package handler import ( "bytes" "encoding/json" "fmt" "io" "log" "net/http" "time" "github.com/gin-gonic/gin" "github.com/gorilla/websocket" "github.com/yourname/cyrene-ai/gateway/internal/config" "github.com/yourname/cyrene-ai/gateway/internal/middleware" "github.com/yourname/cyrene-ai/gateway/internal/ws" ) // ChatHandler 聊天处理器 type ChatHandler struct { cfg *config.Config hub *ws.Hub upgrader websocket.Upgrader } // NewChatHandler 创建聊天处理器 func NewChatHandler(cfg *config.Config, hub *ws.Hub) *ChatHandler { return &ChatHandler{ cfg: cfg, hub: hub, upgrader: websocket.Upgrader{ ReadBufferSize: 1024, WriteBufferSize: 1024, CheckOrigin: func(r *http.Request) bool { return true // 开发阶段允许所有来源 }, }, } } // HandleWebSocket 处理WebSocket升级和消息路由 func (h *ChatHandler) HandleWebSocket(c *gin.Context) { // 从query参数获取token和session_id token := c.Query("token") sessionID := c.Query("session_id") if token == "" { // 也尝试从Authorization头读取 authHeader := c.GetHeader("Authorization") if len(authHeader) > 7 && authHeader[:7] == "Bearer " { token = authHeader[7:] } } if token == "" { c.JSON(http.StatusUnauthorized, gin.H{"error": "需要认证令牌"}) return } // 验证token userID, err := h.cfg.ValidateToken(token) if err != nil { c.JSON(http.StatusUnauthorized, gin.H{"error": "认证令牌无效"}) return } if sessionID == "" { sessionID = "session_" + generateID() } // 升级WebSocket连接 conn, err := h.upgrader.Upgrade(c.Writer, c.Request, nil) if err != nil { log.Printf("[WS] 升级连接失败: %v", err) return } // 创建客户端 client := ws.NewClient(h.hub, conn, userID, sessionID) // 注册到Hub h.hub.Register(client) // 启动读写协程 go client.WritePump() go client.ReadPump(func(client *ws.Client, msg ws.ClientMessage) { h.handleMessage(client, msg) }) } // handleMessage 处理WebSocket消息 func (h *ChatHandler) handleMessage(client *ws.Client, msg ws.ClientMessage) { switch msg.Type { case "message": h.handleChatMessage(client, msg) case "voice_input": h.handleVoiceInput(client, msg) default: log.Printf("[WS] 未知消息类型: %s from user=%s", msg.Type, client.UserID) } } // handleChatMessage 处理文字聊天消息 - 转发到 AI-Core func (h *ChatHandler) handleChatMessage(client *ws.Client, msg ws.ClientMessage) { mode := msg.Mode if mode == "" { mode = "text" } // 记录用户消息 h.hub.RecordMessage(client.SessionID, "user", msg.Content) // 设置会话状态为 thinking h.hub.UpdateSessionState(client.SessionID, "thinking") // 构建 AI-Core 请求 aiReq := map[string]string{ "user_id": client.UserID, "session_id": client.SessionID, "message": msg.Content, "mode": mode, } reqBody, err := json.Marshal(aiReq) if err != nil { log.Printf("[chat] 序列化请求失败: %v", err) h.hub.UpdateSessionState(client.SessionID, "error") client.SendMessage(ws.ServerMessage{ Type: "error", MessageID: "msg_" + generateID(), Error: "内部错误,请稍后重试", Timestamp: time.Now().UnixMilli(), }) return } // 调用 AI-Core aiCoreURL := h.cfg.AICoreURL + "/api/v1/chat" httpReq, err := http.NewRequest("POST", aiCoreURL, bytes.NewReader(reqBody)) if err != nil { log.Printf("[chat] 创建 AI-Core 请求失败: %v", err) h.hub.UpdateSessionState(client.SessionID, "error") client.SendMessage(ws.ServerMessage{ Type: "error", MessageID: "msg_" + generateID(), Error: "服务暂不可用", Timestamp: time.Now().UnixMilli(), }) return } httpReq.Header.Set("Content-Type", "application/json") httpClient := &http.Client{Timeout: 120 * time.Second} resp, err := httpClient.Do(httpReq) if err != nil { log.Printf("[chat] AI-Core 调用失败: %v", err) h.hub.UpdateSessionState(client.SessionID, "error") client.SendMessage(ws.ServerMessage{ Type: "error", MessageID: "msg_" + generateID(), Error: fmt.Sprintf("AI-Core 调用失败: %v", err), Timestamp: time.Now().UnixMilli(), }) return } defer resp.Body.Close() body, err := io.ReadAll(resp.Body) if err != nil { log.Printf("[chat] 读取 AI-Core 响应失败: %v", err) h.hub.UpdateSessionState(client.SessionID, "error") client.SendMessage(ws.ServerMessage{ Type: "error", MessageID: "msg_" + generateID(), Error: "读取响应失败", Timestamp: time.Now().UnixMilli(), }) return } if resp.StatusCode != http.StatusOK { log.Printf("[chat] AI-Core 返回错误 [%d]: %s", resp.StatusCode, string(body)) h.hub.UpdateSessionState(client.SessionID, "error") client.SendMessage(ws.ServerMessage{ Type: "error", MessageID: "msg_" + generateID(), Error: fmt.Sprintf("AI-Core 错误 (%d)", resp.StatusCode), Timestamp: time.Now().UnixMilli(), }) return } // 解析 AI-Core 响应 var aiResp struct { Text string `json:"text"` Mode string `json:"mode"` MessageID string `json:"message_id"` } if err := json.Unmarshal(body, &aiResp); err != nil { log.Printf("[chat] 解析 AI-Core 响应失败: %v", err) h.hub.UpdateSessionState(client.SessionID, "error") client.SendMessage(ws.ServerMessage{ Type: "error", MessageID: "msg_" + generateID(), Error: "解析响应失败", Timestamp: time.Now().UnixMilli(), }) return } // 记录助手响应 h.hub.RecordMessage(client.SessionID, "assistant", aiResp.Text) // 设置会话状态为 idle h.hub.UpdateSessionState(client.SessionID, "idle") // 发送响应给客户端 response := ws.ServerMessage{ Type: "response", MessageID: aiResp.MessageID, Text: aiResp.Text, ResponseMode: mode, Timestamp: time.Now().UnixMilli(), } if err := client.SendMessage(response); err != nil { log.Printf("[WS] 发送响应失败: %v", err) } } // handleVoiceInput 处理语音输入 func (h *ChatHandler) handleVoiceInput(client *ws.Client, msg ws.ClientMessage) { // MVP阶段:返回提示 response := ws.ServerMessage{ Type: "error", MessageID: "msg_" + generateID(), Error: "语音处理功能将在后续版本中启用", Timestamp: time.Now().UnixMilli(), } client.SendMessage(response) } // SendSystemMessage 向用户发送系统消息(用于主动通知) func (h *ChatHandler) SendSystemMessage(userID, sessionID, text string) error { msg := ws.ServerMessage{ Type: "response", MessageID: "sys_" + generateID(), Text: text, Timestamp: time.Now().UnixMilli(), } data, err := json.Marshal(msg) if err != nil { return err } h.hub.SendToSession(userID, sessionID, data) return nil } func generateID() string { return time.Now().Format("20060102150405") + randomStr(6) } func randomStr(n int) string { const letters = "abcdefghijklmnopqrstuvwxyz0123456789" b := make([]byte, n) for i := range b { b[i] = letters[time.Now().UnixNano()%int64(len(letters))] } return string(b) } // 确保未使用变量不报错 var _ = middleware.GetUserID