From 1abcc95ed619f11428cf7afff9c0264ecb9510b2 Mon Sep 17 00:00:00 2001 From: liugangjian Date: Thu, 5 Mar 2026 09:02:19 +0800 Subject: [PATCH] Fix deepseek model max_tokens validation error --- pkg/providers/factory.go | 30 ++++++++++++ pkg/providers/factory_provider.go | 28 ++++++++++-- pkg/providers/factory_test.go | 61 +++++++++++++++++++++++-- pkg/providers/openai_compat/provider.go | 15 +++++- 4 files changed, 127 insertions(+), 7 deletions(-) diff --git a/pkg/providers/factory.go b/pkg/providers/factory.go index 5b3e42b9e..99f154a1e 100644 --- a/pkg/providers/factory.go +++ b/pkg/providers/factory.go @@ -217,7 +217,9 @@ func resolveProviderSelection(cfg *config.Config) (providerSelection, error) { strings.HasPrefix(model, "openai/") || strings.HasPrefix(model, "meta-llama/") || strings.HasPrefix(model, "deepseek/") || + strings.HasPrefix(model, "deepseek-ai/") || strings.HasPrefix(model, "google/"): + sel.apiKey = cfg.Providers.OpenRouter.APIKey sel.proxy = cfg.Providers.OpenRouter.Proxy if cfg.Providers.OpenRouter.APIBase != "" { @@ -279,7 +281,35 @@ func resolveProviderSelection(cfg *config.Config) (providerSelection, error) { if sel.apiBase == "" { sel.apiBase = "https://api.groq.com/openai/v1" } + + case (strings.Contains(lowerModel, "deepseek") || strings.HasPrefix(model, "deepseek/") || strings.HasPrefix(model, "deepseek-ai/")) && cfg.Providers.DeepSeek.APIKey != "": + + sel.apiKey = cfg.Providers.DeepSeek.APIKey + + sel.apiBase = cfg.Providers.DeepSeek.APIBase + + sel.proxy = cfg.Providers.DeepSeek.Proxy + + if sel.apiBase == "" { + + sel.apiBase = "https://api.deepseek.com/v1" + + } + + // Keep original model name logic for backward compatibility + + if !strings.HasPrefix(model, "deepseek/") && !strings.HasPrefix(model, "deepseek-ai/") { + + if model != "deepseek-chat" && model != "deepseek-reasoner" { + + sel.model = "deepseek-chat" + + } + + } + case (strings.Contains(lowerModel, "nvidia") || strings.HasPrefix(model, "nvidia/")) && cfg.Providers.Nvidia.APIKey != "": + sel.apiKey = cfg.Providers.Nvidia.APIKey sel.apiBase = cfg.Providers.Nvidia.APIBase sel.proxy = cfg.Providers.Nvidia.Proxy diff --git a/pkg/providers/factory_provider.go b/pkg/providers/factory_provider.go index 155317a3b..e158fdf72 100644 --- a/pkg/providers/factory_provider.go +++ b/pkg/providers/factory_provider.go @@ -55,6 +55,13 @@ func ExtractProtocol(model string) (protocol, modelID string) { // It uses the protocol prefix in the Model field to determine which provider to create. // Supported protocols: openai, litellm, anthropic, antigravity, claude-cli, codex-cli, github-copilot // Returns the provider, the model ID (without protocol prefix), and any error. +// ExtractProtocol extracts the protocol prefix and model identifier from a model string. +// If no prefix is specified, it defaults to "openai". +// Examples: +// - "openai/gpt-4o" -> ("openai", "gpt-4o") +// - "anthropic/claude-sonnet-4.6" -> ("anthropic", "claude-sonnet-4.6") +// - "gpt-4o" -> ("openai", "gpt-4o") // default protocol +func ExtractProtocol(model string) (protocol, modelID string) { func CreateProviderFromConfig(cfg *config.ModelConfig) (LLMProvider, string, error) { if cfg == nil { return nil, "", fmt.Errorf("config is nil") @@ -84,12 +91,17 @@ func CreateProviderFromConfig(cfg *config.ModelConfig) (LLMProvider, string, err if apiBase == "" { apiBase = getDefaultAPIBase(protocol) } + // For Ollama models, default to higher timeout if none specified + requestTimeout := cfg.RequestTimeout + if protocol == "ollama" && requestTimeout <= 0 { + requestTimeout = 300 // Ollama models often need more time, default to 300 seconds + } return NewHTTPProviderWithMaxTokensFieldAndRequestTimeout( cfg.APIKey, apiBase, cfg.Proxy, cfg.MaxTokensField, - cfg.RequestTimeout, + requestTimeout, ), modelID, nil case "litellm", "openrouter", "groq", "zhipu", "gemini", "nvidia", @@ -103,12 +115,17 @@ func CreateProviderFromConfig(cfg *config.ModelConfig) (LLMProvider, string, err if apiBase == "" { apiBase = getDefaultAPIBase(protocol) } + // For Ollama models, default to higher timeout if none specified + requestTimeout := cfg.RequestTimeout + if protocol == "ollama" && requestTimeout <= 0 { + requestTimeout = 300 // Ollama models often need more time, default to 300 seconds + } return NewHTTPProviderWithMaxTokensFieldAndRequestTimeout( cfg.APIKey, apiBase, cfg.Proxy, cfg.MaxTokensField, - cfg.RequestTimeout, + requestTimeout, ), modelID, nil case "anthropic": @@ -128,12 +145,17 @@ func CreateProviderFromConfig(cfg *config.ModelConfig) (LLMProvider, string, err if cfg.APIKey == "" { return nil, "", fmt.Errorf("api_key is required for anthropic protocol (model: %s)", cfg.Model) } + // For Ollama models, default to higher timeout if none specified + requestTimeout := cfg.RequestTimeout + if protocol == "ollama" && requestTimeout <= 0 { + requestTimeout = 300 // Ollama models often need more time, default to 300 seconds + } return NewHTTPProviderWithMaxTokensFieldAndRequestTimeout( cfg.APIKey, apiBase, cfg.Proxy, cfg.MaxTokensField, - cfg.RequestTimeout, + requestTimeout, ), modelID, nil case "antigravity": diff --git a/pkg/providers/factory_test.go b/pkg/providers/factory_test.go index f7a916d9e..cde265500 100644 --- a/pkg/providers/factory_test.go +++ b/pkg/providers/factory_test.go @@ -179,9 +179,64 @@ func TestResolveProviderSelection(t *testing.T) { setup: func(cfg *config.Config) { cfg.Agents.Defaults.Model = "openrouter/auto" }, - wantErrSubstr: "no API key configured for provider", - }, - } + { + + name: "deepseek-ai model with DeepSeek API key uses DeepSeek API", + + setup: func(cfg *config.Config) { + + cfg.Agents.Defaults.Model = "deepseek-ai/DeepSeek-V3.2" + + cfg.Providers.DeepSeek.APIKey = "deepseek-key" + + cfg.Providers.DeepSeek.APIBase = "https://api.deepseek.com/v1" + + }, + + wantType: providerTypeHTTPCompat, + + wantAPIBase: "https://api.deepseek.com/v1", + + }, + + { + + name: "deepseek-ai prefix model routes to DeepSeek when key configured", + + setup: func(cfg *config.Config) { + + cfg.Agents.Defaults.Model = "deepseek-ai/deepseek-coder" + + cfg.Providers.DeepSeek.APIKey = "deepseek-key" + + }, + + wantType: providerTypeHTTPCompat, + + wantAPIBase: "https://api.deepseek.com/v1", + + }, + + { + + name: "deepseek model still works for backward compatibility", + + setup: func(cfg *config.Config) { + + cfg.Agents.Defaults.Model = "deepseek/deepseek-chat" + + cfg.Providers.DeepSeek.APIKey = "deepseek-key" + + }, + + wantType: providerTypeHTTPCompat, + + wantAPIBase: "https://api.deepseek.com/v1", + + }, + +} + for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { diff --git a/pkg/providers/openai_compat/provider.go b/pkg/providers/openai_compat/provider.go index ff9109e96..41bfb52ae 100644 --- a/pkg/providers/openai_compat/provider.go +++ b/pkg/providers/openai_compat/provider.go @@ -137,7 +137,20 @@ func (p *Provider) Chat( fieldName = "max_tokens" } } - requestBody[fieldName] = maxTokens + // Validate max_tokens range for deepseek model + if strings.Contains(strings.ToLower(model), "deepseek") { + // Clamp max_tokens to the valid range [1, 8192] + adjustedMaxTokens := maxTokens + if maxTokens > 8192 { + adjustedMaxTokens = 8192 + } + if maxTokens < 1 { + adjustedMaxTokens = 1 + } + requestBody[fieldName] = adjustedMaxTokens + } else { + requestBody[fieldName] = maxTokens + } } if temperature, ok := asFloat(options["temperature"]); ok {