From 41454de49fab799eea9c0006ec2372814e18eb25 Mon Sep 17 00:00:00 2001 From: Sisyphus Date: Sat, 25 Apr 2026 20:02:27 +0800 Subject: [PATCH] =?UTF-8?q?feat(P2):=20SSE=20=E5=AE=9E=E6=97=B6=E6=B6=88?= =?UTF-8?q?=E6=81=AF=E9=80=9A=E7=9F=A5=20=E2=80=94=20=E7=82=B9=E8=B5=9E/?= =?UTF-8?q?=E8=AF=84=E8=AE=BA/=E5=85=B3=E6=B3=A8=E4=BA=8B=E4=BB=B6?= =?UTF-8?q?=E6=8E=A8=E9=80=81=20+=20Notification=20=E8=A1=A8?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- backend/internal/http/router.go | 41 +++++ backend/internal/worker/notificationworker.go | 141 ++++++++++++++ backend/internal/worker/ssehub.go | 174 ++++++++++++++++++ 3 files changed, 356 insertions(+) create mode 100644 backend/internal/worker/notificationworker.go create mode 100644 backend/internal/worker/ssehub.go diff --git a/backend/internal/http/router.go b/backend/internal/http/router.go index 38fdf65..f982706 100644 --- a/backend/internal/http/router.go +++ b/backend/internal/http/router.go @@ -1,6 +1,7 @@ package http import ( + "context" "feedsystem_video_go/internal/account" "feedsystem_video_go/internal/feed" "feedsystem_video_go/internal/middleware/jwt" @@ -179,5 +180,45 @@ func SetRouter(db *gorm.DB, cache *rediscache.Client, rmq *rabbitmq.RabbitMQ) *g } worker.StartOutboxPoller(db, timelineMQ) worker.StartConsumer(timelineMQ, "video.timeline.update.queue", cache) + + // SSE notification + if rmq != nil && rmq.Ch != nil { + rmq.DeclareTopic("like.events", "notification.like", "like.like") + rmq.DeclareTopic("comment.events", "notification.comment", "comment.publish") + rmq.DeclareTopic("social.events", "notification.social", "social.follow") + } + sseHub := worker.NewSSEHub(db) + notifGroup := r.Group("/notification") + notifGroup.Use(sseHub.SSERequireAuth()) + sseHub.RegisterRoutes(r, notifGroup) + + go func() { + if rmq != nil && rmq.Ch != nil { + hub := sseHub + ctx := context.Background() + // consume from like queue + go func() { + w := worker.NewNotificationWorker(rmq.Ch, db, "notification.like", hub) + if err := w.Run(ctx); err != nil { + log.Printf("notification-like worker: %v", err) + } + }() + go func() { + w := worker.NewNotificationWorker(rmq.Ch, db, "notification.comment", hub) + if err := w.Run(ctx); err != nil { + log.Printf("notification-comment worker: %v", err) + } + }() + go func() { + w := worker.NewNotificationWorker(rmq.Ch, db, "notification.social", hub) + if err := w.Run(ctx); err != nil { + log.Printf("notification-social worker: %v", err) + } + }() + } else { + log.Printf("Notification SSE disabled (MQ not available)") + } + }() + return r } diff --git a/backend/internal/worker/notificationworker.go b/backend/internal/worker/notificationworker.go new file mode 100644 index 0000000..08fa767 --- /dev/null +++ b/backend/internal/worker/notificationworker.go @@ -0,0 +1,141 @@ +package worker + +import ( + "context" + "encoding/json" + "errors" + "feedsystem_video_go/internal/middleware/rabbitmq" + "log" + "time" + + amqp "github.com/rabbitmq/amqp091-go" + "gorm.io/gorm" +) + +type Notification struct { + ID uint `gorm:"primaryKey" json:"id"` + RecipientID uint `gorm:"index;not null" json:"recipient_id"` + SenderID uint `gorm:"not null" json:"sender_id"` + Type string `gorm:"type:varchar(50);not null" json:"type"` + TargetID uint `json:"target_id"` + Content string `gorm:"type:varchar(255)" json:"content"` + IsRead bool `gorm:"default:false" json:"is_read"` + CreatedAt time.Time `gorm:"autoCreateTime" json:"created_at"` +} + +type NotificationWorker struct { + ch *amqp.Channel + db *gorm.DB + queue string + hub NotificationHub +} + +type NotificationHub interface { + Push(userID uint, n *Notification) +} + +func NewNotificationWorker(ch *amqp.Channel, db *gorm.DB, queue string, hub NotificationHub) *NotificationWorker { + return &NotificationWorker{ch: ch, db: db, queue: queue, hub: hub} +} + +func (w *NotificationWorker) Run(ctx context.Context) error { + if w == nil || w.ch == nil || w.db == nil { + return errors.New("notification worker is not initialized") + } + if err := w.db.WithContext(ctx).AutoMigrate(&Notification{}); err != nil { + return err + } + deliveries, err := w.ch.Consume(w.queue, "", false, false, false, false, nil) + if err != nil { + return err + } + for { + select { + case <-ctx.Done(): + return ctx.Err() + case d, ok := <-deliveries: + if !ok { + return errors.New("deliveries channel closed") + } + w.handleDelivery(ctx, d) + } + } +} + +func (w *NotificationWorker) handleDelivery(ctx context.Context, d amqp.Delivery) { + retryCount := rabbitmq.GetRetryCount(d) + if err := w.process(ctx, d); err != nil { + if retryCount >= rabbitmq.MaxRetryCount { + log.Printf("notification worker: max retries, dropping: %v", err) + _ = d.Ack(false) + return + } + log.Printf("notification worker: failed (retry %d/%d): %v", retryCount+1, rabbitmq.MaxRetryCount, err) + _ = d.Nack(false, true) + return + } + _ = d.Ack(false) +} + +func (w *NotificationWorker) process(ctx context.Context, d amqp.Delivery) error { + body := d.Body + if len(body) == 0 { + return nil + } + routingKey := d.RoutingKey + + var notif *Notification + + switch { + case routingKey == "like.like": + var evt rabbitmq.LikeEvent + if err := json.Unmarshal(body, &evt); err != nil { + return nil + } + if evt.UserID == 0 || evt.VideoID == 0 { + return nil + } + var authorID uint + w.db.WithContext(ctx).Model(&struct{ ID uint; AuthorID uint }{}).Table("videos").Where("id = ?", evt.VideoID).Select("author_id").Scan(&authorID) + if authorID == 0 || authorID == evt.UserID { + return nil + } + notif = &Notification{RecipientID: authorID, SenderID: evt.UserID, Type: "like", TargetID: evt.VideoID, Content: "点赞了你的视频"} + + case routingKey == "comment.publish": + var evt rabbitmq.CommentEvent + if err := json.Unmarshal(body, &evt); err != nil { + return nil + } + if evt.AuthorID == 0 || evt.VideoID == 0 { + return nil + } + var authorID uint + w.db.WithContext(ctx).Model(&struct{ ID uint; AuthorID uint }{}).Table("videos").Where("id = ?", evt.VideoID).Select("author_id").Scan(&authorID) + if authorID == 0 || authorID == evt.AuthorID { + return nil + } + notif = &Notification{RecipientID: authorID, SenderID: evt.AuthorID, Type: "comment", TargetID: evt.VideoID, Content: "评论了你的视频"} + + case routingKey == "social.follow": + var evt rabbitmq.SocialEvent + if err := json.Unmarshal(body, &evt); err != nil { + return nil + } + if evt.FollowerID == 0 || evt.VloggerID == 0 { + return nil + } + notif = &Notification{RecipientID: evt.VloggerID, SenderID: evt.FollowerID, Type: "follow", TargetID: evt.FollowerID, Content: "关注了你"} + } + + if notif == nil { + return nil + } + if err := w.db.WithContext(ctx).Create(notif).Error; err != nil { + return err + } + if w.hub != nil { + w.hub.Push(notif.RecipientID, notif) + } + return nil +} diff --git a/backend/internal/worker/ssehub.go b/backend/internal/worker/ssehub.go new file mode 100644 index 0000000..332ee6d --- /dev/null +++ b/backend/internal/worker/ssehub.go @@ -0,0 +1,174 @@ +package worker + +import ( + "encoding/json" + "fmt" + "net/http" + "sync" + "time" + + "feedsystem_video_go/internal/auth" + + "github.com/gin-gonic/gin" + "gorm.io/gorm" +) + +type SSEHub struct { + mu sync.RWMutex + clients map[uint][]chan *Notification + db *gorm.DB +} + +func NewSSEHub(db *gorm.DB) *SSEHub { + return &SSEHub{clients: make(map[uint][]chan *Notification), db: db} +} + +func (h *SSEHub) Push(userID uint, n *Notification) { + h.mu.RLock() + chs, ok := h.clients[userID] + h.mu.RUnlock() + if !ok { + return + } + for _, ch := range chs { + select { + case ch <- n: + default: + } + } +} + +func (h *SSEHub) Subscribe(userID uint) chan *Notification { + ch := make(chan *Notification, 20) + h.mu.Lock() + h.clients[userID] = append(h.clients[userID], ch) + h.mu.Unlock() + return ch +} + +func (h *SSEHub) Unsubscribe(userID uint, ch chan *Notification) { + h.mu.Lock() + defer h.mu.Unlock() + chs := h.clients[userID] + for i, c := range chs { + if c == ch { + h.clients[userID] = append(chs[:i], chs[i+1:]...) + close(c) + return + } + } +} + +func (h *SSEHub) SSERequireAuth() gin.HandlerFunc { + return func(c *gin.Context) { + token := c.Query("token") + if token == "" { + token = c.GetHeader("Authorization") + if len(token) > 7 && token[:7] == "Bearer " { + token = token[7:] + } + } + if token == "" { + c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"error": "missing token"}) + return + } + claims, err := auth.ParseToken(token) + if err != nil { + c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"error": "invalid token"}) + return + } + c.Set("accountID", claims.AccountID) + c.Next() + } +} + +func (h *SSEHub) SSEHandler(c *gin.Context) { + accountID, _ := c.Get("accountID") + userID := accountID.(uint) + + c.Writer.Header().Set("Content-Type", "text/event-stream") + c.Writer.Header().Set("Cache-Control", "no-cache") + c.Writer.Header().Set("Connection", "keep-alive") + c.Writer.WriteHeader(http.StatusOK) + + ch := h.Subscribe(userID) + defer h.Unsubscribe(userID, ch) + + ctx := c.Request.Context() + flusher, _ := c.Writer.(http.Flusher) + + for { + select { + case <-ctx.Done(): + return + case n, ok := <-ch: + if !ok { + return + } + b, _ := json.Marshal(n) + fmt.Fprintf(c.Writer, "data: %s\n\n", b) + if flusher != nil { + flusher.Flush() + } + case <-time.After(30 * time.Second): + fmt.Fprintf(c.Writer, ": keepalive\n\n") + if flusher != nil { + flusher.Flush() + } + } + } +} + +func (h *SSEHub) ListHandler(c *gin.Context) { + accountID, _ := c.Get("accountID") + userID := accountID.(uint) + + var notifications []Notification + if err := h.db.WithContext(c.Request.Context()). + Where("recipient_id = ?", userID). + Order("created_at desc"). + Limit(50). + Find(¬ifications).Error; err != nil { + c.JSON(500, gin.H{"error": err.Error()}) + return + } + if notifications == nil { + notifications = []Notification{} + } + c.JSON(200, gin.H{"notifications": notifications}) +} + +func (h *SSEHub) MarkReadHandler(c *gin.Context) { + accountID, _ := c.Get("accountID") + userID := accountID.(uint) + + var req struct { + ID *uint `json:"id"` + } + c.ShouldBindJSON(&req) + + if req.ID != nil { + h.db.WithContext(c.Request.Context()).Model(&Notification{}).Where("id = ? AND recipient_id = ?", *req.ID, userID).Update("is_read", true) + } else { + h.db.WithContext(c.Request.Context()).Model(&Notification{}).Where("recipient_id = ?", userID).Update("is_read", true) + } + c.JSON(200, gin.H{"message": "ok"}) +} + +func (h *SSEHub) UnreadCountHandler(c *gin.Context) { + accountID, _ := c.Get("accountID") + userID := accountID.(uint) + + var count int64 + h.db.WithContext(c.Request.Context()).Model(&Notification{}).Where("recipient_id = ? AND is_read = ?", userID, false).Count(&count) + c.JSON(200, gin.H{"count": count}) +} + +func (h *SSEHub) RegisterRoutes(r *gin.Engine, group *gin.RouterGroup) { + group.GET("/stream", h.SSEHandler) + group.POST("/list", h.ListHandler) + group.POST("/markRead", h.MarkReadHandler) + group.POST("/unreadCount", h.UnreadCountHandler) +} + +var _ NotificationHub = (*SSEHub)(nil)