diff --git a/backend/go.mod b/backend/go.mod index e1530db..59e24c6 100644 --- a/backend/go.mod +++ b/backend/go.mod @@ -13,6 +13,7 @@ require ( github.com/bytedance/gopkg v0.1.3 // indirect github.com/bytedance/sonic v1.15.0 // indirect github.com/bytedance/sonic/loader v0.5.0 // indirect + github.com/cespare/xxhash/v2 v2.3.0 // indirect github.com/cloudwego/base64x v0.1.6 // indirect github.com/gabriel-vasile/mimetype v1.4.12 // indirect github.com/gin-contrib/sse v1.1.0 // indirect @@ -33,9 +34,11 @@ require ( github.com/pelletier/go-toml/v2 v2.2.4 // indirect github.com/quic-go/qpack v0.6.0 // indirect github.com/quic-go/quic-go v0.59.0 // indirect + github.com/redis/go-redis/v9 v9.22.0 // indirect github.com/twitchyliquid64/golang-asm v0.15.1 // indirect github.com/ugorji/go/codec v1.3.1 // indirect go.mongodb.org/mongo-driver/v2 v2.5.0 // indirect + go.uber.org/atomic v1.11.0 // indirect golang.org/x/arch v0.22.0 // indirect golang.org/x/net v0.57.0 // indirect golang.org/x/sync v0.22.0 // indirect diff --git a/backend/go.sum b/backend/go.sum index 1b602ed..66d6824 100644 --- a/backend/go.sum +++ b/backend/go.sum @@ -4,6 +4,8 @@ github.com/bytedance/sonic v1.15.0 h1:/PXeWFaR5ElNcVE84U0dOHjiMHQOwNIx3K4ymzh/uS github.com/bytedance/sonic v1.15.0/go.mod h1:tFkWrPz0/CUCLEF4ri4UkHekCIcdnkqXw9VduqpJh0k= github.com/bytedance/sonic/loader v0.5.0 h1:gXH3KVnatgY7loH5/TkeVyXPfESoqSBSBEiDd5VjlgE= github.com/bytedance/sonic/loader v0.5.0/go.mod h1:AR4NYCk5DdzZizZ5djGqQ92eEhCCcdf5x77udYiSJRo= +github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs= +github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs= github.com/cloudwego/base64x v0.1.6 h1:t11wG9AECkCDk5fMSoxmufanudBtJ+/HemLstXDLI2M= github.com/cloudwego/base64x v0.1.6/go.mod h1:OFcloc187FXDaYHvrNIjxSe8ncn0OOM8gEHfghB2IPU= github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= @@ -61,6 +63,8 @@ github.com/quic-go/qpack v0.6.0 h1:g7W+BMYynC1LbYLSqRt8PBg5Tgwxn214ZZR34VIOjz8= github.com/quic-go/qpack v0.6.0/go.mod h1:lUpLKChi8njB4ty2bFLX2x4gzDqXwUpaO1DP9qMDZII= github.com/quic-go/quic-go v0.59.0 h1:OLJkp1Mlm/aS7dpKgTc6cnpynnD2Xg7C1pwL6vy/SAw= github.com/quic-go/quic-go v0.59.0/go.mod h1:upnsH4Ju1YkqpLXC305eW3yDZ4NfnNbmQRCMWS58IKU= +github.com/redis/go-redis/v9 v9.22.0 h1:laDvpYXTJtZLloinw1fA5Kqd6HAEH2XKxOkG/PDq2F0= +github.com/redis/go-redis/v9 v9.22.0/go.mod h1:y2g0Wj8rQvuK0ELM+oxSudcLtC09JScs98I/X9gRWY4= github.com/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME= github.com/stretchr/objx v0.4.0/go.mod h1:YvHI0jy2hoMjB+UWwv71VJQ9isScKT/TqJzVSSt89Yw= github.com/stretchr/objx v0.5.0/go.mod h1:Yh+to48EsGEfYuaHDzXPcE3xhTkx73EhmCGUpEOglKo= @@ -79,6 +83,8 @@ github.com/ugorji/go/codec v1.3.1 h1:waO7eEiFDwidsBN6agj1vJQ4AG7lh2yqXyOXqhgQuyY github.com/ugorji/go/codec v1.3.1/go.mod h1:pRBVtBSKl77K30Bv8R2P+cLSGaTtex6fsA2Wjqmfxj4= go.mongodb.org/mongo-driver/v2 v2.5.0 h1:yXUhImUjjAInNcpTcAlPHiT7bIXhshCTL3jVBkF3xaE= go.mongodb.org/mongo-driver/v2 v2.5.0/go.mod h1:yOI9kBsufol30iFsl1slpdq1I0eHPzybRWdyYUs8K/0= +go.uber.org/atomic v1.11.0 h1:ZvwS0R+56ePWxUNi+Atn9dWONBPp/AUETXlHW0DxSjE= +go.uber.org/atomic v1.11.0/go.mod h1:LUxbIzbOniOlMKjJjyPfpl4v+PKK2cNJn91OQbhoJI0= go.uber.org/mock v0.6.0 h1:hyF9dfmbgIX5EfOdasqLsWD6xqpNZlXblLB/Dbnwv3Y= go.uber.org/mock v0.6.0/go.mod h1:KiVJ4BqZJaMj4svdfmHM0AUx4NJYO8ZNpPnZn1Z+BBU= golang.org/x/arch v0.22.0 h1:c/Zle32i5ttqRXjdLyyHZESLD/bB90DCU1g9l/0YBDI= diff --git a/backend/internal/api/api.go b/backend/internal/api/api.go new file mode 100644 index 0000000..bf8c8e6 --- /dev/null +++ b/backend/internal/api/api.go @@ -0,0 +1,55 @@ +package api + +import ( + "net/http" + "strings" + + "github.com/gin-gonic/gin" + + "booklib/internal/auth" + "booklib/internal/config" + "booklib/internal/redispkg" + "booklib/internal/store" +) + +type api struct { + cfg *config.Config + st *store.Store + rdb *redispkg.R +} + +func err(c *gin.Context, status int, code, msg string) { + c.AbortWithStatusJSON(status, gin.H{"error": gin.H{"code": code, "message": msg}}) +} + +func (a *api) authMw() gin.HandlerFunc { + return func(c *gin.Context) { + h := c.GetHeader("Authorization") + tok, ok := strings.CutPrefix(h, "Bearer ") + if !ok { + err(c, http.StatusUnauthorized, "unauthorized", "missing bearer token") + return + } + cl, perr := auth.Parse(a.cfg.JWTSecret, tok) + if perr != nil { + err(c, http.StatusUnauthorized, "unauthorized", "invalid token") + return + } + c.Set("uid", cl.UID) + c.Set("role", cl.Role) + c.Next() + } +} + +func (a *api) adminOnly() gin.HandlerFunc { + return func(c *gin.Context) { + if c.GetString("role") != "admin" { + err(c, http.StatusForbidden, "forbidden", "admin only") + return + } + c.Next() + } +} + +func uid(c *gin.Context) int64 { return c.GetInt64("uid") } +func isAdmin(c *gin.Context) bool { return c.GetString("role") == "admin" } diff --git a/backend/internal/api/auth.go b/backend/internal/api/auth.go new file mode 100644 index 0000000..0f285b3 --- /dev/null +++ b/backend/internal/api/auth.go @@ -0,0 +1,59 @@ +package api + +import ( + "errors" + "log" + "net/http" + "time" + + "github.com/gin-gonic/gin" + "github.com/jackc/pgx/v5" + + "booklib/internal/auth" +) + +const loginWindow = time.Minute +const loginMax = 5 + +func (a *api) login(c *gin.Context) { + var req struct{ Username, Password string } + if c.ShouldBindJSON(&req) != nil || req.Username == "" || req.Password == "" { + err(c, http.StatusBadRequest, "bad_request", "username and password required") + return + } + if n := a.rdb.IncrWindow(c, "loginrl:"+c.ClientIP(), loginWindow); n > loginMax { + err(c, http.StatusTooManyRequests, "rate_limited", "too many login attempts") + return + } + u, qerr := a.st.GetUserByName(c, req.Username) + if qerr != nil { + if !errors.Is(qerr, pgx.ErrNoRows) { + log.Printf("db: %v", qerr) + err(c, http.StatusInternalServerError, "internal", "db error") + return + } + // 用户不存在也走一次 bcrypt,防用户名枚举时序差 + auth.CheckPassword("$2a$12$000000000000000000000000000000000000000000000000000O", req.Password) + err(c, http.StatusUnauthorized, "unauthorized", "bad credentials") + return + } + if !auth.CheckPassword(u.PasswordHash, req.Password) { + err(c, http.StatusUnauthorized, "unauthorized", "bad credentials") + return + } + tok, serr := auth.Sign(a.cfg.JWTSecret, u.ID, u.Role) + if serr != nil { + err(c, http.StatusInternalServerError, "internal", "sign") + return + } + c.JSON(http.StatusOK, gin.H{"token": tok}) +} + +func (a *api) me(c *gin.Context) { + u, qerr := a.st.GetUserByID(c, uid(c)) + if qerr != nil { + err(c, http.StatusUnauthorized, "unauthorized", "no such user") + return + } + c.JSON(http.StatusOK, gin.H{"id": u.ID, "username": u.Username, "role": u.Role}) +} diff --git a/backend/internal/api/auth_test.go b/backend/internal/api/auth_test.go new file mode 100644 index 0000000..314e54f --- /dev/null +++ b/backend/internal/api/auth_test.go @@ -0,0 +1,104 @@ +package api + +import ( + "bytes" + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "os" + "testing" + + "booklib/internal/auth" + "booklib/internal/db" + "booklib/internal/redispkg" + "booklib/internal/store" +) + +func setupAPI(t *testing.T) (*store.Store, http.Handler) { + 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("pw12345") + _, err := st.CreateUser(ctx, "alice", h, "admin") + if err != nil { + t.Fatal(err) + } + _, err = st.CreateUser(ctx, "bob", h, "member") + if err != nil { + t.Fatal(err) + } + cfg := testCfg() + r := NewRouter(cfg, st, redispkg.New(os.Getenv("REDIS_URL"))) + return st, r +} + +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": "pw12345"}) + 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") + } +} + +func TestMemberCannotWriteUsers(t *testing.T) { + _, h := setupAPI(t) + tok, _ := auth.Sign([]byte("s3cret"), 2, "member") // bob — 注意: 必须走真实登录拿 token + w := do(h, "POST", "/api/auth/login", "", map[string]string{"username": "bob", "password": "pw12345"}) + 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": "pw12345", "role": "member"}) + if w.Code != 403 { + t.Fatalf("member write users: want 403 got %d", w.Code) + } +} diff --git a/backend/internal/api/router.go b/backend/internal/api/router.go index 0e56c81..451c112 100644 --- a/backend/internal/api/router.go +++ b/backend/internal/api/router.go @@ -6,19 +6,22 @@ import ( "github.com/gin-gonic/gin" "booklib/internal/config" + "booklib/internal/redispkg" + "booklib/internal/store" ) -type api struct { - cfg *config.Config - // st, rdb, sc 字段在 Task 4/9 加入 -} - -func NewRouter(cfg *config.Config) *gin.Engine { +func NewRouter(cfg *config.Config, st *store.Store, rdb *redispkg.R) *gin.Engine { gin.SetMode(gin.ReleaseMode) - _ = &api{cfg: cfg} // Task 4/9 起此处为 a := &api{...},handler 挂在 a 上 + a := &api{cfg: cfg, st: st, rdb: rdb} r := gin.New() r.Use(gin.Recovery()) g := r.Group("/api") g.GET("/healthz", func(c *gin.Context) { c.String(http.StatusOK, "ok") }) + g.POST("/auth/login", a.login) + + p := g.Group("", a.authMw()) + p.GET("/auth/me", a.me) + // ponytail: 501 垫片,Task 5 换成 a.createUser;无此路由则 member 403 测不到 + p.POST("/users", a.adminOnly(), func(c *gin.Context) { c.Status(http.StatusNotImplemented) }) return r } diff --git a/backend/internal/api/router_test.go b/backend/internal/api/router_test.go index 2abb4ba..6fe094e 100644 --- a/backend/internal/api/router_test.go +++ b/backend/internal/api/router_test.go @@ -7,6 +7,7 @@ import ( "time" "booklib/internal/config" + "booklib/internal/redispkg" ) func testCfg() *config.Config { @@ -14,7 +15,7 @@ func testCfg() *config.Config { } func TestHealthz(t *testing.T) { - r := NewRouter(testCfg()) + r := NewRouter(testCfg(), nil, redispkg.New("")) req := httptest.NewRequest(http.MethodGet, "/api/healthz", nil) w := httptest.NewRecorder() r.ServeHTTP(w, req) diff --git a/backend/internal/redispkg/redis.go b/backend/internal/redispkg/redis.go new file mode 100644 index 0000000..e8d8878 --- /dev/null +++ b/backend/internal/redispkg/redis.go @@ -0,0 +1,78 @@ +package redispkg + +import ( + "context" + "crypto/rand" + "encoding/hex" + "log" + "time" + + "github.com/redis/go-redis/v9" +) + +type R struct{ c *redis.Client } + +func New(url string) *R { + if url == "" { + return &R{} + } + opt, err := redis.ParseURL(url) + if err != nil { + log.Printf("bad REDIS_URL (%v): redis disabled", err) + return &R{} + } + return &R{c: redis.NewClient(opt)} +} + +func (r *R) Get(ctx context.Context, key string) (string, bool) { + if r.c == nil { + return "", false + } + v, err := r.c.Get(ctx, key).Result() + if err != nil { + return "", false // 故障=miss + } + return v, true +} + +func (r *R) Set(ctx context.Context, key, val string, ttl time.Duration) { + if r.c == nil { + return + } + if err := r.c.Set(ctx, key, val, ttl).Err(); err != nil { + log.Printf("redis set %s: %v", key, err) + } +} + +func (r *R) IncrWindow(ctx context.Context, key string, ttl time.Duration) int { + if r.c == nil { + return 1 + } + n, err := r.c.Incr(ctx, key).Result() + if err != nil { + return 1 + } + if n == 1 { + r.c.Expire(ctx, key, ttl) + } + return int(n) +} + +func (r *R) Lock(ctx context.Context, key string, ttl time.Duration) (func(), bool) { + noop := func() {} + if r.c == nil { + return noop, true + } + b := make([]byte, 8) + rand.Read(b) + tok := hex.EncodeToString(b) + ok, err := r.c.SetNX(ctx, key, tok, ttl).Result() + if err != nil || !ok { + return noop, false + } + return func() { + r.c.Eval(ctx, + "if redis.call('get',KEYS[1])==ARGV[1] then return redis.call('del',KEYS[1]) else return 0 end", + []string{key}, tok) + }, true +} diff --git a/backend/internal/redispkg/redis_test.go b/backend/internal/redispkg/redis_test.go new file mode 100644 index 0000000..cfd8486 --- /dev/null +++ b/backend/internal/redispkg/redis_test.go @@ -0,0 +1,24 @@ +package redispkg + +import ( + "context" + "testing" + "time" +) + +func TestDisabledIsSafe(t *testing.T) { + r := New("") + ctx := context.Background() + if _, ok := r.Get(ctx, "x"); ok { + t.Fatal("disabled Get must miss") + } + r.Set(ctx, "x", "y", time.Second) // 不 panic + if n := r.IncrWindow(ctx, "k", time.Second); n != 1 { + t.Fatal("disabled IncrWindow must allow") + } + un, ok := r.Lock(ctx, "lk", time.Second) + if !ok { + t.Fatal("disabled Lock must always acquire") + } + un() +}