feat(backend): jwt middleware, login with redis rate limit, /auth/me
This commit is contained in:
@@ -13,6 +13,7 @@ require (
|
|||||||
github.com/bytedance/gopkg v0.1.3 // indirect
|
github.com/bytedance/gopkg v0.1.3 // indirect
|
||||||
github.com/bytedance/sonic v1.15.0 // indirect
|
github.com/bytedance/sonic v1.15.0 // indirect
|
||||||
github.com/bytedance/sonic/loader v0.5.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/cloudwego/base64x v0.1.6 // indirect
|
||||||
github.com/gabriel-vasile/mimetype v1.4.12 // indirect
|
github.com/gabriel-vasile/mimetype v1.4.12 // indirect
|
||||||
github.com/gin-contrib/sse v1.1.0 // 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/pelletier/go-toml/v2 v2.2.4 // indirect
|
||||||
github.com/quic-go/qpack v0.6.0 // indirect
|
github.com/quic-go/qpack v0.6.0 // indirect
|
||||||
github.com/quic-go/quic-go v0.59.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/twitchyliquid64/golang-asm v0.15.1 // indirect
|
||||||
github.com/ugorji/go/codec v1.3.1 // indirect
|
github.com/ugorji/go/codec v1.3.1 // indirect
|
||||||
go.mongodb.org/mongo-driver/v2 v2.5.0 // 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/arch v0.22.0 // indirect
|
||||||
golang.org/x/net v0.57.0 // indirect
|
golang.org/x/net v0.57.0 // indirect
|
||||||
golang.org/x/sync v0.22.0 // indirect
|
golang.org/x/sync v0.22.0 // indirect
|
||||||
|
|||||||
@@ -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 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 h1:gXH3KVnatgY7loH5/TkeVyXPfESoqSBSBEiDd5VjlgE=
|
||||||
github.com/bytedance/sonic/loader v0.5.0/go.mod h1:AR4NYCk5DdzZizZ5djGqQ92eEhCCcdf5x77udYiSJRo=
|
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 h1:t11wG9AECkCDk5fMSoxmufanudBtJ+/HemLstXDLI2M=
|
||||||
github.com/cloudwego/base64x v0.1.6/go.mod h1:OFcloc187FXDaYHvrNIjxSe8ncn0OOM8gEHfghB2IPU=
|
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=
|
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/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 h1:OLJkp1Mlm/aS7dpKgTc6cnpynnD2Xg7C1pwL6vy/SAw=
|
||||||
github.com/quic-go/quic-go v0.59.0/go.mod h1:upnsH4Ju1YkqpLXC305eW3yDZ4NfnNbmQRCMWS58IKU=
|
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.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.4.0/go.mod h1:YvHI0jy2hoMjB+UWwv71VJQ9isScKT/TqJzVSSt89Yw=
|
||||||
github.com/stretchr/objx v0.5.0/go.mod h1:Yh+to48EsGEfYuaHDzXPcE3xhTkx73EhmCGUpEOglKo=
|
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=
|
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 h1:yXUhImUjjAInNcpTcAlPHiT7bIXhshCTL3jVBkF3xaE=
|
||||||
go.mongodb.org/mongo-driver/v2 v2.5.0/go.mod h1:yOI9kBsufol30iFsl1slpdq1I0eHPzybRWdyYUs8K/0=
|
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 h1:hyF9dfmbgIX5EfOdasqLsWD6xqpNZlXblLB/Dbnwv3Y=
|
||||||
go.uber.org/mock v0.6.0/go.mod h1:KiVJ4BqZJaMj4svdfmHM0AUx4NJYO8ZNpPnZn1Z+BBU=
|
go.uber.org/mock v0.6.0/go.mod h1:KiVJ4BqZJaMj4svdfmHM0AUx4NJYO8ZNpPnZn1Z+BBU=
|
||||||
golang.org/x/arch v0.22.0 h1:c/Zle32i5ttqRXjdLyyHZESLD/bB90DCU1g9l/0YBDI=
|
golang.org/x/arch v0.22.0 h1:c/Zle32i5ttqRXjdLyyHZESLD/bB90DCU1g9l/0YBDI=
|
||||||
|
|||||||
@@ -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"
|
"github.com/gin-gonic/gin"
|
||||||
|
|
||||||
"booklib/internal/config"
|
"booklib/internal/config"
|
||||||
|
"booklib/internal/redispkg"
|
||||||
|
"booklib/internal/store"
|
||||||
)
|
)
|
||||||
|
|
||||||
type api struct {
|
func NewRouter(cfg *config.Config, st *store.Store, rdb *redispkg.R) *gin.Engine {
|
||||||
cfg *config.Config
|
|
||||||
// st, rdb, sc 字段在 Task 4/9 加入
|
|
||||||
}
|
|
||||||
|
|
||||||
func NewRouter(cfg *config.Config) *gin.Engine {
|
|
||||||
gin.SetMode(gin.ReleaseMode)
|
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 := gin.New()
|
||||||
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") })
|
||||||
|
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
|
return r
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -7,6 +7,7 @@ import (
|
|||||||
"time"
|
"time"
|
||||||
|
|
||||||
"booklib/internal/config"
|
"booklib/internal/config"
|
||||||
|
"booklib/internal/redispkg"
|
||||||
)
|
)
|
||||||
|
|
||||||
func testCfg() *config.Config {
|
func testCfg() *config.Config {
|
||||||
@@ -14,7 +15,7 @@ func testCfg() *config.Config {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestHealthz(t *testing.T) {
|
func TestHealthz(t *testing.T) {
|
||||||
r := NewRouter(testCfg())
|
r := NewRouter(testCfg(), nil, redispkg.New(""))
|
||||||
req := httptest.NewRequest(http.MethodGet, "/api/healthz", nil)
|
req := httptest.NewRequest(http.MethodGet, "/api/healthz", nil)
|
||||||
w := httptest.NewRecorder()
|
w := httptest.NewRecorder()
|
||||||
r.ServeHTTP(w, req)
|
r.ServeHTTP(w, req)
|
||||||
|
|||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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()
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user