feat(backend): jwt middleware, login with redis rate limit, /auth/me
This commit is contained in:
@@ -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" }
|
||||
@@ -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})
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user