- 新增 OpenAIProvider / ClaudeProvider,接入 Deepseek 真实 API - 修复 Deepseek thinking 模式 reasoning_content 回传问题 - 修复 Thinking 阶段过渡指令每轮重复插入导致的死循环 - 新增 maxTurns 保护、dumpMessages/dumpTools 调试输出
193 lines
7.0 KiB
Go
193 lines
7.0 KiB
Go
// internal/provider/openai.go
|
||
package provider
|
||
|
||
import (
|
||
"context"
|
||
"encoding/json"
|
||
"fmt"
|
||
// "os"
|
||
|
||
"github.com/openai/openai-go/v3"
|
||
"github.com/openai/openai-go/v3/option"
|
||
"github.com/openai/openai-go/v3/packages/param"
|
||
"github.com/openai/openai-go/v3/shared"
|
||
"go-tiny-claw/internal/schema"
|
||
)
|
||
|
||
type OpenAIProvider struct {
|
||
client openai.Client // 值类型,非指针
|
||
model string
|
||
}
|
||
|
||
// NewZhipuOpenAIProvider 构造函数:基于 OpenAI V3 SDK,指向智谱底座
|
||
func DeepseekOpenAIProvider(model string) *OpenAIProvider {
|
||
apiKey := "sk-1f44696abe644bd684f09cc43f12c557"
|
||
if apiKey == "" {
|
||
panic("请设置 ZHIPU_API_KEY 环境变量")
|
||
}
|
||
// 核心:将官方 SDK 的地址替换为智谱的兼容端点
|
||
baseURL := "https://api.deepseek.com"
|
||
|
||
return &OpenAIProvider{
|
||
client: openai.NewClient(option.WithAPIKey(apiKey), option.WithBaseURL(baseURL)),
|
||
model: model,
|
||
}
|
||
}
|
||
|
||
func (p *OpenAIProvider) Generate(ctx context.Context, msgs []schema.Message, availableTools []schema.ToolDefinition) (*schema.Message, error) {
|
||
var openaiMsgs []openai.ChatCompletionMessageParamUnion
|
||
|
||
// 1. 翻译上下文消息
|
||
for _, msg := range msgs {
|
||
switch msg.Role {
|
||
case schema.RoleSystem:
|
||
openaiMsgs = append(openaiMsgs, openai.SystemMessage(msg.Content))
|
||
|
||
case schema.RoleUser:
|
||
if msg.ToolCallID != "" {
|
||
// 注意:v3 新版参数顺序是 (content, toolCallID)
|
||
openaiMsgs = append(openaiMsgs, openai.ToolMessage(msg.Content, msg.ToolCallID))
|
||
} else {
|
||
openaiMsgs = append(openaiMsgs, openai.UserMessage(msg.Content))
|
||
}
|
||
|
||
case schema.RoleAssistant:
|
||
// Deepseek thinking mode: reasoning_content 必须回传
|
||
if msg.ReasoningContent != "" || len(msg.ToolCalls) > 0 {
|
||
msgMap := map[string]interface{}{
|
||
"role": "assistant",
|
||
"content": msg.Content,
|
||
}
|
||
if msg.ReasoningContent != "" {
|
||
msgMap["reasoning_content"] = msg.ReasoningContent
|
||
}
|
||
if len(msg.ToolCalls) > 0 {
|
||
var toolCalls []map[string]interface{}
|
||
for _, tc := range msg.ToolCalls {
|
||
toolCalls = append(toolCalls, map[string]interface{}{
|
||
"id": tc.ID,
|
||
"type": "function",
|
||
"function": map[string]interface{}{
|
||
"name": tc.Name,
|
||
"arguments": string(tc.Arguments),
|
||
},
|
||
})
|
||
}
|
||
msgMap["tool_calls"] = toolCalls
|
||
}
|
||
rawJSON, _ := json.Marshal(msgMap)
|
||
astParam := param.Override[openai.ChatCompletionAssistantMessageParam](json.RawMessage(rawJSON))
|
||
openaiMsgs = append(openaiMsgs, openai.ChatCompletionMessageParamUnion{
|
||
OfAssistant: &astParam,
|
||
})
|
||
break
|
||
}
|
||
|
||
astParam := openai.ChatCompletionAssistantMessageParam{}
|
||
|
||
if msg.Content != "" {
|
||
astParam.Content = openai.ChatCompletionAssistantMessageParamContentUnion{
|
||
OfString: openai.String(msg.Content),
|
||
}
|
||
}
|
||
|
||
// 【重要】如果历史包含 ToolCalls,必须原样放回,以维系大模型的逻辑链
|
||
if len(msg.ToolCalls) > 0 {
|
||
var toolCalls []openai.ChatCompletionMessageToolCallUnionParam
|
||
for _, tc := range msg.ToolCalls {
|
||
// OfFunction 对应 GetFunction(),字段类型严格要求为指针
|
||
toolCalls = append(toolCalls, openai.ChatCompletionMessageToolCallUnionParam{
|
||
OfFunction: &openai.ChatCompletionMessageFunctionToolCallParam{
|
||
ID: tc.ID,
|
||
Type: "function",
|
||
Function: openai.ChatCompletionMessageFunctionToolCallFunctionParam{
|
||
Name: tc.Name,
|
||
Arguments: string(tc.Arguments),
|
||
},
|
||
},
|
||
})
|
||
}
|
||
astParam.ToolCalls = toolCalls
|
||
}
|
||
|
||
openaiMsgs = append(openaiMsgs, openai.ChatCompletionMessageParamUnion{
|
||
OfAssistant: &astParam,
|
||
})
|
||
}
|
||
}
|
||
|
||
// 2. 翻译工具定义 (v3 新 API 特性适配)
|
||
var openaiTools []openai.ChatCompletionToolUnionParam
|
||
for _, toolDef := range availableTools {
|
||
var params shared.FunctionParameters
|
||
|
||
// 尝试直接断言,如果不成功则通过 JSON 往返序列化来保证类型匹配
|
||
if m, ok := toolDef.InputSchema.(map[string]interface{}); ok {
|
||
params = shared.FunctionParameters(m)
|
||
} else {
|
||
// fallback:JSON 往返序列化
|
||
b, _ := json.Marshal(toolDef.InputSchema)
|
||
_ = json.Unmarshal(b, ¶ms)
|
||
}
|
||
|
||
openaiTools = append(openaiTools, openai.ChatCompletionFunctionTool(
|
||
shared.FunctionDefinitionParam{
|
||
Name: toolDef.Name,
|
||
Description: openai.String(toolDef.Description),
|
||
Parameters: params,
|
||
},
|
||
))
|
||
}
|
||
|
||
// 3. 构建请求并发送
|
||
params := openai.ChatCompletionNewParams{
|
||
Model: p.model,
|
||
Messages: openaiMsgs,
|
||
}
|
||
|
||
// 【慢思考机制支撑】仅当 availableTools 存在时才挂载 Tools
|
||
if len(openaiTools) > 0 {
|
||
params.Tools = openaiTools
|
||
}
|
||
|
||
resp, err := p.client.Chat.Completions.New(ctx, params)
|
||
if err != nil {
|
||
return nil, fmt.Errorf("OpenAI/Zhipu API 请求失败: %w", err)
|
||
}
|
||
if len(resp.Choices) == 0 {
|
||
return nil, fmt.Errorf("API 返回了空的 Choices")
|
||
}
|
||
|
||
// 4. 将 API Response 反向翻译为内部 schema.Message
|
||
choice := resp.Choices[0].Message
|
||
resultMsg := &schema.Message{
|
||
Role: schema.RoleAssistant,
|
||
Content: choice.Content,
|
||
}
|
||
|
||
// Deepseek thinking mode: 从原始响应中提取 reasoning_content
|
||
rawMsg := choice.RawJSON()
|
||
if rawMsg != "" {
|
||
var rawMap map[string]json.RawMessage
|
||
if err := json.Unmarshal([]byte(rawMsg), &rawMap); err == nil {
|
||
if rc, ok := rawMap["reasoning_content"]; ok {
|
||
var s string
|
||
if json.Unmarshal(rc, &s) == nil {
|
||
resultMsg.ReasoningContent = s
|
||
}
|
||
}
|
||
}
|
||
}
|
||
|
||
for _, tc := range choice.ToolCalls {
|
||
if tc.Type == "function" {
|
||
resultMsg.ToolCalls = append(resultMsg.ToolCalls, schema.ToolCall{
|
||
ID: tc.ID,
|
||
Name: tc.Function.Name,
|
||
Arguments: []byte(tc.Function.Arguments), // 提取 JSON 字符串字节
|
||
})
|
||
}
|
||
}
|
||
|
||
return resultMsg, nil
|
||
} |