feat(P2): SSE 实时消息通知 — 点赞/评论/关注事件推送 + Notification 表
This commit is contained in:
@@ -1,6 +1,7 @@
|
|||||||
package http
|
package http
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
"feedsystem_video_go/internal/account"
|
"feedsystem_video_go/internal/account"
|
||||||
"feedsystem_video_go/internal/feed"
|
"feedsystem_video_go/internal/feed"
|
||||||
"feedsystem_video_go/internal/middleware/jwt"
|
"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.StartOutboxPoller(db, timelineMQ)
|
||||||
worker.StartConsumer(timelineMQ, "video.timeline.update.queue", cache)
|
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
|
return r
|
||||||
}
|
}
|
||||||
|
|||||||
141
backend/internal/worker/notificationworker.go
Normal file
141
backend/internal/worker/notificationworker.go
Normal file
@@ -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
|
||||||
|
}
|
||||||
174
backend/internal/worker/ssehub.go
Normal file
174
backend/internal/worker/ssehub.go
Normal file
@@ -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)
|
||||||
Reference in New Issue
Block a user