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:
Jlypx
2026-07-20 00:08:26 +08:00
co-authored by Sisyphus
parent f69042ca6e
commit 041db5d824
2 changed files with 208 additions and 16 deletions
+86 -16
View File
@@ -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")
}
+122
View File
@@ -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{})