mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-04 07:06:36 +08:00
fix lint
This commit is contained in:
@@ -5,21 +5,30 @@ import (
|
||||
"net"
|
||||
)
|
||||
|
||||
// EmptyProvider is a no-op GeoIP backend used before a real provider is configured.
|
||||
type EmptyProvider struct{}
|
||||
|
||||
// Name returns the provider identifier.
|
||||
func (e *EmptyProvider) Name() string {
|
||||
return "EmptyProvider"
|
||||
}
|
||||
|
||||
// Initialize prepares the empty provider for use.
|
||||
func (e *EmptyProvider) Initialize() error {
|
||||
return nil
|
||||
}
|
||||
func (e *EmptyProvider) GetGeoInfo(ip net.IP) (*GeoInfo, error) {
|
||||
|
||||
// GetGeoInfo reports that no GeoIP provider has been configured.
|
||||
func (e *EmptyProvider) GetGeoInfo(_ net.IP) (*GeoInfo, error) {
|
||||
return nil, fmt.Errorf("you are using an empty GeoIP provider, please set a valid provider")
|
||||
}
|
||||
|
||||
// UpdateDatabase reports that no GeoIP provider has been configured.
|
||||
func (e *EmptyProvider) UpdateDatabase() error {
|
||||
return fmt.Errorf("you are using an empty GeoIP provider, please set a valid provider")
|
||||
}
|
||||
|
||||
// Close releases resources held by the empty provider.
|
||||
func (e *EmptyProvider) Close() error {
|
||||
return nil
|
||||
}
|
||||
|
||||
+29
-12
@@ -1,3 +1,4 @@
|
||||
// Package geoip resolves geographic information for IP addresses.
|
||||
package geoip
|
||||
|
||||
import (
|
||||
@@ -12,19 +13,27 @@ import (
|
||||
ristretto "github.com/dgraph-io/ristretto/v2"
|
||||
)
|
||||
|
||||
var CurrentProvider GeoIPService
|
||||
// CurrentProvider is the active GeoIP backend used by package-level helpers.
|
||||
var CurrentProvider Service
|
||||
var geoCache *providerCache
|
||||
var providerMutex sync.RWMutex
|
||||
var providerFactory = newProvider
|
||||
|
||||
// Supported GeoIP provider identifiers.
|
||||
const (
|
||||
// ProviderDisabled disables GeoIP lookups.
|
||||
ProviderDisabled = "disabled"
|
||||
ProviderMaxMind = "mmdb"
|
||||
ProviderIPAPI = "ip-api"
|
||||
ProviderGeoJS = "geojs"
|
||||
ProviderIPInfo = "ipinfo"
|
||||
|
||||
defaultGeoCacheDuration = 48 * time.Hour
|
||||
isoCountryCodeLength = 2
|
||||
geoipDataDirPerm = 0o750
|
||||
)
|
||||
|
||||
// GeoInfo contains normalized geographic metadata for an IP address.
|
||||
type GeoInfo struct {
|
||||
ISOCode string
|
||||
Name string
|
||||
@@ -34,11 +43,11 @@ type GeoInfo struct {
|
||||
|
||||
func init() {
|
||||
CurrentProvider = &EmptyProvider{}
|
||||
geoCache = newProviderCache(48 * time.Hour)
|
||||
geoCache = newProviderCache(defaultGeoCacheDuration)
|
||||
}
|
||||
|
||||
// GeoIPService 接口定义了获取地理位置信息的核心方法。
|
||||
type GeoIPService interface {
|
||||
// Service defines the core GeoIP lookup operations implemented by providers.
|
||||
type Service interface {
|
||||
Name() string
|
||||
GetGeoInfo(ip net.IP) (*GeoInfo, error)
|
||||
UpdateDatabase() error
|
||||
@@ -94,8 +103,9 @@ func (c *providerCache) Flush() {
|
||||
c.items.Clear()
|
||||
}
|
||||
|
||||
// GetRegionUnicodeEmoji returns the regional indicator emoji for a two-letter ISO code.
|
||||
func GetRegionUnicodeEmoji(isoCode string) string {
|
||||
if len(isoCode) != 2 {
|
||||
if len(isoCode) != isoCountryCodeLength {
|
||||
return ""
|
||||
}
|
||||
isoCode = strings.ToUpper(isoCode)
|
||||
@@ -104,11 +114,12 @@ func GetRegionUnicodeEmoji(isoCode string) string {
|
||||
return ""
|
||||
}
|
||||
|
||||
rune1 := rune(0x1F1E6 + (rune(isoCode[0]) - 'A'))
|
||||
rune2 := rune(0x1F1E6 + (rune(isoCode[1]) - 'A'))
|
||||
rune1 := 0x1F1E6 + (rune(isoCode[0]) - 'A')
|
||||
rune2 := 0x1F1E6 + (rune(isoCode[1]) - 'A')
|
||||
return string(rune1) + string(rune2)
|
||||
}
|
||||
|
||||
// InitGeoIP configures the active GeoIP provider.
|
||||
func InitGeoIP(provider string) {
|
||||
providerName := normalizeProvider(provider)
|
||||
nextProvider, err := providerFactory(providerName)
|
||||
@@ -124,6 +135,7 @@ func InitGeoIP(provider string) {
|
||||
slog.Info("GeoIP provider configured", "provider", CurrentProvider.Name())
|
||||
}
|
||||
|
||||
// GetGeoInfo looks up geographic information for ip using the active provider.
|
||||
func GetGeoInfo(ip net.IP) (*GeoInfo, error) {
|
||||
if ip == nil {
|
||||
return nil, fmt.Errorf("IP address cannot be nil")
|
||||
@@ -142,6 +154,7 @@ func GetGeoInfo(ip net.IP) (*GeoInfo, error) {
|
||||
return info, err
|
||||
}
|
||||
|
||||
// LookupGeoInfoWithProvider looks up geographic information using a temporary provider.
|
||||
func LookupGeoInfoWithProvider(providerName string, ip net.IP) (*GeoInfo, error) {
|
||||
if ip == nil {
|
||||
return nil, fmt.Errorf("IP address cannot be nil")
|
||||
@@ -160,6 +173,7 @@ func LookupGeoInfoWithProvider(providerName string, ip net.IP) (*GeoInfo, error)
|
||||
return provider.GetGeoInfo(ip)
|
||||
}
|
||||
|
||||
// UpdateDatabase refreshes the active provider database and clears cached lookups.
|
||||
func UpdateDatabase() error {
|
||||
err := getProvider().UpdateDatabase()
|
||||
if err == nil {
|
||||
@@ -169,6 +183,7 @@ func UpdateDatabase() error {
|
||||
return err
|
||||
}
|
||||
|
||||
// IsValidProvider reports whether provider names a supported GeoIP backend.
|
||||
func IsValidProvider(provider string) bool {
|
||||
switch normalizeProvider(provider) {
|
||||
case ProviderDisabled, ProviderMaxMind, ProviderIPAPI, ProviderGeoJS, ProviderIPInfo:
|
||||
@@ -186,7 +201,7 @@ func normalizeProvider(provider string) string {
|
||||
return normalized
|
||||
}
|
||||
|
||||
func newProvider(provider string) (GeoIPService, error) {
|
||||
func newProvider(provider string) (Service, error) {
|
||||
switch provider {
|
||||
case ProviderDisabled:
|
||||
return &EmptyProvider{}, nil
|
||||
@@ -203,7 +218,7 @@ func newProvider(provider string) (GeoIPService, error) {
|
||||
}
|
||||
}
|
||||
|
||||
func setProvider(provider GeoIPService) {
|
||||
func setProvider(provider Service) {
|
||||
providerMutex.Lock()
|
||||
previous := CurrentProvider
|
||||
CurrentProvider = provider
|
||||
@@ -216,7 +231,7 @@ func setProvider(provider GeoIPService) {
|
||||
}
|
||||
}
|
||||
|
||||
func getProvider() GeoIPService {
|
||||
func getProvider() Service {
|
||||
providerMutex.RLock()
|
||||
defer providerMutex.RUnlock()
|
||||
if CurrentProvider == nil {
|
||||
@@ -229,11 +244,13 @@ func float64Pointer(value float64) *float64 {
|
||||
return &value
|
||||
}
|
||||
|
||||
func ProviderFactoryForTest() func(string) (GeoIPService, error) {
|
||||
// ProviderFactoryForTest returns the provider factory used by package helpers.
|
||||
func ProviderFactoryForTest() func(string) (Service, error) {
|
||||
return providerFactory
|
||||
}
|
||||
|
||||
func SetProviderFactoryForTest(factory func(string) (GeoIPService, error)) {
|
||||
// SetProviderFactoryForTest replaces the provider factory for tests.
|
||||
func SetProviderFactoryForTest(factory func(string) (Service, error)) {
|
||||
if factory == nil {
|
||||
providerFactory = newProvider
|
||||
return
|
||||
|
||||
@@ -82,7 +82,7 @@ func TestIsValidProvider(t *testing.T) {
|
||||
|
||||
func TestLookupGeoInfoWithProviderUsesTemporaryProvider(t *testing.T) {
|
||||
previousFactory := providerFactory
|
||||
providerFactory = func(provider string) (GeoIPService, error) {
|
||||
providerFactory = func(provider string) (Service, error) {
|
||||
return &fakeProvider{}, nil
|
||||
}
|
||||
defer func() {
|
||||
|
||||
+8
-3
@@ -1,6 +1,7 @@
|
||||
package geoip
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net"
|
||||
@@ -8,7 +9,7 @@ import (
|
||||
"time"
|
||||
)
|
||||
|
||||
// GeoJSService 使用 geojs.io 服务实现 GeoIPService 接口。
|
||||
// GeoJSService resolves geographic information using the geojs.io service.
|
||||
type GeoJSService struct {
|
||||
Client *http.Client
|
||||
}
|
||||
@@ -44,11 +45,15 @@ func (s *GeoJSService) GetGeoInfo(ip net.IP) (*GeoInfo, error) {
|
||||
// GeoJS 的 API 端点
|
||||
apiURL := fmt.Sprintf("https://get.geojs.io/v1/ip/geo/%s.json", ip.String())
|
||||
|
||||
resp, err := s.Client.Get(apiURL)
|
||||
req, err := http.NewRequestWithContext(context.Background(), http.MethodGet, apiURL, nil)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create request for geojs.io: %w", err)
|
||||
}
|
||||
resp, err := s.Client.Do(req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to get geo info from geojs.io: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
defer func() { _ = resp.Body.Close() }()
|
||||
|
||||
// 检查响应状态码
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
|
||||
+9
-3
@@ -1,6 +1,7 @@
|
||||
package geoip
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net"
|
||||
@@ -8,7 +9,7 @@ import (
|
||||
"time"
|
||||
)
|
||||
|
||||
// IPAPIService 使用 ip-api.com 服务实现 GeoIPService 接口。
|
||||
// IPAPIService resolves geographic information using the ip-api.com service.
|
||||
type IPAPIService struct {
|
||||
Client *http.Client
|
||||
}
|
||||
@@ -32,6 +33,7 @@ type ipAPIResponse struct {
|
||||
Query string `json:"query"`
|
||||
}
|
||||
|
||||
// Name returns the provider identifier for the ip-api.com service.
|
||||
func (s *IPAPIService) Name() string {
|
||||
return "ip-api.com"
|
||||
}
|
||||
@@ -50,11 +52,15 @@ func (s *IPAPIService) GetGeoInfo(ip net.IP) (*GeoInfo, error) {
|
||||
// API URL, 使用 fields 参数来仅请求需要的字段
|
||||
apiURL := fmt.Sprintf("http://ip-api.com/json/%s?fields=status,message,country,countryCode", ip.String())
|
||||
|
||||
resp, err := s.Client.Get(apiURL)
|
||||
req, err := http.NewRequestWithContext(context.Background(), http.MethodGet, apiURL, nil)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create request for ip-api.com: %w", err)
|
||||
}
|
||||
resp, err := s.Client.Do(req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to get geo info from ip-api.com: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
defer func() { _ = resp.Body.Close() }()
|
||||
|
||||
var apiResp ipAPIResponse
|
||||
if err := json.NewDecoder(resp.Body).Decode(&apiResp); err != nil {
|
||||
|
||||
+9
-4
@@ -1,6 +1,7 @@
|
||||
package geoip
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net"
|
||||
@@ -10,7 +11,7 @@ import (
|
||||
"time"
|
||||
)
|
||||
|
||||
// IPInfoService 使用 ipinfo.io 服务实现 GeoIPService 接口。
|
||||
// IPInfoService resolves geographic information using the ipinfo.io service.
|
||||
type IPInfoService struct {
|
||||
Client *http.Client
|
||||
// 每天 1000 次请求,限制由 IP 地址的所有人共享。
|
||||
@@ -52,11 +53,15 @@ func (s *IPInfoService) GetGeoInfo(ip net.IP) (*GeoInfo, error) {
|
||||
// API URL: https://ipinfo.io/json (查询自身IP) 或 https://ipinfo.io/YOUR_IP/json
|
||||
apiURL := fmt.Sprintf("https://ipinfo.io/%s/json", ip.String())
|
||||
|
||||
resp, err := s.Client.Get(apiURL)
|
||||
req, err := http.NewRequestWithContext(context.Background(), http.MethodGet, apiURL, nil)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create request for ipinfo.io: %w", err)
|
||||
}
|
||||
resp, err := s.Client.Do(req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to get geo info from ipinfo.io: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
defer func() { _ = resp.Body.Close() }()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return nil, fmt.Errorf("ipinfo.io returned non-200 status: %d %s", resp.StatusCode, resp.Status)
|
||||
@@ -98,7 +103,7 @@ func (s *IPInfoService) Close() error {
|
||||
|
||||
func parseIPInfoCoordinates(value string) (*float64, *float64) {
|
||||
parts := strings.Split(strings.TrimSpace(value), ",")
|
||||
if len(parts) != 2 {
|
||||
if len(parts) != isoCountryCodeLength {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
// Package iputil provides helpers for parsing, normalizing, and scoring IP addresses.
|
||||
package iputil
|
||||
|
||||
import (
|
||||
@@ -5,6 +6,12 @@ import (
|
||||
"strings"
|
||||
)
|
||||
|
||||
const (
|
||||
scorePublic = 2
|
||||
scorePrivate = 1
|
||||
)
|
||||
|
||||
// NormalizeIP parses and normalizes a raw IP address string, preferring the IPv4 form for mapped addresses.
|
||||
func NormalizeIP(raw string) string {
|
||||
trimmed := strings.TrimSpace(raw)
|
||||
if trimmed == "" {
|
||||
@@ -20,6 +27,7 @@ func NormalizeIP(raw string) string {
|
||||
return ip.String()
|
||||
}
|
||||
|
||||
// NormalizeRemoteAddr extracts and normalizes the IP address from a host:port remote address string.
|
||||
func NormalizeRemoteAddr(remoteAddr string) string {
|
||||
trimmed := strings.TrimSpace(remoteAddr)
|
||||
if trimmed == "" {
|
||||
@@ -31,6 +39,7 @@ func NormalizeRemoteAddr(remoteAddr string) string {
|
||||
return NormalizeIP(trimmed)
|
||||
}
|
||||
|
||||
// IsPublic reports whether the given IP address is a publicly routable unicast address.
|
||||
func IsPublic(ip net.IP) bool {
|
||||
if ip == nil {
|
||||
return false
|
||||
@@ -44,11 +53,13 @@ func IsPublic(ip net.IP) bool {
|
||||
return true
|
||||
}
|
||||
|
||||
// IsPublicString parses raw and reports whether it represents a publicly routable IP address.
|
||||
func IsPublicString(raw string) bool {
|
||||
ip := net.ParseIP(strings.TrimSpace(raw))
|
||||
return IsPublic(ip)
|
||||
}
|
||||
|
||||
// Score returns a preference score for the IP address: 2 for public, 1 for private, -1 for invalid or non-unicast.
|
||||
func Score(ip net.IP) int {
|
||||
if ip == nil {
|
||||
return -1
|
||||
@@ -60,7 +71,7 @@ func Score(ip net.IP) int {
|
||||
return -1
|
||||
}
|
||||
if IsPublic(ip) {
|
||||
return 2
|
||||
return scorePublic
|
||||
}
|
||||
return 1
|
||||
return scorePrivate
|
||||
}
|
||||
|
||||
+34
-17
@@ -1,6 +1,7 @@
|
||||
package geoip
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
@@ -12,47 +13,55 @@ import (
|
||||
"github.com/oschwald/maxminddb-golang"
|
||||
)
|
||||
|
||||
var GeoIpUrl = "https://raw.githubusercontent.com/Loyalsoldier/geoip/release/GeoLite2-Country.mmdb"
|
||||
var GeoIpFilePath = "./data/GeoLite2-Country.mmdb"
|
||||
// GeoIPURL is the default download URL for the MaxMind GeoLite2 country database.
|
||||
var GeoIPURL = "https://raw.githubusercontent.com/Loyalsoldier/geoip/release/GeoLite2-Country.mmdb"
|
||||
|
||||
type GeoIpRecord struct {
|
||||
// GeoIPFilePath is the default local path for the MaxMind country database file.
|
||||
var GeoIPFilePath = "./data/GeoLite2-Country.mmdb"
|
||||
|
||||
// Record is the MaxMind database record structure for country lookups.
|
||||
type Record struct {
|
||||
Country struct {
|
||||
ISOCode string `maxminddb:"iso_code"`
|
||||
Names map[string]string `maxminddb:"names"`
|
||||
} `maxminddb:"country"`
|
||||
}
|
||||
|
||||
// MaxMindGeoIPService resolves geographic information using a local MaxMind MMDB database.
|
||||
type MaxMindGeoIPService struct {
|
||||
maxMindDBReader *maxminddb.Reader
|
||||
dbFilePath string
|
||||
mu sync.RWMutex
|
||||
}
|
||||
|
||||
// Name returns the provider identifier for the MaxMind database service.
|
||||
func (s *MaxMindGeoIPService) Name() string {
|
||||
return "MaxMind"
|
||||
}
|
||||
|
||||
// NewMaxMindGeoIPService creates a MaxMind service using the default database path and URL.
|
||||
func NewMaxMindGeoIPService() (*MaxMindGeoIPService, error) {
|
||||
return NewMaxMindGeoIPServiceWithConfig(GeoIpFilePath, GeoIpUrl)
|
||||
return NewMaxMindGeoIPServiceWithConfig(GeoIPFilePath, GeoIPURL)
|
||||
}
|
||||
|
||||
// NewMaxMindGeoIPServiceWithConfig creates a MaxMind service with custom database path and download URL.
|
||||
func NewMaxMindGeoIPServiceWithConfig(dbFilePath string, downloadURL string) (*MaxMindGeoIPService, error) {
|
||||
if dbFilePath == "" {
|
||||
dbFilePath = GeoIpFilePath
|
||||
dbFilePath = GeoIPFilePath
|
||||
}
|
||||
if downloadURL == "" {
|
||||
downloadURL = GeoIpUrl
|
||||
downloadURL = GeoIPURL
|
||||
}
|
||||
service := &MaxMindGeoIPService{
|
||||
dbFilePath: dbFilePath,
|
||||
}
|
||||
|
||||
if err := os.MkdirAll(filepath.Dir(service.dbFilePath), os.ModePerm); err != nil {
|
||||
if err := os.MkdirAll(filepath.Dir(service.dbFilePath), geoipDataDirPerm); err != nil {
|
||||
return nil, fmt.Errorf("failed to create data directory for MaxMind database: %w", err)
|
||||
}
|
||||
|
||||
if _, err := os.Stat(service.dbFilePath); os.IsNotExist(err) {
|
||||
if err := DownloadMaxMindDatabase(service.dbFilePath, downloadURL); err != nil {
|
||||
if err := DownloadMaxMindDatabase(context.Background(), service.dbFilePath, downloadURL); err != nil {
|
||||
return nil, fmt.Errorf("failed to download initial MaxMind database: %w", err)
|
||||
}
|
||||
}
|
||||
@@ -81,6 +90,7 @@ func (s *MaxMindGeoIPService) initialize() error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetGeoInfo looks up geographic information for ip in the MaxMind database.
|
||||
func (s *MaxMindGeoIPService) GetGeoInfo(ip net.IP) (*GeoInfo, error) {
|
||||
s.mu.RLock()
|
||||
defer s.mu.RUnlock()
|
||||
@@ -92,7 +102,7 @@ func (s *MaxMindGeoIPService) GetGeoInfo(ip net.IP) (*GeoInfo, error) {
|
||||
return nil, fmt.Errorf("IP address cannot be nil")
|
||||
}
|
||||
|
||||
var record GeoIpRecord
|
||||
var record Record
|
||||
if err := s.maxMindDBReader.Lookup(ip, &record); err != nil {
|
||||
return nil, fmt.Errorf("error looking up IP %s in MaxMind database: %w", ip.String(), err)
|
||||
}
|
||||
@@ -108,36 +118,42 @@ func (s *MaxMindGeoIPService) GetGeoInfo(ip net.IP) (*GeoInfo, error) {
|
||||
return geoInfo, nil
|
||||
}
|
||||
|
||||
// UpdateDatabase downloads the latest MaxMind database and reloads the reader.
|
||||
func (s *MaxMindGeoIPService) UpdateDatabase() error {
|
||||
if err := DownloadMaxMindDatabase(s.dbFilePath, GeoIpUrl); err != nil {
|
||||
if err := DownloadMaxMindDatabase(context.Background(), s.dbFilePath, GeoIPURL); err != nil {
|
||||
return err
|
||||
}
|
||||
return s.initialize()
|
||||
}
|
||||
|
||||
func DownloadMaxMindDatabase(dbFilePath string, downloadURL string) error {
|
||||
// DownloadMaxMindDatabase downloads the MaxMind database from downloadURL to dbFilePath.
|
||||
func DownloadMaxMindDatabase(ctx context.Context, dbFilePath string, downloadURL string) error {
|
||||
if dbFilePath == "" {
|
||||
dbFilePath = GeoIpFilePath
|
||||
dbFilePath = GeoIPFilePath
|
||||
}
|
||||
if downloadURL == "" {
|
||||
downloadURL = GeoIpUrl
|
||||
downloadURL = GeoIPURL
|
||||
}
|
||||
resp, err := http.Get(downloadURL)
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, downloadURL, nil) //nolint:gosec // URL from trusted GeoIP provider config
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to initiate MaxMind database download: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
resp, err := http.DefaultClient.Do(req)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to initiate MaxMind database download: %w", err)
|
||||
}
|
||||
defer func() { _ = resp.Body.Close() }()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return fmt.Errorf("failed to download MaxMind database: HTTP status %s", resp.Status)
|
||||
}
|
||||
|
||||
if err := os.MkdirAll(filepath.Dir(dbFilePath), os.ModePerm); err != nil {
|
||||
if err := os.MkdirAll(filepath.Dir(dbFilePath), geoipDataDirPerm); err != nil {
|
||||
return fmt.Errorf("failed to create data directory for MaxMind database update: %w", err)
|
||||
}
|
||||
|
||||
tempPath := dbFilePath + ".download"
|
||||
out, err := os.Create(tempPath)
|
||||
out, err := os.Create(tempPath) //nolint:gosec // tempPath is derived from configured dbFilePath
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to create MaxMind database file at %s: %w", tempPath, err)
|
||||
}
|
||||
@@ -157,6 +173,7 @@ func DownloadMaxMindDatabase(dbFilePath string, downloadURL string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// Close closes the MaxMind database reader.
|
||||
func (s *MaxMindGeoIPService) Close() error {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
|
||||
+13
-2
@@ -29,11 +29,13 @@ type OutboundIPAPIAdapter interface {
|
||||
DecodeIP(io.Reader) (net.IP, error)
|
||||
}
|
||||
|
||||
// HTTPOutboundIPStrategy resolves the public egress IP via an HTTP API adapter.
|
||||
type HTTPOutboundIPStrategy struct {
|
||||
Client *http.Client
|
||||
Adapter OutboundIPAPIAdapter
|
||||
}
|
||||
|
||||
// NewHTTPOutboundIPStrategy creates a strategy that queries adapter over HTTP.
|
||||
func NewHTTPOutboundIPStrategy(adapter OutboundIPAPIAdapter, client *http.Client) *HTTPOutboundIPStrategy {
|
||||
if client == nil {
|
||||
client = &http.Client{Timeout: defaultOutboundIPLookupTimeout}
|
||||
@@ -44,6 +46,7 @@ func NewHTTPOutboundIPStrategy(adapter OutboundIPAPIAdapter, client *http.Client
|
||||
}
|
||||
}
|
||||
|
||||
// Name returns the strategy or adapter identifier.
|
||||
func (s *HTTPOutboundIPStrategy) Name() string {
|
||||
if s == nil || s.Adapter == nil {
|
||||
return "http-outbound-ip"
|
||||
@@ -51,12 +54,13 @@ func (s *HTTPOutboundIPStrategy) Name() string {
|
||||
return s.Adapter.Name()
|
||||
}
|
||||
|
||||
// GetOutboundIP queries the configured HTTP endpoint for the current public IP.
|
||||
func (s *HTTPOutboundIPStrategy) GetOutboundIP(ctx context.Context) (net.IP, error) {
|
||||
if s == nil || s.Adapter == nil {
|
||||
return nil, errors.New("outbound IP adapter is nil")
|
||||
}
|
||||
if ctx == nil {
|
||||
ctx = context.Background()
|
||||
return nil, errors.New("context is required")
|
||||
}
|
||||
client := s.Client
|
||||
if client == nil {
|
||||
@@ -70,7 +74,7 @@ func (s *HTTPOutboundIPStrategy) GetOutboundIP(ctx context.Context) (net.IP, err
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%s request failed: %w", s.Name(), err)
|
||||
}
|
||||
defer response.Body.Close()
|
||||
defer func() { _ = response.Body.Close() }()
|
||||
if response.StatusCode != http.StatusOK {
|
||||
return nil, fmt.Errorf("%s returned non-200 status: %d %s", s.Name(), response.StatusCode, response.Status)
|
||||
}
|
||||
@@ -84,6 +88,7 @@ func (s *HTTPOutboundIPStrategy) GetOutboundIP(ctx context.Context) (net.IP, err
|
||||
return ip, nil
|
||||
}
|
||||
|
||||
// RealIPCCAdapter decodes public IP responses from realip.cc.
|
||||
type RealIPCCAdapter struct {
|
||||
URL string
|
||||
}
|
||||
@@ -92,14 +97,17 @@ type realIPCCResponse struct {
|
||||
IP string `json:"ip"`
|
||||
}
|
||||
|
||||
// NewRealIPCCOutboundIPStrategy creates the default realip.cc lookup strategy.
|
||||
func NewRealIPCCOutboundIPStrategy() *HTTPOutboundIPStrategy {
|
||||
return NewHTTPOutboundIPStrategy(RealIPCCAdapter{}, nil)
|
||||
}
|
||||
|
||||
// Name returns the realip.cc adapter identifier.
|
||||
func (a RealIPCCAdapter) Name() string {
|
||||
return "realip.cc"
|
||||
}
|
||||
|
||||
// Endpoint returns the realip.cc API URL.
|
||||
func (a RealIPCCAdapter) Endpoint() string {
|
||||
if strings.TrimSpace(a.URL) != "" {
|
||||
return strings.TrimSpace(a.URL)
|
||||
@@ -107,6 +115,7 @@ func (a RealIPCCAdapter) Endpoint() string {
|
||||
return "https://realip.cc"
|
||||
}
|
||||
|
||||
// DecodeIP parses a realip.cc JSON response into a public IP address.
|
||||
func (a RealIPCCAdapter) DecodeIP(reader io.Reader) (net.IP, error) {
|
||||
var payload realIPCCResponse
|
||||
if err := json.NewDecoder(reader).Decode(&payload); err != nil {
|
||||
@@ -122,12 +131,14 @@ func (a RealIPCCAdapter) DecodeIP(reader io.Reader) (net.IP, error) {
|
||||
return ip, nil
|
||||
}
|
||||
|
||||
// DefaultOutboundIPStrategies returns the built-in public egress IP lookup strategies.
|
||||
func DefaultOutboundIPStrategies() []OutboundIPStrategy {
|
||||
return []OutboundIPStrategy{
|
||||
NewRealIPCCOutboundIPStrategy(),
|
||||
}
|
||||
}
|
||||
|
||||
// GetOutboundIP tries each strategy until one returns a public egress IP.
|
||||
func GetOutboundIP(ctx context.Context, strategies ...OutboundIPStrategy) (net.IP, error) {
|
||||
if len(strategies) == 0 {
|
||||
strategies = DefaultOutboundIPStrategies()
|
||||
|
||||
+25
-1
@@ -1,24 +1,29 @@
|
||||
// Package protocol defines the communication protocol between OpenFlare server, agent, and relay components.
|
||||
package protocol
|
||||
|
||||
import "encoding/json"
|
||||
|
||||
// APIResponse is a generic API response wrapper.
|
||||
type APIResponse[T any] struct {
|
||||
ErrorMsg string `json:"error_msg"`
|
||||
Data T `json:"data"`
|
||||
}
|
||||
|
||||
// HeartbeatData is the heartbeat request payload from agent.
|
||||
type HeartbeatData struct {
|
||||
AgentSettings *AgentSettings `json:"agent_settings"`
|
||||
ActiveConfig *ActiveConfigMeta `json:"active_config"`
|
||||
WAFIPGroups []WAFIPGroup `json:"waf_ip_groups,omitempty"`
|
||||
}
|
||||
|
||||
// HeartbeatResult is the heartbeat response payload.
|
||||
type HeartbeatResult struct {
|
||||
AgentSettings *AgentSettings
|
||||
ActiveConfig *ActiveConfigMeta
|
||||
WAFIPGroups []WAFIPGroup
|
||||
}
|
||||
|
||||
// AgentSettings holds agent configuration settings.
|
||||
type AgentSettings struct {
|
||||
HeartbeatInterval int `json:"heartbeat_interval"`
|
||||
WebsocketUpgradeEnabled bool `json:"websocket_upgrade_enabled"`
|
||||
@@ -30,6 +35,7 @@ type AgentSettings struct {
|
||||
RestartOpenrestyNow bool `json:"restart_openresty_now"`
|
||||
}
|
||||
|
||||
// WSMessageType constants define WebSocket message types.
|
||||
const (
|
||||
WSMessageTypeStatus = "status"
|
||||
WSMessageTypeSettings = "settings"
|
||||
@@ -40,16 +46,19 @@ const (
|
||||
WSMessageTypePong = "pong"
|
||||
)
|
||||
|
||||
// WSMessage represents a WebSocket message.
|
||||
type WSMessage struct {
|
||||
Type string `json:"type"`
|
||||
Payload json.RawMessage `json:"payload,omitempty"`
|
||||
}
|
||||
|
||||
// WSOutboundMessage represents an outbound WebSocket message.
|
||||
type WSOutboundMessage struct {
|
||||
Type string `json:"type"`
|
||||
Payload any `json:"payload,omitempty"`
|
||||
}
|
||||
|
||||
// WebSocketConnection defines the WebSocket connection interface.
|
||||
type WebSocketConnection interface {
|
||||
URL() string
|
||||
SendStatus(payload NodePayload) error
|
||||
@@ -58,12 +67,14 @@ type WebSocketConnection interface {
|
||||
Close() error
|
||||
}
|
||||
|
||||
// OpenrestyStatus constants define OpenResty health status values.
|
||||
const (
|
||||
OpenrestyStatusHealthy = "healthy"
|
||||
OpenrestyStatusUnhealthy = "unhealthy"
|
||||
OpenrestyStatusUnknown = "unknown"
|
||||
)
|
||||
|
||||
// NodePayload is the agent node registration payload.
|
||||
type NodePayload struct {
|
||||
NodeID string `json:"node_id"`
|
||||
Name string `json:"name"`
|
||||
@@ -84,6 +95,7 @@ type NodePayload struct {
|
||||
WAFIPGroupChecksums map[string]string `json:"waf_ip_group_checksums,omitempty"`
|
||||
}
|
||||
|
||||
// NodeSystemProfile describes the system profile of a node.
|
||||
type NodeSystemProfile struct {
|
||||
Hostname string `json:"hostname"`
|
||||
OSName string `json:"os_name"`
|
||||
@@ -98,6 +110,7 @@ type NodeSystemProfile struct {
|
||||
ReportedAtUnix int64 `json:"reported_at_unix"`
|
||||
}
|
||||
|
||||
// NodeMetricSnapshot is a metric snapshot of a node.
|
||||
type NodeMetricSnapshot struct {
|
||||
CapturedAtUnix int64 `json:"captured_at_unix"`
|
||||
CPUUsagePercent float64 `json:"cpu_usage_percent"`
|
||||
@@ -111,6 +124,7 @@ type NodeMetricSnapshot struct {
|
||||
NetworkTxBytes int64 `json:"network_tx_bytes"`
|
||||
}
|
||||
|
||||
// NodeOpenrestyObservation holds OpenResty observation data.
|
||||
type NodeOpenrestyObservation struct {
|
||||
CapturedAtUnix int64 `json:"captured_at_unix"`
|
||||
OpenrestyRxBytes int64 `json:"openresty_rx_bytes"`
|
||||
@@ -118,6 +132,7 @@ type NodeOpenrestyObservation struct {
|
||||
OpenrestyConnections int64 `json:"openresty_connections"`
|
||||
}
|
||||
|
||||
// NodeTrafficReport is a traffic report from agent.
|
||||
type NodeTrafficReport struct {
|
||||
WindowStartedAtUnix int64 `json:"window_started_at_unix"`
|
||||
WindowEndedAtUnix int64 `json:"window_ended_at_unix"`
|
||||
@@ -129,6 +144,7 @@ type NodeTrafficReport struct {
|
||||
SourceCountries map[string]int64 `json:"source_countries"`
|
||||
}
|
||||
|
||||
// NodeAccessLog is an access log entry from agent.
|
||||
type NodeAccessLog struct {
|
||||
LoggedAtUnix int64 `json:"logged_at_unix"`
|
||||
RemoteAddr string `json:"remote_addr"`
|
||||
@@ -137,6 +153,7 @@ type NodeAccessLog struct {
|
||||
StatusCode int `json:"status_code"`
|
||||
}
|
||||
|
||||
// BufferedObservabilityRecord is a buffered observability record.
|
||||
type BufferedObservabilityRecord struct {
|
||||
WindowStartedAtUnix int64 `json:"window_started_at_unix"`
|
||||
Snapshot *NodeMetricSnapshot `json:"snapshot,omitempty"`
|
||||
@@ -145,6 +162,7 @@ type BufferedObservabilityRecord struct {
|
||||
AccessLogs []NodeAccessLog `json:"access_logs,omitempty"`
|
||||
}
|
||||
|
||||
// NodeHealthEvent represents a node health event.
|
||||
type NodeHealthEvent struct {
|
||||
EventType string `json:"event_type"`
|
||||
Severity string `json:"severity"`
|
||||
@@ -153,12 +171,14 @@ type NodeHealthEvent struct {
|
||||
Metadata map[string]string `json:"metadata,omitempty"`
|
||||
}
|
||||
|
||||
// RegisterNodeResponse is the node registration response.
|
||||
type RegisterNodeResponse struct {
|
||||
NodeID string `json:"node_id"`
|
||||
AccessToken string `json:"agent_token"`
|
||||
Name string `json:"name"`
|
||||
}
|
||||
|
||||
// ActiveConfigResponse is the active configuration response.
|
||||
type ActiveConfigResponse struct {
|
||||
Version string `json:"version"`
|
||||
Checksum string `json:"checksum"`
|
||||
@@ -167,6 +187,7 @@ type ActiveConfigResponse struct {
|
||||
CreatedAt string `json:"created_at"`
|
||||
}
|
||||
|
||||
// WAFIPGroup defines a WAF IP group.
|
||||
type WAFIPGroup struct {
|
||||
ID uint `json:"id"`
|
||||
Name string `json:"name"`
|
||||
@@ -176,16 +197,19 @@ type WAFIPGroup struct {
|
||||
Checksum string `json:"checksum"`
|
||||
}
|
||||
|
||||
// WAFIPGroupSyncRequest is a WAF IP group sync request.
|
||||
type WAFIPGroupSyncRequest struct {
|
||||
IDs []uint `json:"ids,omitempty"`
|
||||
Checksums map[string]string `json:"checksums,omitempty"`
|
||||
}
|
||||
|
||||
// WAFIPGroupSyncResponse is a WAF IP group sync response.
|
||||
type WAFIPGroupSyncResponse struct {
|
||||
Groups []WAFIPGroup `json:"groups"`
|
||||
}
|
||||
|
||||
// SupportFile represents a support file for relay.
|
||||
type SupportFile struct {
|
||||
Path string `json:"path"`
|
||||
Content string `json:"content"`
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,9 +1,15 @@
|
||||
package protocol
|
||||
|
||||
// AgentNodeSystemProfile is an alias for NodeSystemProfile used by server.
|
||||
type AgentNodeSystemProfile = NodeSystemProfile
|
||||
|
||||
// AgentNodeMetricSnapshot is an alias for NodeMetricSnapshot used by server.
|
||||
type AgentNodeMetricSnapshot = NodeMetricSnapshot
|
||||
|
||||
// AgentNodeHealthEvent is an alias for NodeHealthEvent used by server.
|
||||
type AgentNodeHealthEvent = NodeHealthEvent
|
||||
|
||||
// RelayProxyStat holds relay proxy statistics.
|
||||
type RelayProxyStat struct {
|
||||
Name string `json:"name"`
|
||||
Type string `json:"type"`
|
||||
@@ -14,6 +20,7 @@ type RelayProxyStat struct {
|
||||
ClientAddr string `json:"client_addr"`
|
||||
}
|
||||
|
||||
// RelayHeartbeatPayload is the relay heartbeat payload.
|
||||
type RelayHeartbeatPayload struct {
|
||||
Version string `json:"version"`
|
||||
ExtVersion string `json:"frp_version"`
|
||||
@@ -29,6 +36,7 @@ type RelayHeartbeatPayload struct {
|
||||
HealthEvents []AgentNodeHealthEvent `json:"health_events,omitempty"`
|
||||
}
|
||||
|
||||
// RelayConfig holds relay configuration.
|
||||
type RelayConfig struct {
|
||||
BindPort int `json:"bind_port"`
|
||||
VhostHTTPPort int `json:"vhost_http_port"`
|
||||
@@ -37,6 +45,7 @@ type RelayConfig struct {
|
||||
WebServerEnabled bool `json:"web_server_enabled"`
|
||||
}
|
||||
|
||||
// RelaySettings holds relay runtime settings.
|
||||
type RelaySettings struct {
|
||||
HeartbeatInterval int `json:"heartbeat_interval"`
|
||||
WebsocketUpgradeEnabled bool `json:"websocket_upgrade_enabled"`
|
||||
@@ -47,22 +56,26 @@ type RelaySettings struct {
|
||||
UpdateTag string `json:"update_tag"`
|
||||
}
|
||||
|
||||
// RelayHeartbeatResponse is the relay heartbeat response.
|
||||
type RelayHeartbeatResponse struct {
|
||||
RelayConfig *RelayConfig `json:"relay_config"`
|
||||
RelaySettings *RelaySettings `json:"relay_settings"`
|
||||
}
|
||||
|
||||
// ActiveConfigMeta holds active configuration metadata.
|
||||
type ActiveConfigMeta struct {
|
||||
Version string `json:"version"`
|
||||
Checksum string `json:"checksum"`
|
||||
}
|
||||
|
||||
// FlaredConnectedRelay describes a connected relay info for flared.
|
||||
type FlaredConnectedRelay struct {
|
||||
RelayNodeID string `json:"relay_node_id"`
|
||||
Status string `json:"status"`
|
||||
ProxyCount int `json:"proxy_count"`
|
||||
}
|
||||
|
||||
// FlaredHeartbeatPayload is the flared heartbeat payload.
|
||||
type FlaredHeartbeatPayload struct {
|
||||
ClientVersion string `json:"client_version"`
|
||||
FrpVersion string `json:"frp_version"`
|
||||
@@ -73,11 +86,13 @@ type FlaredHeartbeatPayload struct {
|
||||
CurrentChecksum string `json:"current_checksum"`
|
||||
}
|
||||
|
||||
// FlaredHeartbeatResponse is the flared heartbeat response.
|
||||
type FlaredHeartbeatResponse struct {
|
||||
ActiveConfig *ActiveConfigMeta `json:"active_config"`
|
||||
TunnelSettings *RelaySettings `json:"tunnel_settings"`
|
||||
}
|
||||
|
||||
// FlaredTunnelConfigResponse is the flared tunnel configuration response.
|
||||
type FlaredTunnelConfigResponse struct {
|
||||
Version string `json:"version"`
|
||||
Checksum string `json:"checksum"`
|
||||
@@ -85,6 +100,7 @@ type FlaredTunnelConfigResponse struct {
|
||||
Proxies []FlaredProxyEntry `json:"proxies"`
|
||||
}
|
||||
|
||||
// FlaredRelayInfo holds flared relay information.
|
||||
type FlaredRelayInfo struct {
|
||||
RelayNodeID string `json:"relay_node_id"`
|
||||
Address string `json:"address"`
|
||||
@@ -92,6 +108,7 @@ type FlaredRelayInfo struct {
|
||||
ProxyURL string `json:"proxy_url"`
|
||||
}
|
||||
|
||||
// FlaredProxyEntry represents a flared proxy entry.
|
||||
type FlaredProxyEntry struct {
|
||||
Name string `json:"name"`
|
||||
Type string `json:"type"`
|
||||
@@ -100,6 +117,7 @@ type FlaredProxyEntry struct {
|
||||
CustomDomains []string `json:"custom_domains"`
|
||||
}
|
||||
|
||||
// ApplyLogPayload is the apply log payload for flared.
|
||||
type ApplyLogPayload struct {
|
||||
NodeID string `json:"node_id"`
|
||||
Version string `json:"version"`
|
||||
|
||||
+51
-126
@@ -1,3 +1,4 @@
|
||||
// Package openresty renders OpenResty configuration from proxy route definitions.
|
||||
package openresty
|
||||
|
||||
import (
|
||||
@@ -16,6 +17,13 @@ import (
|
||||
"strings"
|
||||
)
|
||||
|
||||
const (
|
||||
routeUpstreamTypePages = "pages"
|
||||
indexHTML = "/index.html"
|
||||
)
|
||||
|
||||
// RenderJSON parses the given JSON string as a Document and renders the full
|
||||
// OpenResty configuration bundle, injecting the provided certificate support files.
|
||||
func RenderJSON(sourceJSON string, certificateFiles []SupportFile) (*Result, error) {
|
||||
var doc Document
|
||||
if err := json.Unmarshal([]byte(strings.TrimSpace(sourceJSON)), &doc); err != nil {
|
||||
@@ -24,6 +32,8 @@ func RenderJSON(sourceJSON string, certificateFiles []SupportFile) (*Result, err
|
||||
return Render(doc, certificateFiles)
|
||||
}
|
||||
|
||||
// Render produces a complete OpenResty configuration Result from a Document and
|
||||
// a set of certificate support files.
|
||||
func Render(doc Document, certificateFiles []SupportFile) (*Result, error) {
|
||||
mainConfig := RenderMainConfig(doc.OpenRestyConfig)
|
||||
routeConfig, err := RenderRouteConfig(doc, certificateFiles)
|
||||
@@ -45,6 +55,8 @@ func Render(doc Document, certificateFiles []SupportFile) (*Result, error) {
|
||||
}, nil
|
||||
}
|
||||
|
||||
// RenderMainConfig renders the nginx main configuration string from the given
|
||||
// ConfigSnapshot, falling back to the built-in default template when none is set.
|
||||
func RenderMainConfig(cfg ConfigSnapshot) string {
|
||||
templateText := cfg.MainConfigTemplate
|
||||
if strings.TrimSpace(templateText) == "" {
|
||||
@@ -53,6 +65,8 @@ func RenderMainConfig(cfg ConfigSnapshot) string {
|
||||
return renderMainConfigTemplate(templateText, cfg)
|
||||
}
|
||||
|
||||
// ValidateMainConfigTemplate checks that the provided template text is non-empty
|
||||
// and contains all required OpenResty placeholder tokens.
|
||||
func ValidateMainConfigTemplate(templateText string) error {
|
||||
trimmed := strings.TrimSpace(templateText)
|
||||
if trimmed == "" {
|
||||
@@ -66,6 +80,8 @@ func ValidateMainConfigTemplate(templateText string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// RenderRouteConfig generates the nginx server-block configuration for all
|
||||
// routes in the Document, resolving certificate files as needed.
|
||||
func RenderRouteConfig(doc Document, certificateFiles []SupportFile) (string, error) {
|
||||
var builder strings.Builder
|
||||
builder.WriteString("# This file is generated by OpenFlare. Do not edit manually.\n")
|
||||
@@ -83,120 +99,21 @@ func RenderRouteConfig(doc Document, certificateFiles []SupportFile) (string, er
|
||||
cacheConfig := routeCacheConfig{Enabled: route.CacheEnabled, Policy: route.CachePolicy, Rules: route.CacheRules}
|
||||
limitConfig := routeLimitConfig{LimitConnPerServer: route.LimitConnPerServer, LimitConnPerIP: route.LimitConnPerIP, LimitRate: route.LimitRate}
|
||||
powEnabled, _ := getPoWConfigForRoute(route.ID, doc.WAF)
|
||||
if normalizeRouteUpstreamType(route.UpstreamType) == "pages" {
|
||||
if route.PagesDeployment == nil {
|
||||
return "", fmt.Errorf("route %s pages deployment is missing", route.Domain)
|
||||
}
|
||||
if !route.EnableHTTPS {
|
||||
builder.WriteString(renderHTTPPagesServer(serverNames, displayName, route.PagesDeployment, limitConfig, powEnabled, route.BasicAuthEnabled, route.BasicAuthUsername, route.BasicAuthPassword))
|
||||
continue
|
||||
}
|
||||
certIDs := normalizeCertIDs(route.CertID, route.CertIDs)
|
||||
domainCertIDs := normalizeDomainCertIDs(domains, certIDs, route.DomainCertIDs)
|
||||
if len(certIDs) == 0 {
|
||||
return "", fmt.Errorf("路由 %s 未配置证书", route.Domain)
|
||||
}
|
||||
httpOnlyDomains := make([]string, 0, len(domains))
|
||||
domainsByCertID := make(map[uint][]string, len(certIDs))
|
||||
for index, domain := range domains {
|
||||
if index >= len(domainCertIDs) || domainCertIDs[index] == 0 {
|
||||
httpOnlyDomains = append(httpOnlyDomains, domain)
|
||||
continue
|
||||
}
|
||||
domainsByCertID[domainCertIDs[index]] = append(domainsByCertID[domainCertIDs[index]], domain)
|
||||
}
|
||||
for _, certID := range certIDs {
|
||||
assignedDomains := domainsByCertID[certID]
|
||||
if len(assignedDomains) == 0 {
|
||||
continue
|
||||
}
|
||||
certPEM, ok := certificates[certID]
|
||||
if !ok {
|
||||
return "", fmt.Errorf("route %s certificate %d does not exist", route.Domain, certID)
|
||||
}
|
||||
if err := validateCertificateCoverage(certPEM, assignedDomains); err != nil {
|
||||
return "", fmt.Errorf("site %s certificate validation failed: %w", displayName, err)
|
||||
}
|
||||
}
|
||||
if route.RedirectHTTP {
|
||||
if len(httpOnlyDomains) > 0 {
|
||||
builder.WriteString(renderHTTPPagesServer(renderServerNames(httpOnlyDomains), displayName, route.PagesDeployment, limitConfig, powEnabled, route.BasicAuthEnabled, route.BasicAuthUsername, route.BasicAuthPassword))
|
||||
}
|
||||
for _, certID := range certIDs {
|
||||
if assignedDomains := domainsByCertID[certID]; len(assignedDomains) > 0 {
|
||||
builder.WriteString(renderHTTPRedirectServer(renderServerNames(assignedDomains)))
|
||||
}
|
||||
}
|
||||
} else {
|
||||
builder.WriteString(renderHTTPPagesServer(serverNames, displayName, route.PagesDeployment, limitConfig, powEnabled, route.BasicAuthEnabled, route.BasicAuthUsername, route.BasicAuthPassword))
|
||||
}
|
||||
for _, certID := range certIDs {
|
||||
if assignedDomains := domainsByCertID[certID]; len(assignedDomains) > 0 {
|
||||
builder.WriteString(renderHTTPSPagesServer(renderServerNames(assignedDomains), displayName, certID, route.PagesDeployment, limitConfig, powEnabled, route.BasicAuthEnabled, route.BasicAuthUsername, route.BasicAuthPassword, doc.OpenRestyConfig))
|
||||
}
|
||||
if normalizeRouteUpstreamType(route.UpstreamType) == routeUpstreamTypePages {
|
||||
if err := renderPagesRoute(&builder, route, displayName, serverNames, certificates, limitConfig, powEnabled, doc.OpenRestyConfig); err != nil {
|
||||
return "", err
|
||||
}
|
||||
continue
|
||||
}
|
||||
upstreams := route.Upstreams
|
||||
if len(upstreams) == 0 && strings.TrimSpace(route.OriginURL) != "" {
|
||||
upstreams = []string{route.OriginURL}
|
||||
}
|
||||
upstreamConfig := buildRouteUpstreamConfig(route, upstreams)
|
||||
if upstreamConfig.UsesNamedUpstream {
|
||||
builder.WriteString(renderNamedUpstreamBlock(upstreamConfig))
|
||||
}
|
||||
if !route.EnableHTTPS {
|
||||
builder.WriteString(renderHTTPProxyServer(serverNames, displayName, route.OriginURL, route.OriginHost, route.CustomHeaders, cacheConfig, limitConfig, upstreamConfig, powEnabled, route.BasicAuthEnabled, route.BasicAuthUsername, route.BasicAuthPassword, doc.OpenRestyConfig))
|
||||
continue
|
||||
}
|
||||
certIDs := normalizeCertIDs(route.CertID, route.CertIDs)
|
||||
domainCertIDs := normalizeDomainCertIDs(domains, certIDs, route.DomainCertIDs)
|
||||
if len(certIDs) == 0 {
|
||||
return "", fmt.Errorf("路由 %s 未配置证书", route.Domain)
|
||||
}
|
||||
httpOnlyDomains := make([]string, 0, len(domains))
|
||||
domainsByCertID := make(map[uint][]string, len(certIDs))
|
||||
for index, domain := range domains {
|
||||
if index >= len(domainCertIDs) || domainCertIDs[index] == 0 {
|
||||
httpOnlyDomains = append(httpOnlyDomains, domain)
|
||||
continue
|
||||
}
|
||||
domainsByCertID[domainCertIDs[index]] = append(domainsByCertID[domainCertIDs[index]], domain)
|
||||
}
|
||||
for _, certID := range certIDs {
|
||||
assignedDomains := domainsByCertID[certID]
|
||||
if len(assignedDomains) == 0 {
|
||||
continue
|
||||
}
|
||||
certPEM, ok := certificates[certID]
|
||||
if !ok {
|
||||
return "", fmt.Errorf("route %s certificate %d does not exist", route.Domain, certID)
|
||||
}
|
||||
if err := validateCertificateCoverage(certPEM, assignedDomains); err != nil {
|
||||
return "", fmt.Errorf("site %s certificate validation failed: %w", displayName, err)
|
||||
}
|
||||
}
|
||||
if route.RedirectHTTP {
|
||||
if len(httpOnlyDomains) > 0 {
|
||||
builder.WriteString(renderHTTPProxyServer(renderServerNames(httpOnlyDomains), displayName, route.OriginURL, route.OriginHost, route.CustomHeaders, cacheConfig, limitConfig, upstreamConfig, powEnabled, route.BasicAuthEnabled, route.BasicAuthUsername, route.BasicAuthPassword, doc.OpenRestyConfig))
|
||||
}
|
||||
for _, certID := range certIDs {
|
||||
if assignedDomains := domainsByCertID[certID]; len(assignedDomains) > 0 {
|
||||
builder.WriteString(renderHTTPRedirectServer(renderServerNames(assignedDomains)))
|
||||
}
|
||||
}
|
||||
} else {
|
||||
builder.WriteString(renderHTTPProxyServer(serverNames, displayName, route.OriginURL, route.OriginHost, route.CustomHeaders, cacheConfig, limitConfig, upstreamConfig, powEnabled, route.BasicAuthEnabled, route.BasicAuthUsername, route.BasicAuthPassword, doc.OpenRestyConfig))
|
||||
}
|
||||
for _, certID := range certIDs {
|
||||
if assignedDomains := domainsByCertID[certID]; len(assignedDomains) > 0 {
|
||||
builder.WriteString(renderHTTPSServer(renderServerNames(assignedDomains), displayName, route.OriginURL, route.OriginHost, certID, route.CustomHeaders, cacheConfig, limitConfig, upstreamConfig, powEnabled, route.BasicAuthEnabled, route.BasicAuthUsername, route.BasicAuthPassword, doc.OpenRestyConfig))
|
||||
}
|
||||
if err := renderProxyRoute(&builder, route, displayName, serverNames, certificates, cacheConfig, limitConfig, powEnabled, doc.OpenRestyConfig); err != nil {
|
||||
return "", err
|
||||
}
|
||||
}
|
||||
return builder.String(), nil
|
||||
}
|
||||
|
||||
// RenderPoWConfig serialises the Proof-of-Work configuration for all enabled
|
||||
// routes as a JSON string consumed by the OpenResty Lua runtime.
|
||||
func RenderPoWConfig(doc Document) (string, error) {
|
||||
type domainEntry struct {
|
||||
Domains []string `json:"domains"`
|
||||
@@ -218,6 +135,8 @@ func RenderPoWConfig(doc Document) (string, error) {
|
||||
return string(data), err
|
||||
}
|
||||
|
||||
// RenderWAFConfig serialises the WAF runtime configuration (rule groups and
|
||||
// per-site bindings) as a JSON string consumed by the OpenResty Lua runtime.
|
||||
func RenderWAFConfig(snapshot WAFDocument) (string, error) {
|
||||
type wafRuntimeRuleGroup struct {
|
||||
ID uint `json:"id"`
|
||||
@@ -312,6 +231,9 @@ func sortedUniqueUintIDs(values []uint) []uint {
|
||||
return items
|
||||
}
|
||||
|
||||
// ChecksumBundle returns a stable SHA-256 hex digest over the combined content
|
||||
// of the main config, route config, and deduplicated support files, excluding
|
||||
// the source config JSON file itself.
|
||||
func ChecksumBundle(mainConfig string, routeConfig string, supportFiles []SupportFile) string {
|
||||
var builder strings.Builder
|
||||
builder.WriteString(mainConfig)
|
||||
@@ -333,6 +255,8 @@ func ChecksumBundle(mainConfig string, routeConfig string, supportFiles []Suppor
|
||||
return hex.EncodeToString(sum[:])
|
||||
}
|
||||
|
||||
// DedupeSupportFiles returns a new slice with duplicate paths removed, keeping
|
||||
// the last occurrence of each path.
|
||||
func DedupeSupportFiles(files []SupportFile) []SupportFile {
|
||||
if len(files) == 0 {
|
||||
return nil
|
||||
@@ -437,21 +361,22 @@ func renderPagesAPIProxyLocationBlock(deployment *PagesDeployment) string {
|
||||
cleanPath := strings.TrimSuffix(path, "/")
|
||||
|
||||
var builder strings.Builder
|
||||
builder.WriteString(fmt.Sprintf("\n location %s {\n", cleanPath))
|
||||
// 使用 fmt.Fprintf 替代 WriteString(fmt.Sprintf(...))(QF1012)
|
||||
fmt.Fprintf(&builder, "\n location %s {\n", cleanPath)
|
||||
if rewrite != "" {
|
||||
if !strings.HasPrefix(rewrite, "/") {
|
||||
rewrite = "/" + rewrite
|
||||
}
|
||||
cleanRewrite := strings.TrimSuffix(rewrite, "/")
|
||||
if cleanRewrite == "" {
|
||||
builder.WriteString(fmt.Sprintf(" rewrite ^%s/(.*)$ /$1 break;\n", regexp.QuoteMeta(cleanPath)))
|
||||
builder.WriteString(fmt.Sprintf(" rewrite ^%s$ / break;\n", regexp.QuoteMeta(cleanPath)))
|
||||
fmt.Fprintf(&builder, " rewrite ^%s/(.*)$ /$1 break;\n", regexp.QuoteMeta(cleanPath))
|
||||
fmt.Fprintf(&builder, " rewrite ^%s$ / break;\n", regexp.QuoteMeta(cleanPath))
|
||||
} else {
|
||||
builder.WriteString(fmt.Sprintf(" rewrite ^%s/(.*)$ %s/$1 break;\n", regexp.QuoteMeta(cleanPath), cleanRewrite))
|
||||
builder.WriteString(fmt.Sprintf(" rewrite ^%s$ %s break;\n", regexp.QuoteMeta(cleanPath), cleanRewrite))
|
||||
fmt.Fprintf(&builder, " rewrite ^%s/(.*)$ %s/$1 break;\n", regexp.QuoteMeta(cleanPath), cleanRewrite)
|
||||
fmt.Fprintf(&builder, " rewrite ^%s$ %s break;\n", regexp.QuoteMeta(cleanPath), cleanRewrite)
|
||||
}
|
||||
}
|
||||
builder.WriteString(fmt.Sprintf(" proxy_pass %s;\n", pass))
|
||||
fmt.Fprintf(&builder, " proxy_pass %s;\n", pass)
|
||||
builder.WriteString(" proxy_http_version 1.1;\n")
|
||||
builder.WriteString(" proxy_set_header Host $http_host;\n")
|
||||
builder.WriteString(" proxy_set_header X-Real-IP $remote_addr;\n")
|
||||
@@ -499,7 +424,7 @@ func renderPagesLocationBlock(deployment *PagesDeployment, limitConfig routeLimi
|
||||
var builder strings.Builder
|
||||
builder.WriteString(renderRouteLimitBlock(limitConfig))
|
||||
if deployment != nil && deployment.SPAFallbackEnabled {
|
||||
builder.WriteString(fmt.Sprintf(" try_files $uri $uri/ %s;\n", pagesFallbackPath(deployment)))
|
||||
fmt.Fprintf(&builder, " try_files $uri $uri/ %s;\n", pagesFallbackPath(deployment))
|
||||
} else {
|
||||
builder.WriteString(" try_files $uri $uri/ =404;\n")
|
||||
}
|
||||
@@ -522,18 +447,18 @@ func pagesEntryFile(deployment *PagesDeployment) string {
|
||||
|
||||
func pagesFallbackPath(deployment *PagesDeployment) string {
|
||||
if deployment == nil || strings.TrimSpace(deployment.SPAFallbackPath) == "" {
|
||||
return "/index.html"
|
||||
return indexHTML
|
||||
}
|
||||
value := filepathToNginxPath(strings.TrimSpace(deployment.SPAFallbackPath))
|
||||
if !strings.HasPrefix(value, "/") {
|
||||
value = "/" + value
|
||||
}
|
||||
if value == "/" || strings.HasSuffix(value, "/") || strings.Contains(value, "\\") || strings.ContainsAny(value, "\"';") || strings.ContainsAny(value, " \t\r\n") {
|
||||
return "/index.html"
|
||||
return indexHTML
|
||||
}
|
||||
for _, segment := range strings.Split(value, "/") {
|
||||
if segment == "." || segment == ".." {
|
||||
return "/index.html"
|
||||
return indexHTML
|
||||
}
|
||||
}
|
||||
cleaned := path.Clean(value)
|
||||
@@ -546,13 +471,13 @@ func pagesFallbackPath(deployment *PagesDeployment) string {
|
||||
func renderProxyHeaderBlock(originURL string, originHost string, customHeaders []CustomHeader, upstreamConfig routeUpstreamConfig, cfg ConfigSnapshot) string {
|
||||
var builder strings.Builder
|
||||
if strings.TrimSpace(originHost) != "" {
|
||||
builder.WriteString(fmt.Sprintf(" proxy_set_header Host %s;\n", quoteNginxStringLiteral(originHost)))
|
||||
fmt.Fprintf(&builder, " proxy_set_header Host %s;\n", quoteNginxStringLiteral(originHost))
|
||||
} else {
|
||||
builder.WriteString(" proxy_set_header Host $host;\n")
|
||||
}
|
||||
if upstreamServerName := resolveUpstreamServerName(originURL, originHost); upstreamServerName != "" {
|
||||
builder.WriteString(" proxy_ssl_server_name on;\n")
|
||||
builder.WriteString(fmt.Sprintf(" proxy_ssl_name %s;\n", quoteNginxStringLiteral(upstreamServerName)))
|
||||
fmt.Fprintf(&builder, " proxy_ssl_name %s;\n", quoteNginxStringLiteral(upstreamServerName))
|
||||
}
|
||||
builder.WriteString(" proxy_set_header X-Real-IP $remote_addr;\n")
|
||||
builder.WriteString(" proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;\n")
|
||||
@@ -566,7 +491,7 @@ func renderProxyHeaderBlock(originURL string, originHost string, customHeaders [
|
||||
builder.WriteString(" proxy_set_header Connection \"\";\n")
|
||||
}
|
||||
for _, header := range customHeaders {
|
||||
builder.WriteString(fmt.Sprintf(" proxy_set_header %s %s;\n", header.Key, quoteNginxStringLiteral(header.Value)))
|
||||
fmt.Fprintf(&builder, " proxy_set_header %s %s;\n", header.Key, quoteNginxStringLiteral(header.Value))
|
||||
}
|
||||
return builder.String()
|
||||
}
|
||||
@@ -635,13 +560,13 @@ func renderRouteCacheBlock(cacheConfig routeCacheConfig, cfg ConfigSnapshot) str
|
||||
func renderRouteLimitBlock(limitConfig routeLimitConfig) string {
|
||||
var builder strings.Builder
|
||||
if limitConfig.LimitConnPerServer > 0 {
|
||||
builder.WriteString(fmt.Sprintf(" limit_conn openflare_conn_per_server %d;\n", limitConfig.LimitConnPerServer))
|
||||
fmt.Fprintf(&builder, " limit_conn openflare_conn_per_server %d;\n", limitConfig.LimitConnPerServer)
|
||||
}
|
||||
if limitConfig.LimitConnPerIP > 0 {
|
||||
builder.WriteString(fmt.Sprintf(" limit_conn openflare_conn_per_ip %d;\n", limitConfig.LimitConnPerIP))
|
||||
fmt.Fprintf(&builder, " limit_conn openflare_conn_per_ip %d;\n", limitConfig.LimitConnPerIP)
|
||||
}
|
||||
if strings.TrimSpace(limitConfig.LimitRate) != "" {
|
||||
builder.WriteString(fmt.Sprintf(" limit_rate %s;\n", limitConfig.LimitRate))
|
||||
fmt.Fprintf(&builder, " limit_rate %s;\n", limitConfig.LimitRate)
|
||||
}
|
||||
return builder.String()
|
||||
}
|
||||
@@ -700,8 +625,8 @@ func buildRouteUpstreamConfig(route Route, upstreams []string) routeUpstreamConf
|
||||
|
||||
func normalizeRouteUpstreamType(raw string) string {
|
||||
switch strings.ToLower(strings.TrimSpace(raw)) {
|
||||
case "pages":
|
||||
return "pages"
|
||||
case routeUpstreamTypePages:
|
||||
return routeUpstreamTypePages
|
||||
default:
|
||||
return "direct"
|
||||
}
|
||||
@@ -709,9 +634,9 @@ func normalizeRouteUpstreamType(raw string) string {
|
||||
|
||||
func renderNamedUpstreamBlock(upstreamConfig routeUpstreamConfig) string {
|
||||
var builder strings.Builder
|
||||
builder.WriteString(fmt.Sprintf("upstream %s {\n", upstreamConfig.Name))
|
||||
fmt.Fprintf(&builder, "upstream %s {\n", upstreamConfig.Name)
|
||||
for _, server := range upstreamConfig.Servers {
|
||||
builder.WriteString(fmt.Sprintf(" server %s max_fails=3 fail_timeout=10s;\n", server))
|
||||
fmt.Fprintf(&builder, " server %s max_fails=3 fail_timeout=10s;\n", server)
|
||||
}
|
||||
builder.WriteString(" keepalive 128;\n}\n\n")
|
||||
return builder.String()
|
||||
|
||||
@@ -0,0 +1,151 @@
|
||||
package openresty
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strings"
|
||||
)
|
||||
|
||||
type routeCertPartition struct {
|
||||
httpOnlyDomains []string
|
||||
domainsByCertID map[uint][]string
|
||||
}
|
||||
|
||||
func partitionRouteDomainsByCert(domains []string, certIDs, domainCertIDs []uint) routeCertPartition {
|
||||
httpOnlyDomains := make([]string, 0, len(domains))
|
||||
domainsByCertID := make(map[uint][]string, len(certIDs))
|
||||
for index, domain := range domains {
|
||||
if index >= len(domainCertIDs) || domainCertIDs[index] == 0 {
|
||||
httpOnlyDomains = append(httpOnlyDomains, domain)
|
||||
continue
|
||||
}
|
||||
domainsByCertID[domainCertIDs[index]] = append(domainsByCertID[domainCertIDs[index]], domain)
|
||||
}
|
||||
return routeCertPartition{
|
||||
httpOnlyDomains: httpOnlyDomains,
|
||||
domainsByCertID: domainsByCertID,
|
||||
}
|
||||
}
|
||||
|
||||
func validateRouteCertificates(route Route, displayName string, certIDs []uint, partition routeCertPartition, certificates map[uint]string) error {
|
||||
for _, certID := range certIDs {
|
||||
assignedDomains := partition.domainsByCertID[certID]
|
||||
if len(assignedDomains) == 0 {
|
||||
continue
|
||||
}
|
||||
certPEM, ok := certificates[certID]
|
||||
if !ok {
|
||||
return fmt.Errorf("route %s certificate %d does not exist", route.Domain, certID)
|
||||
}
|
||||
if err := validateCertificateCoverage(certPEM, assignedDomains); err != nil {
|
||||
return fmt.Errorf("site %s certificate validation failed: %w", displayName, err)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func renderPagesRouteHTTPS(
|
||||
builder *strings.Builder,
|
||||
serverNames, displayName string,
|
||||
route Route,
|
||||
partition routeCertPartition,
|
||||
certIDs []uint,
|
||||
limitConfig routeLimitConfig,
|
||||
powEnabled bool,
|
||||
cfg ConfigSnapshot,
|
||||
) {
|
||||
if route.RedirectHTTP {
|
||||
if len(partition.httpOnlyDomains) > 0 {
|
||||
builder.WriteString(renderHTTPPagesServer(renderServerNames(partition.httpOnlyDomains), displayName, route.PagesDeployment, limitConfig, powEnabled, route.BasicAuthEnabled, route.BasicAuthUsername, route.BasicAuthPassword))
|
||||
}
|
||||
for _, certID := range certIDs {
|
||||
if assignedDomains := partition.domainsByCertID[certID]; len(assignedDomains) > 0 {
|
||||
builder.WriteString(renderHTTPRedirectServer(renderServerNames(assignedDomains)))
|
||||
}
|
||||
}
|
||||
} else {
|
||||
builder.WriteString(renderHTTPPagesServer(serverNames, displayName, route.PagesDeployment, limitConfig, powEnabled, route.BasicAuthEnabled, route.BasicAuthUsername, route.BasicAuthPassword))
|
||||
}
|
||||
for _, certID := range certIDs {
|
||||
if assignedDomains := partition.domainsByCertID[certID]; len(assignedDomains) > 0 {
|
||||
builder.WriteString(renderHTTPSPagesServer(renderServerNames(assignedDomains), displayName, certID, route.PagesDeployment, limitConfig, powEnabled, route.BasicAuthEnabled, route.BasicAuthUsername, route.BasicAuthPassword, cfg))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func renderProxyRouteHTTPS(
|
||||
builder *strings.Builder,
|
||||
serverNames, displayName string,
|
||||
route Route,
|
||||
partition routeCertPartition,
|
||||
certIDs []uint,
|
||||
cacheConfig routeCacheConfig,
|
||||
limitConfig routeLimitConfig,
|
||||
upstreamConfig routeUpstreamConfig,
|
||||
powEnabled bool,
|
||||
cfg ConfigSnapshot,
|
||||
) {
|
||||
if route.RedirectHTTP {
|
||||
if len(partition.httpOnlyDomains) > 0 {
|
||||
builder.WriteString(renderHTTPProxyServer(renderServerNames(partition.httpOnlyDomains), displayName, route.OriginURL, route.OriginHost, route.CustomHeaders, cacheConfig, limitConfig, upstreamConfig, powEnabled, route.BasicAuthEnabled, route.BasicAuthUsername, route.BasicAuthPassword, cfg))
|
||||
}
|
||||
for _, certID := range certIDs {
|
||||
if assignedDomains := partition.domainsByCertID[certID]; len(assignedDomains) > 0 {
|
||||
builder.WriteString(renderHTTPRedirectServer(renderServerNames(assignedDomains)))
|
||||
}
|
||||
}
|
||||
} else {
|
||||
builder.WriteString(renderHTTPProxyServer(serverNames, displayName, route.OriginURL, route.OriginHost, route.CustomHeaders, cacheConfig, limitConfig, upstreamConfig, powEnabled, route.BasicAuthEnabled, route.BasicAuthUsername, route.BasicAuthPassword, cfg))
|
||||
}
|
||||
for _, certID := range certIDs {
|
||||
if assignedDomains := partition.domainsByCertID[certID]; len(assignedDomains) > 0 {
|
||||
builder.WriteString(renderHTTPSServer(renderServerNames(assignedDomains), displayName, route.OriginURL, route.OriginHost, certID, route.CustomHeaders, cacheConfig, limitConfig, upstreamConfig, powEnabled, route.BasicAuthEnabled, route.BasicAuthUsername, route.BasicAuthPassword, cfg))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func renderPagesRoute(builder *strings.Builder, route Route, displayName, serverNames string, certificates map[uint]string, limitConfig routeLimitConfig, powEnabled bool, cfg ConfigSnapshot) error {
|
||||
if route.PagesDeployment == nil {
|
||||
return fmt.Errorf("route %s pages deployment is missing", route.Domain)
|
||||
}
|
||||
if !route.EnableHTTPS {
|
||||
builder.WriteString(renderHTTPPagesServer(serverNames, displayName, route.PagesDeployment, limitConfig, powEnabled, route.BasicAuthEnabled, route.BasicAuthUsername, route.BasicAuthPassword))
|
||||
return nil
|
||||
}
|
||||
certIDs := normalizeCertIDs(route.CertID, route.CertIDs)
|
||||
domainCertIDs := normalizeDomainCertIDs(normalizedRouteDomains(route), certIDs, route.DomainCertIDs)
|
||||
if len(certIDs) == 0 {
|
||||
return fmt.Errorf("路由 %s 未配置证书", route.Domain)
|
||||
}
|
||||
partition := partitionRouteDomainsByCert(normalizedRouteDomains(route), certIDs, domainCertIDs)
|
||||
if err := validateRouteCertificates(route, displayName, certIDs, partition, certificates); err != nil {
|
||||
return err
|
||||
}
|
||||
renderPagesRouteHTTPS(builder, serverNames, displayName, route, partition, certIDs, limitConfig, powEnabled, cfg)
|
||||
return nil
|
||||
}
|
||||
|
||||
func renderProxyRoute(builder *strings.Builder, route Route, displayName, serverNames string, certificates map[uint]string, cacheConfig routeCacheConfig, limitConfig routeLimitConfig, powEnabled bool, cfg ConfigSnapshot) error {
|
||||
upstreams := route.Upstreams
|
||||
if len(upstreams) == 0 && strings.TrimSpace(route.OriginURL) != "" {
|
||||
upstreams = []string{route.OriginURL}
|
||||
}
|
||||
upstreamConfig := buildRouteUpstreamConfig(route, upstreams)
|
||||
if upstreamConfig.UsesNamedUpstream {
|
||||
builder.WriteString(renderNamedUpstreamBlock(upstreamConfig))
|
||||
}
|
||||
if !route.EnableHTTPS {
|
||||
builder.WriteString(renderHTTPProxyServer(serverNames, displayName, route.OriginURL, route.OriginHost, route.CustomHeaders, cacheConfig, limitConfig, upstreamConfig, powEnabled, route.BasicAuthEnabled, route.BasicAuthUsername, route.BasicAuthPassword, cfg))
|
||||
return nil
|
||||
}
|
||||
certIDs := normalizeCertIDs(route.CertID, route.CertIDs)
|
||||
domainCertIDs := normalizeDomainCertIDs(normalizedRouteDomains(route), certIDs, route.DomainCertIDs)
|
||||
if len(certIDs) == 0 {
|
||||
return fmt.Errorf("路由 %s 未配置证书", route.Domain)
|
||||
}
|
||||
partition := partitionRouteDomainsByCert(normalizedRouteDomains(route), certIDs, domainCertIDs)
|
||||
if err := validateRouteCertificates(route, displayName, certIDs, partition, certificates); err != nil {
|
||||
return err
|
||||
}
|
||||
renderProxyRouteHTTPS(builder, serverNames, displayName, route, partition, certIDs, cacheConfig, limitConfig, upstreamConfig, powEnabled, cfg)
|
||||
return nil
|
||||
}
|
||||
@@ -1,5 +1,7 @@
|
||||
package openresty
|
||||
|
||||
// Placeholder constants used as sentinel values in rendered OpenResty config
|
||||
// files; the deploy process replaces them with real paths before reload.
|
||||
const (
|
||||
CertDirPlaceholder = "__OPENFLARE_CERT_DIR__"
|
||||
RouteConfigPlaceholder = "__OPENFLARE_ROUTE_CONFIG__"
|
||||
@@ -63,16 +65,22 @@ http {
|
||||
}
|
||||
`
|
||||
|
||||
// SupportFile represents an auxiliary file (certificate, WAF config, etc.)
|
||||
// that is written alongside the main OpenResty configuration.
|
||||
type SupportFile struct {
|
||||
Path string `json:"path"`
|
||||
Content string `json:"content"`
|
||||
}
|
||||
|
||||
// CustomHeader is a key/value pair injected as an additional proxy_set_header
|
||||
// directive for a specific route.
|
||||
type CustomHeader struct {
|
||||
Key string `json:"key"`
|
||||
Value string `json:"value"`
|
||||
}
|
||||
|
||||
// PoWListConfig holds the IP, CIDR, path, and user-agent lists used by the
|
||||
// Proof-of-Work whitelist or blacklist filter.
|
||||
type PoWListConfig struct {
|
||||
IPs []string `json:"ips"`
|
||||
IPCidrs []string `json:"ip_cidrs"`
|
||||
@@ -81,6 +89,8 @@ type PoWListConfig struct {
|
||||
UserAgents []string `json:"user_agents"`
|
||||
}
|
||||
|
||||
// PoWConfig holds the full Proof-of-Work challenge parameters for a route,
|
||||
// including difficulty, algorithm, TTLs, and allow/block lists.
|
||||
type PoWConfig struct {
|
||||
Difficulty int `json:"difficulty"`
|
||||
Algorithm string `json:"algorithm"`
|
||||
@@ -90,6 +100,8 @@ type PoWConfig struct {
|
||||
Blacklist PoWListConfig `json:"blacklist"`
|
||||
}
|
||||
|
||||
// Route describes a single proxy or pages site entry in the OpenFlare config
|
||||
// document, including upstream, TLS, caching, rate-limiting and WAF settings.
|
||||
type Route struct {
|
||||
ID uint `json:"id,omitempty"`
|
||||
SiteName string `json:"site_name,omitempty"`
|
||||
@@ -121,6 +133,8 @@ type Route struct {
|
||||
PagesDeployment *PagesDeployment `json:"pages_deployment,omitempty"`
|
||||
}
|
||||
|
||||
// PagesDeployment holds the static-site deployment parameters for a Pages-type
|
||||
// route, including local root, entry file, SPA fallback, and API proxy options.
|
||||
type PagesDeployment struct {
|
||||
ProjectID uint `json:"project_id"`
|
||||
ProjectSlug string `json:"project_slug"`
|
||||
@@ -137,6 +151,8 @@ type PagesDeployment struct {
|
||||
LocalRoot string `json:"local_root"`
|
||||
}
|
||||
|
||||
// WAFRuleGroup defines a WAF rule group with IP/country/region lists, PoW
|
||||
// integration, and per-group block status configuration.
|
||||
type WAFRuleGroup struct {
|
||||
ID uint `json:"id"`
|
||||
Name string `json:"name"`
|
||||
@@ -156,6 +172,8 @@ type WAFRuleGroup struct {
|
||||
PoWConfig *PoWConfig `json:"pow_config,omitempty"`
|
||||
}
|
||||
|
||||
// WAFIPGroup is a named, reusable list of IP addresses or CIDRs that can be
|
||||
// referenced by multiple WAF rule groups as a whitelist or blacklist.
|
||||
type WAFIPGroup struct {
|
||||
ID uint `json:"id"`
|
||||
Name string `json:"name"`
|
||||
@@ -164,18 +182,24 @@ type WAFIPGroup struct {
|
||||
IPList []string `json:"ip_list,omitempty"`
|
||||
}
|
||||
|
||||
// WAFBinding associates a route (by site name) with the WAF rule groups that
|
||||
// should be enforced for that site.
|
||||
type WAFBinding struct {
|
||||
RouteID uint `json:"route_id"`
|
||||
SiteName string `json:"site_name"`
|
||||
RuleGroupIDs []uint `json:"rule_group_ids"`
|
||||
}
|
||||
|
||||
// WAFDocument is the top-level WAF configuration snapshot containing rule
|
||||
// groups, IP groups, and per-site bindings.
|
||||
type WAFDocument struct {
|
||||
RuleGroups []WAFRuleGroup `json:"rule_groups"`
|
||||
IPGroups []WAFIPGroup `json:"ip_groups,omitempty"`
|
||||
Bindings []WAFBinding `json:"bindings"`
|
||||
}
|
||||
|
||||
// ConfigSnapshot holds the full set of OpenResty tuning parameters that are
|
||||
// rendered into the nginx main configuration template.
|
||||
type ConfigSnapshot struct {
|
||||
DefaultServerReturnStatus int `json:"default_server_return_status"`
|
||||
WorkerProcesses string `json:"worker_processes"`
|
||||
@@ -216,12 +240,16 @@ type ConfigSnapshot struct {
|
||||
MainConfigTemplate string `json:"main_config_template,omitempty"`
|
||||
}
|
||||
|
||||
// Document is the top-level input structure for the OpenResty renderer,
|
||||
// combining routes, OpenResty tuning, and WAF configuration.
|
||||
type Document struct {
|
||||
Routes []Route `json:"routes"`
|
||||
OpenRestyConfig ConfigSnapshot `json:"openresty_config"`
|
||||
WAF WAFDocument `json:"waf"`
|
||||
}
|
||||
|
||||
// Result is the output produced by Render, containing the rendered main
|
||||
// config, route config, support files, and a content checksum.
|
||||
type Result struct {
|
||||
MainConfig string
|
||||
RouteConfig string
|
||||
|
||||
+31
-19
@@ -1,3 +1,4 @@
|
||||
// Package utils provides shared formatting and string helper functions.
|
||||
package utils
|
||||
|
||||
import (
|
||||
@@ -5,48 +6,59 @@ import (
|
||||
"strconv"
|
||||
)
|
||||
|
||||
const (
|
||||
secondsPerYear = 31104000 // 360 days
|
||||
secondsPerMonth = 2592000 // 30 days
|
||||
secondsPerDay = 86400
|
||||
secondsPerHour = 3600
|
||||
secondsPerMinute = 60
|
||||
)
|
||||
|
||||
var sizeKB = 1024
|
||||
var sizeMB = sizeKB * 1024
|
||||
var sizeGB = sizeMB * 1024
|
||||
|
||||
// Bytes2Size converts a byte count to a human-readable string with unit (B, KB, MB, GB).
|
||||
func Bytes2Size(num int64) string {
|
||||
numStr := ""
|
||||
unit := "B"
|
||||
if num/int64(sizeGB) > 1 {
|
||||
switch {
|
||||
case num/int64(sizeGB) > 1:
|
||||
numStr = fmt.Sprintf("%.2f", float64(num)/float64(sizeGB))
|
||||
unit = "GB"
|
||||
} else if num/int64(sizeMB) > 1 {
|
||||
case num/int64(sizeMB) > 1:
|
||||
numStr = fmt.Sprintf("%d", int(float64(num)/float64(sizeMB)))
|
||||
unit = "MB"
|
||||
} else if num/int64(sizeKB) > 1 {
|
||||
case num/int64(sizeKB) > 1:
|
||||
numStr = fmt.Sprintf("%d", int(float64(num)/float64(sizeKB)))
|
||||
unit = "KB"
|
||||
} else {
|
||||
default:
|
||||
numStr = fmt.Sprintf("%d", num)
|
||||
}
|
||||
return numStr + " " + unit
|
||||
}
|
||||
|
||||
// Seconds2Time converts a number of seconds to a human-readable Chinese duration string.
|
||||
func Seconds2Time(num int) (time string) {
|
||||
if num/31104000 > 0 {
|
||||
time += strconv.Itoa(num/31104000) + " 年 "
|
||||
num %= 31104000
|
||||
if num/secondsPerYear > 0 {
|
||||
time += strconv.Itoa(num/secondsPerYear) + " 年 "
|
||||
num %= secondsPerYear
|
||||
}
|
||||
if num/2592000 > 0 {
|
||||
time += strconv.Itoa(num/2592000) + " 个月 "
|
||||
num %= 2592000
|
||||
if num/secondsPerMonth > 0 {
|
||||
time += strconv.Itoa(num/secondsPerMonth) + " 个月 "
|
||||
num %= secondsPerMonth
|
||||
}
|
||||
if num/86400 > 0 {
|
||||
time += strconv.Itoa(num/86400) + " 天 "
|
||||
num %= 86400
|
||||
if num/secondsPerDay > 0 {
|
||||
time += strconv.Itoa(num/secondsPerDay) + " 天 "
|
||||
num %= secondsPerDay
|
||||
}
|
||||
if num/3600 > 0 {
|
||||
time += strconv.Itoa(num/3600) + " 小时 "
|
||||
num %= 3600
|
||||
if num/secondsPerHour > 0 {
|
||||
time += strconv.Itoa(num/secondsPerHour) + " 小时 "
|
||||
num %= secondsPerHour
|
||||
}
|
||||
if num/60 > 0 {
|
||||
time += strconv.Itoa(num/60) + " 分钟 "
|
||||
num %= 60
|
||||
if num/secondsPerMinute > 0 {
|
||||
time += strconv.Itoa(num/secondsPerMinute) + " 分钟 "
|
||||
num %= secondsPerMinute
|
||||
}
|
||||
time += strconv.Itoa(num) + " 秒"
|
||||
return
|
||||
|
||||
+22
-15
@@ -6,7 +6,8 @@ import (
|
||||
"strings"
|
||||
)
|
||||
|
||||
func GetIp() (ip string) {
|
||||
// GetIP returns the first private IPv4 address found on the local network interfaces.
|
||||
func GetIP() (ip string) {
|
||||
ips, err := net.InterfaceAddrs()
|
||||
if err != nil {
|
||||
slog.Error("get interface addresses failed", "error", err)
|
||||
@@ -14,21 +15,27 @@ func GetIp() (ip string) {
|
||||
}
|
||||
|
||||
for _, a := range ips {
|
||||
if ipNet, ok := a.(*net.IPNet); ok && !ipNet.IP.IsLoopback() {
|
||||
if ipNet.IP.To4() != nil {
|
||||
ip = ipNet.IP.String()
|
||||
if strings.HasPrefix(ip, "10") {
|
||||
return
|
||||
}
|
||||
if strings.HasPrefix(ip, "172") {
|
||||
return
|
||||
}
|
||||
if strings.HasPrefix(ip, "192.168") {
|
||||
return
|
||||
}
|
||||
ip = ""
|
||||
}
|
||||
if candidate, ok := privateIPv4FromAddr(a); ok {
|
||||
return candidate
|
||||
}
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
func privateIPv4FromAddr(addr net.Addr) (string, bool) {
|
||||
ipNet, ok := addr.(*net.IPNet)
|
||||
if !ok || ipNet.IP.IsLoopback() || ipNet.IP.To4() == nil {
|
||||
return "", false
|
||||
}
|
||||
ip := ipNet.IP.String()
|
||||
if isPrivateIPv4(ip) {
|
||||
return ip, true
|
||||
}
|
||||
return "", false
|
||||
}
|
||||
|
||||
func isPrivateIPv4(ip string) bool {
|
||||
return strings.HasPrefix(ip, "10") ||
|
||||
strings.HasPrefix(ip, "172") ||
|
||||
strings.HasPrefix(ip, "192.168")
|
||||
}
|
||||
|
||||
+5
-4
@@ -2,14 +2,15 @@ package utils
|
||||
|
||||
import "fmt"
|
||||
|
||||
// Interface2String converts a string, int, or float64 value to its string representation.
|
||||
func Interface2String(inter interface{}) string {
|
||||
switch inter.(type) {
|
||||
switch v := inter.(type) {
|
||||
case string:
|
||||
return inter.(string)
|
||||
return v
|
||||
case int:
|
||||
return fmt.Sprintf("%d", inter.(int))
|
||||
return fmt.Sprintf("%d", v)
|
||||
case float64:
|
||||
return fmt.Sprintf("%f", inter.(float64))
|
||||
return fmt.Sprintf("%f", v)
|
||||
}
|
||||
return "Not Implemented"
|
||||
}
|
||||
|
||||
+11
-93
@@ -5,6 +5,9 @@ import (
|
||||
"strings"
|
||||
)
|
||||
|
||||
const gitDescribeMinIdentifiers = 2
|
||||
|
||||
// VersionInfo holds the parsed components of a semantic version string.
|
||||
type VersionInfo struct {
|
||||
Valid bool
|
||||
IsDev bool
|
||||
@@ -14,6 +17,7 @@ type VersionInfo struct {
|
||||
GitDescribeTail []string
|
||||
}
|
||||
|
||||
// ParseVersionInfo parses a version string into a structured VersionInfo.
|
||||
func ParseVersionInfo(version string) VersionInfo {
|
||||
normalized := strings.TrimSpace(strings.TrimPrefix(version, "v"))
|
||||
if normalized == "" || normalized == "dev" {
|
||||
@@ -66,7 +70,7 @@ func ParseVersionInfo(version string) VersionInfo {
|
||||
}
|
||||
|
||||
func parseGitDescribeIdentifiers(identifiers []string) (int, []string, bool) {
|
||||
if len(identifiers) < 2 {
|
||||
if len(identifiers) < gitDescribeMinIdentifiers {
|
||||
return 0, nil, false
|
||||
}
|
||||
distance, err := strconv.Atoi(strings.TrimSpace(identifiers[0]))
|
||||
@@ -109,100 +113,14 @@ func CompareVersions(local, remote string) int {
|
||||
return 0
|
||||
}
|
||||
|
||||
maxLen := len(left.Numbers)
|
||||
if len(right.Numbers) > maxLen {
|
||||
maxLen = len(right.Numbers)
|
||||
if result := compareVersionNumbers(left, right); result != 0 {
|
||||
return result
|
||||
}
|
||||
for index := 0; index < maxLen; index++ {
|
||||
leftValue := 0
|
||||
rightValue := 0
|
||||
if index < len(left.Numbers) {
|
||||
leftValue = left.Numbers[index]
|
||||
}
|
||||
if index < len(right.Numbers) {
|
||||
rightValue = right.Numbers[index]
|
||||
}
|
||||
if leftValue < rightValue {
|
||||
return -1
|
||||
}
|
||||
if leftValue > rightValue {
|
||||
return 1
|
||||
}
|
||||
}
|
||||
|
||||
if left.GitDescribeDistance != right.GitDescribeDistance {
|
||||
if left.GitDescribeDistance < right.GitDescribeDistance {
|
||||
return -1
|
||||
}
|
||||
return 1
|
||||
if result := compareGitDescribeDistance(left, right); result != 0 {
|
||||
return result
|
||||
}
|
||||
if left.GitDescribeDistance > 0 || right.GitDescribeDistance > 0 {
|
||||
maxLen = len(left.GitDescribeTail)
|
||||
if len(right.GitDescribeTail) > maxLen {
|
||||
maxLen = len(right.GitDescribeTail)
|
||||
}
|
||||
for index := 0; index < maxLen; index++ {
|
||||
if index >= len(left.GitDescribeTail) {
|
||||
return -1
|
||||
}
|
||||
if index >= len(right.GitDescribeTail) {
|
||||
return 1
|
||||
}
|
||||
if left.GitDescribeTail[index] < right.GitDescribeTail[index] {
|
||||
return -1
|
||||
}
|
||||
if left.GitDescribeTail[index] > right.GitDescribeTail[index] {
|
||||
return 1
|
||||
}
|
||||
}
|
||||
return 0
|
||||
return compareGitDescribeTails(left, right)
|
||||
}
|
||||
|
||||
if len(left.Prerelease) == 0 && len(right.Prerelease) == 0 {
|
||||
return 0
|
||||
}
|
||||
if len(left.Prerelease) == 0 {
|
||||
return 1
|
||||
}
|
||||
if len(right.Prerelease) == 0 {
|
||||
return -1
|
||||
}
|
||||
|
||||
maxLen = len(left.Prerelease)
|
||||
if len(right.Prerelease) > maxLen {
|
||||
maxLen = len(right.Prerelease)
|
||||
}
|
||||
for index := 0; index < maxLen; index++ {
|
||||
if index >= len(left.Prerelease) {
|
||||
return -1
|
||||
}
|
||||
if index >= len(right.Prerelease) {
|
||||
return 1
|
||||
}
|
||||
leftPart := left.Prerelease[index]
|
||||
rightPart := right.Prerelease[index]
|
||||
leftNumber, leftErr := strconv.Atoi(leftPart)
|
||||
rightNumber, rightErr := strconv.Atoi(rightPart)
|
||||
switch {
|
||||
case leftErr == nil && rightErr == nil:
|
||||
if leftNumber < rightNumber {
|
||||
return -1
|
||||
}
|
||||
if leftNumber > rightNumber {
|
||||
return 1
|
||||
}
|
||||
case leftErr == nil:
|
||||
return -1
|
||||
case rightErr == nil:
|
||||
return 1
|
||||
default:
|
||||
if leftPart < rightPart {
|
||||
return -1
|
||||
}
|
||||
if leftPart > rightPart {
|
||||
return 1
|
||||
}
|
||||
}
|
||||
}
|
||||
return 0
|
||||
return comparePrereleaseIdentifiers(left, right)
|
||||
}
|
||||
|
||||
@@ -0,0 +1,114 @@
|
||||
package utils
|
||||
|
||||
import "strconv"
|
||||
|
||||
func compareVersionNumbers(left, right VersionInfo) int {
|
||||
maxLen := len(left.Numbers)
|
||||
if len(right.Numbers) > maxLen {
|
||||
maxLen = len(right.Numbers)
|
||||
}
|
||||
for index := 0; index < maxLen; index++ {
|
||||
leftValue := 0
|
||||
rightValue := 0
|
||||
if index < len(left.Numbers) {
|
||||
leftValue = left.Numbers[index]
|
||||
}
|
||||
if index < len(right.Numbers) {
|
||||
rightValue = right.Numbers[index]
|
||||
}
|
||||
if leftValue < rightValue {
|
||||
return -1
|
||||
}
|
||||
if leftValue > rightValue {
|
||||
return 1
|
||||
}
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
func compareGitDescribeDistance(left, right VersionInfo) int {
|
||||
if left.GitDescribeDistance == right.GitDescribeDistance {
|
||||
return 0
|
||||
}
|
||||
if left.GitDescribeDistance < right.GitDescribeDistance {
|
||||
return -1
|
||||
}
|
||||
return 1
|
||||
}
|
||||
|
||||
func compareGitDescribeTails(left, right VersionInfo) int {
|
||||
maxLen := len(left.GitDescribeTail)
|
||||
if len(right.GitDescribeTail) > maxLen {
|
||||
maxLen = len(right.GitDescribeTail)
|
||||
}
|
||||
for index := 0; index < maxLen; index++ {
|
||||
if index >= len(left.GitDescribeTail) {
|
||||
return -1
|
||||
}
|
||||
if index >= len(right.GitDescribeTail) {
|
||||
return 1
|
||||
}
|
||||
if left.GitDescribeTail[index] < right.GitDescribeTail[index] {
|
||||
return -1
|
||||
}
|
||||
if left.GitDescribeTail[index] > right.GitDescribeTail[index] {
|
||||
return 1
|
||||
}
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
func comparePrereleaseIdentifiers(left, right VersionInfo) int {
|
||||
if len(left.Prerelease) == 0 && len(right.Prerelease) == 0 {
|
||||
return 0
|
||||
}
|
||||
if len(left.Prerelease) == 0 {
|
||||
return 1
|
||||
}
|
||||
if len(right.Prerelease) == 0 {
|
||||
return -1
|
||||
}
|
||||
|
||||
maxLen := len(left.Prerelease)
|
||||
if len(right.Prerelease) > maxLen {
|
||||
maxLen = len(right.Prerelease)
|
||||
}
|
||||
for index := 0; index < maxLen; index++ {
|
||||
if index >= len(left.Prerelease) {
|
||||
return -1
|
||||
}
|
||||
if index >= len(right.Prerelease) {
|
||||
return 1
|
||||
}
|
||||
if result := comparePrereleasePart(left.Prerelease[index], right.Prerelease[index]); result != 0 {
|
||||
return result
|
||||
}
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
func comparePrereleasePart(leftPart, rightPart string) int {
|
||||
leftNumber, leftErr := strconv.Atoi(leftPart)
|
||||
rightNumber, rightErr := strconv.Atoi(rightPart)
|
||||
switch {
|
||||
case leftErr == nil && rightErr == nil:
|
||||
if leftNumber < rightNumber {
|
||||
return -1
|
||||
}
|
||||
if leftNumber > rightNumber {
|
||||
return 1
|
||||
}
|
||||
case leftErr == nil:
|
||||
return -1
|
||||
case rightErr == nil:
|
||||
return 1
|
||||
default:
|
||||
if leftPart < rightPart {
|
||||
return -1
|
||||
}
|
||||
if leftPart > rightPart {
|
||||
return 1
|
||||
}
|
||||
}
|
||||
return 0
|
||||
}
|
||||
+23
-3
@@ -1,3 +1,4 @@
|
||||
// Package wsclient provides a WebSocket client for agent/server communication.
|
||||
package wsclient
|
||||
|
||||
import (
|
||||
@@ -14,6 +15,12 @@ import (
|
||||
"golang.org/x/net/websocket"
|
||||
)
|
||||
|
||||
const (
|
||||
writeDeadlineSecs = 5
|
||||
defaultReadDeadlineSecs = 75
|
||||
)
|
||||
|
||||
// Config holds the configuration for a WebSocket client connection.
|
||||
type Config struct {
|
||||
BaseURL string
|
||||
Token string
|
||||
@@ -22,27 +29,32 @@ type Config struct {
|
||||
WSPath string // e.g. "/api/relay/ws", "/api/agent/ws", "/api/flared/ws"
|
||||
}
|
||||
|
||||
// Client provides methods to connect and communicate over WebSocket.
|
||||
type Client struct {
|
||||
cfg Config
|
||||
}
|
||||
|
||||
// WSMessage represents a typed WebSocket message with an optional JSON payload.
|
||||
type WSMessage struct {
|
||||
Type string `json:"type"`
|
||||
Payload json.RawMessage `json:"payload,omitempty"`
|
||||
}
|
||||
|
||||
// MessageHandler handles WebSocket connection lifecycle and incoming messages.
|
||||
type MessageHandler interface {
|
||||
OnConnect(ctx context.Context) error
|
||||
HandleMessage(ctx context.Context, msg WSMessage) error
|
||||
OnClose(err error)
|
||||
}
|
||||
|
||||
// Connection represents an active WebSocket connection.
|
||||
type Connection struct {
|
||||
Conn *websocket.Conn
|
||||
URL string
|
||||
ReadTimeout time.Duration
|
||||
}
|
||||
|
||||
// New creates a new WebSocket client with the given configuration.
|
||||
func New(cfg Config) *Client {
|
||||
cfg.BaseURL = strings.TrimRight(cfg.BaseURL, "/")
|
||||
cfg.Token = strings.TrimSpace(cfg.Token)
|
||||
@@ -53,10 +65,12 @@ func New(cfg Config) *Client {
|
||||
}
|
||||
}
|
||||
|
||||
// SetToken updates the authentication token used for the WebSocket connection.
|
||||
func (c *Client) SetToken(token string) {
|
||||
c.cfg.Token = strings.TrimSpace(token)
|
||||
}
|
||||
|
||||
// URL returns the WebSocket URL for the configured endpoint, or empty string on error.
|
||||
func (c *Client) URL() string {
|
||||
wsURL, err := c.BuildWebsocketURL()
|
||||
if err != nil {
|
||||
@@ -65,6 +79,7 @@ func (c *Client) URL() string {
|
||||
return wsURL
|
||||
}
|
||||
|
||||
// BuildWebsocketURL constructs the WebSocket URL by converting the base URL scheme and appending the WS path.
|
||||
func (c *Client) BuildWebsocketURL() (string, error) {
|
||||
parsed, err := url.Parse(c.cfg.BaseURL)
|
||||
if err != nil {
|
||||
@@ -90,6 +105,7 @@ func (c *Client) BuildWebsocketURL() (string, error) {
|
||||
return parsed.String(), nil
|
||||
}
|
||||
|
||||
// Connect establishes a new WebSocket connection to the configured server.
|
||||
func (c *Client) Connect(ctx context.Context) (*Connection, error) {
|
||||
wsURL, err := c.BuildWebsocketURL()
|
||||
if err != nil {
|
||||
@@ -122,6 +138,7 @@ func (c *Client) Connect(ctx context.Context) (*Connection, error) {
|
||||
return &Connection{Conn: conn, URL: wsURL, ReadTimeout: websocketReadTimeout(c.cfg.Timeout)}, nil
|
||||
}
|
||||
|
||||
// SendMessage sends a typed message with an optional payload over the WebSocket connection.
|
||||
func (conn *Connection) SendMessage(msgType string, payload any) error {
|
||||
if conn == nil || conn.Conn == nil {
|
||||
return errors.New("ws connection is nil")
|
||||
@@ -137,10 +154,11 @@ func (conn *Connection) SendMessage(msgType string, payload any) error {
|
||||
Payload: payload,
|
||||
}
|
||||
|
||||
_ = conn.Conn.SetWriteDeadline(time.Now().Add(5 * time.Second))
|
||||
_ = conn.Conn.SetWriteDeadline(time.Now().Add(writeDeadlineSecs * time.Second))
|
||||
return websocket.JSON.Send(conn.Conn, message)
|
||||
}
|
||||
|
||||
// Receive reads a single message from the WebSocket connection into target.
|
||||
func (conn *Connection) Receive(target any) error {
|
||||
if conn == nil || conn.Conn == nil {
|
||||
return errors.New("ws connection is nil")
|
||||
@@ -161,12 +179,13 @@ func (conn *Connection) Receive(target any) error {
|
||||
|
||||
func websocketReadTimeout(requestTimeout time.Duration) time.Duration {
|
||||
timeout := requestTimeout * 6
|
||||
if timeout < 75*time.Second {
|
||||
return 75 * time.Second
|
||||
if timeout < defaultReadDeadlineSecs*time.Second {
|
||||
return defaultReadDeadlineSecs * time.Second
|
||||
}
|
||||
return timeout
|
||||
}
|
||||
|
||||
// RunReceiveLoop continuously receives messages and dispatches them to the handler until the context is cancelled.
|
||||
func (conn *Connection) RunReceiveLoop(ctx context.Context, handler MessageHandler) error {
|
||||
doneChan := make(chan struct{})
|
||||
defer close(doneChan)
|
||||
@@ -214,6 +233,7 @@ func (conn *Connection) RunReceiveLoop(ctx context.Context, handler MessageHandl
|
||||
}
|
||||
}
|
||||
|
||||
// Close gracefully closes the WebSocket connection.
|
||||
func (conn *Connection) Close() error {
|
||||
if conn == nil || conn.Conn == nil {
|
||||
return nil
|
||||
|
||||
Reference in New Issue
Block a user