fix(security): trusted-proxy-scoped login rate limit

This commit is contained in:
2026-09-07 12:44:47 +08:00
parent f7c27272ba
commit 6adbab44c3
5 changed files with 71 additions and 21 deletions
+21
View File
@@ -4,6 +4,7 @@ import (
"bytes" "bytes"
"context" "context"
"encoding/json" "encoding/json"
"fmt"
"net/http" "net/http"
"net/http/httptest" "net/http/httptest"
"os" "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) { func TestMemberCannotWriteUsers(t *testing.T) {
_, _, h, _ := setupAPI(t) _, _, h, _ := setupAPI(t)
w := do(h, "POST", "/api/auth/login", "", map[string]string{"username": "bob", "password": testPW}) w := do(h, "POST", "/api/auth/login", "", map[string]string{"username": "bob", "password": testPW})
+3
View File
@@ -15,6 +15,9 @@ func NewRouter(cfg *config.Config, st *store.Store, rdb *redispkg.R, sc *scanner
gin.SetMode(gin.ReleaseMode) gin.SetMode(gin.ReleaseMode)
a := &api{cfg: cfg, st: st, rdb: rdb, sc: sc} a := &api{cfg: cfg, st: st, rdb: rdb, sc: sc}
r := gin.New() r := gin.New()
if e := r.SetTrustedProxies(cfg.TrustedProxies); e != nil {
panic(e)
}
r.Use(gin.Recovery()) r.Use(gin.Recovery())
g := r.Group("/api") g := r.Group("/api")
g.GET("/healthz", func(c *gin.Context) { c.String(http.StatusOK, "ok") }) g.GET("/healthz", func(c *gin.Context) { c.String(http.StatusOK, "ok") })
+2 -1
View File
@@ -18,7 +18,8 @@ import (
) )
func testCfg() *config.Config { 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) { func TestHealthz(t *testing.T) {
+13
View File
@@ -5,6 +5,7 @@ import (
"os" "os"
"path/filepath" "path/filepath"
"strconv" "strconv"
"strings"
"time" "time"
) )
@@ -19,6 +20,7 @@ type Config struct {
CacheDir string CacheDir string
ScanInterval time.Duration ScanInterval time.Duration
UploadMaxMB int64 UploadMaxMB int64
TrustedProxies []string
} }
func Load() (*Config, error) { func Load() (*Config, error) {
@@ -28,6 +30,16 @@ func Load() (*Config, error) {
} }
return def 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")) scanSec, err := strconv.Atoi(env("SCAN_INTERVAL_SEC", "60"))
if err != nil { if err != nil {
return nil, fmt.Errorf("SCAN_INTERVAL_SEC: %w", err) return nil, fmt.Errorf("SCAN_INTERVAL_SEC: %w", err)
@@ -59,5 +71,6 @@ func Load() (*Config, error) {
CacheDir: resolveDir(env("CACHE_DIR", "/data/cache")), CacheDir: resolveDir(env("CACHE_DIR", "/data/cache")),
ScanInterval: time.Duration(scanSec) * time.Second, ScanInterval: time.Duration(scanSec) * time.Second,
UploadMaxMB: uploadMB, UploadMaxMB: uploadMB,
TrustedProxies: envList("TRUSTED_PROXY_CIDRS", "172.16.0.0/12"),
}, nil }, nil
} }
+12
View File
@@ -22,6 +22,7 @@ func TestLoad(t *testing.T) {
} }
t.Setenv("SCAN_INTERVAL_SEC", "30") t.Setenv("SCAN_INTERVAL_SEC", "30")
t.Setenv("DATABASE_URL", "postgres://x") t.Setenv("DATABASE_URL", "postgres://x")
t.Setenv("TRUSTED_PROXY_CIDRS", "")
c, err := Load() c, err := Load()
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
@@ -34,4 +35,15 @@ func TestLoad(t *testing.T) {
if c.ScanInterval != 30*time.Second || c.BooksDir != wantBooks || c.Addr != ":8080" { if c.ScanInterval != 30*time.Second || c.BooksDir != wantBooks || c.Addr != ":8080" {
t.Fatalf("%+v", c) 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)
}
} }