diff --git a/backend/internal/server/middleware/session_binding.go b/backend/internal/server/middleware/session_binding.go index 3f4d04b93..19fca2203 100644 --- a/backend/internal/server/middleware/session_binding.go +++ b/backend/internal/server/middleware/session_binding.go @@ -17,12 +17,12 @@ import ( // 接管解析,关闭时使用 Gin 的 server.trusted_proxies 可信代理链。 func SessionBindingContext(cfg *config.Config) gin.HandlerFunc { return func(c *gin.Context) { - trustForwarded := cfg.TrustForwardedIPForAPIKeyACL() - ip.SetLegacyForwardedIPTrust(c, trustForwarded) + forwardedIPSettings := cfg.ForwardedClientIPSettings() + ip.SetForwardedIPSettings(c, forwardedIPSettings.TrustForwardedIP, forwardedIPSettings.Headers) userAgent := normalizePersistentText(c.Request.UserAgent(), maxPersistentUserAgentBytes) c.Request.Header.Set("User-Agent", userAgent) binding := &service.SessionBinding{ - IP: ip.GetSecurityClientIP(c, trustForwarded), + IP: ip.GetSecurityClientIP(c, forwardedIPSettings.TrustForwardedIP), UserAgent: userAgent, } c.Request = c.Request.WithContext(service.WithSessionBinding(c.Request.Context(), binding)) diff --git a/backend/internal/server/middleware/session_binding_test.go b/backend/internal/server/middleware/session_binding_test.go index 11a098d9f..08e944878 100644 --- a/backend/internal/server/middleware/session_binding_test.go +++ b/backend/internal/server/middleware/session_binding_test.go @@ -8,6 +8,7 @@ import ( "testing" "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/Wei-Shaw/sub2api/internal/pkg/ip" "github.com/Wei-Shaw/sub2api/internal/service" "github.com/gin-gonic/gin" @@ -55,6 +56,39 @@ func TestSessionBindingContextFollowsForwardedIPSwitch(t *testing.T) { } } +func TestSessionBindingContextSnapshotsForwardedModeAndHeaders(t *testing.T) { + gin.SetMode(gin.TestMode) + + cfg := &config.Config{} + cfg.SetForwardedClientIPSettings(true, []string{"X-Initial-IP"}) + + r := gin.New() + require.NoError(t, r.SetTrustedProxies(nil)) + r.Use(SessionBindingContext(cfg)) + r.GET("/t", func(c *gin.Context) { + binding := service.SessionBindingFromContext(c.Request.Context()) + require.NotNil(t, binding) + require.Equal(t, "1.2.3.4", binding.IP) + + cfg.SetForwardedClientIPSettings(false, []string{"X-Changed-IP"}) + require.Equal(t, "1.2.3.4", ip.GetSecurityClientIP(c, false)) + c.Status(200) + }) + + w := httptest.NewRecorder() + req := httptest.NewRequest("GET", "/t", nil) + req.RemoteAddr = "9.9.9.9:12345" + req.Header.Set("X-Initial-IP", "1.2.3.4") + req.Header.Set("X-Changed-IP", "4.4.4.4") + req.Header.Set("X-Real-IP", "8.8.8.8") + r.ServeHTTP(w, req) + + require.Equal(t, 200, w.Code) + runtimeSettings := cfg.ForwardedClientIPSettings() + require.False(t, runtimeSettings.TrustForwardedIP) + require.Equal(t, []string{"X-Changed-IP"}, runtimeSettings.Headers) +} + func TestSessionBindingContextBoundsPersistedUserAgent(t *testing.T) { cfg := &config.Config{} r := gin.New()