feat: 建立可靠通知、吊销与跨实例状态

This commit is contained in:
weijuesen
2026-08-14 01:41:14 +08:00
parent a6a996a518
commit a25bc7abc6
22 changed files with 1017 additions and 245 deletions
+2
View File
@@ -36,6 +36,8 @@ func Init(cfg *config.Config) error {
&model.Permission{}, &model.RolePermission{},
&model.Disease{}, &model.KnowledgeArticle{},
&model.InspectionRecord{},
&model.OutboxEvent{},
&model.Notification{},
&model.Tray{}, &model.Batch{}, &model.RearingRecord{},
&model.WechatBinding{},
&model.WeatherAlert{},
+1 -1
View File
@@ -13,7 +13,7 @@ import (
)
// CurrentSchemaVersion 是当前后端代码期望的迁移版本。
const CurrentSchemaVersion = "2"
const CurrentSchemaVersion = "4"
// RunMigrations 使用嵌入式 SQL 迁移文件将数据库升级到最新版本。
func RunMigrations(db *gorm.DB) error {
@@ -104,4 +104,8 @@ func TestEmbeddedMigrationsIncludeBaseline(t *testing.T) {
if err != nil || next != 2 {
t.Fatalf("expected risk assessment migration version 2, got %d (err %v)", next, err)
}
next, err = driver.Next(next)
if err != nil || next != 4 {
t.Fatalf("expected notifications/outbox migration version 4, got %d (err %v)", next, err)
}
}
+27 -7
View File
@@ -72,7 +72,7 @@ func registerHandler(db *gorm.DB, cfg *config.Config) gin.HandlerFunc {
}
// 强制角色为 viewer,防止垂直越权(注册接口不允许自选角色)
role := model.RoleViewer
role := model.RoleViewer
user := model.User{
Username: body.Username,
Email: body.Email,
@@ -123,13 +123,19 @@ func loginHandler(db *gorm.DB, cfg *config.Config) gin.HandlerFunc {
var user model.User
if db.Where("username = ? OR email = ?", body.Username, body.Username).First(&user).Error != nil {
middleware.RecordLoginFail(c, body.Username)
if err := middleware.RecordLoginFail(c, body.Username); err != nil {
c.JSON(http.StatusServiceUnavailable, gin.H{"error": err.Error()})
return
}
c.JSON(http.StatusUnauthorized, gin.H{"error": "用户名或密码错误"})
return
}
if err := bcrypt.CompareHashAndPassword([]byte(user.PasswordHash), []byte(body.Password)); err != nil {
middleware.RecordLoginFail(c, body.Username)
if err := middleware.RecordLoginFail(c, body.Username); err != nil {
c.JSON(http.StatusServiceUnavailable, gin.H{"error": err.Error()})
return
}
c.JSON(http.StatusUnauthorized, gin.H{"error": "用户名或密码错误"})
return
}
@@ -140,7 +146,10 @@ func loginHandler(db *gorm.DB, cfg *config.Config) gin.HandlerFunc {
}
// 登录成功,清空失败计数
middleware.RecordLoginSuccess(c, body.Username)
if err := middleware.RecordLoginSuccess(c, body.Username); err != nil {
c.JSON(http.StatusServiceUnavailable, gin.H{"error": err.Error()})
return
}
// 记录审计日志
uid := user.ID
@@ -184,7 +193,12 @@ func refreshHandler(db *gorm.DB, cfg *config.Config) gin.HandlerFunc {
c.JSON(http.StatusUnauthorized, gin.H{"error": "刷新令牌无效"})
return
}
if middleware.IsRevoked(claims) {
revoked, err := middleware.IsRevoked(claims)
if err != nil {
c.JSON(http.StatusServiceUnavailable, gin.H{"error": "认证状态服务不可用"})
return
}
if revoked {
c.JSON(http.StatusUnauthorized, gin.H{"error": "刷新令牌已注销"})
return
}
@@ -201,7 +215,10 @@ func refreshHandler(db *gorm.DB, cfg *config.Config) gin.HandlerFunc {
}
// 吊销旧刷新令牌(一次性使用),签发新令牌对
middleware.RevokeToken(claims, body.RefreshToken, claims.ExpiresAt.Time)
if err := middleware.RevokeToken(claims, body.RefreshToken, claims.ExpiresAt.Time); err != nil {
c.JSON(http.StatusServiceUnavailable, gin.H{"error": err.Error()})
return
}
c.JSON(http.StatusOK, buildLoginPayload(db, user, cfg))
}
}
@@ -214,7 +231,10 @@ func logoutHandler(cfg *config.Config) gin.HandlerFunc {
if len(parts) == 2 && strings.EqualFold(parts[0], "Bearer") {
if claims, token, err := middleware.ExtractClaims(parts[1], cfg.JWTSecret); err == nil && token.Valid {
if claims.ExpiresAt != nil {
middleware.RevokeToken(claims, parts[1], claims.ExpiresAt.Time)
if err := middleware.RevokeToken(claims, parts[1], claims.ExpiresAt.Time); err != nil {
c.JSON(http.StatusServiceUnavailable, gin.H{"error": err.Error()})
return
}
}
}
}
+35 -33
View File
@@ -2,8 +2,8 @@ package handler
import (
"bytes"
"context"
"encoding/json"
"fmt"
"io"
"log/slog"
"net/http"
@@ -27,8 +27,8 @@ func isUUID(s string) bool {
}
// RegisterInspectionRoutes 注册 AI 巡检路由
func RegisterInspectionRoutes(rg *gin.RouterGroup, db *gorm.DB, s3 *service.S3Service, ai *service.AIClient, imageBucket string, wechat *service.WechatService, inspectionTemplateID string, appEnv string) {
rg.POST("/inspections", middleware.RequirePermission(db, "inspection:create"), createInspection(db, s3, ai, imageBucket, wechat, inspectionTemplateID, appEnv))
func RegisterInspectionRoutes(rg *gin.RouterGroup, db *gorm.DB, s3 *service.S3Service, ai *service.AIClient, imageBucket string, outbox *service.Outbox, inspectionTemplateID string, appEnv string) {
rg.POST("/inspections", middleware.RequirePermission(db, "inspection:create"), createInspection(db, s3, ai, imageBucket, outbox, inspectionTemplateID, appEnv))
rg.GET("/inspections", middleware.RequirePermission(db, "inspection:read"), listInspections(db))
}
@@ -66,7 +66,7 @@ func buildRiskInput(detRes *service.AIDetectResponse) service.RiskInput {
// createInspection 拍照巡检:图片存 S3 → 调 AI /detect → 写记录。
// 幂等:客户端传 Idempotency-Key 头时,重复请求返回已有记录。
func createInspection(db *gorm.DB, s3 *service.S3Service, ai *service.AIClient, bucket string, wechat *service.WechatService, inspectionTemplateID string, appEnv string) gin.HandlerFunc {
func createInspection(db *gorm.DB, s3 *service.S3Service, ai *service.AIClient, bucket string, outbox *service.Outbox, inspectionTemplateID string, appEnv string) gin.HandlerFunc {
return func(c *gin.Context) {
idemKey := strings.TrimSpace(c.GetHeader("Idempotency-Key"))
roomID := strings.TrimSpace(c.PostForm("roomId"))
@@ -155,38 +155,40 @@ func createInspection(db *gorm.DB, s3 *service.S3Service, ai *service.AIClient,
rawRisk, _ := json.Marshal(assessment)
rec.RiskAssessment = rawRisk
// 微信订阅消息(#11 骨架):Mock 结果不进入告警,风险非绿且用户已授权时异步推送
if !isMock {
if key := service.WechatTemplateKey(assessment.Level); key != "" {
go func(uid *string, lv string, sc float64) {
if uid == nil || !wechat.Configured() || inspectionTemplateID == "" {
return
}
var binding model.WechatBinding
if db.Where("user_id = ?", *uid).First(&binding).Error != nil {
return
}
var authorized []string
if len(binding.AuthorizedTemplates) > 0 {
_ = json.Unmarshal(binding.AuthorizedTemplates, &authorized)
}
if !service.IsAuthorized(authorized, key) {
return
}
_ = wechat.SendSubscribe(
context.Background(),
binding.OpenID,
inspectionTemplateID,
service.BuildSubscribeData(lv, sc),
"pages/inspection/index",
)
}(rec.UserID, assessment.Level, assessment.Score)
}
}
}
}
if err := db.Create(&rec).Error; err != nil {
// 微信订阅消息(#11 骨架):Mock 不进入告警;业务事务内写 outbox,重启后仍可重试
txErr := db.Transaction(func(tx *gorm.DB) error {
if err := tx.Create(&rec).Error; err != nil {
return err
}
if rec.AIStatus == "done" && rec.UserID != nil && rec.IsMock != nil && !*rec.IsMock {
if key := service.WechatTemplateKey(*rec.RiskLevel); key != "" {
payload, _ := json.Marshal(service.WechatSubscribePayload{
UserID: *rec.UserID,
TemplateKey: key,
TemplateID: inspectionTemplateID,
Page: "pages/inspection/index",
Title: "巡检风险提醒",
Body: fmt.Sprintf("风险等级 %s,风险分 %.0f", *rec.RiskLevel, *rec.RiskScore),
Data: service.BuildSubscribeData(*rec.RiskLevel, *rec.RiskScore),
BusinessType: "inspection",
BusinessID: rec.ID,
})
event := service.Event{
ID: "inspection-" + rec.ID + "-" + *rec.RiskLevel,
Type: service.OutboxEventWechatSubscribe,
AggregateType: "inspection",
AggregateID: rec.ID,
Payload: payload,
}
return outbox.PublishTx(tx, event)
}
}
return nil
})
if txErr != nil {
// 并发幂等:唯一索引冲突时返回已有记录
if idemKey != "" {
var exist model.InspectionRecord
+25 -51
View File
@@ -1,61 +1,34 @@
package handler
import (
"fmt"
"math/rand"
"net/http"
"sync"
"time"
"silk-server-go/internal/middleware"
"silk-server-go/internal/model"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
)
// notificationItem 通知项
type notificationItem struct {
ID string `json:"id"`
Channel string `json:"channel"`
Target string `json:"target"`
Title string `json:"title"`
Body string `json:"body"`
CreatedAt string `json:"createdAt"`
}
// 通知内存存储(后续可替换为 Redis)
var (
notificationStore []notificationItem
notificationMu sync.Mutex
)
// RegisterNotificationRoutes 注册通知路由
func RegisterNotificationRoutes(rg *gin.RouterGroup, db *gorm.DB) {
readPerm := middleware.RequirePermission(db, "alarm:read")
rg.GET("/notifications", readPerm, listNotifications())
rg.POST("/notifications", readPerm, createNotification())
rg.GET("/notifications", readPerm, listNotifications(db))
rg.POST("/notifications", readPerm, createNotification(db))
}
// listNotifications 通知列表(上限200
func listNotifications() gin.HandlerFunc {
func listNotifications(db *gorm.DB) gin.HandlerFunc {
return func(c *gin.Context) {
notificationMu.Lock()
defer notificationMu.Unlock()
limit := 200
if len(notificationStore) < limit {
limit = len(notificationStore)
}
// 返回最新的 limit 条(存储已按新到旧排序)
result := make([]notificationItem, limit)
copy(result, notificationStore[:limit])
c.JSON(http.StatusOK, result)
var list []model.Notification
db.Order("created_at DESC").Limit(200).Find(&list)
c.JSON(http.StatusOK, list)
}
}
// createNotification 手动发通知
func createNotification() gin.HandlerFunc {
func createNotification(db *gorm.DB) gin.HandlerFunc {
return func(c *gin.Context) {
var body struct {
Channel string `json:"channel"`
@@ -68,24 +41,25 @@ func createNotification() gin.HandlerFunc {
return
}
ntf := notificationItem{
ID: fmt.Sprintf("%d%d", time.Now().UnixNano(), rand.Intn(1000000)),
Channel: body.Channel,
Target: body.Target,
Title: body.Title,
Body: body.Body,
CreatedAt: time.Now().Format(time.RFC3339),
ntf := model.Notification{
UserID: currentUserID(c),
Channel: body.Channel,
Target: body.Target,
Title: body.Title,
Body: body.Body,
Status: "sent",
NextAttemptAt: time.Now(),
}
notificationMu.Lock()
// 插入到头部(最新在前)
notificationStore = append([]notificationItem{ntf}, notificationStore...)
// 保留最近 500 条
if len(notificationStore) > 500 {
notificationStore = notificationStore[:500]
if ntf.Channel == "" {
ntf.Channel = "manual"
}
if ntf.Target == "" {
ntf.Target = "all"
}
if err := db.Create(&ntf).Error; err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "创建通知失败"})
return
}
notificationMu.Unlock()
c.JSON(http.StatusCreated, ntf)
}
}
+16 -11
View File
@@ -20,16 +20,16 @@ type JWTClaims struct {
// 白名单路径,无需鉴权
var whitelist = map[string]bool{
"/api/v1/health": true,
"/health": true,
"/api/health": true,
"/api/v1/auth/login": true,
"/auth/login": true,
"/api/v1/auth/register": true,
"/auth/register": true,
"/api/v1/auth/refresh": true,
"/auth/refresh": true,
"/api/v1/video/clips/internal": true,
"/api/v1/health": true,
"/health": true,
"/api/health": true,
"/api/v1/auth/login": true,
"/auth/login": true,
"/api/v1/auth/register": true,
"/auth/register": true,
"/api/v1/auth/refresh": true,
"/auth/refresh": true,
"/api/v1/video/clips/internal": true,
"/api/v1/video/recordings/internal/end": true,
}
@@ -83,7 +83,12 @@ func Auth(cfg *config.Config) gin.HandlerFunc {
}
// 校验令牌是否已被登出吊销
if IsRevoked(claims) {
revoked, err := IsRevoked(claims)
if err != nil {
c.AbortWithStatusJSON(http.StatusServiceUnavailable, gin.H{"error": "认证状态服务不可用"})
return
}
if revoked {
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"error": "令牌已注销,请重新登录"})
return
}
+26 -89
View File
@@ -1,110 +1,41 @@
package middleware
import (
"context"
"net/http"
"strings"
"sync"
"time"
"github.com/gin-gonic/gin"
)
// loginAttempt 登录失败计数(按 IP + 用户名维度)
type loginAttempt struct {
failures int
lockUntil time.Time
lastFail time.Time
}
type loginLimiter struct {
mu sync.Mutex
seen map[string]*loginAttempt
}
const (
maxFailures = 5 // 连续失败 5 次后锁定
lockDuration = 15 * time.Minute
failureWindow = 10 * time.Minute // 失败计数窗口
cleanupInterval = 5 * time.Minute
maxFailures = 5 // 连续失败 5 次后锁定
lockDuration = 15 * time.Minute
failureWindow = 10 * time.Minute // 失败计数窗口
)
var defaultLoginLimiter = newLoginLimiter()
func newLoginLimiter() *loginLimiter {
l := &loginLimiter{seen: make(map[string]*loginAttempt)}
go l.cleanupLoop()
return l
}
func (l *loginLimiter) cleanupLoop() {
t := time.NewTicker(cleanupInterval)
defer t.Stop()
for range t.C {
l.mu.Lock()
now := time.Now()
for k, v := range l.seen {
if now.After(v.lockUntil) && now.Sub(v.lastFail) > failureWindow {
delete(l.seen, k)
}
}
l.mu.Unlock()
}
}
// key = ip + "|" + username(小写)
func limiterKey(c *gin.Context, username string) string {
return c.ClientIP() + "|" + strings.ToLower(strings.TrimSpace(username))
}
// checkLock 返回是否被锁定及剩余锁定时间
func (l *loginLimiter) checkLock(key string) (bool, time.Duration) {
l.mu.Lock()
defer l.mu.Unlock()
a, ok := l.seen[key]
if !ok {
return false, 0
}
if time.Now().Before(a.lockUntil) {
return true, time.Until(a.lockUntil)
}
return false, 0
}
// recordFailure 记录一次失败,达到阈值则锁定
func (l *loginLimiter) recordFailure(key string) {
l.mu.Lock()
defer l.mu.Unlock()
a, ok := l.seen[key]
if !ok {
a = &loginAttempt{}
l.seen[key] = a
}
now := time.Now()
// 窗口外重置
if now.Sub(a.lastFail) > failureWindow {
a.failures = 0
}
a.failures++
a.lastFail = now
if a.failures >= maxFailures {
a.lockUntil = now.Add(lockDuration)
}
}
// recordSuccess 登录成功后清空计数
func (l *loginLimiter) recordSuccess(key string) {
l.mu.Lock()
delete(l.seen, key)
l.mu.Unlock()
}
// CheckLoginLock 检查是否被锁定,被锁定则写 429 并返回 true(在 handler 解析 body 后调用)
func CheckLoginLock(c *gin.Context, username string) bool {
if authState == nil {
c.AbortWithStatusJSON(http.StatusServiceUnavailable, gin.H{"error": stateUnavailable("执行登录限流").Error()})
return true
}
key := limiterKey(c, username)
if locked, remain := defaultLoginLimiter.checkLock(key); locked {
locked, remain, err := authState.CheckLoginLock(context.Background(), key)
if err != nil {
c.AbortWithStatusJSON(http.StatusServiceUnavailable, gin.H{"error": "登录状态服务不可用"})
return true
}
if locked {
c.AbortWithStatusJSON(http.StatusTooManyRequests, gin.H{
"error": "登录尝试过多,已锁定,请稍后再试",
"retry": int(remain.Minutes()) + 1,
"error": "登录尝试过多,已锁定,请稍后再试",
"retry": int(remain.Minutes()) + 1,
})
return true
}
@@ -112,11 +43,17 @@ func CheckLoginLock(c *gin.Context, username string) bool {
}
// RecordLoginFail 记录登录失败
func RecordLoginFail(c *gin.Context, username string) {
defaultLoginLimiter.recordFailure(limiterKey(c, username))
func RecordLoginFail(c *gin.Context, username string) error {
if authState == nil {
return stateUnavailable("记录登录失败")
}
return authState.RecordLoginFailure(context.Background(), limiterKey(c, username))
}
// RecordLoginSuccess 登录成功后清空计数
func RecordLoginSuccess(c *gin.Context, username string) {
defaultLoginLimiter.recordSuccess(limiterKey(c, username))
func RecordLoginSuccess(c *gin.Context, username string) error {
if authState == nil {
return stateUnavailable("清空登录失败计数")
}
return authState.RecordLoginSuccess(context.Background(), limiterKey(c, username))
}
+104
View File
@@ -0,0 +1,104 @@
package middleware
import (
"context"
"fmt"
"log/slog"
"time"
"github.com/redis/go-redis/v9"
)
const statePrefix = "silk:auth:"
// StateStore 跨实例认证状态存储。
type StateStore interface {
RevokeToken(ctx context.Context, id string, exp time.Time) error
IsTokenRevoked(ctx context.Context, id string) (bool, error)
CheckLoginLock(ctx context.Context, key string) (bool, time.Duration, error)
RecordLoginFailure(ctx context.Context, key string) error
RecordLoginSuccess(ctx context.Context, key string) error
}
var (
authState StateStore
appEnv string
)
// InitState 设置认证状态存储;store 为 nil 时认证相关接口保守失败。
func InitState(store StateStore, env string) {
authState = store
appEnv = env
}
// RedisState Redis 实现。
type RedisState struct {
rdb *redis.Client
prefix string
}
// NewRedisState 创建 Redis 状态存储。
func NewRedisState(rdb *redis.Client) *RedisState {
return &RedisState{rdb: rdb, prefix: statePrefix}
}
func (s *RedisState) RevokeToken(ctx context.Context, id string, exp time.Time) error {
ttl := time.Until(exp)
if ttl <= 0 {
return nil
}
return s.rdb.Set(ctx, s.prefix+"revoked:"+id, "1", ttl).Err()
}
func (s *RedisState) IsTokenRevoked(ctx context.Context, id string) (bool, error) {
count, err := s.rdb.Exists(ctx, s.prefix+"revoked:"+id).Result()
if err != nil {
return false, err
}
return count > 0, nil
}
func (s *RedisState) CheckLoginLock(ctx context.Context, key string) (bool, time.Duration, error) {
lockKey := s.prefix + "login-lock:" + key
if _, err := s.rdb.Get(ctx, lockKey).Result(); err == redis.Nil {
return false, 0, nil
} else if err != nil {
return false, 0, err
}
ttl, err := s.rdb.TTL(ctx, lockKey).Result()
if err != nil {
return false, 0, err
}
return true, ttl, nil
}
var loginFailureScript = redis.NewScript(`
local count = redis.call('INCR', KEYS[1])
redis.call('EXPIRE', KEYS[1], ARGV[1])
if tonumber(count) >= tonumber(ARGV[2]) then
redis.call('SET', KEYS[2], '1', 'PX', ARGV[3])
end
return count
`)
func (s *RedisState) RecordLoginFailure(ctx context.Context, key string) error {
return loginFailureScript.Run(ctx, s.rdb,
[]string{s.prefix + "login-failures:" + key, s.prefix + "login-lock:" + key},
int(failureWindow.Seconds()), maxFailures, int(lockDuration.Milliseconds()),
).Err()
}
func (s *RedisState) RecordLoginSuccess(ctx context.Context, key string) error {
pipe := s.rdb.Pipeline()
pipe.Del(ctx, s.prefix+"login-failures:"+key, s.prefix+"login-lock:"+key)
_, err := pipe.Exec(ctx)
return err
}
func stateUnavailable(operation string) error {
msg := "Redis 状态服务不可用,无法" + operation
if appEnv == "production" {
slog.Error(msg)
}
return fmt.Errorf("%s", msg)
}
+152
View File
@@ -0,0 +1,152 @@
package middleware
import (
"context"
"net/http"
"net/http/httptest"
"sync"
"testing"
"time"
"github.com/gin-gonic/gin"
)
type memoryState struct {
mu sync.Mutex
revoked map[string]time.Time
failures map[string]int
locked map[string]time.Time
}
func newMemoryState() *memoryState {
return &memoryState{
revoked: map[string]time.Time{},
failures: map[string]int{},
locked: map[string]time.Time{},
}
}
func (m *memoryState) RevokeToken(_ context.Context, id string, exp time.Time) error {
m.mu.Lock()
defer m.mu.Unlock()
m.revoked[id] = exp
return nil
}
func (m *memoryState) IsTokenRevoked(_ context.Context, id string) (bool, error) {
m.mu.Lock()
defer m.mu.Unlock()
exp, ok := m.revoked[id]
if !ok {
return false, nil
}
if time.Now().After(exp) {
delete(m.revoked, id)
return false, nil
}
return true, nil
}
func (m *memoryState) CheckLoginLock(_ context.Context, key string) (bool, time.Duration, error) {
m.mu.Lock()
defer m.mu.Unlock()
until, ok := m.locked[key]
if !ok {
return false, 0, nil
}
if time.Now().After(until) {
delete(m.locked, key)
return false, 0, nil
}
return true, time.Until(until), nil
}
func (m *memoryState) RecordLoginFailure(_ context.Context, key string) error {
m.mu.Lock()
defer m.mu.Unlock()
m.failures[key]++
if m.failures[key] >= maxFailures {
m.locked[key] = time.Now().Add(lockDuration)
}
return nil
}
func (m *memoryState) RecordLoginSuccess(_ context.Context, key string) error {
m.mu.Lock()
defer m.mu.Unlock()
delete(m.failures, key)
delete(m.locked, key)
return nil
}
func withMemoryState(t *testing.T) *memoryState {
t.Helper()
store := newMemoryState()
InitState(store, "test")
t.Cleanup(func() { InitState(nil, "test") })
return store
}
func TestRevokeTokenIsCrossInstanceState(t *testing.T) {
withMemoryState(t)
claims := &JWTClaims{}
claims.ID = "token-1"
if err := RevokeToken(claims, "token", time.Now().Add(time.Hour)); err != nil {
t.Fatalf("RevokeToken failed: %v", err)
}
revoked, err := IsRevoked(claims)
if err != nil {
t.Fatalf("IsRevoked failed: %v", err)
}
if !revoked {
t.Fatal("token should be revoked")
}
}
func TestRevokeTokenUnavailableFailsClosed(t *testing.T) {
InitState(nil, "production")
t.Cleanup(func() { InitState(nil, "test") })
claims := &JWTClaims{}
claims.ID = "x"
if err := RevokeToken(claims, "token", time.Now().Add(time.Hour)); err == nil {
t.Fatal("state unavailable should fail closed")
}
}
func TestLoginLockStateSemantics(t *testing.T) {
store := withMemoryState(t)
for i := 0; i < maxFailures; i++ {
if err := store.RecordLoginFailure(context.Background(), "1.1.1.1|user"); err != nil {
t.Fatalf("record failure failed: %v", err)
}
}
locked, _, err := store.CheckLoginLock(context.Background(), "1.1.1.1|user")
if err != nil {
t.Fatalf("check lock failed: %v", err)
}
if !locked {
t.Fatal("expected login lock after max failures")
}
if err := store.RecordLoginSuccess(context.Background(), "1.1.1.1|user"); err != nil {
t.Fatalf("record success failed: %v", err)
}
locked, _, _ = store.CheckLoginLock(context.Background(), "1.1.1.1|user")
if locked {
t.Fatal("login lock should be cleared after success")
}
}
func TestCheckLoginLockUnavailableReturns503(t *testing.T) {
InitState(nil, "test")
t.Cleanup(func() { InitState(nil, "test") })
gin.SetMode(gin.TestMode)
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
c.Request = httptest.NewRequest(http.MethodPost, "/auth/login", nil)
if !CheckLoginLock(c, "user") {
t.Fatal("state unavailable should abort login")
}
if rec.Code != http.StatusServiceUnavailable {
t.Fatalf("status = %d, want 503", rec.Code)
}
}
@@ -1,45 +1,26 @@
package middleware
import (
"sync"
"context"
"time"
"github.com/golang-jwt/jwt/v5"
)
// tokenBlacklist 登出令牌黑名单(内存版,进程重启后失效,令牌自然过期兜底)
type tokenBlacklist struct {
mu sync.RWMutex
revoked map[string]time.Time // tokenID(jti) -> 过期时间
}
var defaultBlacklist = &tokenBlacklist{revoked: make(map[string]time.Time)}
// RevokeToken 将令牌加入黑名单(按 jti,若无 jti 则按 subject+签发时间)
func RevokeToken(claims *JWTClaims, tokenStr string, exp time.Time) {
id := tokenIdentifier(claims)
defaultBlacklist.mu.Lock()
defaultBlacklist.revoked[id] = exp
defaultBlacklist.mu.Unlock()
func RevokeToken(claims *JWTClaims, tokenStr string, exp time.Time) error {
if authState == nil {
return stateUnavailable("吊销令牌")
}
return authState.RevokeToken(context.Background(), tokenIdentifier(claims), exp)
}
// IsRevoked 判断令牌是否已被吊销
func IsRevoked(claims *JWTClaims) bool {
id := tokenIdentifier(claims)
defaultBlacklist.mu.RLock()
exp, ok := defaultBlacklist.revoked[id]
defaultBlacklist.mu.RUnlock()
if !ok {
return false
func IsRevoked(claims *JWTClaims) (bool, error) {
if authState == nil {
return false, stateUnavailable("校验令牌吊销状态")
}
// 已过期的黑名单项自动清理
if time.Now().After(exp) {
defaultBlacklist.mu.Lock()
delete(defaultBlacklist.revoked, id)
defaultBlacklist.mu.Unlock()
return false
}
return true
return authState.IsTokenRevoked(context.Background(), tokenIdentifier(claims))
}
// ExtractClaims 从 token 字符串解析 claims(供 logout handler 使用)
+51
View File
@@ -0,0 +1,51 @@
package model
import (
"encoding/json"
"time"
)
// OutboxEvent 可靠事件箱;业务事务内写入,worker 发送并重试。
type OutboxEvent struct {
ID string `gorm:"type:uuid;primaryKey;default:gen_random_uuid()" json:"id"`
EventID string `gorm:"column:event_id;size:128;uniqueIndex" json:"eventId"`
EventType string `gorm:"column:event_type;size:64" json:"eventType"`
AggregateType *string `gorm:"column:aggregate_type;size:64" json:"aggregateType,omitempty"`
AggregateID *string `gorm:"column:aggregate_id;size:128" json:"aggregateId,omitempty"`
Payload json.RawMessage `gorm:"type:jsonb;default:'{}'" json:"payload"`
Status string `gorm:"size:16;default:pending;index:idx_outbox_events_status_next,priority:1" json:"status"`
Attempts int `gorm:"default:0" json:"attempts"`
MaxAttempts int `gorm:"column:max_attempts;default:5" json:"maxAttempts"`
NextAttemptAt time.Time `gorm:"column:next_attempt_at;type:timestamptz;default:now();index:idx_outbox_events_status_next,priority:2" json:"nextAttemptAt"`
LastError *string `gorm:"column:last_error;type:text" json:"lastError,omitempty"`
CreatedAt time.Time `gorm:"type:timestamptz" json:"createdAt"`
UpdatedAt time.Time `gorm:"type:timestamptz" json:"updatedAt"`
}
func (OutboxEvent) TableName() string { return "outbox_events" }
// Notification 持久化通知记录;状态 pending/sending/sent/retry/failed/cancelled。
type Notification struct {
ID string `gorm:"type:uuid;primaryKey;default:gen_random_uuid()" json:"id"`
UserID *string `gorm:"column:user_id;type:uuid;index:idx_notifications_user_created,priority:1" json:"userId,omitempty"`
Channel string `gorm:"size:32" json:"channel"`
Target string `gorm:"size:128" json:"target"`
Title string `gorm:"size:128" json:"title"`
Body string `gorm:"type:text" json:"body"`
Status string `gorm:"size:16;default:pending;index:idx_notifications_status_next,priority:1" json:"status"`
Attempts int `gorm:"default:0" json:"attempts"`
NextAttemptAt time.Time `gorm:"column:next_attempt_at;type:timestamptz;default:now();index:idx_notifications_status_next,priority:2" json:"nextAttemptAt"`
LastError *string `gorm:"column:last_error;type:text" json:"lastError,omitempty"`
TemplateID *string `gorm:"column:template_id;size:128" json:"templateId,omitempty"`
OpenID *string `gorm:"column:open_id;size:128" json:"openId,omitempty"`
Authorized bool `gorm:"default:false" json:"authorized"`
ProviderResponse json.RawMessage `gorm:"column:provider_response;type:jsonb" json:"providerResponse,omitempty"`
Data json.RawMessage `gorm:"type:jsonb" json:"data,omitempty"`
BusinessType *string `gorm:"column:business_type;size:64" json:"businessType,omitempty"`
BusinessID *string `gorm:"column:business_id;size:128" json:"businessId,omitempty"`
EventID *string `gorm:"column:event_id;size:128;uniqueIndex" json:"eventId,omitempty"`
CreatedAt time.Time `gorm:"type:timestamptz" json:"createdAt"`
UpdatedAt time.Time `gorm:"type:timestamptz" json:"updatedAt"`
}
func (Notification) TableName() string { return "notifications" }
+216
View File
@@ -0,0 +1,216 @@
package service
import (
"context"
"crypto/rand"
"encoding/hex"
"encoding/json"
"errors"
"log/slog"
"time"
"silk-server-go/internal/model"
"gorm.io/gorm"
"gorm.io/gorm/clause"
)
const (
OutboxStatusPending = "pending"
OutboxStatusSending = "sending"
OutboxStatusSent = "sent"
OutboxStatusRetry = "retry"
OutboxStatusFailed = "failed"
OutboxStatusCancelled = "cancelled"
OutboxEventWechatSubscribe = "wechat.subscribe"
)
// Event 写入 outbox 的领域事件。
type Event struct {
ID string
Type string
AggregateType string
AggregateID string
Payload json.RawMessage
}
// EventHandler 处理一条已领取的 outbox 事件。
type EventHandler func(ctx context.Context, event model.OutboxEvent) error
// Outbox 基于 PostgreSQL 的可靠事件箱。
type Outbox struct {
db *gorm.DB
handler EventHandler
interval time.Duration
batchSize int
maxAttempts int
backoff time.Duration
}
// NewOutbox 创建默认配置的 Outbox。
func NewOutbox(db *gorm.DB) *Outbox {
return &Outbox{
db: db,
interval: 5 * time.Second,
batchSize: 20,
maxAttempts: 5,
backoff: 15 * time.Second,
}
}
// SetHandler 设置事件处理器。
func (o *Outbox) SetHandler(handler EventHandler) {
o.handler = handler
}
// PublishTx 在业务事务内写入事件;重复 eventId 通过唯一索引幂等跳过。
func (o *Outbox) PublishTx(tx *gorm.DB, event Event) error {
if event.Type == "" {
return errors.New("outbox event type is required")
}
if event.ID == "" {
event.ID = randomEventID()
}
payload := event.Payload
if len(payload) == 0 {
payload = json.RawMessage(`{}`)
}
record := model.OutboxEvent{
EventID: event.ID,
EventType: event.Type,
AggregateType: strPtrOrNil(event.AggregateType),
AggregateID: strPtrOrNil(event.AggregateID),
Payload: payload,
Status: OutboxStatusPending,
MaxAttempts: o.maxAttempts,
NextAttemptAt: time.Now(),
}
// PublishTx 必须复用调用方事务,避免 GORM 对单条 Create 再开嵌套事务。
return tx.Session(&gorm.Session{SkipDefaultTransaction: true}).Clauses(clause.OnConflict{
Columns: []clause.Column{{Name: "event_id"}},
DoNothing: true,
}).Create(&record).Error
}
// ClaimNext 领取一批到期待处理事件。
func (o *Outbox) ClaimNext(ctx context.Context, limit int) ([]model.OutboxEvent, error) {
if limit <= 0 {
limit = o.batchSize
}
var events []model.OutboxEvent
err := o.db.WithContext(ctx).
Where("status IN ? AND next_attempt_at <= ?", []string{OutboxStatusPending, OutboxStatusRetry}, time.Now()).
Order("created_at ASC").
Limit(limit).
Find(&events).Error
return events, err
}
// MarkSending 原子标记领取状态,避免多 worker 重复处理。
func (o *Outbox) MarkSending(ctx context.Context, id string) (bool, error) {
result := o.db.WithContext(ctx).
Model(&model.OutboxEvent{}).
Where("id = ? AND status IN ?", id, []string{OutboxStatusPending, OutboxStatusRetry}).
Update("status", OutboxStatusSending)
return result.RowsAffected > 0, result.Error
}
// MarkResult 写入成功/失败并计算下一次重试时间。
func (o *Outbox) MarkResult(ctx context.Context, id string, processErr error) error {
var event model.OutboxEvent
if err := o.db.WithContext(ctx).First(&event, "id = ?", id).Error; err != nil {
return err
}
if event.MaxAttempts <= 0 {
event.MaxAttempts = o.maxAttempts
}
event.Attempts++
now := time.Now()
updates := map[string]interface{}{
"attempts": event.Attempts,
"updated_at": now,
"next_attempt_at": now,
}
if processErr == nil {
updates["status"] = OutboxStatusSent
updates["last_error"] = nil
} else {
message := processErr.Error()
updates["last_error"] = message
if event.Attempts >= event.MaxAttempts {
updates["status"] = OutboxStatusFailed
} else {
updates["status"] = OutboxStatusRetry
updates["next_attempt_at"] = now.Add(o.backoff * time.Duration(event.Attempts))
}
}
return o.db.WithContext(ctx).
Model(&model.OutboxEvent{}).
Where("id = ?", id).
Updates(updates).Error
}
// ProcessPending 领取并处理到期事件,返回本轮处理数量。
func (o *Outbox) ProcessPending(ctx context.Context) (int, error) {
if o.handler == nil {
return 0, errors.New("outbox handler is not set")
}
processed := 0
for {
events, err := o.ClaimNext(ctx, o.batchSize)
if err != nil {
return processed, err
}
if len(events) == 0 {
return processed, nil
}
for _, event := range events {
claimed, err := o.MarkSending(ctx, event.ID)
if err != nil {
slog.Warn("outbox mark sending failed", "eventId", event.EventID, "err", err)
continue
}
if !claimed {
continue
}
processErr := o.handler(ctx, event)
if err := o.MarkResult(ctx, event.ID, processErr); err != nil {
slog.Warn("outbox mark result failed", "eventId", event.EventID, "err", err)
continue
}
processed++
}
}
}
// Start 启动后台 worker。
func (o *Outbox) Start(ctx context.Context) {
ticker := time.NewTicker(o.interval)
defer ticker.Stop()
for {
select {
case <-ctx.Done():
return
case <-ticker.C:
if _, err := o.ProcessPending(ctx); err != nil {
slog.Warn("outbox worker failed", "err", err)
}
}
}
}
func strPtrOrNil(value string) *string {
if value == "" {
return nil
}
return &value
}
func randomEventID() string {
b := make([]byte, 16)
if _, err := rand.Read(b); err != nil {
return hex.EncodeToString([]byte(time.Now().Format(time.RFC3339Nano)))
}
return hex.EncodeToString(b)
}
+100
View File
@@ -0,0 +1,100 @@
package service
import (
"context"
"encoding/json"
"testing"
"time"
"github.com/DATA-DOG/go-sqlmock"
"gorm.io/driver/postgres"
"gorm.io/gorm"
"silk-server-go/internal/model"
)
func newOutboxMock(t *testing.T) (*Outbox, sqlmock.Sqlmock) {
t.Helper()
sqlDB, mock, err := sqlmock.New()
if err != nil {
t.Fatalf("create sqlmock: %v", err)
}
gdb, err := gorm.Open(postgres.New(postgres.Config{Conn: sqlDB}), &gorm.Config{})
if err != nil {
t.Fatalf("open gorm: %v", err)
}
o := NewOutbox(gdb)
return o, mock
}
func TestPublishTxUsesSameEventIDIdempotently(t *testing.T) {
o, mock := newOutboxMock(t)
event := Event{
ID: "event-1",
Type: OutboxEventWechatSubscribe,
Payload: json.RawMessage(`{"userId":"u1"}`),
}
mock.ExpectQuery(`INSERT INTO "outbox_events".*ON CONFLICT \("event_id"\) DO NOTHING`).
WithArgs(
sqlmock.AnyArg(), sqlmock.AnyArg(), sqlmock.AnyArg(), sqlmock.AnyArg(),
sqlmock.AnyArg(), sqlmock.AnyArg(), sqlmock.AnyArg(), sqlmock.AnyArg(),
sqlmock.AnyArg(), sqlmock.AnyArg(), sqlmock.AnyArg(), sqlmock.AnyArg(),
).
WillReturnRows(sqlmock.NewRows([]string{"id", "payload", "next_attempt_at"}).
AddRow("00000000-0000-0000-0000-000000000001", json.RawMessage(`{"userId":"u1"}`), time.Now()))
if err := o.PublishTx(o.db, event); err != nil {
t.Fatalf("first publish failed: %v", err)
}
// 第二次相同 eventId 仍走 ON CONFLICT DO NOTHING,不产生重复发送。
mock.ExpectQuery(`INSERT INTO "outbox_events".*ON CONFLICT \("event_id"\) DO NOTHING`).
WithArgs(
sqlmock.AnyArg(), sqlmock.AnyArg(), sqlmock.AnyArg(), sqlmock.AnyArg(),
sqlmock.AnyArg(), sqlmock.AnyArg(), sqlmock.AnyArg(), sqlmock.AnyArg(),
sqlmock.AnyArg(), sqlmock.AnyArg(), sqlmock.AnyArg(), sqlmock.AnyArg(),
).
WillReturnRows(sqlmock.NewRows([]string{"id", "payload", "next_attempt_at"}))
if err := o.PublishTx(o.db, event); err != nil {
t.Fatalf("second publish failed: %v", err)
}
}
func TestProcessPendingWithNoEventsIsNoop(t *testing.T) {
o, mock := newOutboxMock(t)
o.handler = func(ctx context.Context, event model.OutboxEvent) error { return nil }
mock.ExpectQuery(`SELECT .* FROM "outbox_events" WHERE status IN \(.*\)`).
WillReturnRows(sqlmock.NewRows([]string{
"id", "event_id", "event_type", "aggregate_type", "aggregate_id",
"payload", "status", "attempts", "max_attempts", "next_attempt_at",
"last_error", "created_at", "updated_at",
}))
processed, err := o.ProcessPending(context.Background())
if err != nil {
t.Fatalf("ProcessPending failed: %v", err)
}
if processed != 0 {
t.Fatalf("processed = %d, want 0", processed)
}
}
func TestClaimNextReturnsPendingAfterRestart(t *testing.T) {
o, mock := newOutboxMock(t)
now := time.Now()
mock.ExpectQuery(`SELECT .* FROM "outbox_events" WHERE status IN \(.*\)`).
WithArgs(OutboxStatusPending, OutboxStatusRetry, sqlmock.AnyArg(), sqlmock.AnyArg()).
WillReturnRows(sqlmock.NewRows([]string{
"id", "event_id", "event_type", "aggregate_type", "aggregate_id",
"payload", "status", "attempts", "max_attempts", "next_attempt_at",
"last_error", "created_at", "updated_at",
}).AddRow(
"id-1", "event-1", OutboxEventWechatSubscribe, nil, nil,
json.RawMessage(`{}`), OutboxStatusPending, 0, 5, now,
nil, now, now,
))
events, err := o.ClaimNext(context.Background(), 1)
if err != nil {
t.Fatalf("ClaimNext failed: %v", err)
}
if len(events) != 1 || events[0].EventID != "event-1" {
t.Fatalf("events = %+v", events)
}
}
+139 -12
View File
@@ -10,6 +10,11 @@ import (
"net/url"
"sync"
"time"
"silk-server-go/internal/model"
"gorm.io/gorm"
"gorm.io/gorm/clause"
)
const wechatAPIBase = "https://api.weixin.qq.com"
@@ -168,9 +173,21 @@ func (s *WechatService) getAccessToken(ctx context.Context) (string, error) {
// SendSubscribe 发送订阅消息
func (s *WechatService) SendSubscribe(ctx context.Context, openid, templateID string, data map[string]map[string]string, page string) error {
_, err := s.SendSubscribeResult(ctx, openid, templateID, data, page)
return err
}
// WechatSendResult 微信订阅消息响应。
type WechatSendResult struct {
ErrCode int `json:"errcode"`
ErrMsg string `json:"errmsg"`
}
// SendSubscribeResult 发送订阅消息并返回微信响应码。
func (s *WechatService) SendSubscribeResult(ctx context.Context, openid, templateID string, data map[string]map[string]string, page string) (WechatSendResult, error) {
token, err := s.getAccessToken(ctx)
if err != nil {
return err
return WechatSendResult{ErrCode: -1, ErrMsg: err.Error()}, err
}
payload := map[string]any{
"touser": openid,
@@ -182,32 +199,142 @@ func (s *WechatService) SendSubscribe(ctx context.Context, openid, templateID st
}
raw, err := json.Marshal(payload)
if err != nil {
return err
return WechatSendResult{ErrCode: -1, ErrMsg: err.Error()}, err
}
u := s.baseURL + "/cgi-bin/message/subscribe/send?access_token=" + url.QueryEscape(token)
req, err := http.NewRequestWithContext(ctx, http.MethodPost, u, bytes.NewReader(raw))
if err != nil {
return err
return WechatSendResult{ErrCode: -1, ErrMsg: err.Error()}, err
}
req.Header.Set("Content-Type", "application/json")
resp, err := s.httpClient.Do(req)
if err != nil {
return err
return WechatSendResult{ErrCode: -1, ErrMsg: err.Error()}, err
}
defer resp.Body.Close()
body, err := io.ReadAll(resp.Body)
if err != nil {
return err
}
var out struct {
ErrCode int `json:"errcode"`
ErrMsg string `json:"errmsg"`
return WechatSendResult{ErrCode: -1, ErrMsg: err.Error()}, err
}
var out WechatSendResult
if err := json.Unmarshal(body, &out); err != nil {
return err
return WechatSendResult{ErrCode: -1, ErrMsg: err.Error()}, err
}
if out.ErrCode != 0 {
return fmt.Errorf("微信订阅消息发送失败 (%d): %s", out.ErrCode, out.ErrMsg)
return out, fmt.Errorf("微信订阅消息发送失败 (%d): %s", out.ErrCode, out.ErrMsg)
}
return nil
return out, nil
}
// WechatSubscribePayload outbox 微信订阅事件载荷。
type WechatSubscribePayload struct {
UserID string `json:"userId"`
OpenID string `json:"openId"`
TemplateKey string `json:"templateKey"`
TemplateID string `json:"templateId"`
Page string `json:"page"`
Title string `json:"title"`
Body string `json:"body"`
Data map[string]map[string]string `json:"data"`
BusinessType string `json:"businessType"`
BusinessID string `json:"businessId"`
}
// NewWechatOutboxHandler 创建微信订阅 outbox 处理器。
func NewWechatOutboxHandler(db *gorm.DB, wechat *WechatService) EventHandler {
return func(ctx context.Context, event model.OutboxEvent) error {
if event.EventType != OutboxEventWechatSubscribe {
return nil
}
var payload WechatSubscribePayload
if err := json.Unmarshal(event.Payload, &payload); err != nil {
return err
}
return handleWechatSubscribe(ctx, db, wechat, event, payload)
}
}
func handleWechatSubscribe(ctx context.Context, db *gorm.DB, wechat *WechatService, event model.OutboxEvent, payload WechatSubscribePayload) error {
var binding model.WechatBinding
if err := db.Where("user_id = ?", payload.UserID).First(&binding).Error; err != nil {
message := "微信未绑定"
return saveWechatNotification(ctx, db, event, payload, "", false, WechatSendResult{ErrCode: 40001, ErrMsg: message}, "cancelled", &message)
}
var authorized []string
if len(binding.AuthorizedTemplates) > 0 {
_ = json.Unmarshal(binding.AuthorizedTemplates, &authorized)
}
if !IsAuthorized(authorized, payload.TemplateKey) {
message := "用户未授权订阅模板"
return saveWechatNotification(ctx, db, event, payload, binding.OpenID, false, WechatSendResult{ErrCode: 43101, ErrMsg: message}, "cancelled", &message)
}
if payload.TemplateID == "" {
message := "订阅消息模板 ID 未配置"
if err := saveWechatNotification(ctx, db, event, payload, binding.OpenID, true, WechatSendResult{ErrCode: -1, ErrMsg: message}, "retry", &message); err != nil {
return err
}
return fmt.Errorf("订阅消息模板 ID 未配置")
}
result, sendErr := wechat.SendSubscribeResult(ctx, binding.OpenID, payload.TemplateID, payload.Data, payload.Page)
status := "sent"
var lastErr *string
maxAttempts := event.MaxAttempts
if maxAttempts <= 0 {
maxAttempts = 5
}
if sendErr != nil {
message := sendErr.Error()
lastErr = &message
status = "retry"
if event.Attempts+1 >= maxAttempts {
status = "failed"
}
}
if err := saveWechatNotification(ctx, db, event, payload, binding.OpenID, true, result, status, lastErr); err != nil {
return err
}
return sendErr
}
func saveWechatNotification(ctx context.Context, db *gorm.DB, event model.OutboxEvent, payload WechatSubscribePayload, openID string, authorized bool, response WechatSendResult, status string, lastErr *string) error {
data, _ := json.Marshal(payload.Data)
providerResponse, _ := json.Marshal(response)
userID := payload.UserID
eventID := event.EventID
notification := model.Notification{
UserID: &userID,
Channel: "wechat",
Target: openID,
Title: payload.Title,
Body: payload.Body,
Status: status,
Attempts: event.Attempts + 1,
NextAttemptAt: time.Now(),
LastError: lastErr,
TemplateID: strPtrOrNil(payload.TemplateID),
OpenID: strPtrOrNil(openID),
Authorized: authorized,
ProviderResponse: providerResponse,
Data: data,
BusinessType: strPtrOrNil(payload.BusinessType),
BusinessID: strPtrOrNil(payload.BusinessID),
EventID: &eventID,
}
return db.WithContext(ctx).
Clauses(clause.OnConflict{
Columns: []clause.Column{{Name: "event_id"}},
DoUpdates: clause.Assignments(map[string]interface{}{
"status": status,
"attempts": event.Attempts + 1,
"next_attempt_at": time.Now(),
"last_error": lastErr,
"authorized": authorized,
"provider_response": providerResponse,
"updated_at": time.Now(),
}),
}).
Create(&notification).Error
}
+5 -1
View File
@@ -106,7 +106,11 @@ func TestWechatSendSubscribeErrorCode(t *testing.T) {
s.baseURL = srv.URL
s.accessToken = "fake-token"
s.tokenExpire = time.Now().Add(time.Hour)
if err := s.SendSubscribe(context.Background(), "bad", "tmpl1", nil, ""); err == nil {
result, err := s.SendSubscribeResult(context.Background(), "bad", "tmpl1", nil, "")
if err == nil {
t.Error("errcode!=0 应返回错误")
}
if result.ErrCode != 40003 || result.ErrMsg != "invalid openid" {
t.Errorf("response code = %+v", result)
}
}