fix(security): trusted-proxy-scoped login rate limit
This commit is contained in:
@@ -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})
|
||||||
|
|||||||
@@ -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") })
|
||||||
|
|||||||
@@ -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) {
|
||||||
|
|||||||
@@ -5,20 +5,22 @@ import (
|
|||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
"strconv"
|
"strconv"
|
||||||
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
)
|
)
|
||||||
|
|
||||||
type Config struct {
|
type Config struct {
|
||||||
Addr string
|
Addr string
|
||||||
DatabaseURL string
|
DatabaseURL string
|
||||||
RedisURL string
|
RedisURL string
|
||||||
JWTSecret []byte
|
JWTSecret []byte
|
||||||
AdminUser string
|
AdminUser string
|
||||||
AdminPassword string
|
AdminPassword string
|
||||||
BooksDir string
|
BooksDir string
|
||||||
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)
|
||||||
@@ -49,15 +61,16 @@ func Load() (*Config, error) {
|
|||||||
return dir
|
return dir
|
||||||
}
|
}
|
||||||
return &Config{
|
return &Config{
|
||||||
Addr: env("ADDR", ":8080"),
|
Addr: env("ADDR", ":8080"),
|
||||||
DatabaseURL: env("DATABASE_URL", ""),
|
DatabaseURL: env("DATABASE_URL", ""),
|
||||||
RedisURL: env("REDIS_URL", ""),
|
RedisURL: env("REDIS_URL", ""),
|
||||||
JWTSecret: []byte(secret),
|
JWTSecret: []byte(secret),
|
||||||
AdminUser: env("ADMIN_USER", ""),
|
AdminUser: env("ADMIN_USER", ""),
|
||||||
AdminPassword: env("ADMIN_PASSWORD", ""),
|
AdminPassword: env("ADMIN_PASSWORD", ""),
|
||||||
BooksDir: resolveDir(env("BOOKS_DIR", "/data/books")),
|
BooksDir: resolveDir(env("BOOKS_DIR", "/data/books")),
|
||||||
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
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user