Files
yunzerwebsiteallinone/go/controllers/backend_ai_chat.go
T
2026-09-03 17:54:49 +08:00

795 lines
21 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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()
}
}