Files
xk-ai-agent/internal/llm/openai.go
2026-08-15 17:05:22 +08:00

358 lines
9.6 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"
"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), &params)
// 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
}