更新ai广场bug
This commit is contained in:
@@ -54,10 +54,20 @@ func (c *BackendAiChatController) chatOk(data interface{}) {
|
||||
}
|
||||
|
||||
type chatSendPayload struct {
|
||||
SessionID uint64 `json:"session_id"`
|
||||
Content string `json:"content"`
|
||||
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"`
|
||||
}
|
||||
|
||||
// mcpEnabled 解析本轮是否启用 MCP(未传默认开启)
|
||||
func (p *chatSendPayload) mcpEnabled() bool {
|
||||
return p.McpEnabled == nil || *p.McpEnabled
|
||||
}
|
||||
|
||||
// ============ 通用消息/工具结构 ============
|
||||
@@ -124,6 +134,28 @@ type mcpToolSummary struct {
|
||||
|
||||
const maxToolRounds = 6
|
||||
|
||||
// aiUsage 一次 AI 调用的 token 用量
|
||||
type aiUsage struct {
|
||||
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
|
||||
}
|
||||
|
||||
// sanitizeToolName 工具名规范化为 OpenAI 允许的字符集
|
||||
func sanitizeToolName(s string) string {
|
||||
var sb strings.Builder
|
||||
@@ -150,16 +182,23 @@ func mcpToolKey(serverID uint64, name string) string {
|
||||
return key
|
||||
}
|
||||
|
||||
// collectMcpTools 汇总当前用户「已启用」的 MCP 服务器工具,转换为 LLM 工具列表
|
||||
func collectMcpTools(claims *jwtutil.Claims) ([]openaiTool, map[string]toolRef, []mcpToolSummary, error) {
|
||||
// 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
|
||||
_, err := models.Orm.QueryTable(new(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).
|
||||
OrderBy("id").
|
||||
All(&servers)
|
||||
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
|
||||
}
|
||||
@@ -265,12 +304,13 @@ func buildMcpAssistantMessage(text string, pending []pendingToolCall) openaiMess
|
||||
// runToolLoop 运行带已启用 MCP 工具的完整对话循环(非流式),返回最终文本。
|
||||
// 供聊天非流式接口与智能生成等模块复用。
|
||||
func runToolLoop(claims *jwtutil.Claims, provider models.BackendAiProvider, model string, systemPrompt string, messages []openaiMessage) (string, error) {
|
||||
llmTools, refMap, _, _ := collectMcpTools(claims)
|
||||
// 智能生成等内部调用:默认注入全部已启用的 MCP 工具
|
||||
llmTools, refMap, _, _ := collectMcpTools(claims, true, nil)
|
||||
current := messages
|
||||
rounds := 0
|
||||
for {
|
||||
rounds++
|
||||
text, pending, err := callAITools(provider, model, systemPrompt, current, llmTools)
|
||||
text, pending, _, err := callAITools(provider, model, systemPrompt, current, llmTools)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
@@ -542,19 +582,23 @@ func (c *BackendAiChatController) Send() {
|
||||
}
|
||||
_, _ = models.Orm.Insert(&userMsg)
|
||||
|
||||
// 汇总启用的 MCP 工具
|
||||
llmTools, refMap, _, _ := collectMcpTools(claims)
|
||||
// 汇总启用的 MCP 工具(受前端开关与所选服务控制)
|
||||
llmTools, refMap, _, _ := collectMcpTools(claims, p.mcpEnabled(), p.McpServerIDs)
|
||||
|
||||
current := messages
|
||||
reply := ""
|
||||
rounds := 0
|
||||
startTime := time.Now()
|
||||
var totalUsage aiUsage
|
||||
for {
|
||||
rounds++
|
||||
text, pending, callErr := callAITools(provider, useModel, systemPrompt, current, llmTools)
|
||||
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
|
||||
@@ -568,12 +612,16 @@ func (c *BackendAiChatController) Send() {
|
||||
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,
|
||||
CreateTime: time.Now(),
|
||||
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)
|
||||
|
||||
@@ -581,9 +629,12 @@ func (c *BackendAiChatController) Send() {
|
||||
_, _ = models.Orm.Update(&session, "update_time")
|
||||
|
||||
c.chatOk(map[string]interface{}{
|
||||
"session_id": session.ID,
|
||||
"reply": reply,
|
||||
"message_id": assistantMsg.ID,
|
||||
"session_id": session.ID,
|
||||
"reply": reply,
|
||||
"message_id": assistantMsg.ID,
|
||||
"tokens": totalUsage.total(),
|
||||
"duration_ms": durationMs,
|
||||
"model": useModel,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -645,8 +696,8 @@ func (c *BackendAiChatController) SendStream() {
|
||||
}
|
||||
_, _ = models.Orm.Insert(&userMsg)
|
||||
|
||||
// 汇总启用的 MCP 工具
|
||||
llmTools, refMap, toolSummaries, _ := collectMcpTools(claims)
|
||||
// 汇总启用的 MCP 工具(受前端开关与所选服务控制)
|
||||
llmTools, refMap, toolSummaries, _ := collectMcpTools(claims, p.mcpEnabled(), p.McpServerIDs)
|
||||
|
||||
// 设置SSE响应头
|
||||
rw := c.Ctx.ResponseWriter.ResponseWriter
|
||||
@@ -672,8 +723,11 @@ func (c *BackendAiChatController) SendStream() {
|
||||
writeSSE(event, string(b))
|
||||
}
|
||||
|
||||
// session 事件
|
||||
writeJSON("session", map[string]interface{}{"session_id": session.ID})
|
||||
// session 事件(带上用户消息ID,便于前端后续编辑/删除该条消息)
|
||||
writeJSON("session", map[string]interface{}{
|
||||
"session_id": session.ID,
|
||||
"user_message_id": userMsg.ID,
|
||||
})
|
||||
|
||||
// 当前启用的工具列表(前端展示)
|
||||
if len(toolSummaries) > 0 {
|
||||
@@ -684,17 +738,21 @@ func (c *BackendAiChatController) SendStream() {
|
||||
rounds := 0
|
||||
fullReply := ""
|
||||
streamErr := error(nil)
|
||||
startTime := time.Now()
|
||||
var totalUsage aiUsage
|
||||
|
||||
for {
|
||||
rounds++
|
||||
var text string
|
||||
var pending []pendingToolCall
|
||||
text, pending, streamErr = callAIStreamTools(provider, useModel, systemPrompt, current, llmTools, func(chunk string) {
|
||||
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
|
||||
@@ -736,12 +794,16 @@ func (c *BackendAiChatController) SendStream() {
|
||||
}
|
||||
|
||||
// 保存AI回复
|
||||
durationMs := int(time.Since(startTime).Milliseconds())
|
||||
assistantMsg := models.BackendAiChatMessage{
|
||||
TenantID: fmt.Sprintf("%d", claims.TenantId),
|
||||
SessionID: session.ID,
|
||||
Role: "assistant",
|
||||
Content: fullReply,
|
||||
CreateTime: time.Now(),
|
||||
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)
|
||||
|
||||
@@ -749,8 +811,11 @@ func (c *BackendAiChatController) SendStream() {
|
||||
_, _ = models.Orm.Update(&session, "update_time")
|
||||
|
||||
writeJSON("done", map[string]interface{}{
|
||||
"session_id": session.ID,
|
||||
"message_id": assistantMsg.ID,
|
||||
"session_id": session.ID,
|
||||
"message_id": assistantMsg.ID,
|
||||
"tokens": totalUsage.total(),
|
||||
"duration_ms": durationMs,
|
||||
"model": useModel,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -758,7 +823,7 @@ func (c *BackendAiChatController) SendStream() {
|
||||
|
||||
// callAI 兼容入口(无 MCP 工具),供智能生成等模块使用
|
||||
func callAI(provider models.BackendAiProvider, model string, systemPrompt string, messages []openaiMessage) (string, error) {
|
||||
text, pending, err := callAITools(provider, model, systemPrompt, messages, nil)
|
||||
text, pending, _, err := callAITools(provider, model, systemPrompt, messages, nil)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
@@ -767,7 +832,7 @@ func callAI(provider models.BackendAiProvider, model string, systemPrompt string
|
||||
return text, nil
|
||||
}
|
||||
|
||||
func callAITools(provider models.BackendAiProvider, model string, systemPrompt string, messages []openaiMessage, tools []openaiTool) (string, []pendingToolCall, error) {
|
||||
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" {
|
||||
@@ -790,20 +855,20 @@ func callAITools(provider models.BackendAiProvider, model string, systemPrompt s
|
||||
}
|
||||
req, err := http.NewRequest("POST", url, bytes.NewBuffer(jsonData))
|
||||
if err != nil {
|
||||
return "", nil, err
|
||||
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, err
|
||||
return "", nil, aiUsage{}, err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
bodyBytes, _ := io.ReadAll(resp.Body)
|
||||
if resp.StatusCode != 200 {
|
||||
return "", nil, fmt.Errorf("API返回错误状态 %d: %s", resp.StatusCode, string(bodyBytes))
|
||||
return "", nil, aiUsage{}, fmt.Errorf("API返回错误状态 %d: %s", resp.StatusCode, string(bodyBytes))
|
||||
}
|
||||
var result struct {
|
||||
Choices []struct {
|
||||
@@ -812,25 +877,38 @@ func callAITools(provider models.BackendAiProvider, model string, systemPrompt s
|
||||
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, fmt.Errorf("解析响应失败: %s", string(bodyBytes))
|
||||
return "", nil, aiUsage{}, fmt.Errorf("解析响应失败: %s", string(bodyBytes))
|
||||
}
|
||||
if result.Error != nil {
|
||||
return "", nil, fmt.Errorf(result.Error.Message)
|
||||
return "", nil, aiUsage{}, fmt.Errorf(result.Error.Message)
|
||||
}
|
||||
if len(result.Choices) == 0 {
|
||||
return "", nil, fmt.Errorf("AI未返回内容")
|
||||
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,
|
||||
}
|
||||
}
|
||||
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, nil
|
||||
return content, pending, usage, nil
|
||||
}
|
||||
|
||||
// Anthropic 非流式
|
||||
@@ -853,7 +931,7 @@ func callAITools(provider models.BackendAiProvider, model string, systemPrompt s
|
||||
}
|
||||
req, err := http.NewRequest("POST", url, bytes.NewBuffer(jsonData))
|
||||
if err != nil {
|
||||
return "", nil, err
|
||||
return "", nil, aiUsage{}, err
|
||||
}
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("x-api-key", provider.ApiKey)
|
||||
@@ -861,13 +939,13 @@ func callAITools(provider models.BackendAiProvider, model string, systemPrompt s
|
||||
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
return "", nil, err
|
||||
return "", nil, aiUsage{}, err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
bodyBytes, _ := io.ReadAll(resp.Body)
|
||||
if resp.StatusCode != 200 {
|
||||
return "", nil, fmt.Errorf("API返回错误状态 %d: %s", resp.StatusCode, string(bodyBytes))
|
||||
return "", nil, aiUsage{}, fmt.Errorf("API返回错误状态 %d: %s", resp.StatusCode, string(bodyBytes))
|
||||
}
|
||||
var result struct {
|
||||
Content []struct {
|
||||
@@ -877,15 +955,27 @@ func callAITools(provider models.BackendAiProvider, model string, systemPrompt s
|
||||
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, fmt.Errorf("解析响应失败: %s", string(bodyBytes))
|
||||
return "", nil, aiUsage{}, fmt.Errorf("解析响应失败: %s", string(bodyBytes))
|
||||
}
|
||||
if result.Error != nil {
|
||||
return "", nil, fmt.Errorf(result.Error.Message)
|
||||
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,
|
||||
}
|
||||
}
|
||||
var text strings.Builder
|
||||
var pending []pendingToolCall
|
||||
@@ -897,7 +987,7 @@ func callAITools(provider models.BackendAiProvider, model string, systemPrompt s
|
||||
pending = append(pending, pendingToolCall{ID: block.ID, Name: block.Name, Args: block.Input})
|
||||
}
|
||||
}
|
||||
return text.String(), pending, nil
|
||||
return text.String(), pending, usage, nil
|
||||
}
|
||||
|
||||
// parseOpenAIToolCall 解析非流式 OpenAI 工具调用
|
||||
@@ -917,7 +1007,7 @@ func parseOpenAIToolCall(tc openaiToolCall) pendingToolCall {
|
||||
|
||||
// ============ AI 调用(流式,含工具) ============
|
||||
|
||||
func callAIStreamTools(provider models.BackendAiProvider, model string, systemPrompt string, messages []openaiMessage, tools []openaiTool, onChunk func(string)) (string, []pendingToolCall, error) {
|
||||
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)
|
||||
}
|
||||
@@ -925,7 +1015,7 @@ func callAIStreamTools(provider models.BackendAiProvider, model string, systemPr
|
||||
}
|
||||
|
||||
// callOpenAIStreamTools OpenAI 兼容流式工具调用
|
||||
func callOpenAIStreamTools(provider models.BackendAiProvider, model string, systemPrompt string, messages []openaiMessage, tools []openaiTool, onChunk func(string)) (string, []pendingToolCall, error) {
|
||||
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
|
||||
@@ -940,6 +1030,8 @@ func callOpenAIStreamTools(provider models.BackendAiProvider, model string, syst
|
||||
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, "/")
|
||||
@@ -948,23 +1040,24 @@ func callOpenAIStreamTools(provider models.BackendAiProvider, model string, syst
|
||||
}
|
||||
req, err := http.NewRequest("POST", url, bytes.NewBuffer(jsonData))
|
||||
if err != nil {
|
||||
return "", nil, err
|
||||
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, err
|
||||
return "", nil, aiUsage{}, err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != 200 {
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
return "", nil, fmt.Errorf("API返回错误状态 %d: %s", resp.StatusCode, string(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
|
||||
@@ -991,6 +1084,11 @@ func callOpenAIStreamTools(provider models.BackendAiProvider, model string, syst
|
||||
} `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"`
|
||||
@@ -999,7 +1097,14 @@ func callOpenAIStreamTools(provider models.BackendAiProvider, model string, syst
|
||||
continue
|
||||
}
|
||||
if chunk.Error != nil {
|
||||
return "", nil, fmt.Errorf(chunk.Error.Message)
|
||||
return "", nil, aiUsage{}, fmt.Errorf(chunk.Error.Message)
|
||||
}
|
||||
if chunk.Usage != nil {
|
||||
usage = aiUsage{
|
||||
PromptTokens: chunk.Usage.PromptTokens,
|
||||
CompletionTokens: chunk.Usage.CompletionTokens,
|
||||
TotalTokens: chunk.Usage.TotalTokens,
|
||||
}
|
||||
}
|
||||
if len(chunk.Choices) == 0 {
|
||||
continue
|
||||
@@ -1052,11 +1157,11 @@ func callOpenAIStreamTools(provider models.BackendAiProvider, model string, syst
|
||||
}
|
||||
pending = append(pending, pc)
|
||||
}
|
||||
return text.String(), pending, nil
|
||||
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, error) {
|
||||
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{}{
|
||||
@@ -1079,7 +1184,7 @@ func callAnthropicStreamTools(provider models.BackendAiProvider, model string, s
|
||||
}
|
||||
req, err := http.NewRequest("POST", url, bytes.NewBuffer(jsonData))
|
||||
if err != nil {
|
||||
return "", nil, err
|
||||
return "", nil, aiUsage{}, err
|
||||
}
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("x-api-key", provider.ApiKey)
|
||||
@@ -1087,16 +1192,17 @@ func callAnthropicStreamTools(provider models.BackendAiProvider, model string, s
|
||||
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
return "", nil, err
|
||||
return "", nil, aiUsage{}, err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != 200 {
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
return "", nil, fmt.Errorf("API返回错误状态 %d: %s", resp.StatusCode, string(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
|
||||
@@ -1127,10 +1233,16 @@ func callAnthropicStreamTools(provider models.BackendAiProvider, model string, s
|
||||
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{}
|
||||
@@ -1169,9 +1281,15 @@ func callAnthropicStreamTools(provider models.BackendAiProvider, model string, s
|
||||
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":
|
||||
// 结束
|
||||
@@ -1198,7 +1316,8 @@ func callAnthropicStreamTools(provider models.BackendAiProvider, model string, s
|
||||
}
|
||||
pending = append(pending, pc)
|
||||
}
|
||||
return text.String(), pending, nil
|
||||
usage.TotalTokens = usage.PromptTokens + usage.CompletionTokens
|
||||
return text.String(), pending, usage, nil
|
||||
}
|
||||
|
||||
// ============ Anthropic 消息/工具格式转换 ============
|
||||
|
||||
Reference in New Issue
Block a user