package controllers import ( "bufio" "bytes" "encoding/json" "fmt" "io" "net/http" "strconv" "strings" "time" "server/models" "server/pkg/jwtutil" beego "github.com/beego/beego/v2/server/web" ) // BackendAiChatController AI聊天消息控制器 type BackendAiChatController struct { beego.Controller } func (c *BackendAiChatController) chatClaims() (*jwtutil.Claims, error) { auth := c.Ctx.Request.Header.Get("Authorization") if auth == "" { return nil, fmt.Errorf("未登录") } parts := strings.SplitN(auth, " ", 2) if len(parts) != 2 || parts[0] != "Bearer" { return nil, fmt.Errorf("认证信息格式错误") } claims, err := jwtutil.ParseToken(parts[1]) if err != nil { return nil, fmt.Errorf("无效的token") } if claims.UserType != "backend" { return nil, fmt.Errorf("无权访问") } return claims, nil } func (c *BackendAiChatController) chatJsonErr(httpStatus, bizCode int, msg string) { c.Ctx.Output.SetStatus(httpStatus) c.Data["json"] = map[string]interface{}{"code": bizCode, "msg": msg} _ = c.ServeJSON() } func (c *BackendAiChatController) chatOk(data interface{}) { c.Data["json"] = map[string]interface{}{"code": 200, "msg": "success", "data": data} _ = c.ServeJSON() } type chatSendPayload struct { SessionID uint64 `json:"session_id"` Content string `json:"content"` ProviderID uint64 `json:"provider_id"` Model string `json:"model"` } type openaiMessage struct { Role string `json:"role"` Content string `json:"content"` } type openaiRequest struct { Model string `json:"model"` Messages []openaiMessage `json:"messages"` } type openaiResponse struct { Choices []struct { Message struct { Content string `json:"content"` } `json:"message"` } `json:"choices"` Error *struct { Message string `json:"message"` } `json:"error"` } type anthropicRequest struct { Model string `json:"model"` MaxTokens int `json:"max_tokens"` System string `json:"system,omitempty"` Messages []openaiMessage `json:"messages"` } type anthropicResponse struct { Content []struct { Text string `json:"text"` } `json:"content"` Error *struct { Message string `json:"message"` } `json:"error"` } // MessageList GET /backend/ai/chat/message/list?session_id=xxx func (c *BackendAiChatController) MessageList() { claims, err := c.chatClaims() if err != nil { c.chatJsonErr(401, 401, err.Error()) return } sessionIDStr := strings.TrimSpace(c.GetString("session_id")) if sessionIDStr == "" { c.chatJsonErr(400, 400, "缺少会话ID") return } sessionID, err := strconv.ParseUint(sessionIDStr, 10, 64) if err != nil { c.chatJsonErr(400, 400, "会话ID格式错误") return } // 验证会话归属 session := models.BackendAiChatSession{ID: sessionID} if err := models.Orm.Read(&session); err != nil { c.chatJsonErr(404, 404, "会话不存在") return } if session.TenantID != fmt.Sprintf("%d", claims.TenantId) || session.UserID != uint64(claims.UserID) { c.chatJsonErr(403, 403, "无权访问") return } var list []models.BackendAiChatMessage _, err = models.Orm.QueryTable(new(models.BackendAiChatMessage)). Filter("session_id", sessionID). OrderBy("id"). All(&list) if err != nil { c.chatJsonErr(500, 500, "查询失败: "+err.Error()) return } c.chatOk(map[string]interface{}{"list": list, "title": session.Title}) } // DeleteMessage DELETE /backend/ai/chat/message/:id func (c *BackendAiChatController) DeleteMessage() { claims, err := c.chatClaims() if err != nil { c.chatJsonErr(401, 401, err.Error()) return } idStr := c.Ctx.Input.Param(":id") id, err := strconv.ParseUint(idStr, 10, 64) if err != nil { c.chatJsonErr(400, 400, "消息ID格式错误") return } msg := models.BackendAiChatMessage{ID: id} if err := models.Orm.Read(&msg); err != nil { c.chatJsonErr(404, 404, "消息不存在") return } // 验证会话归属 session := models.BackendAiChatSession{ID: msg.SessionID} if err := models.Orm.Read(&session); err != nil { c.chatJsonErr(404, 404, "会话不存在") return } if session.TenantID != fmt.Sprintf("%d", claims.TenantId) || session.UserID != uint64(claims.UserID) { c.chatJsonErr(403, 403, "无权操作") return } if _, err := models.Orm.Delete(&msg); err != nil { c.chatJsonErr(500, 500, "删除失败: "+err.Error()) return } c.chatOk(nil) } // Send POST /backend/ai/chat/send func (c *BackendAiChatController) Send() { claims, err := c.chatClaims() if err != nil { c.chatJsonErr(401, 401, err.Error()) return } body, err := io.ReadAll(c.Ctx.Request.Body) if err != nil { c.chatJsonErr(400, 400, "读取请求体失败") return } var p chatSendPayload if err := json.Unmarshal(body, &p); err != nil { c.chatJsonErr(400, 400, "参数格式错误") return } if strings.TrimSpace(p.Content) == "" { c.chatJsonErr(400, 400, "消息内容不能为空") return } // 查找或创建会话 var session models.BackendAiChatSession if p.SessionID > 0 { session = models.BackendAiChatSession{ID: p.SessionID} if err := models.Orm.Read(&session); err != nil { c.chatJsonErr(404, 404, "会话不存在") return } if session.TenantID != fmt.Sprintf("%d", claims.TenantId) || session.UserID != uint64(claims.UserID) { c.chatJsonErr(403, 403, "无权操作") return } } else { // 新建会话,标题取消息前20字 title := p.Content if len([]rune(title)) > 20 { title = string([]rune(title)[:20]) + "..." } session = models.BackendAiChatSession{ TenantID: fmt.Sprintf("%d", claims.TenantId), UserID: uint64(claims.UserID), ProviderID: p.ProviderID, Title: title, CreateTime: time.Now(), UpdateTime: time.Now(), } id, err := models.Orm.Insert(&session) if err != nil { c.chatJsonErr(500, 500, "创建会话失败: "+err.Error()) return } session.ID = uint64(id) } // 查找AI接入配置 providerID := p.ProviderID if providerID == 0 { providerID = session.ProviderID } var provider models.BackendAiProvider if providerID > 0 { provider = models.BackendAiProvider{ID: providerID} if err := models.Orm.Read(&provider); err != nil { c.chatJsonErr(400, 400, "指定的AI接入配置不存在,请先在设置中配置") return } if provider.TenantID != fmt.Sprintf("%d", claims.TenantId) || provider.UserID != uint64(claims.UserID) { c.chatJsonErr(403, 403, "无权使用该配置") return } if provider.Status != 1 { c.chatJsonErr(400, 400, "该AI接入配置已被禁用") return } } else { // 找第一个启用的配置 var providers []models.BackendAiProvider _, err := models.Orm.QueryTable(new(models.BackendAiProvider)). Filter("tenant_id", fmt.Sprintf("%d", claims.TenantId)). Filter("user_id", uint64(claims.UserID)). Filter("status", 1). Filter("delete_time__isnull", true). OrderBy("-id"). Limit(1). All(&providers) if err != nil || len(providers) == 0 { c.chatJsonErr(400, 400, "尚未配置AI接入,请先在设置中配置OpenAI或Anthropic接入") return } provider = providers[0] } // 加载历史消息(最近20条) var history []models.BackendAiChatMessage _, _ = models.Orm.QueryTable(new(models.BackendAiChatMessage)). Filter("session_id", session.ID). OrderBy("-id"). Limit(20). All(&history) // 反转顺序 for i, j := 0, len(history)-1; i < j; i, j = i+1, j-1 { history[i], history[j] = history[j], history[i] } // 构建消息列表 messages := make([]openaiMessage, 0, len(history)+1) for _, m := range history { messages = append(messages, openaiMessage{Role: m.Role, Content: m.Content}) } messages = append(messages, openaiMessage{Role: "user", Content: p.Content}) // 保存用户消息 userMsg := models.BackendAiChatMessage{ TenantID: fmt.Sprintf("%d", claims.TenantId), SessionID: session.ID, Role: "user", Content: p.Content, CreateTime: time.Now(), } _, _ = models.Orm.Insert(&userMsg) // 确定使用的模型:用户指定 > provider第一个模型 useModel := strings.TrimSpace(p.Model) if useModel == "" { providerModels := parseProviderModels(provider.Models) if len(providerModels) > 0 { useModel = providerModels[0] } } if useModel == "" { c.chatJsonErr(400, 400, "未指定模型且该接入配置无可用模型") return } // 查询用户默认角色预设 var systemPrompt string var presets []models.BackendAiChatPreset _, _ = models.Orm.QueryTable(new(models.BackendAiChatPreset)). Filter("tenant_id", fmt.Sprintf("%d", claims.TenantId)). Filter("user_id", uint64(claims.UserID)). Filter("is_default", 1). Filter("delete_time__isnull", true). Limit(1). All(&presets) if len(presets) > 0 { systemPrompt = presets[0].Content } // 调用AI reply, err := callAI(provider, useModel, systemPrompt, messages) if err != nil { c.chatJsonErr(500, 500, "AI调用失败: "+err.Error()) return } // 保存AI回复 assistantMsg := models.BackendAiChatMessage{ TenantID: fmt.Sprintf("%d", claims.TenantId), SessionID: session.ID, Role: "assistant", Content: reply, CreateTime: time.Now(), } _, _ = models.Orm.Insert(&assistantMsg) // 更新会话时间 session.UpdateTime = time.Now() _, _ = models.Orm.Update(&session, "update_time") c.chatOk(map[string]interface{}{ "session_id": session.ID, "reply": reply, "message_id": assistantMsg.ID, }) } func callAI(provider models.BackendAiProvider, model string, systemPrompt string, messages []openaiMessage) (string, error) { client := &http.Client{Timeout: 120 * time.Second} if provider.ProviderType == "openai" { // OpenAI 兼容格式 reqMessages := messages if systemPrompt != "" { reqMessages = append([]openaiMessage{{Role: "system", Content: systemPrompt}}, messages...) } reqBody := openaiRequest{ Model: model, Messages: reqMessages, } jsonData, _ := json.Marshal(reqBody) // 处理URL:如果已经包含/chat/completions,直接使用;否则拼接 url := strings.TrimRight(provider.ApiBase, "/") if !strings.HasSuffix(url, "/chat/completions") { url = url + "/chat/completions" } req, err := http.NewRequest("POST", url, bytes.NewBuffer(jsonData)) if err != nil { return "", err } req.Header.Set("Content-Type", "application/json") req.Header.Set("Authorization", "Bearer "+provider.ApiKey) resp, err := client.Do(req) if err != nil { return "", err } defer resp.Body.Close() body, _ := io.ReadAll(resp.Body) var result openaiResponse if err := json.Unmarshal(body, &result); err != nil { return "", fmt.Errorf("解析响应失败: %s", string(body)) } if result.Error != nil { return "", fmt.Errorf(result.Error.Message) } if len(result.Choices) == 0 { return "", fmt.Errorf("AI未返回内容") } return result.Choices[0].Message.Content, nil } else { // Anthropic 格式 reqBody := anthropicRequest{ Model: model, MaxTokens: 4096, System: systemPrompt, Messages: messages, } jsonData, _ := json.Marshal(reqBody) // 处理URL:如果已经包含/v1/messages,直接使用;否则拼接 url := strings.TrimRight(provider.ApiBase, "/") if !strings.HasSuffix(url, "/v1/messages") { url = url + "/v1/messages" } req, err := http.NewRequest("POST", url, bytes.NewBuffer(jsonData)) if err != nil { return "", err } req.Header.Set("Content-Type", "application/json") req.Header.Set("x-api-key", provider.ApiKey) req.Header.Set("anthropic-version", "2023-06-01") resp, err := client.Do(req) if err != nil { return "", err } defer resp.Body.Close() body, _ := io.ReadAll(resp.Body) var result anthropicResponse if err := json.Unmarshal(body, &result); err != nil { return "", fmt.Errorf("解析响应失败: %s", string(body)) } if result.Error != nil { return "", fmt.Errorf(result.Error.Message) } if len(result.Content) == 0 { return "", fmt.Errorf("AI未返回内容") } return result.Content[0].Text, nil } } // SendStream POST /backend/ai/chat/send-stream // 流式响应(SSE),逐块返回AI回复 func (c *BackendAiChatController) SendStream() { claims, err := c.chatClaims() if err != nil { c.chatJsonErr(401, 401, err.Error()) return } body, err := io.ReadAll(c.Ctx.Request.Body) if err != nil { c.chatJsonErr(400, 400, "读取请求体失败") return } var p chatSendPayload if err := json.Unmarshal(body, &p); err != nil { c.chatJsonErr(400, 400, "参数格式错误") return } if strings.TrimSpace(p.Content) == "" { c.chatJsonErr(400, 400, "消息内容不能为空") return } // 查找或创建会话 var session models.BackendAiChatSession if p.SessionID > 0 { session = models.BackendAiChatSession{ID: p.SessionID} if err := models.Orm.Read(&session); err != nil { c.chatJsonErr(404, 404, "会话不存在") return } if session.TenantID != fmt.Sprintf("%d", claims.TenantId) || session.UserID != uint64(claims.UserID) { c.chatJsonErr(403, 403, "无权操作") return } } else { title := p.Content if len([]rune(title)) > 20 { title = string([]rune(title)[:20]) + "..." } session = models.BackendAiChatSession{ TenantID: fmt.Sprintf("%d", claims.TenantId), UserID: uint64(claims.UserID), ProviderID: p.ProviderID, Title: title, CreateTime: time.Now(), UpdateTime: time.Now(), } id, err := models.Orm.Insert(&session) if err != nil { c.chatJsonErr(500, 500, "创建会话失败: "+err.Error()) return } session.ID = uint64(id) } // 查找AI接入配置 providerID := p.ProviderID if providerID == 0 { providerID = session.ProviderID } var provider models.BackendAiProvider if providerID > 0 { provider = models.BackendAiProvider{ID: providerID} if err := models.Orm.Read(&provider); err != nil { c.chatJsonErr(400, 400, "指定的AI接入配置不存在") return } if provider.TenantID != fmt.Sprintf("%d", claims.TenantId) || provider.UserID != uint64(claims.UserID) { c.chatJsonErr(403, 403, "无权使用该配置") return } if provider.Status != 1 { c.chatJsonErr(400, 400, "该AI接入配置已被禁用") return } } else { var providers []models.BackendAiProvider _, err := models.Orm.QueryTable(new(models.BackendAiProvider)). Filter("tenant_id", fmt.Sprintf("%d", claims.TenantId)). Filter("user_id", uint64(claims.UserID)). Filter("status", 1). Filter("delete_time__isnull", true). OrderBy("-id"). Limit(1). All(&providers) if err != nil || len(providers) == 0 { c.chatJsonErr(400, 400, "尚未配置AI接入,请先在设置中配置") return } provider = providers[0] } // 确定模型 useModel := strings.TrimSpace(p.Model) if useModel == "" { providerModels := parseProviderModels(provider.Models) if len(providerModels) > 0 { useModel = providerModels[0] } } if useModel == "" { c.chatJsonErr(400, 400, "未指定模型且该接入配置无可用模型") return } // 查询默认预设 var systemPrompt string var presets []models.BackendAiChatPreset _, _ = models.Orm.QueryTable(new(models.BackendAiChatPreset)). Filter("tenant_id", fmt.Sprintf("%d", claims.TenantId)). Filter("user_id", uint64(claims.UserID)). Filter("is_default", 1). Filter("delete_time__isnull", true). Limit(1). All(&presets) if len(presets) > 0 { systemPrompt = presets[0].Content } // 加载历史消息 var history []models.BackendAiChatMessage _, _ = models.Orm.QueryTable(new(models.BackendAiChatMessage)). Filter("session_id", session.ID). OrderBy("-id"). Limit(20). All(&history) for i, j := 0, len(history)-1; i < j; i, j = i+1, j-1 { history[i], history[j] = history[j], history[i] } messages := make([]openaiMessage, 0, len(history)+1) for _, m := range history { messages = append(messages, openaiMessage{Role: m.Role, Content: m.Content}) } messages = append(messages, openaiMessage{Role: "user", Content: p.Content}) // 保存用户消息 userMsg := models.BackendAiChatMessage{ TenantID: fmt.Sprintf("%d", claims.TenantId), SessionID: session.ID, Role: "user", Content: p.Content, CreateTime: time.Now(), } _, _ = models.Orm.Insert(&userMsg) // 设置SSE响应头 - 直接操作底层ResponseWriter,绕过beego包装层 rw := c.Ctx.ResponseWriter.ResponseWriter rw.Header().Set("Content-Type", "text/event-stream") rw.Header().Set("Cache-Control", "no-cache") rw.Header().Set("Connection", "keep-alive") rw.Header().Set("X-Accel-Buffering", "no") rw.WriteHeader(200) var flusher http.Flusher if f, ok := rw.(http.Flusher); ok { flusher = f } writeSSE := func(event string, data string) { rw.Write([]byte(fmt.Sprintf("event: %s\ndata: %s\n\n", event, data))) if flusher != nil { flusher.Flush() } } // 发送session_id事件 writeSSE("session", fmt.Sprintf("{\"session_id\":%d}", session.ID)) // 流式调用AI fullReply := "" streamErr := callAIStream(provider, useModel, systemPrompt, messages, func(chunk string) { fullReply += chunk chunkJSON, _ := json.Marshal(map[string]string{"content": chunk}) writeSSE("content", string(chunkJSON)) }) if streamErr != nil { errJSON, _ := json.Marshal(map[string]string{"error": streamErr.Error()}) writeSSE("error", string(errJSON)) return } // 保存AI回复 assistantMsg := models.BackendAiChatMessage{ TenantID: fmt.Sprintf("%d", claims.TenantId), SessionID: session.ID, Role: "assistant", Content: fullReply, CreateTime: time.Now(), } _, _ = models.Orm.Insert(&assistantMsg) // 更新会话时间 session.UpdateTime = time.Now() _, _ = models.Orm.Update(&session, "update_time") // 发送done事件 doneJSON, _ := json.Marshal(map[string]interface{}{ "session_id": session.ID, "message_id": assistantMsg.ID, }) writeSSE("done", string(doneJSON)) } // callAIStream 流式调用AI,onChunk回调每收到一个内容块就调用 func callAIStream(provider models.BackendAiProvider, model string, systemPrompt string, messages []openaiMessage, onChunk func(string)) error { client := &http.Client{Timeout: 120 * time.Second} if provider.ProviderType == "openai" { // OpenAI 兼容格式 - 流式 reqMessages := messages if systemPrompt != "" { reqMessages = append([]openaiMessage{{Role: "system", Content: systemPrompt}}, messages...) } reqBody := map[string]interface{}{ "model": model, "messages": reqMessages, "stream": true, } jsonData, _ := json.Marshal(reqBody) // 处理URL:如果已经包含/chat/completions,直接使用;否则拼接 url := strings.TrimRight(provider.ApiBase, "/") if !strings.HasSuffix(url, "/chat/completions") { url = url + "/chat/completions" } req, err := http.NewRequest("POST", url, bytes.NewBuffer(jsonData)) if err != nil { return err } req.Header.Set("Content-Type", "application/json") req.Header.Set("Authorization", "Bearer "+provider.ApiKey) resp, err := client.Do(req) if err != nil { return err } defer resp.Body.Close() if resp.StatusCode != 200 { body, _ := io.ReadAll(resp.Body) return fmt.Errorf("API返回错误状态 %d: %s", resp.StatusCode, string(body)) } scanner := bufio.NewScanner(resp.Body) scanner.Buffer(make([]byte, 1024*1024), 1024*1024) for scanner.Scan() { line := scanner.Text() if !strings.HasPrefix(line, "data: ") { continue } data := strings.TrimPrefix(line, "data: ") if data == "[DONE]" { break } var chunk struct { Choices []struct { Delta struct { Content string `json:"content"` } `json:"delta"` } `json:"choices"` } if err := json.Unmarshal([]byte(data), &chunk); err != nil { continue } if len(chunk.Choices) > 0 && chunk.Choices[0].Delta.Content != "" { onChunk(chunk.Choices[0].Delta.Content) } } return scanner.Err() } else { // Anthropic 格式 - 流式 reqBody := map[string]interface{}{ "model": model, "max_tokens": 4096, "messages": messages, "stream": true, } if systemPrompt != "" { reqBody["system"] = systemPrompt } jsonData, _ := json.Marshal(reqBody) // 处理URL:如果已经包含/v1/messages,直接使用;否则拼接 url := strings.TrimRight(provider.ApiBase, "/") if !strings.HasSuffix(url, "/v1/messages") { url = url + "/v1/messages" } req, err := http.NewRequest("POST", url, bytes.NewBuffer(jsonData)) if err != nil { return err } req.Header.Set("Content-Type", "application/json") req.Header.Set("x-api-key", provider.ApiKey) req.Header.Set("anthropic-version", "2023-06-01") resp, err := client.Do(req) if err != nil { return err } defer resp.Body.Close() if resp.StatusCode != 200 { body, _ := io.ReadAll(resp.Body) return fmt.Errorf("API返回错误状态 %d: %s", resp.StatusCode, string(body)) } scanner := bufio.NewScanner(resp.Body) scanner.Buffer(make([]byte, 1024*1024), 1024*1024) currentEvent := "" for scanner.Scan() { line := scanner.Text() if strings.HasPrefix(line, "event: ") { currentEvent = strings.TrimPrefix(line, "event: ") continue } if strings.HasPrefix(line, "data: ") { data := strings.TrimPrefix(line, "data: ") if currentEvent == "content_block_delta" { var chunk struct { Delta struct { Text string `json:"text"` } `json:"delta"` } if err := json.Unmarshal([]byte(data), &chunk); err == nil && chunk.Delta.Text != "" { onChunk(chunk.Delta.Text) } } if currentEvent == "message_stop" { break } } } return scanner.Err() } }