feat(backend): jwt middleware, login with redis rate limit, /auth/me

This commit is contained in:
2026-09-04 23:58:26 +08:00
parent a2f202ada6
commit bc88cabc47
9 changed files with 341 additions and 8 deletions
+55
View File
@@ -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" }
+59
View File
@@ -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})
}
+104
View File
@@ -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)
}
}
+10 -7
View File
@@ -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
}
+2 -1
View File
@@ -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)