Files
xk-ai-agent/internal/llm/spark.go
2026-08-14 21:50:48 +08:00

468 lines
17 KiB
Go
Raw Permalink 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 llm
import (
"bytes"
"context"
"encoding/json"
"fmt"
"io"
"log"
"net/http"
"strings"
"time"
"tcm-agent/internal/agentcfg"
"tcm-agent/internal/config"
"tcm-agent/internal/types"
)
// ========================================================================
// 讯飞星火 (Xunfei Spark) 客户端
// ========================================================================
// 对接方式:
// - 讯飞星火当前提供两套对外协议:
// 1) 老协议:基于 APIKey + APISecret 做 JWT 签名,WebSocket 长连接(v1/v2/v3/v4)
// 2) 新协议 OpenAPI:HTTPS POST /v1/chat/completions,Authorization: Bearer APIPassword
// - PHP 端 SparkAiAgent 已经走的是新协议 OpenAPI(与 OpenAI 协议一致),
// 为了保持多端行为一致、降低维护成本,Go 端也走 OpenAPI 协议。
//
// 适用场景:
// - 国内合规部署(讯飞数据不出境)
// - 中文医疗场景(讯飞有医学大模型经验:spark-medicine)
//
// 推荐模型:
// - spark-lite 轻量,免费额度大,适合分类/简单问答
// - spark-pro 中等档位,性价比高
// - spark-max 最强档位,适合复杂医学推理
// - spark-medicine 医学专用模型(需开通对应授权)
// ========================================================================
// SparkClient 讯飞星火客户端(OpenAPI 兼容协议)
type SparkClient struct {
apiKey string // Bearer 鉴权用的 APIPassword
baseURL string // OpenAPI 完整地址(含 /v1/chat/completions)
model string // 默认模型名(spark-pro / spark-max 等)
client *http.Client // 复用的 HTTP 连接
capabilities map[string]bool // 能力声明
lastResult *types.ChatResult // 最近一次 Chat 的 token/finish_reason 快照
}
// createSparkClient 工厂方法
//
// 参数:
// cfg - 从 config.yaml / DB 读到的供应商配置
//
// 默认值约定(与 PHP 端 AiRuntimeConfigService 一致):
// - BaseURL 未配置 → 回落到讯飞官方域名(与 .env 中 SPARK_API_URL 同源)
// - Model 未配置 → spark-lite(轻量档位,便于联调)
func createSparkClient(cfg *config.LLMConfigEx) (LLMClient, error) {
baseURL := cfg.BaseURL
if baseURL == "" {
// 讯飞星火 OpenAPI 默认域名(与 PHP 端 .env 默认值保持一致)
baseURL = "https://spark-api-open.xf-yun.com/v1/chat/completions"
}
model := cfg.Model
if model == "" {
model = "lite" // 默认轻量档位
}
// 按版本检测 function calling 能力(文档说仅 Pro/Max/Ultra 支持)
// capability map 在客户端初始化时就标好,运行时 Supports() 查这个 map
supportsFC := supportsSparkFunctionCalling(model)
client := &SparkClient{
apiKey: cfg.APIKey,
baseURL: baseURL,
model: model,
client: &http.Client{Timeout: time.Duration(cfg.Timeout) * time.Second},
capabilities: map[string]bool{
// FunctionCalling:按 model 版本动态判断(lite/general* 不支持,pro/max/ultra 支持)
CapFunctionCalling: supportsFC,
CapStreaming: true,
CapJSONMode: true,
// spark-max 支持 8K~32K,spark-pro 一般 8K,统一标 false,按需打开
CapLongContext: false,
// 讯飞 Embedding 需要走独立接口,本客户端暂不实现
CapEmbedding: false,
},
}
log.Printf("[Spark] 初始化完成 | 模型: %s | 地址: %s | function_calling=%v",
model, baseURL, supportsFC)
return client, nil
}
// supportsSparkFunctionCalling 独立函数版(供 createSparkClient 在构造 capabilities 前调用)
//
// 与 (*SparkClient).supportsFunctionCalling 是同一份白名单,
// 拆成独立函数是为了在客户端尚未构造完成时也能调用
func supportsSparkFunctionCalling(model string) bool {
m := strings.ToLower(strings.TrimSpace(model))
switch m {
case "generalv3", "pro-128k",
"generalv3.5", "max-32k",
"4.0ultra":
return true
}
return false
}
func (c *SparkClient) Name() string { return c.model }
func (c *SparkClient) Provider() string { return "spark" }
func (c *SparkClient) Supports(cap string) bool { return c.capabilities[cap] }
// supportsFunctionCalling 检测当前 model 是否支持 OpenAI 风格 function calling
//
// 讯飞 OpenAPI 各版本 model 名(来自官方文档):
// - lite → Lite 版(不支持 function calling)
// - general → Spark V1.5(不支持)
// - generalv2 → Spark V2(不支持)
// - generalv3 → Pro 版(支持)
// - pro-128k → Pro-128K(支持)
// - generalv3.5 → Max 版(支持)
// - max-32k → Max-32K(支持)
// - 4.0Ultra → 4.0 Ultra(支持)
//
// 判断规则:model 名命中以下白名单才算支持
// 用白名单而非黑名单,是为了未来新增 model 时默认保守(不允许),
// 避免新版本突然能传 tools 时反而不被识别
func (c *SparkClient) supportsFunctionCalling() bool {
m := strings.ToLower(strings.TrimSpace(c.model))
switch m {
case "generalv3", "pro-128k", // Pro 系列
"generalv3.5", "max-32k", // Max 系列
"4.0ultra": // 4.0 Ultra(已经 ToLower)
return true
}
return false
}
// Chat 发起一次非流式对话(与 OpenAI 协议一致)
//
// 本方法使用客户端默认参数(temperature=0.3, max_tokens=4096);
// 需要运行时覆盖参数请用 ChatWithOpts(实现 OptAwareClient)。
//
// 返回的 Message 可能包含 ToolCall(触发工具调用),由上层 Runner 决定下一步。
func (c *SparkClient) Chat(ctx context.Context, messages []types.Message, tools []types.Tool) (*types.Message, error) {
return c.ChatWithOpts(ctx, messages, tools, ChatOpts{})
}
// ChatWithOpts 带运行时参数的 Chat(实现 OptAwareClient 接口)
func (c *SparkClient) ChatWithOpts(ctx context.Context, messages []types.Message, tools []types.Tool, opts ChatOpts) (*types.Message, error) {
// 组装请求体(与 DeepSeek/OpenAI 一致)
temperature := 0.3
if opts.Temperature > 0 {
temperature = opts.Temperature
}
maxTokens := 4096
if opts.MaxTokens > 0 {
maxTokens = opts.MaxTokens
}
reqBody := map[string]any{
"model": c.model,
"messages": c.convertMessages(messages),
// 默认关闭流式,业务侧解析整包 JSON 更简单(与 PHP 端默认一致)
"stream": false,
"temperature": temperature,
"max_tokens": maxTokens,
}
// 函数调用(仅 Pro/Max/Ultra 支持,Lite 不支持)
//
// 讯飞 OpenAPI 文档明确说明:
// - Lite 版不支持 function calling
// - Pro/Max/4.0 Ultra 支持 function calling
// 若给 Lite 版传 tools 字段,讯飞会返回 10003 "用户的消息格式有错误"
// 这里按 model 名自动降级,避免低版本模型直接报错
if len(tools) > 0 {
if !c.supportsFunctionCalling() {
log.Printf("[Spark] ⚠️ 当前 model=%s 不支持 function calling,已自动降级为纯文本(忽略 %d 个工具)",
c.model, len(tools))
} else {
reqBody["tools"] = c.buildToolDefs(tools)
reqBody["tool_choice"] = "auto"
}
}
buf, _ := json.Marshal(reqBody)
req, err := http.NewRequestWithContext(ctx, "POST", c.baseURL, bytes.NewReader(buf))
if err != nil {
return nil, fmt.Errorf("[Spark] 构造请求失败: %w", err)
}
// 讯飞 OpenAPI 用 Bearer + APIPassword(与 PHP 端 Authorization 头一致)
req.Header.Set("Authorization", "Bearer "+c.apiKey)
req.Header.Set("Content-Type", "application/json")
// ===== 调试日志:请求体打印受 ai_agent_debug_log 开关控制 =====
// 请求体包含完整患者病历(PHI 隐私数据),生产环境默认只打不含内容的摘要行;
// 排查厂商 API 错误(如 10003 消息格式错误)时,把 xk_system_config 的
// ai_agent_debug_log 置 1 即可看到完整请求体(60s 内生效,无需重启)
if agentcfg.Get().Debug.LogRequestBody {
reqSnippet := string(buf)
if len(reqSnippet) > 1500 {
reqSnippet = reqSnippet[:1500] + "...(截断)"
}
log.Printf("[Spark] → POST %s | model=%s | msgs=%d | tools=%d | body=%s",
c.baseURL, c.model, len(messages), len(tools), reqSnippet)
} else {
// 摘要行:不含消息内容,只有规模信息,便于观察调用频率与请求大小
log.Printf("[Spark] → POST | model=%s | msgs=%d | tools=%d | bytes=%d",
c.model, len(messages), len(tools), len(buf))
}
startedAt := time.Now()
resp, err := c.client.Do(req)
if err != nil {
return nil, fmt.Errorf("[Spark] 请求失败(请检查网络或密钥): %w", err)
}
defer resp.Body.Close()
data, _ := io.ReadAll(resp.Body)
if resp.StatusCode != 200 {
// 截取前 1000 字符,包含完整错误信息(讯飞的 10003 错误响应通常 < 300 字符)
snippet := string(data)
if len(snippet) > 1000 {
snippet = snippet[:1000]
}
log.Printf("[Spark] ← HTTP %d | resp=%s", resp.StatusCode, snippet)
return nil, fmt.Errorf("[Spark] API 返回 %d: %s", resp.StatusCode, snippet)
}
// 解析响应并提取 usage / finish_reason
msg, chatResult := c.parseResponseWithMeta(data)
if chatResult != nil {
chatResult.Provider = "spark"
chatResult.Model = c.model
chatResult.DurationMs = int(time.Since(startedAt).Milliseconds())
c.lastResult = chatResult
}
// 记录耗时(仅日志,token 计费统计在 KnowledgeEnhancer 层做)
log.Printf("[Spark] 调用耗时 %v | finish_reason=%s tokens=%d/%d",
time.Since(startedAt), chatResult.FinishReason,
chatResult.PromptTokens, chatResult.CompletionTokens)
return msg, nil
}
// LastChatResult 实现 TokenAwareClient 接口
func (c *SparkClient) LastChatResult() *types.ChatResult {
return c.lastResult
}
// StreamChat 流式对话(暂未实现)
//
// 讯飞星火支持 SSE 流式输出,等业务侧确实需要"边生成边展示"时再补:
// - 设置 stream=true
// - 用 bufio.Scanner 逐行解析 data: {...} 块
func (c *SparkClient) StreamChat(ctx context.Context, messages []types.Message, tools []types.Tool) (<-chan string, error) {
ch := make(chan string, 10)
close(ch)
return ch, fmt.Errorf("[Spark] 流式输出暂未实现")
}
// Embed 文本向量化(讯飞 Embedding 走独立接口,本客户端暂不实现)
//
// 若后续需要把处方/病历向量化做相似度匹配,可在此实现:
// - 调用讯飞独立 embedding 接口
// - 或回落到本地 Ollama embedding(避免重复请求)
func (c *SparkClient) Embed(ctx context.Context, texts []string) ([][]float32, error) {
return nil, fmt.Errorf("[Spark] Embedding 暂未实现,请使用 Ollama/Qwen 的 Embed 能力")
}
// Close 释放资源(关闭底层连接池)
func (c *SparkClient) Close() error {
c.client.CloseIdleConnections()
return nil
}
// ========================================================================
// 内部辅助方法
// ========================================================================
// convertMessages 把 types.Message 转成讯飞 OpenAPI 接受的格式
//
// 讯飞接口与 OpenAI 协议一致,role + content 即可。
// 工具调用相关帧必须严格遵守 OpenAI Function Calling 协议:
// - assistant 帧的 tool_calls 必须带 id,arguments 必须是 JSON 字符串(不是对象)
// - tool 结果帧必须带 tool_call_id 与上面的 id 对应
// 否则讯飞会报 10003 消息格式错误(曾因 arguments 传对象踩过坑)
func (c *SparkClient) convertMessages(messages []types.Message) []map[string]any {
out := make([]map[string]any, 0, len(messages))
for _, m := range messages {
item := map[string]any{
"role": m.Role,
"content": m.Content,
}
// tool 结果帧:带上与 assistant 帧 tool_calls[].id 对应的 tool_call_id
if m.Role == "tool" && m.ToolCallID != "" {
item["tool_call_id"] = m.ToolCallID
}
// assistant 工具调用帧回放
if m.ToolCall != nil {
// arguments 必须是 JSON 字符串(OpenAI 协议),不能直接传 map 对象
args, _ := json.Marshal(m.ToolCall.Params)
// id 优先用厂商返回的原始 id;缺失时生成一个(保证协议完整性)
id := m.ToolCall.ID
if id == "" {
id = fmt.Sprintf("call_%d", time.Now().UnixNano())
}
item["tool_calls"] = []map[string]any{
{
"id": id,
"type": "function",
"function": map[string]any{
"name": m.ToolCall.ToolName,
"arguments": string(args),
},
},
}
}
out = append(out, item)
}
return out
}
// buildToolDefs 构造 tool 列表(OpenAI 风格的 function 定义)
func (c *SparkClient) buildToolDefs(tools []types.Tool) []map[string]any {
defs := make([]map[string]any, 0, len(tools))
for _, t := range tools {
defs = append(defs, map[string]any{
"type": "function",
"function": map[string]any{
"name": t.Name(),
"description": t.Description(),
"parameters": map[string]any{
"type": "object",
"properties": map[string]any{
"query": map[string]any{
"type": "string",
"description": "检索/查询关键词",
},
},
"required": []string{"query"},
},
},
})
}
return defs
}
// parseResponse 解析讯飞返回的 JSON(保留旧 API)
func (c *SparkClient) parseResponse(data []byte) *types.Message {
msg, _ := c.parseResponseWithMeta(data)
return msg
}
// parseResponseWithMeta 解析响应并附带 token/finish_reason 等元数据
//
// 返回:
// - *types.Message:assistant 消息(可能含 tool_calls)
// - *types.ChatResult:token 用量 + finish_reason
func (c *SparkClient) parseResponseWithMeta(data []byte) (*types.Message, *types.ChatResult) {
var result struct {
Choices []struct {
Message struct {
Role string `json:"role"`
Content string `json:"content"`
ToolCalls []struct {
ID string `json:"id"` // 厂商生成的调用 id(回放时须原样带回)
Function struct {
Name string `json:"name"`
Arguments string `json:"arguments"` // JSON 字符串
} `json:"function"`
} `json:"tool_calls"`
} `json:"message"`
FinishReason string `json:"finish_reason"` // stop/length/content_filter/tool_calls
} `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"`
}
chatResult := &types.ChatResult{}
if err := json.Unmarshal(data, &result); err != nil {
// JSON 解析失败也要返回一个 Message,避免上层 nil panic
return &types.Message{
Role: "assistant",
Content: fmt.Sprintf("[Spark] 响应解析失败: %v | 原始: %s", err, string(data[:min(len(data), 200)])),
Timestamp: time.Now().Unix(),
}, chatResult
}
// 错误响应(如额度耗尽 / 模型未授权)
if result.Error.Message != "" {
return &types.Message{
Role: "assistant",
Content: fmt.Sprintf("[Spark] 调用失败: %s", result.Error.Message),
Timestamp: time.Now().Unix(),
}, chatResult
}
if len(result.Choices) == 0 {
return &types.Message{
Role: "assistant",
Content: "[Spark] 返回空响应",
Timestamp: time.Now().Unix(),
}, chatResult
}
// 解析 finish_reason
chatResult.FinishReason = result.Choices[0].FinishReason
// 解析 token 用量
chatResult.PromptTokens = result.Usage.PromptTokens
chatResult.CompletionTokens = result.Usage.CompletionTokens
chatResult.TotalTokens = result.Usage.TotalTokens
choice := result.Choices[0].Message
msg := &types.Message{
Role: "assistant",
Content: choice.Content,
Timestamp: time.Now().Unix(),
}
// 若触发工具调用,构造 ToolCallInfo
if len(choice.ToolCalls) > 0 {
tc := choice.ToolCalls[0]
// arguments 是 JSON 字符串,解析成 map,解析失败则原样塞进 query
params := map[string]any{}
if tc.Function.Arguments != "" {
if err := json.Unmarshal([]byte(tc.Function.Arguments), &params); err != nil {
params["query"] = tc.Function.Arguments
}
}
// id 厂商未返回时生成一个,保证回放帧协议完整(tool_call_id 有对应目标)
id := tc.ID
if id == "" {
id = fmt.Sprintf("call_%d", time.Now().UnixNano())
}
msg.ToolCall = &types.ToolCallInfo{
ID: id,
ToolName: tc.Function.Name,
Params: params,
}
chatResult.ToolCall = msg.ToolCall
}
// finish_reason 回落:讯飞非流式成功响应本身不带 finish_reason 字段
// (官方文档的成功响应示例里没有该字段),为空不是解析错误。
// 这里按语义回落:有工具调用 → tool_calls,否则 → stop,避免日志误导
if chatResult.FinishReason == "" {
if msg.ToolCall != nil {
chatResult.FinishReason = "tool_calls"
} else {
chatResult.FinishReason = "stop"
}
}
return msg, chatResult
}