mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 15:38:22 +08:00
feat: 支持自定义客户端 IP 请求头配置
Ultraworked with [Sisyphus](https://github.com/code-yeongyu/oh-my-openagent) Co-authored-by: Sisyphus <clio-agent@sisyphuslabs.ai>
This commit is contained in:
@@ -7,6 +7,7 @@ import (
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"math"
|
||||
"net/textproto"
|
||||
"net/url"
|
||||
"os"
|
||||
"strings"
|
||||
@@ -14,6 +15,7 @@ import (
|
||||
"time"
|
||||
|
||||
"github.com/spf13/viper"
|
||||
"golang.org/x/net/http/httpguts"
|
||||
)
|
||||
|
||||
const (
|
||||
@@ -668,6 +670,13 @@ type CORSConfig struct {
|
||||
AllowCredentials bool `mapstructure:"allow_credentials"`
|
||||
}
|
||||
|
||||
const MaxForwardedClientIPHeaders = 16
|
||||
|
||||
type ForwardedClientIPSettings struct {
|
||||
TrustForwardedIP bool
|
||||
Headers []string
|
||||
}
|
||||
|
||||
type SecurityConfig struct {
|
||||
URLAllowlist URLAllowlistConfig `mapstructure:"url_allowlist"`
|
||||
ResponseHeaders ResponseHeaderConfig `mapstructure:"response_headers"`
|
||||
@@ -676,19 +685,58 @@ type SecurityConfig struct {
|
||||
ProxyProbe ProxyProbeConfig `mapstructure:"proxy_probe"`
|
||||
// TrustForwardedIPForAPIKeyACL enables legacy raw forwarded-header takeover.
|
||||
// When disabled, server.trusted_proxies is authoritative for all client-IP consumers.
|
||||
TrustForwardedIPForAPIKeyACL bool `mapstructure:"trust_forwarded_ip_for_api_key_acl"`
|
||||
trustForwardedIPForAPIKeyACLLive *atomic.Bool `mapstructure:"-"`
|
||||
TrustForwardedIPForAPIKeyACL bool `mapstructure:"trust_forwarded_ip_for_api_key_acl"`
|
||||
ForwardedClientIPHeaders []string `mapstructure:"forwarded_client_ip_headers" json:"forwarded_client_ip_headers" yaml:"forwarded_client_ip_headers"`
|
||||
forwardedClientIPSettingsLive atomic.Pointer[ForwardedClientIPSettings] `mapstructure:"-" json:"-" yaml:"-"`
|
||||
}
|
||||
|
||||
func NormalizeForwardedClientIPHeaders(headers []string) ([]string, error) {
|
||||
normalized := make([]string, 0, len(headers))
|
||||
seen := make(map[string]struct{}, len(headers))
|
||||
for _, header := range headers {
|
||||
header = strings.TrimSpace(header)
|
||||
if !httpguts.ValidHeaderFieldName(header) {
|
||||
return nil, fmt.Errorf("invalid HTTP header field name %q", header)
|
||||
}
|
||||
canonical := textproto.CanonicalMIMEHeaderKey(header)
|
||||
key := strings.ToLower(canonical)
|
||||
if _, exists := seen[key]; exists {
|
||||
continue
|
||||
}
|
||||
if len(normalized) == MaxForwardedClientIPHeaders {
|
||||
return nil, fmt.Errorf("forwarded client IP headers must contain at most %d unique names", MaxForwardedClientIPHeaders)
|
||||
}
|
||||
seen[key] = struct{}{}
|
||||
normalized = append(normalized, canonical)
|
||||
}
|
||||
return normalized, nil
|
||||
}
|
||||
|
||||
func cloneForwardedClientIPHeaders(headers []string) []string {
|
||||
if len(headers) == 0 {
|
||||
return []string{}
|
||||
}
|
||||
return append([]string(nil), headers...)
|
||||
}
|
||||
|
||||
func (c *Config) ForwardedClientIPSettings() ForwardedClientIPSettings {
|
||||
if c == nil {
|
||||
return ForwardedClientIPSettings{Headers: []string{}}
|
||||
}
|
||||
if snapshot := c.Security.forwardedClientIPSettingsLive.Load(); snapshot != nil {
|
||||
return ForwardedClientIPSettings{
|
||||
TrustForwardedIP: snapshot.TrustForwardedIP,
|
||||
Headers: cloneForwardedClientIPHeaders(snapshot.Headers),
|
||||
}
|
||||
}
|
||||
return ForwardedClientIPSettings{
|
||||
TrustForwardedIP: c.Security.TrustForwardedIPForAPIKeyACL,
|
||||
Headers: cloneForwardedClientIPHeaders(c.Security.ForwardedClientIPHeaders),
|
||||
}
|
||||
}
|
||||
|
||||
func (c *Config) TrustForwardedIPForAPIKeyACL() bool {
|
||||
if c == nil {
|
||||
return false
|
||||
}
|
||||
live := c.Security.trustForwardedIPForAPIKeyACLLive
|
||||
if live == nil {
|
||||
return c.Security.TrustForwardedIPForAPIKeyACL
|
||||
}
|
||||
return live.Load()
|
||||
return c.ForwardedClientIPSettings().TrustForwardedIP
|
||||
}
|
||||
|
||||
// ForwardedClientIPTrustEnabled reports whether the legacy forwarded-header
|
||||
@@ -697,15 +745,22 @@ func (c *Config) ForwardedClientIPTrustEnabled() bool {
|
||||
return c != nil && c.TrustForwardedIPForAPIKeyACL()
|
||||
}
|
||||
|
||||
func (c *Config) SetForwardedClientIPSettings(enabled bool, headers []string) {
|
||||
if c == nil {
|
||||
return
|
||||
}
|
||||
headers = cloneForwardedClientIPHeaders(headers)
|
||||
c.Security.forwardedClientIPSettingsLive.Store(&ForwardedClientIPSettings{
|
||||
TrustForwardedIP: enabled,
|
||||
Headers: headers,
|
||||
})
|
||||
}
|
||||
|
||||
func (c *Config) SetTrustForwardedIPForAPIKeyACL(enabled bool) {
|
||||
if c == nil {
|
||||
return
|
||||
}
|
||||
c.Security.TrustForwardedIPForAPIKeyACL = enabled
|
||||
if c.Security.trustForwardedIPForAPIKeyACLLive == nil {
|
||||
c.Security.trustForwardedIPForAPIKeyACLLive = &atomic.Bool{}
|
||||
}
|
||||
c.Security.trustForwardedIPForAPIKeyACLLive.Store(enabled)
|
||||
c.SetForwardedClientIPSettings(enabled, c.ForwardedClientIPSettings().Headers)
|
||||
}
|
||||
|
||||
type URLAllowlistConfig struct {
|
||||
@@ -1574,6 +1629,7 @@ func load(allowMissingJWTSecret bool) (*Config, error) {
|
||||
// 配置文件不存在时使用默认值
|
||||
}
|
||||
trustedProxiesEnv, trustedProxiesEnvConfigured := os.LookupEnv("SERVER_TRUSTED_PROXIES")
|
||||
forwardedClientIPHeadersEnv, forwardedClientIPHeadersEnvConfigured := os.LookupEnv("SECURITY_FORWARDED_CLIENT_IP_HEADERS")
|
||||
trustedProxiesConfigured := viper.InConfig("server.trusted_proxies") ||
|
||||
viper.IsSet("server.trusted_proxies") || trustedProxiesEnvConfigured
|
||||
|
||||
@@ -1584,6 +1640,9 @@ func load(allowMissingJWTSecret bool) (*Config, error) {
|
||||
if trustedProxiesEnvConfigured {
|
||||
cfg.Server.TrustedProxies = normalizeStringSlice(strings.Split(trustedProxiesEnv, ","))
|
||||
}
|
||||
if forwardedClientIPHeadersEnvConfigured {
|
||||
cfg.Security.ForwardedClientIPHeaders = normalizeStringSlice(strings.Split(forwardedClientIPHeadersEnv, ","))
|
||||
}
|
||||
cfg.Server.TrustedProxiesConfigured = trustedProxiesConfigured
|
||||
if cfg.Gateway.OpenAIScheduler.StickyEscapeTTFTMs == 0 {
|
||||
cfg.Gateway.OpenAIScheduler.StickyEscapeTTFTMs = 15000
|
||||
@@ -1640,7 +1699,12 @@ func load(allowMissingJWTSecret bool) (*Config, error) {
|
||||
cfg.Security.ResponseHeaders.AdditionalAllowed = normalizeStringSlice(cfg.Security.ResponseHeaders.AdditionalAllowed)
|
||||
cfg.Security.ResponseHeaders.ForceRemove = normalizeStringSlice(cfg.Security.ResponseHeaders.ForceRemove)
|
||||
cfg.Security.CSP.Policy = strings.TrimSpace(cfg.Security.CSP.Policy)
|
||||
cfg.SetTrustForwardedIPForAPIKeyACL(cfg.Security.TrustForwardedIPForAPIKeyACL)
|
||||
forwardedClientIPHeaders, err := NormalizeForwardedClientIPHeaders(cfg.Security.ForwardedClientIPHeaders)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("security.forwarded_client_ip_headers: %w", err)
|
||||
}
|
||||
cfg.Security.ForwardedClientIPHeaders = forwardedClientIPHeaders
|
||||
cfg.SetForwardedClientIPSettings(cfg.Security.TrustForwardedIPForAPIKeyACL, forwardedClientIPHeaders)
|
||||
cfg.Log.Level = strings.ToLower(strings.TrimSpace(cfg.Log.Level))
|
||||
cfg.Log.Format = strings.ToLower(strings.TrimSpace(cfg.Log.Format))
|
||||
cfg.Log.ServiceName = strings.TrimSpace(cfg.Log.ServiceName)
|
||||
@@ -2244,6 +2308,12 @@ func setDefaults() {
|
||||
}
|
||||
|
||||
func (c *Config) Validate() error {
|
||||
forwardedClientIPHeaders, err := NormalizeForwardedClientIPHeaders(c.Security.ForwardedClientIPHeaders)
|
||||
if err != nil {
|
||||
return fmt.Errorf("security.forwarded_client_ip_headers: %w", err)
|
||||
}
|
||||
c.Security.ForwardedClientIPHeaders = forwardedClientIPHeaders
|
||||
c.SetForwardedClientIPSettings(c.Security.TrustForwardedIPForAPIKeyACL, forwardedClientIPHeaders)
|
||||
if c.Server.ReadHeaderTimeout < 1 || c.Server.ReadHeaderTimeout > 60 {
|
||||
return fmt.Errorf("server.read_header_timeout must be between 1 and 60 seconds")
|
||||
}
|
||||
|
||||
@@ -1,10 +1,12 @@
|
||||
package config
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"math"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
@@ -50,6 +52,126 @@ func TestLoadHTTPIngressSafetyDefaults(t *testing.T) {
|
||||
require.Equal(t, 16384, cfg.APIKeyAuth.InvalidAbuse.Capacity)
|
||||
}
|
||||
|
||||
func TestNormalizeForwardedClientIPHeaders(t *testing.T) {
|
||||
headers, err := NormalizeForwardedClientIPHeaders([]string{
|
||||
" x-cdn-client-ip ",
|
||||
"X-CDN-CLIENT-IP",
|
||||
"true-client-ip",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, []string{"X-Cdn-Client-Ip", "True-Client-Ip"}, headers)
|
||||
|
||||
_, err = NormalizeForwardedClientIPHeaders([]string{"X Invalid"})
|
||||
require.ErrorContains(t, err, "invalid HTTP header field name")
|
||||
}
|
||||
|
||||
func TestNormalizeForwardedClientIPHeadersLimit(t *testing.T) {
|
||||
headers := make([]string, 0, MaxForwardedClientIPHeaders+1)
|
||||
for i := 0; i <= MaxForwardedClientIPHeaders; i++ {
|
||||
headers = append(headers, fmt.Sprintf("X-CDN-IP-%d", i))
|
||||
}
|
||||
|
||||
_, err := NormalizeForwardedClientIPHeaders(headers)
|
||||
require.ErrorContains(t, err, "at most 16 unique names")
|
||||
}
|
||||
|
||||
func TestLoadForwardedClientIPHeadersNormalizesAndSnapshots(t *testing.T) {
|
||||
resetViperWithJWTSecret(t)
|
||||
viper.Set("security.forwarded_client_ip_headers", []string{" x-cdn-ip ", "X-CDN-IP", "true-client-ip"})
|
||||
|
||||
cfg, err := Load()
|
||||
require.NoError(t, err)
|
||||
snapshot := cfg.ForwardedClientIPSettings()
|
||||
require.Equal(t, []string{"X-Cdn-Ip", "True-Client-Ip"}, snapshot.Headers)
|
||||
|
||||
snapshot.Headers[0] = "X-Mutated"
|
||||
require.Equal(t, []string{"X-Cdn-Ip", "True-Client-Ip"}, cfg.ForwardedClientIPSettings().Headers)
|
||||
}
|
||||
|
||||
func TestForwardedClientIPSettingsConcurrentPublication(t *testing.T) {
|
||||
cfg := &Config{}
|
||||
cfg.SetForwardedClientIPSettings(true, []string{"X-Public-A"})
|
||||
|
||||
const iterations = 2000
|
||||
start := make(chan struct{})
|
||||
errCh := make(chan error, 8)
|
||||
var wg sync.WaitGroup
|
||||
|
||||
for _, settings := range []ForwardedClientIPSettings{
|
||||
{TrustForwardedIP: true, Headers: []string{"X-Public-A"}},
|
||||
{TrustForwardedIP: false, Headers: []string{"X-Public-B"}},
|
||||
} {
|
||||
settings := settings
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
<-start
|
||||
for i := 0; i < iterations; i++ {
|
||||
cfg.SetForwardedClientIPSettings(settings.TrustForwardedIP, settings.Headers)
|
||||
}
|
||||
}()
|
||||
}
|
||||
|
||||
for i := 0; i < cap(errCh); i++ {
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
<-start
|
||||
for j := 0; j < iterations; j++ {
|
||||
snapshot := cfg.ForwardedClientIPSettings()
|
||||
validA := snapshot.TrustForwardedIP && len(snapshot.Headers) == 1 && snapshot.Headers[0] == "X-Public-A"
|
||||
validB := !snapshot.TrustForwardedIP && len(snapshot.Headers) == 1 && snapshot.Headers[0] == "X-Public-B"
|
||||
if !validA && !validB {
|
||||
errCh <- fmt.Errorf("observed inconsistent forwarded IP settings: %+v", snapshot)
|
||||
return
|
||||
}
|
||||
}
|
||||
}()
|
||||
}
|
||||
|
||||
close(start)
|
||||
wg.Wait()
|
||||
close(errCh)
|
||||
for err := range errCh {
|
||||
require.NoError(t, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadForwardedClientIPHeadersFromEnvironment(t *testing.T) {
|
||||
resetViperWithJWTSecret(t)
|
||||
t.Setenv("SECURITY_FORWARDED_CLIENT_IP_HEADERS", " x-cdn-ip , X-CDN-IP, true-client-ip ")
|
||||
|
||||
cfg, err := Load()
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, []string{"X-Cdn-Ip", "True-Client-Ip"}, cfg.ForwardedClientIPSettings().Headers)
|
||||
}
|
||||
|
||||
func TestLoadExplicitEmptyForwardedClientIPHeadersFromEnvironment(t *testing.T) {
|
||||
resetViperWithJWTSecret(t)
|
||||
viper.Set("security.forwarded_client_ip_headers", []string{"X-Yaml-IP"})
|
||||
t.Setenv("SECURITY_FORWARDED_CLIENT_IP_HEADERS", "")
|
||||
|
||||
cfg, err := Load()
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, cfg.ForwardedClientIPSettings().Headers)
|
||||
}
|
||||
|
||||
func TestLoadRejectsInvalidForwardedClientIPHeaderFromEnvironment(t *testing.T) {
|
||||
resetViperWithJWTSecret(t)
|
||||
t.Setenv("SECURITY_FORWARDED_CLIENT_IP_HEADERS", "X-Valid-IP, X Invalid")
|
||||
|
||||
_, err := Load()
|
||||
require.ErrorContains(t, err, "security.forwarded_client_ip_headers")
|
||||
}
|
||||
|
||||
func TestLoadRejectsInvalidForwardedClientIPHeader(t *testing.T) {
|
||||
resetViperWithJWTSecret(t)
|
||||
viper.Set("security.forwarded_client_ip_headers", []string{"X Invalid"})
|
||||
|
||||
_, err := Load()
|
||||
require.ErrorContains(t, err, "security.forwarded_client_ip_headers")
|
||||
}
|
||||
|
||||
func TestLoadExplicitEmptyTrustedProxiesEnablesConfiguredMode(t *testing.T) {
|
||||
resetViperWithJWTSecret(t)
|
||||
viper.Set("server.trusted_proxies", []string{})
|
||||
|
||||
Reference in New Issue
Block a user