This commit is contained in:
ryan
2026-06-19 15:13:24 +08:00
parent 0b34792709
commit 32861c5db9
376 changed files with 3648 additions and 19957 deletions
+10 -1
View File
@@ -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
View File
@@ -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
+1 -1
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
}
+13 -2
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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"`
}
}
+18
View File
@@ -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
View File
@@ -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()
+151
View File
@@ -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
}
+28
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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)
}
+114
View File
@@ -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
View File
@@ -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