package api import ( "bytes" "context" "encoding/json" "fmt" "net/http" "net/http/httptest" "os" "path/filepath" "testing" "github.com/redis/go-redis/v9" "booklib/internal/auth" "booklib/internal/db" "booklib/internal/redispkg" "booklib/internal/scanner" "booklib/internal/store" ) func setupAPI(t *testing.T) (*store.Store, *scanner.Scanner, http.Handler, string) { t.Helper() url := os.Getenv("DATABASE_URL") if url == "" { t.Skip("DATABASE_URL not set") } ctx := context.Background() p, _ := db.Connect(ctx, url) if err := db.Migrate(ctx, p); err != nil { t.Fatal(err) } st := store.New(p) p.Exec(ctx, "DELETE FROM reading_progress; DELETE FROM books; DELETE FROM libraries; DELETE FROM users") h, _ := auth.HashPassword(testPW) if _, err := st.CreateUser(ctx, "alice", h, "admin"); err != nil { t.Fatal(err) } if _, err := st.CreateUser(ctx, "bob", h, "member"); err != nil { t.Fatal(err) } booksParent := t.TempDir() booksDir, err := filepath.EvalSymlinks(booksParent) // macOS 上 /var→/private,root 校验要用真实路径 if err != nil { t.Fatal(err) } cfg := testCfg() cfg.BooksDir = booksDir cfg.CacheDir = t.TempDir() rdb := redispkg.New(os.Getenv("REDIS_URL")) sc := scanner.New(st, cfg, rdb) r := NewRouter(cfg, st, rdb, sc) if u := os.Getenv("REDIS_URL"); u != "" { // 测试卫生: 共享 redis 上重置登录限流桶, 防跨测试累计 429 if opt, e := redis.ParseURL(u); e == nil { rc := redis.NewClient(opt) rc.Del(ctx, "loginrl:192.0.2.1") rc.Close() } } return st, sc, r, booksDir } // testPW 是唯一的 fixture 口令常量: 直接种子 (HashPassword) 与所有登录/建户必须同值, 且 >=8 位 const testPW = "pw123456" func do(h http.Handler, method, path, token string, body any) *httptest.ResponseRecorder { var r *bytes.Reader if body != nil { b, _ := json.Marshal(body) r = bytes.NewReader(b) } else { r = bytes.NewReader(nil) } req := httptest.NewRequest(method, path, r) if token != "" { req.Header.Set("Authorization", "Bearer "+token) } w := httptest.NewRecorder() h.ServeHTTP(w, req) return w } func TestLoginMe(t *testing.T) { _, _, h, _ := setupAPI(t) w := do(h, "POST", "/api/auth/login", "", map[string]string{"username": "alice", "password": testPW}) if w.Code != 200 { t.Fatalf("login %d %s", w.Code, w.Body) } var tok struct{ Token string } json.Unmarshal(w.Body.Bytes(), &tok) if tok.Token == "" { t.Fatal("no token") } w = do(h, "GET", "/api/auth/me", tok.Token, nil) var me map[string]any json.Unmarshal(w.Body.Bytes(), &me) if w.Code != 200 || me["username"] != "alice" || me["role"] != "admin" { t.Fatalf("me %d %s", w.Code, w.Body) } // 错密码 → 401 统一错误体 w = do(h, "POST", "/api/auth/login", "", map[string]string{"username": "alice", "password": "nope"}) if w.Code != 401 { t.Fatalf("want 401 got %d", w.Code) } // 无 token / 坏 token 访问受保护端点 → 401(me 已注册;books 路由 Task 10 才有) if w = do(h, "GET", "/api/auth/me", "", nil); w.Code != 401 { t.Fatal("me without token must 401") } if w = do(h, "GET", "/api/auth/me", "garbage", nil); w.Code != 401 { t.Fatal("me with bad token must 401") } } // 伪造 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}) var v struct{ Token string } json.Unmarshal(w.Body.Bytes(), &v) tok := v.Token w = do(h, "POST", "/api/users", tok, map[string]string{"username": "eve", "password": testPW, "role": "member"}) if w.Code != 403 { t.Fatalf("member write users: want 403 got %d", w.Code) } }