From 6adbab44c37bc34e4133303769074436b7d6e62a Mon Sep 17 00:00:00 2001 From: Fendy Date: Mon, 7 Sep 2026 12:44:47 +0800 Subject: [PATCH] fix(security): trusted-proxy-scoped login rate limit --- backend/internal/api/auth_test.go | 21 ++++++++++ backend/internal/api/router.go | 3 ++ backend/internal/api/router_test.go | 3 +- backend/internal/config/config.go | 53 ++++++++++++++++---------- backend/internal/config/config_test.go | 12 ++++++ 5 files changed, 71 insertions(+), 21 deletions(-) diff --git a/backend/internal/api/auth_test.go b/backend/internal/api/auth_test.go index 7708cda..dbbd9bc 100644 --- a/backend/internal/api/auth_test.go +++ b/backend/internal/api/auth_test.go @@ -4,6 +4,7 @@ import ( "bytes" "context" "encoding/json" + "fmt" "net/http" "net/http/httptest" "os" @@ -111,6 +112,26 @@ func TestLoginMe(t *testing.T) { } } +// 伪造 XFF 换不了限流桶: peer(192.0.2.1)不在可信代理段 → ClientIP 取 peer,XFF 忽略 +func TestLoginRateLimitResistsXFFSpoof(t *testing.T) { + if os.Getenv("REDIS_URL") == "" { + t.Skip("REDIS_URL not set (no redis → IncrWindow always allows)") + } + _, _, h, _ := setupAPI(t) // setupAPI 已重置 loginrl:192.0.2.1 + for i := 1; i <= 6; i++ { + req := httptest.NewRequest("POST", "/api/auth/login", bytes.NewBufferString(`{"username":"alice","password":"nope"}`)) + req.Header.Set("X-Forwarded-For", fmt.Sprintf("203.0.113.%d", i)) + w := httptest.NewRecorder() + h.ServeHTTP(w, req) + if i < 6 && w.Code != http.StatusUnauthorized { + t.Fatalf("attempt %d: want 401 got %d %s", i, w.Code, w.Body) + } + if i == 6 && w.Code != http.StatusTooManyRequests { + t.Fatalf("attempt 6: spoofed XFF escaped per-peer limit: want 429 got %d", w.Code) + } + } +} + func TestMemberCannotWriteUsers(t *testing.T) { _, _, h, _ := setupAPI(t) w := do(h, "POST", "/api/auth/login", "", map[string]string{"username": "bob", "password": testPW}) diff --git a/backend/internal/api/router.go b/backend/internal/api/router.go index 58c26a6..8a023a4 100644 --- a/backend/internal/api/router.go +++ b/backend/internal/api/router.go @@ -15,6 +15,9 @@ func NewRouter(cfg *config.Config, st *store.Store, rdb *redispkg.R, sc *scanner gin.SetMode(gin.ReleaseMode) a := &api{cfg: cfg, st: st, rdb: rdb, sc: sc} r := gin.New() + if e := r.SetTrustedProxies(cfg.TrustedProxies); e != nil { + panic(e) + } r.Use(gin.Recovery()) g := r.Group("/api") g.GET("/healthz", func(c *gin.Context) { c.String(http.StatusOK, "ok") }) diff --git a/backend/internal/api/router_test.go b/backend/internal/api/router_test.go index 6eb9c1f..f36db66 100644 --- a/backend/internal/api/router_test.go +++ b/backend/internal/api/router_test.go @@ -18,7 +18,8 @@ import ( ) func testCfg() *config.Config { - return &config.Config{Addr: ":8080", JWTSecret: []byte("s3cret"), ScanInterval: time.Minute, UploadMaxMB: 200} + return &config.Config{Addr: ":8080", JWTSecret: []byte("s3cret"), ScanInterval: time.Minute, UploadMaxMB: 200, + TrustedProxies: []string{"172.16.0.0/12"}} // 与 prod 默认一致: 只有 compose 网段内代理才可信 } func TestHealthz(t *testing.T) { diff --git a/backend/internal/config/config.go b/backend/internal/config/config.go index 236be76..8e06393 100644 --- a/backend/internal/config/config.go +++ b/backend/internal/config/config.go @@ -5,20 +5,22 @@ import ( "os" "path/filepath" "strconv" + "strings" "time" ) type Config struct { - Addr string - DatabaseURL string - RedisURL string - JWTSecret []byte - AdminUser string - AdminPassword string - BooksDir string - CacheDir string - ScanInterval time.Duration - UploadMaxMB int64 + Addr string + DatabaseURL string + RedisURL string + JWTSecret []byte + AdminUser string + AdminPassword string + BooksDir string + CacheDir string + ScanInterval time.Duration + UploadMaxMB int64 + TrustedProxies []string } func Load() (*Config, error) { @@ -28,6 +30,16 @@ func Load() (*Config, error) { } return def } + envList := func(k, def string) []string { + v := env(k, def) + out := make([]string, 0, len(strings.Split(v, ","))) + for _, s := range strings.Split(v, ",") { + if s = strings.TrimSpace(s); s != "" { + out = append(out, s) + } + } + return out + } scanSec, err := strconv.Atoi(env("SCAN_INTERVAL_SEC", "60")) if err != nil { return nil, fmt.Errorf("SCAN_INTERVAL_SEC: %w", err) @@ -49,15 +61,16 @@ func Load() (*Config, error) { return dir } return &Config{ - Addr: env("ADDR", ":8080"), - DatabaseURL: env("DATABASE_URL", ""), - RedisURL: env("REDIS_URL", ""), - JWTSecret: []byte(secret), - AdminUser: env("ADMIN_USER", ""), - AdminPassword: env("ADMIN_PASSWORD", ""), - BooksDir: resolveDir(env("BOOKS_DIR", "/data/books")), - CacheDir: resolveDir(env("CACHE_DIR", "/data/cache")), - ScanInterval: time.Duration(scanSec) * time.Second, - UploadMaxMB: uploadMB, + Addr: env("ADDR", ":8080"), + DatabaseURL: env("DATABASE_URL", ""), + RedisURL: env("REDIS_URL", ""), + JWTSecret: []byte(secret), + AdminUser: env("ADMIN_USER", ""), + AdminPassword: env("ADMIN_PASSWORD", ""), + BooksDir: resolveDir(env("BOOKS_DIR", "/data/books")), + CacheDir: resolveDir(env("CACHE_DIR", "/data/cache")), + ScanInterval: time.Duration(scanSec) * time.Second, + UploadMaxMB: uploadMB, + TrustedProxies: envList("TRUSTED_PROXY_CIDRS", "172.16.0.0/12"), }, nil } diff --git a/backend/internal/config/config_test.go b/backend/internal/config/config_test.go index 84e4d72..c6fd5c0 100644 --- a/backend/internal/config/config_test.go +++ b/backend/internal/config/config_test.go @@ -22,6 +22,7 @@ func TestLoad(t *testing.T) { } t.Setenv("SCAN_INTERVAL_SEC", "30") t.Setenv("DATABASE_URL", "postgres://x") + t.Setenv("TRUSTED_PROXY_CIDRS", "") c, err := Load() if err != nil { t.Fatal(err) @@ -34,4 +35,15 @@ func TestLoad(t *testing.T) { if c.ScanInterval != 30*time.Second || c.BooksDir != wantBooks || c.Addr != ":8080" { t.Fatalf("%+v", c) } + if len(c.TrustedProxies) != 1 || c.TrustedProxies[0] != "172.16.0.0/12" { + t.Fatalf("trusted proxies default: %+v", c.TrustedProxies) + } + t.Setenv("TRUSTED_PROXY_CIDRS", "10.0.0.0/8, 1.2.3.4") + c, err = Load() + if err != nil { + t.Fatal(err) + } + if len(c.TrustedProxies) != 2 || c.TrustedProxies[0] != "10.0.0.0/8" || c.TrustedProxies[1] != "1.2.3.4" { + t.Fatalf("trusted proxies override: %+v", c.TrustedProxies) + } }