Files
yunzerwebsiteallinone/go/controllers/backend_ai_chat.go
T
2026-09-11 00:22:49 +08:00

1489 lines
43 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"
"server/services"
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/Model 由前端选择
ProviderID uint64 `json:"provider_id"`
Model string `json:"model"`
// McpEnabled 是否为本轮对话注入 MCP 工具;不传(nil)视为开启,兼容旧客户端
McpEnabled *bool `json:"mcp_enabled"`
// McpServerIDs 限定参与本轮对话的 MCP 服务器;为空表示使用全部已启用的服务器
McpServerIDs []uint64 `json:"mcp_server_ids"`
// Images 用户上传的图片(base64 dataURL),仅多模态模型可用
Images []string `json:"images"`
}
// mcpEnabled 解析本轮是否启用 MCP(未传默认开启)
func (p *chatSendPayload) mcpEnabled() bool {
return p.McpEnabled == nil || *p.McpEnabled
}
// ============ 通用消息/工具结构 ============
// openaiMessage 通用消息(OpenAI 与 Anthropic 共用内部表示)
type openaiMessage struct {
Role string `json:"role"`
Content string `json:"content,omitempty"`
ToolCalls []openaiToolCall `json:"tool_calls,omitempty"`
ToolCallID string `json:"tool_call_id,omitempty"`
// Images 仅 user 消息使用,存放 base64 dataURL,用于多模态输入
Images []string `json:"-"`
}
// openAIMessages 将内部消息转换为 OpenAI 兼容的请求结构。
// 当且仅当 user 消息携带图片时,content 会被构造为多模态 parts(text + image_url)。
func openAIMessages(messages []openaiMessage) []map[string]interface{} {
out := make([]map[string]interface{}, 0, len(messages))
for _, m := range messages {
item := map[string]interface{}{"role": m.Role}
if len(m.Images) > 0 {
parts := []map[string]interface{}{}
if strings.TrimSpace(m.Content) != "" {
parts = append(parts, map[string]interface{}{"type": "text", "text": m.Content})
}
for _, img := range m.Images {
parts = append(parts, map[string]interface{}{
"type": "image_url",
"image_url": map[string]interface{}{"url": img},
})
}
item["content"] = parts
} else {
item["content"] = m.Content
}
if len(m.ToolCalls) > 0 {
item["tool_calls"] = m.ToolCalls
}
if m.ToolCallID != "" {
item["tool_call_id"] = m.ToolCallID
}
out = append(out, item)
}
return out
}
// marshalImages 将图片列表序列化为可存储的 JSON 字符串(无图片返回空串)
func marshalImages(imgs []string) string {
if len(imgs) == 0 {
return ""
}
b, _ := json.Marshal(imgs)
return string(b)
}
type openaiToolCall struct {
ID string `json:"id,omitempty"`
Type string `json:"type,omitempty"`
Function struct {
Name string `json:"name,omitempty"`
Arguments string `json:"arguments,omitempty"`
} `json:"function,omitempty"`
}
type openaiTool struct {
Type string `json:"type"`
Function struct {
Name string `json:"name"`
Description string `json:"description"`
Parameters interface{} `json:"parameters"`
} `json:"function"`
}
// openaiDeltaToolCall 流式增量工具调用
type openaiDeltaToolCall struct {
Index int `json:"index"`
ID string `json:"id,omitempty"`
Type string `json:"type,omitempty"`
Function struct {
Name string `json:"name,omitempty"`
Arguments string `json:"arguments,omitempty"`
} `json:"function,omitempty"`
}
// pendingToolCall 一轮中 AI 请求调用的工具
type pendingToolCall struct {
ID string
Name string // 唯一工具key(mcp_{serverID}_{toolName})
Args map[string]interface{}
}
// toolRef 工具key -> 实际 MCP 服务器与工具
type toolRef struct {
ServerID uint64
ServerName string
ToolName string
}
// mcpToolSummary 前端展示用工具摘要
type mcpToolSummary struct {
Key string `json:"key"`
ServerID uint64 `json:"server_id"`
ServerName string `json:"server_name"`
ToolName string `json:"tool_name"`
Description string `json:"description"`
}
const maxToolRounds = 6
// aiUsage 一次 AI 调用的 token 用量与真实模型
type aiUsage struct {
Model string `json:"model,omitempty"` // 接口实际返回的使用模型(如网关把 auto 解析后的真实模型)
PromptTokens int `json:"prompt_tokens"`
CompletionTokens int `json:"completion_tokens"`
TotalTokens int `json:"total_tokens"`
}
// add 累加用量
func (u *aiUsage) add(o aiUsage) {
u.PromptTokens += o.PromptTokens
u.CompletionTokens += o.CompletionTokens
u.TotalTokens += o.TotalTokens
}
// total 返回本次调用的总 token 数(部分厂商不返回 total,则自行求和)
func (u aiUsage) total() int {
if u.TotalTokens > 0 {
return u.TotalTokens
}
return u.PromptTokens + u.CompletionTokens
}
// usedMcpTool 单次调用中实际使用的 MCP 工具(含所属服务器)
type usedMcpTool struct {
ServerID uint64 `json:"server_id"`
ServerName string `json:"server_name"`
ToolName string `json:"tool_name"`
}
// aiToolUsage 一次 AI 工具循环调用的元信息:用了什么模型、调用了哪些 MCP 工具
type aiToolUsage struct {
Model string `json:"model"` // 实际使用的模型
Rounds int `json:"rounds"` // 工具循环轮数
Tools []usedMcpTool `json:"tools"` // 实际被调用的 MCP 工具(去重)
ToolCount int `json:"tool_count"` // 工具调用种类数
}
// sanitizeToolName 工具名规范化为 OpenAI 允许的字符集
func sanitizeToolName(s string) string {
var sb strings.Builder
for _, r := range s {
if (r >= 'a' && r <= 'z') || (r >= 'A' && r <= 'Z') ||
(r >= '0' && r <= '9') || r == '_' || r == '-' {
sb.WriteRune(r)
} else {
sb.WriteRune('_')
}
}
if sb.Len() == 0 {
return "tool"
}
return sb.String()
}
// mcpToolKey 生成 LLM 调用使用的唯一工具名
func mcpToolKey(serverID uint64, name string) string {
key := fmt.Sprintf("mcp_%d_%s", serverID, sanitizeToolName(name))
if len(key) > 64 {
key = key[:64]
}
return key
}
// collectMcpTools 汇总「已启用」的 MCP 服务器工具,转换为 LLM 工具列表
// enabled=false 时不注入任何工具;serverIDs 非空时只取指定的服务器
func collectMcpTools(claims *jwtutil.Claims, enabled bool, serverIDs []uint64) ([]openaiTool, map[string]toolRef, []mcpToolSummary, error) {
if !enabled {
return nil, map[string]toolRef{}, nil, nil
}
var servers []models.BackendMcpServer
qs := models.Orm.QueryTable(new(models.BackendMcpServer)).
Filter("tenant_id", fmt.Sprintf("%d", claims.TenantId)).
Filter("user_id", uint64(claims.UserID)).
Filter("enabled", 1).
Filter("delete_time__isnull", true)
if len(serverIDs) > 0 {
qs = qs.Filter("id__in", serverIDs)
}
_, err := qs.OrderBy("id").All(&servers)
if err != nil {
return nil, nil, nil, err
}
var tools []openaiTool
refMap := make(map[string]toolRef)
var summaries []mcpToolSummary
for i := range servers {
srv := servers[i]
toolInfos, err := services.McpClientManager.EnsureConnected(&srv)
if err != nil {
// 单个服务连接失败不阻断整体,跳过并记录到摘要(标记不可用)
summaries = append(summaries, mcpToolSummary{
ServerID: srv.ID,
ServerName: srv.Name,
Description: "连接失败:" + err.Error(),
})
continue
}
for _, ti := range toolInfos {
key := mcpToolKey(srv.ID, ti.Name)
var fn openaiTool
fn.Type = "function"
fn.Function.Name = key
fn.Function.Description = ti.Description
if ti.InputSchema == nil {
fn.Function.Parameters = map[string]interface{}{
"type": "object",
"properties": map[string]interface{}{},
}
} else {
fn.Function.Parameters = ti.InputSchema
}
tools = append(tools, fn)
refMap[key] = toolRef{ServerID: srv.ID, ServerName: srv.Name, ToolName: ti.Name}
summaries = append(summaries, mcpToolSummary{
Key: key,
ServerID: srv.ID,
ServerName: srv.Name,
ToolName: ti.Name,
Description: ti.Description,
})
}
}
return tools, refMap, summaries, nil
}
// executePendingTools 执行 AI 请求的工具调用,返回工具结果消息
func executePendingTools(claims *jwtutil.Claims, pending []pendingToolCall, refMap map[string]toolRef, onStart func(pc pendingToolCall), onResult func(pc pendingToolCall, text string, ok bool, errMsg string)) []openaiMessage {
var toolMsgs []openaiMessage
for _, pc := range pending {
ref, ok := refMap[pc.Name]
if !ok {
msg := fmt.Sprintf("工具 %s 未找到对应 MCP 服务", pc.Name)
if onStart != nil {
onStart(pc)
}
if onResult != nil {
onResult(pc, "", false, msg)
}
toolMsgs = append(toolMsgs, openaiMessage{Role: "tool", ToolCallID: pc.ID, Content: msg})
continue
}
if onStart != nil {
onStart(pc)
}
text, isErr, err := services.McpClientManager.CallTool(ref.ServerID, ref.ToolName, pc.Args)
if err != nil {
errMsg := "工具调用失败: " + err.Error()
if onResult != nil {
onResult(pc, "", false, errMsg)
}
toolMsgs = append(toolMsgs, openaiMessage{Role: "tool", ToolCallID: pc.ID, Content: errMsg})
continue
}
if onResult != nil {
onResult(pc, text, !isErr, "")
}
toolMsgs = append(toolMsgs, openaiMessage{Role: "tool", ToolCallID: pc.ID, Content: text})
}
return toolMsgs
}
// buildMcpAssistantMessage 构造携带 tool_calls 的 assistant 消息
func buildMcpAssistantMessage(text string, pending []pendingToolCall) openaiMessage {
msg := openaiMessage{Role: "assistant", Content: text}
for _, pc := range pending {
argsStr := "{}"
if pc.Args != nil {
if b, err := json.Marshal(pc.Args); err == nil {
argsStr = string(b)
}
}
tc := openaiToolCall{ID: pc.ID, Type: "function"}
tc.Function.Name = pc.Name
tc.Function.Arguments = argsStr
msg.ToolCalls = append(msg.ToolCalls, tc)
}
return msg
}
// runToolLoop 运行带已启用 MCP 工具的完整对话循环(非流式),返回最终文本与调用元信息。
// 供聊天非流式接口与智能生成等模块复用。
func runToolLoop(claims *jwtutil.Claims, provider models.BackendAiProvider, model string, systemPrompt string, messages []openaiMessage) (string, *aiToolUsage, error) {
// 智能生成等内部调用:默认注入全部已启用的 MCP 工具
llmTools, refMap, _, _ := collectMcpTools(claims, true, nil)
usage := &aiToolUsage{Model: model}
usedSet := make(map[string]bool)
addUsed := func(key string) {
ref, ok := refMap[key]
if !ok {
return
}
id := fmt.Sprintf("%d:%s", ref.ServerID, ref.ToolName)
if usedSet[id] {
return
}
usedSet[id] = true
usage.Tools = append(usage.Tools, usedMcpTool{
ServerID: ref.ServerID,
ServerName: ref.ServerName,
ToolName: ref.ToolName,
})
}
current := messages
rounds := 0
for {
rounds++
text, pending, roundUsage, err := callAITools(provider, model, systemPrompt, current, llmTools)
// 用接口实际返回的真实模型覆盖占位符(如 auto),若网关未回传则保留请求模型
if roundUsage.Model != "" {
usage.Model = roundUsage.Model
}
if err != nil {
usage.Rounds = rounds
usage.ToolCount = len(usage.Tools)
return "", usage, err
}
for _, pc := range pending {
addUsed(pc.Name)
}
if len(pending) == 0 {
usage.Rounds = rounds
usage.ToolCount = len(usage.Tools)
return text, usage, nil
}
if rounds >= maxToolRounds {
usage.Rounds = rounds
usage.ToolCount = len(usage.Tools)
return text, usage, nil
}
current = append(current, buildMcpAssistantMessage(text, pending))
toolMsgs := executePendingTools(claims, pending, refMap, nil, nil)
current = append(current, toolMsgs...)
}
}
// ============ 会话/消息接口 ============
// 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)
}
// ============ 公共:会话/接入/预设解析 ============
// resolveOrCreateSession 查找或创建会话
func resolveOrCreateSession(c *BackendAiChatController, claims *jwtutil.Claims, p *chatSendPayload) (models.BackendAiChatSession, error) {
var session models.BackendAiChatSession
if p.SessionID > 0 {
session = models.BackendAiChatSession{ID: p.SessionID}
if err := models.Orm.Read(&session); err != nil {
return session, fmt.Errorf("会话不存在")
}
if session.TenantID != fmt.Sprintf("%d", claims.TenantId) || session.UserID != uint64(claims.UserID) {
return session, fmt.Errorf("无权操作")
}
} 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 {
return session, fmt.Errorf("创建会话失败: " + err.Error())
}
session.ID = uint64(id)
}
return session, nil
}
// resolveProvider 解析 AI 接入配置
func resolveProvider(claims *jwtutil.Claims, p *chatSendPayload, session *models.BackendAiChatSession) (models.BackendAiProvider, error) {
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 {
return provider, fmt.Errorf("指定的AI接入配置不存在,请先在设置中配置")
}
if provider.TenantID != fmt.Sprintf("%d", claims.TenantId) || provider.UserID != uint64(claims.UserID) {
return provider, fmt.Errorf("无权使用该配置")
}
if provider.Status != 1 {
return provider, fmt.Errorf("该AI接入配置已被禁用")
}
} 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 {
return provider, fmt.Errorf("尚未配置AI接入,请先在设置中配置OpenAI或Anthropic接入")
}
provider = providers[0]
}
return provider, nil
}
// resolveModel 确定使用的模型
func resolveModel(p *chatSendPayload, provider *models.BackendAiProvider) (string, error) {
useModel := strings.TrimSpace(p.Model)
if useModel == "" {
providerModels := parseProviderModels(provider.Models)
if len(providerModels) > 0 {
useModel = providerModels[0]
}
}
if useModel == "" {
return "", fmt.Errorf("未指定模型且该接入配置无可用模型")
}
return useModel, nil
}
// loadSystemPrompt 加载默认角色预设
func loadSystemPrompt(claims *jwtutil.Claims) 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 {
return presets[0].Content
}
return ""
}
// loadHistory 加载会话历史(最近20条,正序)
func loadHistory(sessionID uint64) []openaiMessage {
var history []models.BackendAiChatMessage
_, _ = models.Orm.QueryTable(new(models.BackendAiChatMessage)).
Filter("session_id", sessionID).
OrderBy("-id").
Limit(20).
All(&history)
out := make([]openaiMessage, 0, len(history)+1)
for i := len(history) - 1; i >= 0; i-- {
h := history[i]
imgs := []string{}
if strings.TrimSpace(h.Images) != "" {
_ = json.Unmarshal([]byte(h.Images), &imgs)
}
out = append(out, openaiMessage{Role: h.Role, Content: h.Content, Images: imgs})
}
return out
}
// ============ Send(非流式,含 MCP 工具调用) ============
// 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) == "" && len(p.Images) == 0 {
c.chatJsonErr(400, 400, "消息内容不能为空")
return
}
session, err := resolveOrCreateSession(c, claims, &p)
if err != nil {
c.chatJsonErr(400, 400, err.Error())
return
}
provider, err := resolveProvider(claims, &p, &session)
if err != nil {
c.chatJsonErr(400, 400, err.Error())
return
}
useModel, err := resolveModel(&p, &provider)
if err != nil {
c.chatJsonErr(400, 400, err.Error())
return
}
systemPrompt := loadSystemPrompt(claims)
messages := loadHistory(session.ID)
messages = append(messages, openaiMessage{Role: "user", Content: p.Content, Images: p.Images})
// 保存用户消息
userMsg := models.BackendAiChatMessage{
TenantID: fmt.Sprintf("%d", claims.TenantId),
SessionID: session.ID,
Role: "user",
Content: p.Content,
Images: marshalImages(p.Images),
CreateTime: time.Now(),
}
_, _ = models.Orm.Insert(&userMsg)
// 汇总启用的 MCP 工具(受前端开关与所选服务控制)
llmTools, refMap, _, _ := collectMcpTools(claims, p.mcpEnabled(), p.McpServerIDs)
current := messages
reply := ""
rounds := 0
startTime := time.Now()
var totalUsage aiUsage
for {
rounds++
var roundUsage aiUsage
text, pending, roundUsage, callErr := callAITools(provider, useModel, systemPrompt, current, llmTools)
if callErr != nil {
c.chatJsonErr(500, 500, "AI调用失败: "+callErr.Error())
return
}
totalUsage.add(roundUsage)
if len(pending) == 0 {
reply = text
break
}
if rounds >= maxToolRounds {
reply = text
break
}
current = append(current, buildMcpAssistantMessage(text, pending))
toolMsgs := executePendingTools(claims, pending, refMap, nil, nil)
current = append(current, toolMsgs...)
}
durationMs := int(time.Since(startTime).Milliseconds())
assistantMsg := models.BackendAiChatMessage{
TenantID: fmt.Sprintf("%d", claims.TenantId),
SessionID: session.ID,
Role: "assistant",
Content: reply,
Tokens: totalUsage.total(),
DurationMs: durationMs,
Model: useModel,
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,
"tokens": totalUsage.total(),
"duration_ms": durationMs,
"model": useModel,
})
}
// ============ SendStream(SSE 流式,含 MCP 工具调用) ============
// SendStream POST /backend/ai/chat/send-stream
func (c *BackendAiChatController) SendStream() {
// SSE 手动写入响应体,禁用框架自动渲染模板
c.EnableRender = false
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) == "" && len(p.Images) == 0 {
c.chatJsonErr(400, 400, "消息内容不能为空")
return
}
session, err := resolveOrCreateSession(c, claims, &p)
if err != nil {
c.chatJsonErr(400, 400, err.Error())
return
}
provider, err := resolveProvider(claims, &p, &session)
if err != nil {
c.chatJsonErr(400, 400, err.Error())
return
}
useModel, err := resolveModel(&p, &provider)
if err != nil {
c.chatJsonErr(400, 400, err.Error())
return
}
systemPrompt := loadSystemPrompt(claims)
messages := loadHistory(session.ID)
messages = append(messages, openaiMessage{Role: "user", Content: p.Content, Images: p.Images})
// 保存用户消息
userMsg := models.BackendAiChatMessage{
TenantID: fmt.Sprintf("%d", claims.TenantId),
SessionID: session.ID,
Role: "user",
Content: p.Content,
Images: marshalImages(p.Images),
CreateTime: time.Now(),
}
_, _ = models.Orm.Insert(&userMsg)
// 汇总启用的 MCP 工具(受前端开关与所选服务控制)
llmTools, refMap, toolSummaries, _ := collectMcpTools(claims, p.mcpEnabled(), p.McpServerIDs)
// 设置SSE响应头
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()
}
}
writeJSON := func(event string, v interface{}) {
b, _ := json.Marshal(v)
writeSSE(event, string(b))
}
// session 事件(带上用户消息ID,便于前端后续编辑/删除该条消息)
writeJSON("session", map[string]interface{}{
"session_id": session.ID,
"user_message_id": userMsg.ID,
})
// 当前启用的工具列表(前端展示)
if len(toolSummaries) > 0 {
writeJSON("tool_list", map[string]interface{}{"tools": toolSummaries})
}
current := messages
rounds := 0
fullReply := ""
streamErr := error(nil)
startTime := time.Now()
var totalUsage aiUsage
for {
rounds++
var text string
var pending []pendingToolCall
var roundUsage aiUsage
text, pending, roundUsage, streamErr = callAIStreamTools(provider, useModel, systemPrompt, current, llmTools, func(chunk string) {
writeJSON("content", map[string]string{"content": chunk})
})
if streamErr != nil {
break
}
totalUsage.add(roundUsage)
if len(pending) == 0 {
fullReply += text
break
}
if rounds >= maxToolRounds {
fullReply += text
break
}
current = append(current, buildMcpAssistantMessage(text, pending))
toolMsgs := executePendingTools(claims, pending, refMap,
func(pc pendingToolCall) {
ref, _ := refMap[pc.Name]
writeJSON("tool_start", map[string]interface{}{
"id": pc.ID,
"key": pc.Name,
"tool": ref.ToolName,
"server_id": ref.ServerID,
"server_name": ref.ServerName,
"args": pc.Args,
})
},
func(pc pendingToolCall, result string, ok bool, errMsg string) {
writeJSON("tool_result", map[string]interface{}{
"id": pc.ID,
"key": pc.Name,
"ok": ok,
"result": result,
"error": errMsg,
})
})
current = append(current, toolMsgs...)
}
if streamErr != nil {
writeJSON("error", map[string]string{"error": streamErr.Error()})
return
}
// 保存AI回复
durationMs := int(time.Since(startTime).Milliseconds())
assistantMsg := models.BackendAiChatMessage{
TenantID: fmt.Sprintf("%d", claims.TenantId),
SessionID: session.ID,
Role: "assistant",
Content: fullReply,
Tokens: totalUsage.total(),
DurationMs: durationMs,
Model: useModel,
CreateTime: time.Now(),
}
_, _ = models.Orm.Insert(&assistantMsg)
session.UpdateTime = time.Now()
_, _ = models.Orm.Update(&session, "update_time")
writeJSON("done", map[string]interface{}{
"session_id": session.ID,
"message_id": assistantMsg.ID,
"tokens": totalUsage.total(),
"duration_ms": durationMs,
"model": useModel,
})
}
// ============ AI 调用(非流式,含工具) ============
// callAI 兼容入口(无 MCP 工具),供智能生成等模块使用
func callAI(provider models.BackendAiProvider, model string, systemPrompt string, messages []openaiMessage) (string, error) {
text, pending, _, err := callAITools(provider, model, systemPrompt, messages, nil)
if err != nil {
return "", err
}
// 一般场景不包含工具调用;若模型仍返回工具调用(未注入工具不应发生),仅返回文本
_ = pending
return text, nil
}
func callAITools(provider models.BackendAiProvider, model string, systemPrompt string, messages []openaiMessage, tools []openaiTool) (string, []pendingToolCall, aiUsage, error) {
client := &http.Client{Timeout: 180 * time.Second}
if provider.ProviderType == "openai" {
reqMessages := messages
if systemPrompt != "" {
reqMessages = append([]openaiMessage{{Role: "system", Content: systemPrompt}}, messages...)
}
reqBody := map[string]interface{}{
"model": model,
"messages": openAIMessages(reqMessages),
}
if len(tools) > 0 {
reqBody["tools"] = tools
}
jsonData, _ := json.Marshal(reqBody)
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 "", nil, aiUsage{}, err
}
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Authorization", "Bearer "+provider.ApiKey)
resp, err := client.Do(req)
if err != nil {
return "", nil, aiUsage{}, err
}
defer resp.Body.Close()
bodyBytes, _ := io.ReadAll(resp.Body)
if resp.StatusCode != 200 {
return "", nil, aiUsage{}, fmt.Errorf("API返回错误状态 %d: %s", resp.StatusCode, string(bodyBytes))
}
var result struct {
Model string `json:"model"`
Choices []struct {
Message struct {
Content string `json:"content"`
ToolCalls []openaiToolCall `json:"tool_calls"`
} `json:"message"`
} `json:"choices"`
Usage *struct {
PromptTokens int `json:"prompt_tokens"`
CompletionTokens int `json:"completion_tokens"`
TotalTokens int `json:"total_tokens"`
} `json:"usage"`
Error *struct {
Message string `json:"message"`
} `json:"error"`
}
if err := json.Unmarshal(bodyBytes, &result); err != nil {
return "", nil, aiUsage{}, fmt.Errorf("解析响应失败: %s", string(bodyBytes))
}
if result.Error != nil {
return "", nil, aiUsage{}, fmt.Errorf(result.Error.Message)
}
if len(result.Choices) == 0 {
return "", nil, aiUsage{}, fmt.Errorf("AI未返回内容")
}
var usage aiUsage
if result.Usage != nil {
usage = aiUsage{
PromptTokens: result.Usage.PromptTokens,
CompletionTokens: result.Usage.CompletionTokens,
TotalTokens: result.Usage.TotalTokens,
}
}
// 接口实际返回的使用模型(网关可能将 auto 解析为具体模型)
if result.Model != "" {
usage.Model = result.Model
}
content := result.Choices[0].Message.Content
var pending []pendingToolCall
for _, tc := range result.Choices[0].Message.ToolCalls {
pending = append(pending, parseOpenAIToolCall(tc))
}
return content, pending, usage, nil
}
// Anthropic 非流式
reqBody := map[string]interface{}{
"model": model,
"max_tokens": 4096,
"messages": toAnthropicMessages(messages),
}
if systemPrompt != "" {
reqBody["system"] = systemPrompt
}
if len(tools) > 0 {
reqBody["tools"] = toAnthropicTools(tools)
}
jsonData, _ := json.Marshal(reqBody)
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 "", nil, aiUsage{}, 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 "", nil, aiUsage{}, err
}
defer resp.Body.Close()
bodyBytes, _ := io.ReadAll(resp.Body)
if resp.StatusCode != 200 {
return "", nil, aiUsage{}, fmt.Errorf("API返回错误状态 %d: %s", resp.StatusCode, string(bodyBytes))
}
var result struct {
Content []struct {
Type string `json:"type"`
Text string `json:"text"`
ID string `json:"id"`
Name string `json:"name"`
Input map[string]interface{} `json:"input"`
} `json:"content"`
Usage *struct {
InputTokens int `json:"input_tokens"`
OutputTokens int `json:"output_tokens"`
} `json:"usage"`
Error *struct {
Message string `json:"message"`
} `json:"error"`
}
if err := json.Unmarshal(bodyBytes, &result); err != nil {
return "", nil, aiUsage{}, fmt.Errorf("解析响应失败: %s", string(bodyBytes))
}
if result.Error != nil {
return "", nil, aiUsage{}, fmt.Errorf(result.Error.Message)
}
var usage aiUsage
if result.Usage != nil {
usage = aiUsage{
PromptTokens: result.Usage.InputTokens,
CompletionTokens: result.Usage.OutputTokens,
TotalTokens: result.Usage.InputTokens + result.Usage.OutputTokens,
}
}
// Anthropic 响应不回传模型名,记录请求使用的模型
usage.Model = model
var text strings.Builder
var pending []pendingToolCall
for _, block := range result.Content {
if block.Type == "text" {
text.WriteString(block.Text)
}
if block.Type == "tool_use" {
pending = append(pending, pendingToolCall{ID: block.ID, Name: block.Name, Args: block.Input})
}
}
return text.String(), pending, usage, nil
}
// parseOpenAIToolCall 解析非流式 OpenAI 工具调用
func parseOpenAIToolCall(tc openaiToolCall) pendingToolCall {
pc := pendingToolCall{ID: tc.ID, Name: tc.Function.Name}
if tc.Function.Arguments != "" {
var m map[string]interface{}
if err := json.Unmarshal([]byte(tc.Function.Arguments), &m); err == nil {
pc.Args = m
}
}
if pc.Args == nil {
pc.Args = map[string]interface{}{}
}
return pc
}
// ============ AI 调用(流式,含工具) ============
func callAIStreamTools(provider models.BackendAiProvider, model string, systemPrompt string, messages []openaiMessage, tools []openaiTool, onChunk func(string)) (string, []pendingToolCall, aiUsage, error) {
if provider.ProviderType == "openai" {
return callOpenAIStreamTools(provider, model, systemPrompt, messages, tools, onChunk)
}
return callAnthropicStreamTools(provider, model, systemPrompt, messages, tools, onChunk)
}
// callOpenAIStreamTools OpenAI 兼容流式工具调用
func callOpenAIStreamTools(provider models.BackendAiProvider, model string, systemPrompt string, messages []openaiMessage, tools []openaiTool, onChunk func(string)) (string, []pendingToolCall, aiUsage, error) {
client := &http.Client{Timeout: 180 * time.Second}
reqMessages := messages
if systemPrompt != "" {
reqMessages = append([]openaiMessage{{Role: "system", Content: systemPrompt}}, messages...)
}
reqBody := map[string]interface{}{
"model": model,
"messages": openAIMessages(reqMessages),
"stream": true,
}
if len(tools) > 0 {
reqBody["tools"] = tools
}
// 请求在最后一个 chunk 返回 token 用量统计
reqBody["stream_options"] = map[string]interface{}{"include_usage": true}
jsonData, _ := json.Marshal(reqBody)
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 "", nil, aiUsage{}, err
}
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Authorization", "Bearer "+provider.ApiKey)
resp, err := client.Do(req)
if err != nil {
return "", nil, aiUsage{}, err
}
defer resp.Body.Close()
if resp.StatusCode != 200 {
body, _ := io.ReadAll(resp.Body)
return "", nil, aiUsage{}, fmt.Errorf("API返回错误状态 %d: %s", resp.StatusCode, string(body))
}
var text strings.Builder
var usage aiUsage
acc := make(map[int]*struct {
ID string
Name string
Args string
})
var order []int
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 {
Model string `json:"model"`
Choices []struct {
Delta struct {
Content string `json:"content"`
ToolCalls []openaiDeltaToolCall `json:"tool_calls"`
} `json:"delta"`
FinishReason string `json:"finish_reason"`
} `json:"choices"`
Usage *struct {
PromptTokens int `json:"prompt_tokens"`
CompletionTokens int `json:"completion_tokens"`
TotalTokens int `json:"total_tokens"`
} `json:"usage"`
Error *struct {
Message string `json:"message"`
} `json:"error"`
}
if err := json.Unmarshal([]byte(data), &chunk); err != nil {
continue
}
if chunk.Error != nil {
return "", nil, aiUsage{}, fmt.Errorf(chunk.Error.Message)
}
if chunk.Usage != nil {
usage.PromptTokens = chunk.Usage.PromptTokens
usage.CompletionTokens = chunk.Usage.CompletionTokens
usage.TotalTokens = chunk.Usage.TotalTokens
}
if chunk.Model != "" {
usage.Model = chunk.Model
}
if len(chunk.Choices) == 0 {
continue
}
choice := chunk.Choices[0]
if choice.Delta.Content != "" {
text.WriteString(choice.Delta.Content)
if onChunk != nil {
onChunk(choice.Delta.Content)
}
}
for _, tc := range choice.Delta.ToolCalls {
if _, ok := acc[tc.Index]; !ok {
acc[tc.Index] = &struct {
ID string
Name string
Args string
}{}
order = append(order, tc.Index)
}
cur := acc[tc.Index]
if tc.ID != "" {
cur.ID = tc.ID
}
if tc.Function.Name != "" {
cur.Name += tc.Function.Name
}
if tc.Function.Arguments != "" {
cur.Args += tc.Function.Arguments
}
}
}
var pending []pendingToolCall
for _, idx := range order {
cur := acc[idx]
id := cur.ID
if id == "" {
id = fmt.Sprintf("call_%d_%d", time.Now().UnixNano(), idx)
}
pc := pendingToolCall{ID: id, Name: cur.Name}
if cur.Args != "" {
var m map[string]interface{}
if err := json.Unmarshal([]byte(cur.Args), &m); err == nil {
pc.Args = m
}
}
if pc.Args == nil {
pc.Args = map[string]interface{}{}
}
pending = append(pending, pc)
}
return text.String(), pending, usage, nil
}
// callAnthropicStreamTools Anthropic 流式工具调用
func callAnthropicStreamTools(provider models.BackendAiProvider, model string, systemPrompt string, messages []openaiMessage, tools []openaiTool, onChunk func(string)) (string, []pendingToolCall, aiUsage, error) {
client := &http.Client{Timeout: 180 * time.Second}
reqBody := map[string]interface{}{
"model": model,
"max_tokens": 4096,
"messages": toAnthropicMessages(messages),
"stream": true,
}
if systemPrompt != "" {
reqBody["system"] = systemPrompt
}
if len(tools) > 0 {
reqBody["tools"] = toAnthropicTools(tools)
}
jsonData, _ := json.Marshal(reqBody)
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 "", nil, aiUsage{}, 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 "", nil, aiUsage{}, err
}
defer resp.Body.Close()
if resp.StatusCode != 200 {
body, _ := io.ReadAll(resp.Body)
return "", nil, aiUsage{}, fmt.Errorf("API返回错误状态 %d: %s", resp.StatusCode, string(body))
}
var text strings.Builder
var usage aiUsage
type toolUseAcc struct {
ID string
Name string
Input strings.Builder
}
acc := make(map[int]*toolUseAcc)
var order []int
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: ") {
continue
}
data := strings.TrimPrefix(line, "data: ")
switch currentEvent {
case "content_block_start":
var block struct {
Index int `json:"index"`
ContentBlock struct {
Type string `json:"type"`
ID string `json:"id"`
Name string `json:"name"`
} `json:"content_block"`
Usage *struct {
InputTokens int `json:"input_tokens"`
} `json:"usage"`
}
if err := json.Unmarshal([]byte(data), &block); err != nil {
continue
}
if block.Usage != nil {
usage.PromptTokens = block.Usage.InputTokens
}
if block.ContentBlock.Type == "tool_use" {
if _, ok := acc[block.Index]; !ok {
acc[block.Index] = &toolUseAcc{}
order = append(order, block.Index)
}
acc[block.Index].ID = block.ContentBlock.ID
acc[block.Index].Name = block.ContentBlock.Name
}
case "content_block_delta":
var delta struct {
Index int `json:"index"`
Delta struct {
Type string `json:"type"`
Text string `json:"text"`
PartialJSON string `json:"partial_json"`
} `json:"delta"`
}
if err := json.Unmarshal([]byte(data), &delta); err != nil {
continue
}
if delta.Delta.Type == "text_delta" && delta.Delta.Text != "" {
text.WriteString(delta.Delta.Text)
if onChunk != nil {
onChunk(delta.Delta.Text)
}
}
if delta.Delta.Type == "input_json_delta" && delta.Delta.PartialJSON != "" {
if _, ok := acc[delta.Index]; !ok {
acc[delta.Index] = &toolUseAcc{}
order = append(order, delta.Index)
}
acc[delta.Index].Input.WriteString(delta.Delta.PartialJSON)
}
case "message_delta":
var md struct {
Delta struct {
StopReason string `json:"stop_reason"`
} `json:"delta"`
Usage *struct {
OutputTokens int `json:"output_tokens"`
} `json:"usage"`
}
if err := json.Unmarshal([]byte(data), &md); err == nil {
_ = md.Delta.StopReason
if md.Usage != nil {
usage.CompletionTokens = md.Usage.OutputTokens
}
}
case "message_stop":
// 结束
}
}
var pending []pendingToolCall
for _, idx := range order {
cur := acc[idx]
id := cur.ID
if id == "" {
id = fmt.Sprintf("toolu_%d_%d", time.Now().UnixNano(), idx)
}
pc := pendingToolCall{ID: id, Name: cur.Name}
inputStr := cur.Input.String()
if inputStr != "" {
var m map[string]interface{}
if err := json.Unmarshal([]byte(inputStr), &m); err == nil {
pc.Args = m
}
}
if pc.Args == nil {
pc.Args = map[string]interface{}{}
}
pending = append(pending, pc)
}
usage.TotalTokens = usage.PromptTokens + usage.CompletionTokens
return text.String(), pending, usage, nil
}
// ============ Anthropic 消息/工具格式转换 ============
func toAnthropicTools(tools []openaiTool) []map[string]interface{} {
out := make([]map[string]interface{}, 0, len(tools))
for _, t := range tools {
out = append(out, map[string]interface{}{
"name": t.Function.Name,
"description": t.Function.Description,
"input_schema": t.Function.Parameters,
})
}
return out
}
// toAnthropicMessages 将内部 openaiMessage 列表转为 Anthropic 消息格式
func toAnthropicMessages(messages []openaiMessage) []map[string]interface{} {
out := make([]map[string]interface{}, 0, len(messages))
for _, m := range messages {
switch m.Role {
case "assistant":
if len(m.ToolCalls) > 0 {
blocks := make([]map[string]interface{}, 0, len(m.ToolCalls)+1)
if m.Content != "" {
blocks = append(blocks, map[string]interface{}{"type": "text", "text": m.Content})
}
for _, tc := range m.ToolCalls {
var input map[string]interface{}
if tc.Function.Arguments != "" {
_ = json.Unmarshal([]byte(tc.Function.Arguments), &input)
}
if input == nil {
input = map[string]interface{}{}
}
blocks = append(blocks, map[string]interface{}{
"type": "tool_use",
"id": tc.ID,
"name": tc.Function.Name,
"input": input,
})
}
out = append(out, map[string]interface{}{"role": "assistant", "content": blocks})
} else {
out = append(out, map[string]interface{}{"role": "assistant", "content": m.Content})
}
case "tool":
out = append(out, map[string]interface{}{
"role": "user",
"content": []map[string]interface{}{
{"type": "tool_result", "tool_use_id": m.ToolCallID, "content": m.Content},
},
})
default:
out = append(out, map[string]interface{}{"role": m.Role, "content": m.Content})
}
}
return out
}