358 lines
9.6 KiB
Go
358 lines
9.6 KiB
Go
package llm
|
||
|
||
import (
|
||
"bytes"
|
||
"context"
|
||
"encoding/json"
|
||
"fmt"
|
||
"io"
|
||
"log"
|
||
"net/http"
|
||
"time"
|
||
|
||
"tcm-agent/internal/config"
|
||
"tcm-agent/internal/types"
|
||
)
|
||
|
||
// ========================================================================
|
||
// OpenAI 客户端
|
||
// ========================================================================
|
||
// 支持模型:
|
||
// - gpt-4o / gpt-4o-mini(综合能力最强,适合处方校验)
|
||
// - gpt-4-turbo(长上下文)
|
||
// - o1-preview / o1-mini(深度推理)
|
||
// - text-embedding-3-large / small(向量化)
|
||
//
|
||
// 协议兼容性:
|
||
// OpenAI 的 API 协议已成为行业事实标准,很多国产模型
|
||
// (DeepSeek/通义千问/Moonshot)都兼容此协议。
|
||
// 因此这个客户端稍作修改(换 BaseURL)就能对接很多服务。
|
||
// ========================================================================
|
||
|
||
// OpenAIClient OpenAI 模型客户端
|
||
type OpenAIClient struct {
|
||
apiKey string
|
||
baseURL string
|
||
model string
|
||
embeddingModel string
|
||
client *http.Client
|
||
capabilities map[string]bool
|
||
}
|
||
|
||
// createOpenAIClient 工厂方法:创建 OpenAI 客户端
|
||
func createOpenAIClient(cfg *config.LLMConfigEx) (LLMClient, error) {
|
||
baseURL := cfg.BaseURL
|
||
if baseURL == "" {
|
||
baseURL = "https://api.openai.com/v1"
|
||
}
|
||
model := cfg.Model
|
||
if model == "" {
|
||
model = "gpt-4o" // 默认用 4o
|
||
}
|
||
|
||
// 判断是否为推理模型(o1 系列不支持 temperature 等参数)
|
||
isReasoningModel := len(model) >= 2 && model[:2] == "o1"
|
||
|
||
client := &OpenAIClient{
|
||
apiKey: cfg.APIKey,
|
||
baseURL: baseURL,
|
||
model: model,
|
||
embeddingModel: cfg.EmbeddingModel,
|
||
client: &http.Client{Timeout: time.Duration(cfg.Timeout) * time.Second},
|
||
capabilities: map[string]bool{
|
||
CapFunctionCalling: true,
|
||
CapStreaming: true,
|
||
CapJSONMode: true,
|
||
CapLongContext: model == "gpt-4o" || model == "gpt-4-turbo",
|
||
CapEmbedding: cfg.EmbeddingModel != "",
|
||
"reasoning": isReasoningModel,
|
||
},
|
||
}
|
||
|
||
log.Printf("[OpenAI] 初始化完成 | 模型: %s | 地址: %s | 推理模型: %v",
|
||
model, baseURL, isReasoningModel)
|
||
return client, nil
|
||
}
|
||
|
||
// Name 返回模型名称
|
||
func (c *OpenAIClient) Name() string { return c.model }
|
||
|
||
// Provider 返回供应商名称
|
||
func (c *OpenAIClient) Provider() string { return "openai" }
|
||
|
||
// Supports 查询能力
|
||
func (c *OpenAIClient) Supports(capability string) bool {
|
||
return c.capabilities[capability]
|
||
}
|
||
|
||
// Chat 发起对话请求
|
||
func (c *OpenAIClient) Chat(ctx context.Context, messages []types.Message, tools []types.Tool) (*types.Message, error) {
|
||
// 转换消息(工具帧协议与 Spark/DeepSeek 客户端保持一致):
|
||
// - tool 结果帧带 tool_call_id 与 assistant 帧 tool_calls[].id 对应
|
||
// - assistant 工具帧回放 tool_calls(id + arguments JSON 字符串)
|
||
openAIMsgs := make([]map[string]any, 0, len(messages))
|
||
for _, m := range messages {
|
||
msg := map[string]any{
|
||
"role": m.Role,
|
||
"content": m.Content,
|
||
}
|
||
if m.Role == "tool" && m.ToolCallID != "" {
|
||
msg["tool_call_id"] = m.ToolCallID
|
||
}
|
||
if m.ToolCall != nil {
|
||
args, _ := json.Marshal(m.ToolCall.Params)
|
||
id := m.ToolCall.ID
|
||
if id == "" {
|
||
id = fmt.Sprintf("call_%d", time.Now().UnixNano())
|
||
}
|
||
msg["tool_calls"] = []map[string]any{
|
||
{
|
||
"id": id,
|
||
"type": "function",
|
||
"function": map[string]any{
|
||
"name": m.ToolCall.ToolName,
|
||
"arguments": string(args),
|
||
},
|
||
},
|
||
}
|
||
}
|
||
openAIMsgs = append(openAIMsgs, msg)
|
||
}
|
||
|
||
// 构建请求体
|
||
body := map[string]any{
|
||
"model": c.model,
|
||
"messages": openAIMsgs,
|
||
}
|
||
|
||
// o1 系列模型不支持 temperature 和 max_tokens 参数
|
||
if !c.capabilities["reasoning"] {
|
||
body["temperature"] = 0.3
|
||
body["max_tokens"] = 4096
|
||
} else {
|
||
// o1 使用 max_completion_tokens
|
||
body["max_completion_tokens"] = 4096
|
||
}
|
||
|
||
// 附加工具
|
||
if len(tools) > 0 && !c.capabilities["reasoning"] {
|
||
// o1 系列目前不支持 function calling
|
||
toolDefs := make([]map[string]any, 0, len(tools))
|
||
for _, t := range tools {
|
||
toolDefs = append(toolDefs, 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"},
|
||
},
|
||
},
|
||
})
|
||
}
|
||
body["tools"] = toolDefs
|
||
body["tool_choice"] = "auto"
|
||
}
|
||
|
||
// 发送请求
|
||
buf, _ := json.Marshal(body)
|
||
url := ResolveChatCompletionsURL(c.baseURL)
|
||
req, _ := http.NewRequestWithContext(ctx, "POST", url, bytes.NewReader(buf))
|
||
req.Header.Set("Authorization", "Bearer "+c.apiKey)
|
||
req.Header.Set("Content-Type", "application/json")
|
||
|
||
resp, err := c.client.Do(req)
|
||
if err != nil {
|
||
return nil, fmt.Errorf("[OpenAI] 请求失败: %w", err)
|
||
}
|
||
defer resp.Body.Close()
|
||
|
||
data, _ := io.ReadAll(resp.Body)
|
||
if resp.StatusCode != 200 {
|
||
return nil, fmt.Errorf("[OpenAI] API 返回 %d: %s", resp.StatusCode, string(data[:min(len(data), 500)]))
|
||
}
|
||
|
||
// 解析响应
|
||
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:"function"`
|
||
} `json:"tool_calls"`
|
||
} `json:"message"`
|
||
} `json:"choices"`
|
||
Error *struct {
|
||
Message string `json:"message"`
|
||
} `json:"error"`
|
||
}
|
||
|
||
if err := json.Unmarshal(data, &result); err != nil {
|
||
return nil, fmt.Errorf("[OpenAI] 解析响应失败: %w", err)
|
||
}
|
||
|
||
if result.Error != nil {
|
||
return nil, fmt.Errorf("[OpenAI] API 错误: %s", result.Error.Message)
|
||
}
|
||
|
||
if len(result.Choices) == 0 {
|
||
return nil, fmt.Errorf("[OpenAI] 返回空响应")
|
||
}
|
||
|
||
msg := &types.Message{
|
||
Role: "assistant",
|
||
Content: result.Choices[0].Message.Content,
|
||
Timestamp: time.Now().Unix(),
|
||
}
|
||
|
||
// 工具调用
|
||
if len(result.Choices[0].Message.ToolCalls) > 0 {
|
||
tc := result.Choices[0].Message.ToolCalls[0]
|
||
params := make(map[string]any)
|
||
json.Unmarshal([]byte(tc.Function.Arguments), ¶ms)
|
||
// 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,
|
||
}
|
||
log.Printf("[OpenAI] 模型决定调用工具: %s", tc.Function.Name)
|
||
}
|
||
|
||
return msg, nil
|
||
}
|
||
|
||
// StreamChat 流式对话
|
||
func (c *OpenAIClient) StreamChat(ctx context.Context, messages []types.Message, tools []types.Tool) (<-chan string, error) {
|
||
// 实现与 DeepSeek 类似,设置 stream: true
|
||
// 为节省篇幅,这里用简化的实现
|
||
ch := make(chan string, 10)
|
||
|
||
openAIMsgs := make([]map[string]any, 0, len(messages))
|
||
for _, m := range messages {
|
||
openAIMsgs = append(openAIMsgs, map[string]any{
|
||
"role": m.Role, "content": m.Content,
|
||
})
|
||
}
|
||
|
||
body := map[string]any{
|
||
"model": c.model,
|
||
"messages": openAIMsgs,
|
||
"stream": true,
|
||
}
|
||
|
||
buf, _ := json.Marshal(body)
|
||
url := ResolveChatCompletionsURL(c.baseURL)
|
||
req, _ := http.NewRequestWithContext(ctx, "POST", url, bytes.NewReader(buf))
|
||
req.Header.Set("Authorization", "Bearer "+c.apiKey)
|
||
req.Header.Set("Content-Type", "application/json")
|
||
req.Header.Set("Accept", "text/event-stream")
|
||
|
||
resp, err := c.client.Do(req)
|
||
if err != nil {
|
||
close(ch)
|
||
return nil, fmt.Errorf("[OpenAI] 流式请求失败: %w", err)
|
||
}
|
||
|
||
go func() {
|
||
defer resp.Body.Close()
|
||
defer close(ch)
|
||
buf := make([]byte, 4096)
|
||
for {
|
||
n, err := resp.Body.Read(buf)
|
||
if n > 0 {
|
||
ch <- string(buf[:n])
|
||
}
|
||
if err != nil {
|
||
break
|
||
}
|
||
}
|
||
}()
|
||
|
||
return ch, nil
|
||
}
|
||
|
||
// Embed 生成文本向量
|
||
//
|
||
// OpenAI 的 text-embedding-3 系列是目前质量最高的 Embedding 模型之一。
|
||
// 推荐:
|
||
// - text-embedding-3-large(3072维,精度最高)
|
||
// - text-embedding-3-small(1536维,性价比高)
|
||
func (c *OpenAIClient) Embed(ctx context.Context, texts []string) ([][]float32, error) {
|
||
if c.embeddingModel == "" {
|
||
return nil, fmt.Errorf("[OpenAI] 未配置 Embedding 模型")
|
||
}
|
||
|
||
// OpenAI Embedding API 一次最多传 2048 条文本
|
||
results := make([][]float32, 0, len(texts))
|
||
|
||
// 分批处理
|
||
batchSize := 100
|
||
for i := 0; i < len(texts); i += batchSize {
|
||
end := i + batchSize
|
||
if end > len(texts) {
|
||
end = len(texts)
|
||
}
|
||
batch := texts[i:end]
|
||
|
||
body := map[string]any{
|
||
"model": c.embeddingModel,
|
||
"input": batch,
|
||
}
|
||
|
||
buf, _ := json.Marshal(body)
|
||
url := c.baseURL + "/embeddings"
|
||
req, _ := http.NewRequestWithContext(ctx, "POST", url, bytes.NewReader(buf))
|
||
req.Header.Set("Authorization", "Bearer "+c.apiKey)
|
||
req.Header.Set("Content-Type", "application/json")
|
||
|
||
resp, err := c.client.Do(req)
|
||
if err != nil {
|
||
return nil, fmt.Errorf("[OpenAI] Embedding 请求失败: %w", err)
|
||
}
|
||
|
||
data, _ := io.ReadAll(resp.Body)
|
||
resp.Body.Close()
|
||
|
||
if resp.StatusCode != 200 {
|
||
return nil, fmt.Errorf("[OpenAI] Embedding 返回 %d: %s", resp.StatusCode, string(data[:min(len(data), 300)]))
|
||
}
|
||
|
||
var result struct {
|
||
Data []struct {
|
||
Embedding []float32 `json:"embedding"`
|
||
} `json:"data"`
|
||
}
|
||
json.Unmarshal(data, &result)
|
||
|
||
for _, d := range result.Data {
|
||
results = append(results, d.Embedding)
|
||
}
|
||
}
|
||
|
||
log.Printf("[OpenAI] Embedding 完成 | 模型: %s | 文本数: %d", c.embeddingModel, len(texts))
|
||
return results, nil
|
||
}
|
||
|
||
// Close 释放资源
|
||
func (c *OpenAIClient) Close() error {
|
||
c.client.CloseIdleConnections()
|
||
return nil
|
||
}
|