fix: improve migration logic and reduce code duplication
- Preserve user's configured model during config migration (issue #5) - Simplify ExtractProtocol using strings.Cut - Extract NormalizeToolCall to shared utility, removing ~70 lines of duplicate code - Clean up unused fields in providerMigrationConfig struct Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
parent
09a0d19119
commit
ec86b21d3f
6 changed files with 600 additions and 285 deletions
|
|
@ -607,7 +607,7 @@ func (al *AgentLoop) runLLMIteration(ctx context.Context, messages []providers.M
|
||||||
|
|
||||||
normalizedToolCalls := make([]providers.ToolCall, 0, len(response.ToolCalls))
|
normalizedToolCalls := make([]providers.ToolCall, 0, len(response.ToolCalls))
|
||||||
for _, tc := range response.ToolCalls {
|
for _, tc := range response.ToolCalls {
|
||||||
normalizedToolCalls = append(normalizedToolCalls, normalizeProviderToolCall(tc))
|
normalizedToolCalls = append(normalizedToolCalls, providers.NormalizeToolCall(tc))
|
||||||
}
|
}
|
||||||
|
|
||||||
// Log tool calls
|
// Log tool calls
|
||||||
|
|
@ -715,45 +715,6 @@ func (al *AgentLoop) runLLMIteration(ctx context.Context, messages []providers.M
|
||||||
return finalContent, iteration, nil
|
return finalContent, iteration, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func normalizeProviderToolCall(tc providers.ToolCall) providers.ToolCall {
|
|
||||||
normalized := tc
|
|
||||||
|
|
||||||
if normalized.Name == "" && normalized.Function != nil {
|
|
||||||
normalized.Name = normalized.Function.Name
|
|
||||||
}
|
|
||||||
|
|
||||||
if normalized.Arguments == nil {
|
|
||||||
normalized.Arguments = map[string]interface{}{}
|
|
||||||
}
|
|
||||||
|
|
||||||
if len(normalized.Arguments) == 0 && normalized.Function != nil && normalized.Function.Arguments != "" {
|
|
||||||
var parsed map[string]interface{}
|
|
||||||
if err := json.Unmarshal([]byte(normalized.Function.Arguments), &parsed); err == nil && parsed != nil {
|
|
||||||
normalized.Arguments = parsed
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
argsJSON, _ := json.Marshal(normalized.Arguments)
|
|
||||||
if normalized.Function == nil {
|
|
||||||
normalized.Function = &providers.FunctionCall{
|
|
||||||
Name: normalized.Name,
|
|
||||||
Arguments: string(argsJSON),
|
|
||||||
}
|
|
||||||
} else {
|
|
||||||
if normalized.Function.Name == "" {
|
|
||||||
normalized.Function.Name = normalized.Name
|
|
||||||
}
|
|
||||||
if normalized.Name == "" {
|
|
||||||
normalized.Name = normalized.Function.Name
|
|
||||||
}
|
|
||||||
if normalized.Function.Arguments == "" {
|
|
||||||
normalized.Function.Arguments = string(argsJSON)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
return normalized
|
|
||||||
}
|
|
||||||
|
|
||||||
// updateToolContexts updates the context for tools that need channel/chatID info.
|
// updateToolContexts updates the context for tools that need channel/chatID info.
|
||||||
func (al *AgentLoop) updateToolContexts(channel, chatID string) {
|
func (al *AgentLoop) updateToolContexts(channel, chatID string) {
|
||||||
// Use ContextualTool interface instead of type assertions
|
// Use ContextualTool interface instead of type assertions
|
||||||
|
|
|
||||||
|
|
@ -5,201 +5,326 @@
|
||||||
|
|
||||||
package config
|
package config
|
||||||
|
|
||||||
|
import (
|
||||||
|
"slices"
|
||||||
|
"strings"
|
||||||
|
)
|
||||||
|
|
||||||
|
// providerMigrationConfig defines how to migrate a provider from old config to new format.
|
||||||
|
type providerMigrationConfig struct {
|
||||||
|
// providerNames are the possible names used in agents.defaults.provider
|
||||||
|
providerNames []string
|
||||||
|
// protocol is the protocol prefix for the model field
|
||||||
|
protocol string
|
||||||
|
// buildConfig creates the ModelConfig from ProviderConfig
|
||||||
|
buildConfig func(p ProvidersConfig) (ModelConfig, bool)
|
||||||
|
}
|
||||||
|
|
||||||
// ConvertProvidersToModelList converts the old ProvidersConfig to a slice of ModelConfig.
|
// ConvertProvidersToModelList converts the old ProvidersConfig to a slice of ModelConfig.
|
||||||
// This enables backward compatibility with existing configurations.
|
// This enables backward compatibility with existing configurations.
|
||||||
|
// It preserves the user's configured model from agents.defaults.model when possible.
|
||||||
func ConvertProvidersToModelList(cfg *Config) []ModelConfig {
|
func ConvertProvidersToModelList(cfg *Config) []ModelConfig {
|
||||||
if cfg == nil {
|
if cfg == nil {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Get user's configured provider and model
|
||||||
|
userProvider := strings.ToLower(cfg.Agents.Defaults.Provider)
|
||||||
|
userModel := cfg.Agents.Defaults.Model
|
||||||
|
|
||||||
var result []ModelConfig
|
var result []ModelConfig
|
||||||
p := cfg.Providers
|
p := cfg.Providers
|
||||||
|
|
||||||
// OpenAI
|
// Define migration rules for each provider
|
||||||
if p.OpenAI.APIKey != "" || p.OpenAI.APIBase != "" {
|
migrations := []providerMigrationConfig{
|
||||||
result = append(result, ModelConfig{
|
{
|
||||||
|
providerNames: []string{"openai", "gpt"},
|
||||||
|
protocol: "openai",
|
||||||
|
buildConfig: func(p ProvidersConfig) (ModelConfig, bool) {
|
||||||
|
if p.OpenAI.APIKey == "" && p.OpenAI.APIBase == "" {
|
||||||
|
return ModelConfig{}, false
|
||||||
|
}
|
||||||
|
return ModelConfig{
|
||||||
ModelName: "openai",
|
ModelName: "openai",
|
||||||
Model: "openai/gpt-4o",
|
Model: "openai/gpt-4o",
|
||||||
APIKey: p.OpenAI.APIKey,
|
APIKey: p.OpenAI.APIKey,
|
||||||
APIBase: p.OpenAI.APIBase,
|
APIBase: p.OpenAI.APIBase,
|
||||||
Proxy: p.OpenAI.Proxy,
|
Proxy: p.OpenAI.Proxy,
|
||||||
AuthMethod: p.OpenAI.AuthMethod,
|
AuthMethod: p.OpenAI.AuthMethod,
|
||||||
})
|
}, true
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
providerNames: []string{"anthropic", "claude"},
|
||||||
|
protocol: "anthropic",
|
||||||
|
buildConfig: func(p ProvidersConfig) (ModelConfig, bool) {
|
||||||
|
if p.Anthropic.APIKey == "" && p.Anthropic.APIBase == "" {
|
||||||
|
return ModelConfig{}, false
|
||||||
}
|
}
|
||||||
|
return ModelConfig{
|
||||||
// Anthropic
|
|
||||||
if p.Anthropic.APIKey != "" || p.Anthropic.APIBase != "" {
|
|
||||||
result = append(result, ModelConfig{
|
|
||||||
ModelName: "anthropic",
|
ModelName: "anthropic",
|
||||||
Model: "anthropic/claude-3-sonnet",
|
Model: "anthropic/claude-3-sonnet",
|
||||||
APIKey: p.Anthropic.APIKey,
|
APIKey: p.Anthropic.APIKey,
|
||||||
APIBase: p.Anthropic.APIBase,
|
APIBase: p.Anthropic.APIBase,
|
||||||
Proxy: p.Anthropic.Proxy,
|
Proxy: p.Anthropic.Proxy,
|
||||||
AuthMethod: p.Anthropic.AuthMethod,
|
AuthMethod: p.Anthropic.AuthMethod,
|
||||||
})
|
}, true
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
providerNames: []string{"openrouter"},
|
||||||
|
protocol: "openrouter",
|
||||||
|
buildConfig: func(p ProvidersConfig) (ModelConfig, bool) {
|
||||||
|
if p.OpenRouter.APIKey == "" && p.OpenRouter.APIBase == "" {
|
||||||
|
return ModelConfig{}, false
|
||||||
}
|
}
|
||||||
|
return ModelConfig{
|
||||||
// OpenRouter
|
|
||||||
if p.OpenRouter.APIKey != "" || p.OpenRouter.APIBase != "" {
|
|
||||||
result = append(result, ModelConfig{
|
|
||||||
ModelName: "openrouter",
|
ModelName: "openrouter",
|
||||||
Model: "openrouter/auto",
|
Model: "openrouter/auto",
|
||||||
APIKey: p.OpenRouter.APIKey,
|
APIKey: p.OpenRouter.APIKey,
|
||||||
APIBase: p.OpenRouter.APIBase,
|
APIBase: p.OpenRouter.APIBase,
|
||||||
Proxy: p.OpenRouter.Proxy,
|
Proxy: p.OpenRouter.Proxy,
|
||||||
})
|
}, true
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
providerNames: []string{"groq"},
|
||||||
|
protocol: "groq",
|
||||||
|
buildConfig: func(p ProvidersConfig) (ModelConfig, bool) {
|
||||||
|
if p.Groq.APIKey == "" && p.Groq.APIBase == "" {
|
||||||
|
return ModelConfig{}, false
|
||||||
}
|
}
|
||||||
|
return ModelConfig{
|
||||||
// Groq
|
|
||||||
if p.Groq.APIKey != "" || p.Groq.APIBase != "" {
|
|
||||||
result = append(result, ModelConfig{
|
|
||||||
ModelName: "groq",
|
ModelName: "groq",
|
||||||
Model: "groq/llama-3.1-70b-versatile",
|
Model: "groq/llama-3.1-70b-versatile",
|
||||||
APIKey: p.Groq.APIKey,
|
APIKey: p.Groq.APIKey,
|
||||||
APIBase: p.Groq.APIBase,
|
APIBase: p.Groq.APIBase,
|
||||||
Proxy: p.Groq.Proxy,
|
Proxy: p.Groq.Proxy,
|
||||||
})
|
}, true
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
providerNames: []string{"zhipu", "glm"},
|
||||||
|
protocol: "openai",
|
||||||
|
buildConfig: func(p ProvidersConfig) (ModelConfig, bool) {
|
||||||
|
if p.Zhipu.APIKey == "" && p.Zhipu.APIBase == "" {
|
||||||
|
return ModelConfig{}, false
|
||||||
}
|
}
|
||||||
|
return ModelConfig{
|
||||||
// Zhipu
|
|
||||||
if p.Zhipu.APIKey != "" || p.Zhipu.APIBase != "" {
|
|
||||||
result = append(result, ModelConfig{
|
|
||||||
ModelName: "zhipu",
|
ModelName: "zhipu",
|
||||||
Model: "openai/glm-4",
|
Model: "openai/glm-4",
|
||||||
APIKey: p.Zhipu.APIKey,
|
APIKey: p.Zhipu.APIKey,
|
||||||
APIBase: p.Zhipu.APIBase,
|
APIBase: p.Zhipu.APIBase,
|
||||||
Proxy: p.Zhipu.Proxy,
|
Proxy: p.Zhipu.Proxy,
|
||||||
})
|
}, true
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
providerNames: []string{"vllm"},
|
||||||
|
protocol: "openai",
|
||||||
|
buildConfig: func(p ProvidersConfig) (ModelConfig, bool) {
|
||||||
|
if p.VLLM.APIKey == "" && p.VLLM.APIBase == "" {
|
||||||
|
return ModelConfig{}, false
|
||||||
}
|
}
|
||||||
|
return ModelConfig{
|
||||||
// VLLM
|
|
||||||
if p.VLLM.APIKey != "" || p.VLLM.APIBase != "" {
|
|
||||||
result = append(result, ModelConfig{
|
|
||||||
ModelName: "vllm",
|
ModelName: "vllm",
|
||||||
Model: "openai/auto",
|
Model: "openai/auto",
|
||||||
APIKey: p.VLLM.APIKey,
|
APIKey: p.VLLM.APIKey,
|
||||||
APIBase: p.VLLM.APIBase,
|
APIBase: p.VLLM.APIBase,
|
||||||
Proxy: p.VLLM.Proxy,
|
Proxy: p.VLLM.Proxy,
|
||||||
})
|
}, true
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
providerNames: []string{"gemini", "google"},
|
||||||
|
protocol: "openai",
|
||||||
|
buildConfig: func(p ProvidersConfig) (ModelConfig, bool) {
|
||||||
|
if p.Gemini.APIKey == "" && p.Gemini.APIBase == "" {
|
||||||
|
return ModelConfig{}, false
|
||||||
}
|
}
|
||||||
|
return ModelConfig{
|
||||||
// Gemini
|
|
||||||
if p.Gemini.APIKey != "" || p.Gemini.APIBase != "" {
|
|
||||||
result = append(result, ModelConfig{
|
|
||||||
ModelName: "gemini",
|
ModelName: "gemini",
|
||||||
Model: "openai/gemini-pro",
|
Model: "openai/gemini-pro",
|
||||||
APIKey: p.Gemini.APIKey,
|
APIKey: p.Gemini.APIKey,
|
||||||
APIBase: p.Gemini.APIBase,
|
APIBase: p.Gemini.APIBase,
|
||||||
Proxy: p.Gemini.Proxy,
|
Proxy: p.Gemini.Proxy,
|
||||||
})
|
}, true
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
providerNames: []string{"nvidia"},
|
||||||
|
protocol: "nvidia",
|
||||||
|
buildConfig: func(p ProvidersConfig) (ModelConfig, bool) {
|
||||||
|
if p.Nvidia.APIKey == "" && p.Nvidia.APIBase == "" {
|
||||||
|
return ModelConfig{}, false
|
||||||
}
|
}
|
||||||
|
return ModelConfig{
|
||||||
// Nvidia
|
|
||||||
if p.Nvidia.APIKey != "" || p.Nvidia.APIBase != "" {
|
|
||||||
result = append(result, ModelConfig{
|
|
||||||
ModelName: "nvidia",
|
ModelName: "nvidia",
|
||||||
Model: "nvidia/meta/llama-3.1-8b-instruct",
|
Model: "nvidia/meta/llama-3.1-8b-instruct",
|
||||||
APIKey: p.Nvidia.APIKey,
|
APIKey: p.Nvidia.APIKey,
|
||||||
APIBase: p.Nvidia.APIBase,
|
APIBase: p.Nvidia.APIBase,
|
||||||
Proxy: p.Nvidia.Proxy,
|
Proxy: p.Nvidia.Proxy,
|
||||||
})
|
}, true
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
providerNames: []string{"ollama"},
|
||||||
|
protocol: "ollama",
|
||||||
|
buildConfig: func(p ProvidersConfig) (ModelConfig, bool) {
|
||||||
|
if p.Ollama.APIKey == "" && p.Ollama.APIBase == "" {
|
||||||
|
return ModelConfig{}, false
|
||||||
}
|
}
|
||||||
|
return ModelConfig{
|
||||||
// Ollama
|
|
||||||
if p.Ollama.APIKey != "" || p.Ollama.APIBase != "" {
|
|
||||||
result = append(result, ModelConfig{
|
|
||||||
ModelName: "ollama",
|
ModelName: "ollama",
|
||||||
Model: "ollama/llama3",
|
Model: "ollama/llama3",
|
||||||
APIKey: p.Ollama.APIKey,
|
APIKey: p.Ollama.APIKey,
|
||||||
APIBase: p.Ollama.APIBase,
|
APIBase: p.Ollama.APIBase,
|
||||||
Proxy: p.Ollama.Proxy,
|
Proxy: p.Ollama.Proxy,
|
||||||
})
|
}, true
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
providerNames: []string{"moonshot", "kimi"},
|
||||||
|
protocol: "moonshot",
|
||||||
|
buildConfig: func(p ProvidersConfig) (ModelConfig, bool) {
|
||||||
|
if p.Moonshot.APIKey == "" && p.Moonshot.APIBase == "" {
|
||||||
|
return ModelConfig{}, false
|
||||||
}
|
}
|
||||||
|
return ModelConfig{
|
||||||
// Moonshot
|
|
||||||
if p.Moonshot.APIKey != "" || p.Moonshot.APIBase != "" {
|
|
||||||
result = append(result, ModelConfig{
|
|
||||||
ModelName: "moonshot",
|
ModelName: "moonshot",
|
||||||
Model: "moonshot/kimi",
|
Model: "moonshot/kimi",
|
||||||
APIKey: p.Moonshot.APIKey,
|
APIKey: p.Moonshot.APIKey,
|
||||||
APIBase: p.Moonshot.APIBase,
|
APIBase: p.Moonshot.APIBase,
|
||||||
Proxy: p.Moonshot.Proxy,
|
Proxy: p.Moonshot.Proxy,
|
||||||
})
|
}, true
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
providerNames: []string{"shengsuanyun"},
|
||||||
|
protocol: "openai",
|
||||||
|
buildConfig: func(p ProvidersConfig) (ModelConfig, bool) {
|
||||||
|
if p.ShengSuanYun.APIKey == "" && p.ShengSuanYun.APIBase == "" {
|
||||||
|
return ModelConfig{}, false
|
||||||
}
|
}
|
||||||
|
return ModelConfig{
|
||||||
// ShengSuanYun
|
|
||||||
if p.ShengSuanYun.APIKey != "" || p.ShengSuanYun.APIBase != "" {
|
|
||||||
result = append(result, ModelConfig{
|
|
||||||
ModelName: "shengsuanyun",
|
ModelName: "shengsuanyun",
|
||||||
Model: "openai/auto",
|
Model: "openai/auto",
|
||||||
APIKey: p.ShengSuanYun.APIKey,
|
APIKey: p.ShengSuanYun.APIKey,
|
||||||
APIBase: p.ShengSuanYun.APIBase,
|
APIBase: p.ShengSuanYun.APIBase,
|
||||||
Proxy: p.ShengSuanYun.Proxy,
|
Proxy: p.ShengSuanYun.Proxy,
|
||||||
})
|
}, true
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
providerNames: []string{"deepseek"},
|
||||||
|
protocol: "openai",
|
||||||
|
buildConfig: func(p ProvidersConfig) (ModelConfig, bool) {
|
||||||
|
if p.DeepSeek.APIKey == "" && p.DeepSeek.APIBase == "" {
|
||||||
|
return ModelConfig{}, false
|
||||||
}
|
}
|
||||||
|
return ModelConfig{
|
||||||
// DeepSeek
|
|
||||||
if p.DeepSeek.APIKey != "" || p.DeepSeek.APIBase != "" {
|
|
||||||
result = append(result, ModelConfig{
|
|
||||||
ModelName: "deepseek",
|
ModelName: "deepseek",
|
||||||
Model: "openai/deepseek-chat",
|
Model: "openai/deepseek-chat",
|
||||||
APIKey: p.DeepSeek.APIKey,
|
APIKey: p.DeepSeek.APIKey,
|
||||||
APIBase: p.DeepSeek.APIBase,
|
APIBase: p.DeepSeek.APIBase,
|
||||||
Proxy: p.DeepSeek.Proxy,
|
Proxy: p.DeepSeek.Proxy,
|
||||||
})
|
}, true
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
providerNames: []string{"cerebras"},
|
||||||
|
protocol: "cerebras",
|
||||||
|
buildConfig: func(p ProvidersConfig) (ModelConfig, bool) {
|
||||||
|
if p.Cerebras.APIKey == "" && p.Cerebras.APIBase == "" {
|
||||||
|
return ModelConfig{}, false
|
||||||
}
|
}
|
||||||
|
return ModelConfig{
|
||||||
// Cerebras
|
|
||||||
if p.Cerebras.APIKey != "" || p.Cerebras.APIBase != "" {
|
|
||||||
result = append(result, ModelConfig{
|
|
||||||
ModelName: "cerebras",
|
ModelName: "cerebras",
|
||||||
Model: "cerebras/llama-3.3-70b",
|
Model: "cerebras/llama-3.3-70b",
|
||||||
APIKey: p.Cerebras.APIKey,
|
APIKey: p.Cerebras.APIKey,
|
||||||
APIBase: p.Cerebras.APIBase,
|
APIBase: p.Cerebras.APIBase,
|
||||||
Proxy: p.Cerebras.Proxy,
|
Proxy: p.Cerebras.Proxy,
|
||||||
})
|
}, true
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
providerNames: []string{"volcengine", "doubao"},
|
||||||
|
protocol: "openai",
|
||||||
|
buildConfig: func(p ProvidersConfig) (ModelConfig, bool) {
|
||||||
|
if p.VolcEngine.APIKey == "" && p.VolcEngine.APIBase == "" {
|
||||||
|
return ModelConfig{}, false
|
||||||
}
|
}
|
||||||
|
return ModelConfig{
|
||||||
// VolcEngine (Doubao)
|
|
||||||
if p.VolcEngine.APIKey != "" || p.VolcEngine.APIBase != "" {
|
|
||||||
result = append(result, ModelConfig{
|
|
||||||
ModelName: "volcengine",
|
ModelName: "volcengine",
|
||||||
Model: "openai/doubao-pro",
|
Model: "openai/doubao-pro",
|
||||||
APIKey: p.VolcEngine.APIKey,
|
APIKey: p.VolcEngine.APIKey,
|
||||||
APIBase: p.VolcEngine.APIBase,
|
APIBase: p.VolcEngine.APIBase,
|
||||||
Proxy: p.VolcEngine.Proxy,
|
Proxy: p.VolcEngine.Proxy,
|
||||||
})
|
}, true
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
providerNames: []string{"github_copilot", "copilot"},
|
||||||
|
protocol: "github-copilot",
|
||||||
|
buildConfig: func(p ProvidersConfig) (ModelConfig, bool) {
|
||||||
|
if p.GitHubCopilot.APIKey == "" && p.GitHubCopilot.APIBase == "" && p.GitHubCopilot.ConnectMode == "" {
|
||||||
|
return ModelConfig{}, false
|
||||||
}
|
}
|
||||||
|
return ModelConfig{
|
||||||
// GitHub Copilot
|
|
||||||
if p.GitHubCopilot.APIKey != "" || p.GitHubCopilot.APIBase != "" || p.GitHubCopilot.ConnectMode != "" {
|
|
||||||
result = append(result, ModelConfig{
|
|
||||||
ModelName: "github-copilot",
|
ModelName: "github-copilot",
|
||||||
Model: "github-copilot/gpt-4o",
|
Model: "github-copilot/gpt-4o",
|
||||||
APIBase: p.GitHubCopilot.APIBase,
|
APIBase: p.GitHubCopilot.APIBase,
|
||||||
ConnectMode: p.GitHubCopilot.ConnectMode,
|
ConnectMode: p.GitHubCopilot.ConnectMode,
|
||||||
})
|
}, true
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
providerNames: []string{"antigravity"},
|
||||||
|
protocol: "antigravity",
|
||||||
|
buildConfig: func(p ProvidersConfig) (ModelConfig, bool) {
|
||||||
|
if p.Antigravity.APIKey == "" && p.Antigravity.AuthMethod == "" {
|
||||||
|
return ModelConfig{}, false
|
||||||
}
|
}
|
||||||
|
return ModelConfig{
|
||||||
// Antigravity
|
|
||||||
if p.Antigravity.APIKey != "" || p.Antigravity.AuthMethod != "" {
|
|
||||||
result = append(result, ModelConfig{
|
|
||||||
ModelName: "antigravity",
|
ModelName: "antigravity",
|
||||||
Model: "antigravity/gemini-2.0-flash",
|
Model: "antigravity/gemini-2.0-flash",
|
||||||
APIKey: p.Antigravity.APIKey,
|
APIKey: p.Antigravity.APIKey,
|
||||||
AuthMethod: p.Antigravity.AuthMethod,
|
AuthMethod: p.Antigravity.AuthMethod,
|
||||||
})
|
}, true
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
providerNames: []string{"qwen", "tongyi"},
|
||||||
|
protocol: "qwen",
|
||||||
|
buildConfig: func(p ProvidersConfig) (ModelConfig, bool) {
|
||||||
|
if p.Qwen.APIKey == "" && p.Qwen.APIBase == "" {
|
||||||
|
return ModelConfig{}, false
|
||||||
}
|
}
|
||||||
|
return ModelConfig{
|
||||||
// Qwen
|
|
||||||
if p.Qwen.APIKey != "" || p.Qwen.APIBase != "" {
|
|
||||||
result = append(result, ModelConfig{
|
|
||||||
ModelName: "qwen",
|
ModelName: "qwen",
|
||||||
Model: "qwen/qwen-max",
|
Model: "qwen/qwen-max",
|
||||||
APIKey: p.Qwen.APIKey,
|
APIKey: p.Qwen.APIKey,
|
||||||
APIBase: p.Qwen.APIBase,
|
APIBase: p.Qwen.APIBase,
|
||||||
Proxy: p.Qwen.Proxy,
|
Proxy: p.Qwen.Proxy,
|
||||||
})
|
}, true
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
// Process each provider migration
|
||||||
|
for _, m := range migrations {
|
||||||
|
mc, ok := m.buildConfig(p)
|
||||||
|
if !ok {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check if this is the user's configured provider
|
||||||
|
if slices.Contains(m.providerNames, userProvider) && userModel != "" {
|
||||||
|
// Use the user's configured model instead of default
|
||||||
|
mc.Model = m.protocol + "/" + userModel
|
||||||
|
}
|
||||||
|
|
||||||
|
result = append(result, mc)
|
||||||
}
|
}
|
||||||
|
|
||||||
return result
|
return result
|
||||||
|
|
|
||||||
|
|
@ -6,6 +6,7 @@
|
||||||
package config
|
package config
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
@ -175,3 +176,218 @@ func TestConvertProvidersToModelList_AuthMethod(t *testing.T) {
|
||||||
t.Errorf("len(result) = %d, want 0 (AuthMethod alone should not create entry)", len(result))
|
t.Errorf("len(result) = %d, want 0 (AuthMethod alone should not create entry)", len(result))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Tests for preserving user's configured model during migration
|
||||||
|
|
||||||
|
func TestConvertProvidersToModelList_PreservesUserModel_DeepSeek(t *testing.T) {
|
||||||
|
cfg := &Config{
|
||||||
|
Agents: AgentsConfig{
|
||||||
|
Defaults: AgentDefaults{
|
||||||
|
Provider: "deepseek",
|
||||||
|
Model: "deepseek-reasoner",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
Providers: ProvidersConfig{
|
||||||
|
DeepSeek: ProviderConfig{APIKey: "sk-deepseek"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
result := ConvertProvidersToModelList(cfg)
|
||||||
|
|
||||||
|
if len(result) != 1 {
|
||||||
|
t.Fatalf("len(result) = %d, want 1", len(result))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Should use user's model, not default
|
||||||
|
if result[0].Model != "openai/deepseek-reasoner" {
|
||||||
|
t.Errorf("Model = %q, want %q (user's configured model)", result[0].Model, "openai/deepseek-reasoner")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConvertProvidersToModelList_PreservesUserModel_OpenAI(t *testing.T) {
|
||||||
|
cfg := &Config{
|
||||||
|
Agents: AgentsConfig{
|
||||||
|
Defaults: AgentDefaults{
|
||||||
|
Provider: "openai",
|
||||||
|
Model: "gpt-4-turbo",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
Providers: ProvidersConfig{
|
||||||
|
OpenAI: ProviderConfig{APIKey: "sk-openai"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
result := ConvertProvidersToModelList(cfg)
|
||||||
|
|
||||||
|
if len(result) != 1 {
|
||||||
|
t.Fatalf("len(result) = %d, want 1", len(result))
|
||||||
|
}
|
||||||
|
|
||||||
|
if result[0].Model != "openai/gpt-4-turbo" {
|
||||||
|
t.Errorf("Model = %q, want %q", result[0].Model, "openai/gpt-4-turbo")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConvertProvidersToModelList_PreservesUserModel_Anthropic(t *testing.T) {
|
||||||
|
cfg := &Config{
|
||||||
|
Agents: AgentsConfig{
|
||||||
|
Defaults: AgentDefaults{
|
||||||
|
Provider: "claude", // alternative name
|
||||||
|
Model: "claude-3-opus-20240229",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
Providers: ProvidersConfig{
|
||||||
|
Anthropic: ProviderConfig{APIKey: "sk-ant"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
result := ConvertProvidersToModelList(cfg)
|
||||||
|
|
||||||
|
if len(result) != 1 {
|
||||||
|
t.Fatalf("len(result) = %d, want 1", len(result))
|
||||||
|
}
|
||||||
|
|
||||||
|
if result[0].Model != "anthropic/claude-3-opus-20240229" {
|
||||||
|
t.Errorf("Model = %q, want %q", result[0].Model, "anthropic/claude-3-opus-20240229")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConvertProvidersToModelList_PreservesUserModel_Qwen(t *testing.T) {
|
||||||
|
cfg := &Config{
|
||||||
|
Agents: AgentsConfig{
|
||||||
|
Defaults: AgentDefaults{
|
||||||
|
Provider: "qwen",
|
||||||
|
Model: "qwen-plus",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
Providers: ProvidersConfig{
|
||||||
|
Qwen: ProviderConfig{APIKey: "sk-qwen"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
result := ConvertProvidersToModelList(cfg)
|
||||||
|
|
||||||
|
if len(result) != 1 {
|
||||||
|
t.Fatalf("len(result) = %d, want 1", len(result))
|
||||||
|
}
|
||||||
|
|
||||||
|
if result[0].Model != "qwen/qwen-plus" {
|
||||||
|
t.Errorf("Model = %q, want %q", result[0].Model, "qwen/qwen-plus")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConvertProvidersToModelList_UsesDefaultWhenNoUserModel(t *testing.T) {
|
||||||
|
cfg := &Config{
|
||||||
|
Agents: AgentsConfig{
|
||||||
|
Defaults: AgentDefaults{
|
||||||
|
Provider: "deepseek",
|
||||||
|
Model: "", // no model specified
|
||||||
|
},
|
||||||
|
},
|
||||||
|
Providers: ProvidersConfig{
|
||||||
|
DeepSeek: ProviderConfig{APIKey: "sk-deepseek"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
result := ConvertProvidersToModelList(cfg)
|
||||||
|
|
||||||
|
if len(result) != 1 {
|
||||||
|
t.Fatalf("len(result) = %d, want 1", len(result))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Should use default model
|
||||||
|
if result[0].Model != "openai/deepseek-chat" {
|
||||||
|
t.Errorf("Model = %q, want %q (default)", result[0].Model, "openai/deepseek-chat")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConvertProvidersToModelList_MultipleProviders_PreservesUserModel(t *testing.T) {
|
||||||
|
cfg := &Config{
|
||||||
|
Agents: AgentsConfig{
|
||||||
|
Defaults: AgentDefaults{
|
||||||
|
Provider: "deepseek",
|
||||||
|
Model: "deepseek-reasoner",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
Providers: ProvidersConfig{
|
||||||
|
OpenAI: ProviderConfig{APIKey: "sk-openai"},
|
||||||
|
DeepSeek: ProviderConfig{APIKey: "sk-deepseek"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
result := ConvertProvidersToModelList(cfg)
|
||||||
|
|
||||||
|
if len(result) != 2 {
|
||||||
|
t.Fatalf("len(result) = %d, want 2", len(result))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Find each provider and verify model
|
||||||
|
for _, mc := range result {
|
||||||
|
switch mc.ModelName {
|
||||||
|
case "openai":
|
||||||
|
if mc.Model != "openai/gpt-4o" {
|
||||||
|
t.Errorf("OpenAI Model = %q, want %q (default)", mc.Model, "openai/gpt-4o")
|
||||||
|
}
|
||||||
|
case "deepseek":
|
||||||
|
if mc.Model != "openai/deepseek-reasoner" {
|
||||||
|
t.Errorf("DeepSeek Model = %q, want %q (user's)", mc.Model, "openai/deepseek-reasoner")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConvertProvidersToModelList_ProviderNameAliases(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
providerAlias string
|
||||||
|
expectedModel string
|
||||||
|
provider ProviderConfig
|
||||||
|
}{
|
||||||
|
{"gpt", "openai/gpt-4-custom", ProviderConfig{APIKey: "key"}},
|
||||||
|
{"claude", "anthropic/claude-custom", ProviderConfig{APIKey: "key"}},
|
||||||
|
{"doubao", "openai/doubao-custom", ProviderConfig{APIKey: "key"}},
|
||||||
|
{"tongyi", "qwen/qwen-custom", ProviderConfig{APIKey: "key"}},
|
||||||
|
{"kimi", "moonshot/kimi-custom", ProviderConfig{APIKey: "key"}},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.providerAlias, func(t *testing.T) {
|
||||||
|
cfg := &Config{
|
||||||
|
Agents: AgentsConfig{
|
||||||
|
Defaults: AgentDefaults{
|
||||||
|
Provider: tt.providerAlias,
|
||||||
|
Model: strings.TrimPrefix(tt.expectedModel, tt.expectedModel[:strings.Index(tt.expectedModel, "/")+1]),
|
||||||
|
},
|
||||||
|
},
|
||||||
|
Providers: ProvidersConfig{},
|
||||||
|
}
|
||||||
|
|
||||||
|
// Set the appropriate provider config
|
||||||
|
switch tt.providerAlias {
|
||||||
|
case "gpt":
|
||||||
|
cfg.Providers.OpenAI = tt.provider
|
||||||
|
case "claude":
|
||||||
|
cfg.Providers.Anthropic = tt.provider
|
||||||
|
case "doubao":
|
||||||
|
cfg.Providers.VolcEngine = tt.provider
|
||||||
|
case "tongyi":
|
||||||
|
cfg.Providers.Qwen = tt.provider
|
||||||
|
case "kimi":
|
||||||
|
cfg.Providers.Moonshot = tt.provider
|
||||||
|
}
|
||||||
|
|
||||||
|
// Need to fix the model name in config
|
||||||
|
cfg.Agents.Defaults.Model = strings.TrimPrefix(tt.expectedModel, tt.expectedModel[:strings.Index(tt.expectedModel, "/")+1])
|
||||||
|
|
||||||
|
result := ConvertProvidersToModelList(cfg)
|
||||||
|
if len(result) != 1 {
|
||||||
|
t.Fatalf("len(result) = %d, want 1", len(result))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Extract just the model ID part (after the first /)
|
||||||
|
expectedModelID := tt.expectedModel
|
||||||
|
if result[0].Model != expectedModelID {
|
||||||
|
t.Errorf("Model = %q, want %q", result[0].Model, expectedModelID)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -45,14 +45,12 @@ func createCodexAuthProvider() (LLMProvider, error) {
|
||||||
// - "gpt-4o" -> ("openai", "gpt-4o") // default protocol
|
// - "gpt-4o" -> ("openai", "gpt-4o") // default protocol
|
||||||
func ExtractProtocol(model string) (protocol, modelID string) {
|
func ExtractProtocol(model string) (protocol, modelID string) {
|
||||||
model = strings.TrimSpace(model)
|
model = strings.TrimSpace(model)
|
||||||
for i := 0; i < len(model); i++ {
|
protocol, modelID, found := strings.Cut(model, "/")
|
||||||
if model[i] == '/' {
|
if !found {
|
||||||
return model[:i], model[i+1:]
|
|
||||||
}
|
|
||||||
}
|
|
||||||
// No prefix found, default to openai
|
|
||||||
return "openai", model
|
return "openai", model
|
||||||
}
|
}
|
||||||
|
return protocol, modelID
|
||||||
|
}
|
||||||
|
|
||||||
// CreateProviderFromConfig creates a provider based on the ModelConfig.
|
// CreateProviderFromConfig creates a provider based on the ModelConfig.
|
||||||
// It uses the protocol prefix in the Model field to determine which provider to create.
|
// It uses the protocol prefix in the Model field to determine which provider to create.
|
||||||
|
|
|
||||||
54
pkg/providers/toolcall_utils.go
Normal file
54
pkg/providers/toolcall_utils.go
Normal file
|
|
@ -0,0 +1,54 @@
|
||||||
|
// PicoClaw - Ultra-lightweight personal AI agent
|
||||||
|
// License: MIT
|
||||||
|
//
|
||||||
|
// Copyright (c) 2026 PicoClaw contributors
|
||||||
|
|
||||||
|
package providers
|
||||||
|
|
||||||
|
import "encoding/json"
|
||||||
|
|
||||||
|
// NormalizeToolCall normalizes a ToolCall to ensure all fields are properly populated.
|
||||||
|
// It handles cases where Name/Arguments might be in different locations (top-level vs Function)
|
||||||
|
// and ensures both are populated consistently.
|
||||||
|
func NormalizeToolCall(tc ToolCall) ToolCall {
|
||||||
|
normalized := tc
|
||||||
|
|
||||||
|
// Ensure Name is populated from Function if not set
|
||||||
|
if normalized.Name == "" && normalized.Function != nil {
|
||||||
|
normalized.Name = normalized.Function.Name
|
||||||
|
}
|
||||||
|
|
||||||
|
// Ensure Arguments is not nil
|
||||||
|
if normalized.Arguments == nil {
|
||||||
|
normalized.Arguments = map[string]interface{}{}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Parse Arguments from Function.Arguments if not already set
|
||||||
|
if len(normalized.Arguments) == 0 && normalized.Function != nil && normalized.Function.Arguments != "" {
|
||||||
|
var parsed map[string]interface{}
|
||||||
|
if err := json.Unmarshal([]byte(normalized.Function.Arguments), &parsed); err == nil && parsed != nil {
|
||||||
|
normalized.Arguments = parsed
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Ensure Function is populated with consistent values
|
||||||
|
argsJSON, _ := json.Marshal(normalized.Arguments)
|
||||||
|
if normalized.Function == nil {
|
||||||
|
normalized.Function = &FunctionCall{
|
||||||
|
Name: normalized.Name,
|
||||||
|
Arguments: string(argsJSON),
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
if normalized.Function.Name == "" {
|
||||||
|
normalized.Function.Name = normalized.Name
|
||||||
|
}
|
||||||
|
if normalized.Name == "" {
|
||||||
|
normalized.Name = normalized.Function.Name
|
||||||
|
}
|
||||||
|
if normalized.Function.Arguments == "" {
|
||||||
|
normalized.Function.Arguments = string(argsJSON)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return normalized
|
||||||
|
}
|
||||||
|
|
@ -85,7 +85,7 @@ func RunToolLoop(ctx context.Context, config ToolLoopConfig, messages []provider
|
||||||
|
|
||||||
normalizedToolCalls := make([]providers.ToolCall, 0, len(response.ToolCalls))
|
normalizedToolCalls := make([]providers.ToolCall, 0, len(response.ToolCalls))
|
||||||
for _, tc := range response.ToolCalls {
|
for _, tc := range response.ToolCalls {
|
||||||
normalizedToolCalls = append(normalizedToolCalls, normalizeProviderToolCall(tc))
|
normalizedToolCalls = append(normalizedToolCalls, providers.NormalizeToolCall(tc))
|
||||||
}
|
}
|
||||||
|
|
||||||
// 5. Log tool calls
|
// 5. Log tool calls
|
||||||
|
|
@ -159,42 +159,3 @@ func RunToolLoop(ctx context.Context, config ToolLoopConfig, messages []provider
|
||||||
Iterations: iteration,
|
Iterations: iteration,
|
||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func normalizeProviderToolCall(tc providers.ToolCall) providers.ToolCall {
|
|
||||||
normalized := tc
|
|
||||||
|
|
||||||
if normalized.Name == "" && normalized.Function != nil {
|
|
||||||
normalized.Name = normalized.Function.Name
|
|
||||||
}
|
|
||||||
|
|
||||||
if normalized.Arguments == nil {
|
|
||||||
normalized.Arguments = map[string]interface{}{}
|
|
||||||
}
|
|
||||||
|
|
||||||
if len(normalized.Arguments) == 0 && normalized.Function != nil && normalized.Function.Arguments != "" {
|
|
||||||
var parsed map[string]interface{}
|
|
||||||
if err := json.Unmarshal([]byte(normalized.Function.Arguments), &parsed); err == nil && parsed != nil {
|
|
||||||
normalized.Arguments = parsed
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
argsJSON, _ := json.Marshal(normalized.Arguments)
|
|
||||||
if normalized.Function == nil {
|
|
||||||
normalized.Function = &providers.FunctionCall{
|
|
||||||
Name: normalized.Name,
|
|
||||||
Arguments: string(argsJSON),
|
|
||||||
}
|
|
||||||
} else {
|
|
||||||
if normalized.Function.Name == "" {
|
|
||||||
normalized.Function.Name = normalized.Name
|
|
||||||
}
|
|
||||||
if normalized.Name == "" {
|
|
||||||
normalized.Name = normalized.Function.Name
|
|
||||||
}
|
|
||||||
if normalized.Function.Arguments == "" {
|
|
||||||
normalized.Function.Arguments = string(argsJSON)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
return normalized
|
|
||||||
}
|
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue