From a0029877849cabb1ebe3bf0a92c27695061d1497 Mon Sep 17 00:00:00 2001 From: yml2213 Date: Wed, 27 May 2026 06:18:55 +0800 Subject: [PATCH] =?UTF-8?q?SSE=20=E6=8E=A8=E9=80=81=E6=9B=BF=E4=BB=A3?= =?UTF-8?q?=E8=BD=AE=E8=AE=AD?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- backend/internal/middleware/auth.go | 19 ++- backend/internal/modules/chat/repository.go | 39 ++++- backend/internal/modules/chathub/handler.go | 87 +++++++++++ backend/internal/modules/chathub/hub.go | 153 +++++++++++++++++++ backend/internal/modules/order/repository.go | 22 ++- backend/internal/router/router.go | 20 ++- frontend/src/composables/useChatSSE.ts | 91 +++++++++++ frontend/src/views/admin/AdminChatsView.vue | 49 ++++-- frontend/src/views/mobile/MobileChatView.vue | 55 +++++-- 9 files changed, 493 insertions(+), 42 deletions(-) create mode 100644 backend/internal/modules/chathub/handler.go create mode 100644 backend/internal/modules/chathub/hub.go create mode 100644 frontend/src/composables/useChatSSE.ts diff --git a/backend/internal/middleware/auth.go b/backend/internal/middleware/auth.go index ef2090f..64b2d9c 100644 --- a/backend/internal/middleware/auth.go +++ b/backend/internal/middleware/auth.go @@ -16,11 +16,19 @@ const ( ContextUsername = "username" ) +func extractToken(c *gin.Context) string { + header := c.GetHeader("Authorization") + tokenText := strings.TrimSpace(strings.TrimPrefix(header, "Bearer ")) + if tokenText != "" && tokenText != header { + return tokenText + } + return c.Query("token") +} + func Auth(jwtManager *auth.JWTManager) gin.HandlerFunc { return func(c *gin.Context) { - header := c.GetHeader("Authorization") - tokenText := strings.TrimSpace(strings.TrimPrefix(header, "Bearer ")) - if tokenText == "" || tokenText == header { + tokenText := extractToken(c) + if tokenText == "" { response.Unauthorized(c, "缺少访问令牌") c.Abort() return @@ -41,9 +49,8 @@ func Auth(jwtManager *auth.JWTManager) gin.HandlerFunc { func AdminAuth(jwtManager *auth.JWTManager) gin.HandlerFunc { return func(c *gin.Context) { - header := c.GetHeader("Authorization") - tokenText := strings.TrimSpace(strings.TrimPrefix(header, "Bearer ")) - if tokenText == "" || tokenText == header { + tokenText := extractToken(c) + if tokenText == "" { response.Unauthorized(c, "缺少后台访问令牌") c.Abort() return diff --git a/backend/internal/modules/chat/repository.go b/backend/internal/modules/chat/repository.go index b46fc42..a38dc84 100644 --- a/backend/internal/modules/chat/repository.go +++ b/backend/internal/modules/chat/repository.go @@ -8,6 +8,7 @@ import ( "time" "hfb_sys/backend/internal/model" + "hfb_sys/backend/internal/modules/chathub" "golang.org/x/crypto/bcrypt" "gorm.io/datatypes" @@ -16,7 +17,8 @@ import ( ) type Repository struct { - db *gorm.DB + db *gorm.DB + hub *chathub.Hub } const ( @@ -25,8 +27,8 @@ const ( defaultSupportNickname = "超级管理员" ) -func NewRepository(db *gorm.DB) *Repository { - return &Repository{db: db} +func NewRepository(db *gorm.DB, hub *chathub.Hub) *Repository { + return &Repository{db: db, hub: hub} } func EnsureOrderConversation(tx *gorm.DB, order model.RentalOrder) (*model.ChatConversation, error) { @@ -101,6 +103,17 @@ func EnsureOrderConversation(tx *gorm.DB, order model.RentalOrder) (*model.ChatC return &conversation, nil } +// NotifyNewConversation 推送群聊创建事件给会话参与者。应在事务提交后调用。 +func (r *Repository) NotifyNewConversation(conversationID uint64) { + if r.hub == nil { + return + } + r.hub.NotifyConversation(conversationID, &chathub.ChatEvent{ + Type: "conversation_updated", + ConversationID: conversationID, + }) +} + func (r *Repository) ListConversations(principal Principal, page, pageSize int) (*PaginatedResult, error) { page, pageSize = normalizePagination(page, pageSize) var total int64 @@ -241,6 +254,26 @@ func (r *Repository) SendMessage(principal Principal, conversationID uint64, req if len(items) == 0 { return nil, ErrConversationNotFound } + // 推送新消息事件给会话中的在线参与者 + if r.hub != nil { + msg := items[0] + r.hub.NotifyConversation(conversationID, &chathub.ChatEvent{ + Type: "new_message", + ConversationID: conversationID, + Message: &chathub.MessageData{ + ID: msg.ID, + ConversationID: msg.ConversationID, + SenderType: msg.SenderType, + SenderID: msg.SenderID, + SenderRole: msg.SenderRole, + SenderName: msg.SenderName, + ContentType: msg.ContentType, + Content: msg.Content, + AttachmentURLS: msg.AttachmentURLS, + CreatedAt: msg.CreatedAt.Format(time.RFC3339), + }, + }) + } return &items[0], nil } diff --git a/backend/internal/modules/chathub/handler.go b/backend/internal/modules/chathub/handler.go new file mode 100644 index 0000000..5d9d9a8 --- /dev/null +++ b/backend/internal/modules/chathub/handler.go @@ -0,0 +1,87 @@ +package chathub + +import ( + "fmt" + "time" + + "hfb_sys/backend/internal/middleware" + "hfb_sys/backend/pkg/response" + + "github.com/gin-gonic/gin" +) + +const heartbeatInterval = 30 * time.Second + +type Handler struct { + hub *Hub +} + +func NewHandler(hub *Hub) *Handler { + return &Handler{hub: hub} +} + +// UserEvents 处理用户端 SSE 连接: GET /api/chats/events +func (h *Handler) UserEvents(c *gin.Context) { + value, ok := c.Get(middleware.ContextUserID) + if !ok { + response.Unauthorized(c, "缺少用户上下文") + return + } + userID, ok := value.(uint64) + if !ok { + response.Unauthorized(c, "用户上下文无效") + return + } + h.serveSSE(c, "user", userID) +} + +// AdminEvents 处理管理端 SSE 连接: GET /api/admin/chats/events +func (h *Handler) AdminEvents(c *gin.Context) { + value, ok := c.Get(middleware.ContextAdminID) + if !ok { + response.Unauthorized(c, "缺少管理员上下文") + return + } + adminID, ok := value.(uint64) + if !ok { + response.Unauthorized(c, "管理员上下文无效") + return + } + h.serveSSE(c, "admin", adminID) +} + +func (h *Handler) serveSSE(c *gin.Context, pType string, pID uint64) { + ch := h.hub.Subscribe(pType, pID) + defer h.hub.Unsubscribe(pType, pID, ch) + + c.Header("Content-Type", "text/event-stream") + c.Header("Cache-Control", "no-cache") + c.Header("Connection", "keep-alive") + c.Header("X-Accel-Buffering", "no") + + // 发送初始连接确认 + fmt.Fprintf(c.Writer, "event: connected\ndata: {\"ok\":true}\n\n") + c.Writer.Flush() + + heartbeat := time.NewTicker(heartbeatInterval) + defer heartbeat.Stop() + + clientGone := c.Request.Context().Done() + + for { + select { + case <-clientGone: + return + case <-heartbeat.C: + fmt.Fprintf(c.Writer, ":heartbeat\n\n") + c.Writer.Flush() + case event, ok := <-ch: + if !ok { + return + } + data := MarshalEvent(event) + fmt.Fprintf(c.Writer, "event: %s\ndata: %s\n\n", event.Type, data) + c.Writer.Flush() + } + } +} diff --git a/backend/internal/modules/chathub/hub.go b/backend/internal/modules/chathub/hub.go new file mode 100644 index 0000000..53e8b84 --- /dev/null +++ b/backend/internal/modules/chathub/hub.go @@ -0,0 +1,153 @@ +package chathub + +import ( + "encoding/json" + "sync" + + "gorm.io/gorm" +) + +// ChatEvent 是通过 SSE 推送给客户端的事件。 +type ChatEvent struct { + Type string `json:"type"` // "new_message" | "conversation_updated" + ConversationID uint64 `json:"conversation_id"` + Message *MessageData `json:"message,omitempty"` +} + +// MessageData 是事件中携带的消息数据,与 chat.MessageDTO 对齐。 +type MessageData struct { + ID uint64 `json:"id"` + ConversationID uint64 `json:"conversation_id"` + SenderType string `json:"sender_type"` + SenderID uint64 `json:"sender_id"` + SenderRole string `json:"sender_role"` + SenderName string `json:"sender_name"` + ContentType string `json:"content_type"` + Content string `json:"content"` + AttachmentURLS []string `json:"attachment_urls"` + CreatedAt string `json:"created_at"` +} + +// principal 标识一个连接方。 +type principal struct { + Type string // "user" or "admin" + ID uint64 +} + +// Hub 管理所有 SSE 连接。 +type Hub struct { + mu sync.RWMutex + clients map[principal]map[chan *ChatEvent]struct{} + db *gorm.DB +} + +// NewHub 创建 Hub 实例。 +func NewHub(db *gorm.DB) *Hub { + return &Hub{ + clients: make(map[principal]map[chan *ChatEvent]struct{}), + db: db, + } +} + +// Subscribe 注册一个 SSE 连接,返回事件 channel。 +func (h *Hub) Subscribe(pType string, pID uint64) <-chan *ChatEvent { + ch := make(chan *ChatEvent, 16) + key := principal{Type: pType, ID: pID} + h.mu.Lock() + if h.clients[key] == nil { + h.clients[key] = make(map[chan *ChatEvent]struct{}) + } + h.clients[key][ch] = struct{}{} + h.mu.Unlock() + return ch +} + +// Unsubscribe 注销 SSE 连接。 +func (h *Hub) Unsubscribe(pType string, pID uint64, ch <-chan *ChatEvent) { + key := principal{Type: pType, ID: pID} + h.mu.Lock() + if clients, ok := h.clients[key]; ok { + // 找到对应的发送 channel 并删除 + for sendCh := range clients { + if sendCh == ch { + delete(clients, sendCh) + close(sendCh) + break + } + } + if len(clients) == 0 { + delete(h.clients, key) + } + } + h.mu.Unlock() +} + +// NotifyConversation 查询会话参与者,向所有在线参与者推送事件。 +func (h *Hub) NotifyConversation(conversationID uint64, event *ChatEvent) { + if h.db == nil { + return + } + // 查询会话参与者 + type participantRow struct { + ParticipantType string + ParticipantID uint64 + } + var rows []participantRow + h.db.Table("chat_participants"). + Select("participant_type, participant_id"). + Where("conversation_id = ?", conversationID). + Find(&rows) + + h.mu.RLock() + defer h.mu.RUnlock() + + seen := make(map[principal]struct{}) + for _, row := range rows { + key := principal{Type: row.ParticipantType, ID: row.ParticipantID} + if _, ok := seen[key]; ok { + continue + } + seen[key] = struct{}{} + if clients, ok := h.clients[key]; ok { + for ch := range clients { + select { + case ch <- event: + default: + // channel 满了,丢弃事件避免阻塞 + } + } + } + } +} + +// NotifyUser 直接向指定用户/Admin 推送事件。 +func (h *Hub) NotifyUser(pType string, pID uint64, event *ChatEvent) { + key := principal{Type: pType, ID: pID} + h.mu.RLock() + defer h.mu.RUnlock() + if clients, ok := h.clients[key]; ok { + for ch := range clients { + select { + case ch <- event: + default: + } + } + } +} + +// MarshalEvent 将事件序列化为 SSE data 行。 +func MarshalEvent(event *ChatEvent) string { + raw, _ := json.Marshal(event) + return string(raw) +} + +// OnlineCount 返回当前在线连接数。 +func (h *Hub) OnlineCount() int { + h.mu.RLock() + defer h.mu.RUnlock() + count := 0 + for _, clients := range h.clients { + count += len(clients) + } + return count +} diff --git a/backend/internal/modules/order/repository.go b/backend/internal/modules/order/repository.go index 70392bb..0371e5c 100644 --- a/backend/internal/modules/order/repository.go +++ b/backend/internal/modules/order/repository.go @@ -20,7 +20,8 @@ import ( ) type Repository struct { - db *gorm.DB + db *gorm.DB + chatRepo *chat.Repository } const defaultPendingPaymentTimeoutMinutes = 15 @@ -29,6 +30,10 @@ func NewRepository(db *gorm.DB) *Repository { return &Repository{db: db} } +func (r *Repository) SetChatRepo(cr *chat.Repository) { + r.chatRepo = cr +} + type orderPricing struct { RentAmount float64 OwnerRentAmount float64 @@ -178,7 +183,8 @@ func (r *Repository) Create(renterID uint64, req CreateRequest) (*OrderDTO, erro } func (r *Repository) Pay(userID uint64, orderID uint64) error { - return r.db.Transaction(func(tx *gorm.DB) error { + var newConvID uint64 + err := r.db.Transaction(func(tx *gorm.DB) error { var order model.RentalOrder if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}). Where("id = ? AND renter_id = ?", orderID, userID). @@ -238,9 +244,11 @@ func (r *Repository) Pay(userID uint64, orderID uint64) error { order.HandoffStatus = "pending_owner" listing.Status = "rented" account.Status = "rented" - if _, err := chat.EnsureOrderConversation(tx, order); err != nil { + conv, err := chat.EnsureOrderConversation(tx, order) + if err != nil { return err } + newConvID = conv.ID orderID := order.ID if err := notification.Append(tx, notification.Entry{ @@ -270,6 +278,14 @@ func (r *Repository) Pay(userID uint64, orderID uint64) error { } return tx.Save(&account).Error }) + if err != nil { + return err + } + // 事务成功后推送群聊创建事件 + if newConvID > 0 && r.chatRepo != nil { + r.chatRepo.NotifyNewConversation(newConvID) + } + return nil } func (r *Repository) Cancel(userID uint64, orderID uint64) error { diff --git a/backend/internal/router/router.go b/backend/internal/router/router.go index bb93f35..071237f 100644 --- a/backend/internal/router/router.go +++ b/backend/internal/router/router.go @@ -12,6 +12,7 @@ import ( "hfb_sys/backend/internal/modules/adminuser" "hfb_sys/backend/internal/modules/auth" "hfb_sys/backend/internal/modules/chat" + "hfb_sys/backend/internal/modules/chathub" "hfb_sys/backend/internal/modules/dispute" filemodule "hfb_sys/backend/internal/modules/file" "hfb_sys/backend/internal/modules/listing" @@ -98,12 +99,23 @@ func New(cfg config.Config, deps Dependencies, logger *zap.Logger) *gin.Engine { } notificationService := notification.NewService(notificationRepo) notificationHandler := notification.NewHandler(notificationService) + var chatHub *chathub.Hub + if deps.DB != nil { + chatHub = chathub.NewHub(deps.DB) + } var chatRepo *chat.Repository if deps.DB != nil { - chatRepo = chat.NewRepository(deps.DB) + chatRepo = chat.NewRepository(deps.DB, chatHub) } chatService := chat.NewService(chatRepo) chatHandler := chat.NewHandler(chatService) + var chatHubHandler *chathub.Handler + if chatHub != nil { + chatHubHandler = chathub.NewHandler(chatHub) + } + if orderRepo != nil && chatRepo != nil { + orderRepo.SetChatRepo(chatRepo) + } var disputeRepo *dispute.Repository if deps.DB != nil { disputeRepo = dispute.NewRepository(deps.DB) @@ -232,6 +244,9 @@ func New(cfg config.Config, deps Dependencies, logger *zap.Logger) *gin.Engine { chatRoutes := api.Group("/chats", requireAuth) { + if chatHubHandler != nil { + chatRoutes.GET("/events", chatHubHandler.UserEvents) + } chatRoutes.GET("", chatHandler.List) chatRoutes.GET("/:id", chatHandler.Detail) chatRoutes.GET("/:id/messages", chatHandler.Messages) @@ -280,6 +295,9 @@ func New(cfg config.Config, deps Dependencies, logger *zap.Logger) *gin.Engine { adminRoutes.GET("/system-configs", requirePerm("system_config:view"), systemConfigHandler.List) adminRoutes.PUT("/system-configs/:key", requirePerm("system_config:update"), systemConfigHandler.Update) adminRoutes.GET("/audit-logs", requirePerm("audit_log:view"), adminAuditHandler.List) + if chatHubHandler != nil { + adminRoutes.GET("/chats/events", requirePerm("chat:view"), chatHubHandler.AdminEvents) + } adminRoutes.GET("/chats", requirePerm("chat:view"), chatHandler.AdminList) adminRoutes.GET("/chats/:id", requirePerm("chat:view"), chatHandler.AdminDetail) adminRoutes.GET("/chats/:id/messages", requirePerm("chat:view"), chatHandler.AdminMessages) diff --git a/frontend/src/composables/useChatSSE.ts b/frontend/src/composables/useChatSSE.ts new file mode 100644 index 0000000..625002b --- /dev/null +++ b/frontend/src/composables/useChatSSE.ts @@ -0,0 +1,91 @@ +import { onBeforeUnmount, ref, type Ref } from 'vue' +import { getAccessToken, type AuthScope } from '@/utils/authStorage' + +export interface SSEMessage { + id: number + conversation_id: number + sender_type: string + sender_id: number + sender_role: string + sender_name: string + content_type: string + content: string + attachment_urls: string[] + created_at: string +} + +export interface ChatEvent { + type: 'new_message' | 'conversation_updated' + conversation_id: number + message?: SSEMessage +} + +type EventHandler = (event: ChatEvent) => void + +const reconnectDelay = 3000 + +export function useChatSSE(scope: AuthScope, endpoint: string) { + const connected: Ref = ref(false) + let source: EventSource | null = null + let reconnectTimer: ReturnType | null = null + let stopped = false + const handlers: EventHandler[] = [] + + function onEvent(handler: EventHandler) { + handlers.push(handler) + } + + function connect() { + const token = getAccessToken(scope) + if (!token) return + + const url = `${endpoint}?token=${encodeURIComponent(token)}` + source = new EventSource(url) + + source.addEventListener('connected', () => { + connected.value = true + }) + + source.addEventListener('new_message', (e) => { + try { + const data = JSON.parse((e as MessageEvent).data) as ChatEvent + handlers.forEach(h => h(data)) + } catch { /* ignore */ } + }) + + source.addEventListener('conversation_updated', (e) => { + try { + const data = JSON.parse((e as MessageEvent).data) as ChatEvent + handlers.forEach(h => h(data)) + } catch { /* ignore */ } + }) + + source.onerror = () => { + connected.value = false + source?.close() + source = null + if (!stopped) { + reconnectTimer = setTimeout(connect, reconnectDelay) + } + } + } + + function disconnect() { + stopped = true + if (reconnectTimer) { + clearTimeout(reconnectTimer) + reconnectTimer = null + } + source?.close() + source = null + connected.value = false + } + + onBeforeUnmount(() => { + disconnect() + }) + + connect() + + return { connected, onEvent, disconnect } +} diff --git a/frontend/src/views/admin/AdminChatsView.vue b/frontend/src/views/admin/AdminChatsView.vue index e1b1d88..c09dee7 100644 --- a/frontend/src/views/admin/AdminChatsView.vue +++ b/frontend/src/views/admin/AdminChatsView.vue @@ -1,5 +1,5 @@