package chathub import ( "context" "fmt" "time" "hfb_sys/backend/internal/middleware" "hfb_sys/backend/pkg/response" "github.com/gin-gonic/gin" "go.uber.org/zap" ) const heartbeatInterval = 30 * time.Second type Handler struct { hub *Hub logger *zap.Logger } func NewHandler(hub *Hub, logger *zap.Logger) *Handler { if logger == nil { logger = zap.NewNop() } return &Handler{hub: hub, logger: logger} } // 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) { startedAt := time.Now() disconnectReason := "handler_completed" var disconnectErr error ch := h.hub.Subscribe(pType, pID) h.logger.Info("SSE 连接建立", h.connectionLogFields(c, pType, pID)...) defer func() { h.hub.Unsubscribe(pType, pID, ch) fields := h.connectionLogFields(c, pType, pID) fields = append(fields, zap.String("disconnect_reason", disconnectReason), zap.Float64("lifetime_ms", float64(time.Since(startedAt).Microseconds())/1000), ) if disconnectErr != nil { fields = append(fields, zap.Error(disconnectErr)) } h.logger.Info("SSE 连接断开", fields...) }() c.Header("Content-Type", "text/event-stream") c.Header("Cache-Control", "no-cache") c.Header("Connection", "keep-alive") c.Header("X-Accel-Buffering", "no") // 发送初始连接确认 if _, err := fmt.Fprintf(c.Writer, "event: connected\ndata: {\"ok\":true}\n\n"); err != nil { disconnectReason = "initial_write_error" disconnectErr = err return } c.Writer.Flush() heartbeat := time.NewTicker(heartbeatInterval) defer heartbeat.Stop() clientGone := c.Request.Context().Done() for { select { case <-clientGone: disconnectReason = contextDisconnectReason(c.Request.Context().Err()) return case <-heartbeat.C: if _, err := fmt.Fprintf(c.Writer, ":heartbeat\n\n"); err != nil { disconnectReason = "heartbeat_write_error" disconnectErr = err return } c.Writer.Flush() case event, ok := <-ch: if !ok { disconnectReason = "subscription_closed" return } data := MarshalEvent(event) if _, err := fmt.Fprintf(c.Writer, "event: %s\ndata: %s\n\n", event.Type, data); err != nil { disconnectReason = "event_write_error" disconnectErr = err return } c.Writer.Flush() } } } func (h *Handler) connectionLogFields(c *gin.Context, pType string, pID uint64) []zap.Field { fields := []zap.Field{ zap.String("request_id", middleware.GetRequestID(c)), zap.String("principal_type", pType), zap.Uint64("principal_id", pID), zap.String("client_ip", c.ClientIP()), zap.Int("active_sse_connections", h.hub.OnlineCount()), zap.Int("active_principal_type_connections", h.hub.OnlineCountByType(pType)), } if pType == "admin" { fields = append(fields, zap.Uint64("admin_id", pID)) } else { fields = append(fields, zap.Uint64("user_id", pID)) } return fields } func contextDisconnectReason(err error) string { switch err { case context.Canceled: return "context_canceled" case context.DeadlineExceeded: return "context_deadline_exceeded" default: return "context_closed" } }