468 lines
17 KiB
Go
468 lines
17 KiB
Go
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), ¶ms); 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
|
||
}
|