From 9ee422c12d33fab05f8288a612cc42f7f8a4c3b7 Mon Sep 17 00:00:00 2001 From: lc6464 <64722907+lc6464@users.noreply.github.com> Date: Mon, 30 Mar 2026 21:38:11 +0800 Subject: [PATCH] refactor(provider): consolidate protocol metadata and local tests --- pkg/providers/factory_provider.go | 111 ++++++-------- pkg/providers/factory_provider_test.go | 196 +++++++++---------------- 2 files changed, 120 insertions(+), 187 deletions(-) diff --git a/pkg/providers/factory_provider.go b/pkg/providers/factory_provider.go index fc4c59e08..80b707b9d 100644 --- a/pkg/providers/factory_provider.go +++ b/pkg/providers/factory_provider.go @@ -17,6 +17,48 @@ import ( "github.com/sipeed/picoclaw/pkg/providers/bedrock" ) +type protocolMeta struct { + defaultAPIBase string + emptyAPIKeyAllowed bool +} + +var protocolMetaByName = map[string]protocolMeta{ + "openai": {defaultAPIBase: "https://api.openai.com/v1"}, + "openrouter": {defaultAPIBase: "https://openrouter.ai/api/v1"}, + "litellm": {defaultAPIBase: "http://localhost:4000/v1"}, + "lmstudio": {defaultAPIBase: "http://localhost:1234/v1", emptyAPIKeyAllowed: true}, + "novita": {defaultAPIBase: "https://api.novita.ai/openai"}, + "groq": {defaultAPIBase: "https://api.groq.com/openai/v1"}, + "zhipu": {defaultAPIBase: "https://open.bigmodel.cn/api/paas/v4"}, + "gemini": {defaultAPIBase: "https://generativelanguage.googleapis.com/v1beta"}, + "nvidia": {defaultAPIBase: "https://integrate.api.nvidia.com/v1"}, + "ollama": {defaultAPIBase: "http://localhost:11434/v1", emptyAPIKeyAllowed: true}, + "moonshot": {defaultAPIBase: "https://api.moonshot.cn/v1"}, + "shengsuanyun": {defaultAPIBase: "https://router.shengsuanyun.com/api/v1"}, + "deepseek": {defaultAPIBase: "https://api.deepseek.com/v1"}, + "cerebras": {defaultAPIBase: "https://api.cerebras.ai/v1"}, + "vivgrid": {defaultAPIBase: "https://api.vivgrid.com/v1"}, + "volcengine": {defaultAPIBase: "https://ark.cn-beijing.volces.com/api/v3"}, + "qwen": {defaultAPIBase: "https://dashscope.aliyuncs.com/compatible-mode/v1"}, + "qwen-intl": {defaultAPIBase: "https://dashscope-intl.aliyuncs.com/compatible-mode/v1"}, + "qwen-international": {defaultAPIBase: "https://dashscope-intl.aliyuncs.com/compatible-mode/v1"}, + "dashscope-intl": {defaultAPIBase: "https://dashscope-intl.aliyuncs.com/compatible-mode/v1"}, + "qwen-us": {defaultAPIBase: "https://dashscope-us.aliyuncs.com/compatible-mode/v1"}, + "dashscope-us": {defaultAPIBase: "https://dashscope-us.aliyuncs.com/compatible-mode/v1"}, + "coding-plan": {defaultAPIBase: "https://coding-intl.dashscope.aliyuncs.com/v1"}, + "alibaba-coding": {defaultAPIBase: "https://coding-intl.dashscope.aliyuncs.com/v1"}, + "qwen-coding": {defaultAPIBase: "https://coding-intl.dashscope.aliyuncs.com/v1"}, + "coding-plan-anthropic": {defaultAPIBase: "https://coding-intl.dashscope.aliyuncs.com/apps/anthropic"}, + "alibaba-coding-anthropic": {defaultAPIBase: "https://coding-intl.dashscope.aliyuncs.com/apps/anthropic"}, + "vllm": {defaultAPIBase: "http://localhost:8000/v1", emptyAPIKeyAllowed: true}, + "mistral": {defaultAPIBase: "https://api.mistral.ai/v1"}, + "avian": {defaultAPIBase: "https://api.avian.io/v1"}, + "minimax": {defaultAPIBase: "https://api.minimaxi.com/v1"}, + "longcat": {defaultAPIBase: "https://api.longcat.chat/openai"}, + "modelscope": {defaultAPIBase: "https://api-inference.modelscope.cn/v1"}, + "mimo": {defaultAPIBase: "https://api.xiaomimimo.com/v1"}, +} + // createClaudeAuthProvider creates a Claude provider using OAuth credentials from auth store. func createClaudeAuthProvider() (LLMProvider, error) { cred, err := getCredential("anthropic") @@ -295,74 +337,15 @@ func CreateProviderFromConfig(cfg *config.ModelConfig) (LLMProvider, string, err } func isEmptyAPIKeyAllowed(protocol string) bool { - switch protocol { - case "ollama", "vllm", "lmstudio": - return true - default: - return false - } + meta, ok := protocolMetaByName[protocol] + return ok && meta.emptyAPIKeyAllowed } // getDefaultAPIBase returns the default API base URL for a given protocol. func getDefaultAPIBase(protocol string) string { - switch protocol { - case "openai": - return "https://api.openai.com/v1" - case "openrouter": - return "https://openrouter.ai/api/v1" - case "litellm": - return "http://localhost:4000/v1" - case "lmstudio": - return "http://localhost:1234/v1" - case "novita": - return "https://api.novita.ai/openai" - case "groq": - return "https://api.groq.com/openai/v1" - case "zhipu": - return "https://open.bigmodel.cn/api/paas/v4" - case "gemini": - return "https://generativelanguage.googleapis.com/v1beta" - case "nvidia": - return "https://integrate.api.nvidia.com/v1" - case "ollama": - return "http://localhost:11434/v1" - case "moonshot": - return "https://api.moonshot.cn/v1" - case "shengsuanyun": - return "https://router.shengsuanyun.com/api/v1" - case "deepseek": - return "https://api.deepseek.com/v1" - case "cerebras": - return "https://api.cerebras.ai/v1" - case "vivgrid": - return "https://api.vivgrid.com/v1" - case "volcengine": - return "https://ark.cn-beijing.volces.com/api/v3" - case "qwen": - return "https://dashscope.aliyuncs.com/compatible-mode/v1" - case "qwen-intl", "qwen-international", "dashscope-intl": - return "https://dashscope-intl.aliyuncs.com/compatible-mode/v1" - case "qwen-us", "dashscope-us": - return "https://dashscope-us.aliyuncs.com/compatible-mode/v1" - case "coding-plan", "alibaba-coding", "qwen-coding": - return "https://coding-intl.dashscope.aliyuncs.com/v1" - case "coding-plan-anthropic", "alibaba-coding-anthropic": - return "https://coding-intl.dashscope.aliyuncs.com/apps/anthropic" - case "vllm": - return "http://localhost:8000/v1" - case "mistral": - return "https://api.mistral.ai/v1" - case "avian": - return "https://api.avian.io/v1" - case "minimax": - return "https://api.minimaxi.com/v1" - case "longcat": - return "https://api.longcat.chat/openai" - case "modelscope": - return "https://api-inference.modelscope.cn/v1" - case "mimo": - return "https://api.xiaomimimo.com/v1" - default: + meta, ok := protocolMetaByName[protocol] + if !ok { return "" } + return meta.defaultAPIBase } diff --git a/pkg/providers/factory_provider_test.go b/pkg/providers/factory_provider_test.go index ebd887220..588b81650 100644 --- a/pkg/providers/factory_provider_test.go +++ b/pkg/providers/factory_provider_test.go @@ -180,132 +180,82 @@ func TestCreateProviderFromConfig_LiteLLM(t *testing.T) { } } -func TestCreateProviderFromConfig_LMStudio_WithAPIKey(t *testing.T) { - cfg := &config.ModelConfig{ - ModelName: "test-lmstudio", - Model: "lmstudio/openai/gpt-oss-20b", - } - cfg.SetAPIKey("test-key") - - provider, modelID, err := CreateProviderFromConfig(cfg) - if err != nil { - t.Fatalf("CreateProviderFromConfig() error = %v", err) - } - if provider == nil { - t.Fatal("CreateProviderFromConfig() returned nil provider") - } - if modelID != "openai/gpt-oss-20b" { - t.Errorf("modelID = %q, want %q", modelID, "openai/gpt-oss-20b") - } - if _, ok := provider.(*HTTPProvider); !ok { - t.Fatalf("expected *HTTPProvider, got %T", provider) - } -} - -func TestCreateProviderFromConfig_LMStudio_NoAPIKey(t *testing.T) { - cfg := &config.ModelConfig{ - ModelName: "test-lmstudio", - Model: "lmstudio/openai/gpt-oss-20b", +func TestCreateProviderFromConfig_LocalProviders(t *testing.T) { + tests := []struct { + name string + modelName string + model string + apiKey string + wantModelID string + }{ + { + name: "LMStudio with API key", + modelName: "test-lmstudio", + model: "lmstudio/openai/gpt-oss-20b", + apiKey: "test-key", + wantModelID: "openai/gpt-oss-20b", + }, + { + name: "LMStudio without API key", + modelName: "test-lmstudio", + model: "lmstudio/openai/gpt-oss-20b", + apiKey: "", + wantModelID: "openai/gpt-oss-20b", + }, + { + name: "Ollama with API key", + modelName: "test-ollama", + model: "ollama/llama3.1:8b", + apiKey: "test-key", + wantModelID: "llama3.1:8b", + }, + { + name: "Ollama without API key", + modelName: "test-ollama", + model: "ollama/llama3.1:8b", + apiKey: "", + wantModelID: "llama3.1:8b", + }, + { + name: "VLLM with API key", + modelName: "test-vllm", + model: "vllm/Qwen/Qwen3-8B", + apiKey: "test-key", + wantModelID: "Qwen/Qwen3-8B", + }, + { + name: "VLLM without API key", + modelName: "test-vllm", + model: "vllm/Qwen/Qwen3-8B", + apiKey: "", + wantModelID: "Qwen/Qwen3-8B", + }, } - provider, modelID, err := CreateProviderFromConfig(cfg) - if err != nil { - t.Fatalf("CreateProviderFromConfig() error = %v", err) - } - if provider == nil { - t.Fatal("CreateProviderFromConfig() returned nil provider") - } - if modelID != "openai/gpt-oss-20b" { - t.Errorf("modelID = %q, want %q", modelID, "openai/gpt-oss-20b") - } - if _, ok := provider.(*HTTPProvider); !ok { - t.Fatalf("expected *HTTPProvider, got %T", provider) - } -} + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + cfg := &config.ModelConfig{ + ModelName: tt.modelName, + Model: tt.model, + } + if tt.apiKey != "" { + cfg.SetAPIKey(tt.apiKey) + } -func TestCreateProviderFromConfig_Ollama_WithAPIKey(t *testing.T) { - cfg := &config.ModelConfig{ - ModelName: "test-ollama", - Model: "ollama/llama3.1:8b", - } - cfg.SetAPIKey("test-key") - - provider, modelID, err := CreateProviderFromConfig(cfg) - if err != nil { - t.Fatalf("CreateProviderFromConfig() error = %v", err) - } - if provider == nil { - t.Fatal("CreateProviderFromConfig() returned nil provider") - } - if modelID != "llama3.1:8b" { - t.Errorf("modelID = %q, want %q", modelID, "llama3.1:8b") - } - if _, ok := provider.(*HTTPProvider); !ok { - t.Fatalf("expected *HTTPProvider, got %T", provider) - } -} - -func TestCreateProviderFromConfig_Ollama_NoAPIKey(t *testing.T) { - cfg := &config.ModelConfig{ - ModelName: "test-ollama", - Model: "ollama/llama3.1:8b", - } - - provider, modelID, err := CreateProviderFromConfig(cfg) - if err != nil { - t.Fatalf("CreateProviderFromConfig() error = %v", err) - } - if provider == nil { - t.Fatal("CreateProviderFromConfig() returned nil provider") - } - if modelID != "llama3.1:8b" { - t.Errorf("modelID = %q, want %q", modelID, "llama3.1:8b") - } - if _, ok := provider.(*HTTPProvider); !ok { - t.Fatalf("expected *HTTPProvider, got %T", provider) - } -} - -func TestCreateProviderFromConfig_VLLM_WithAPIKey(t *testing.T) { - cfg := &config.ModelConfig{ - ModelName: "test-vllm", - Model: "vllm/Qwen/Qwen3-8B", - } - cfg.SetAPIKey("test-key") - - provider, modelID, err := CreateProviderFromConfig(cfg) - if err != nil { - t.Fatalf("CreateProviderFromConfig() error = %v", err) - } - if provider == nil { - t.Fatal("CreateProviderFromConfig() returned nil provider") - } - if modelID != "Qwen/Qwen3-8B" { - t.Errorf("modelID = %q, want %q", modelID, "Qwen/Qwen3-8B") - } - if _, ok := provider.(*HTTPProvider); !ok { - t.Fatalf("expected *HTTPProvider, got %T", provider) - } -} - -func TestCreateProviderFromConfig_VLLM_NoAPIKey(t *testing.T) { - cfg := &config.ModelConfig{ - ModelName: "test-vllm", - Model: "vllm/Qwen/Qwen3-8B", - } - - provider, modelID, err := CreateProviderFromConfig(cfg) - if err != nil { - t.Fatalf("CreateProviderFromConfig() error = %v", err) - } - if provider == nil { - t.Fatal("CreateProviderFromConfig() returned nil provider") - } - if modelID != "Qwen/Qwen3-8B" { - t.Errorf("modelID = %q, want %q", modelID, "Qwen/Qwen3-8B") - } - if _, ok := provider.(*HTTPProvider); !ok { - t.Fatalf("expected *HTTPProvider, got %T", provider) + provider, modelID, err := CreateProviderFromConfig(cfg) + if err != nil { + t.Fatalf("CreateProviderFromConfig() error = %v", err) + } + if provider == nil { + t.Fatal("CreateProviderFromConfig() returned nil provider") + } + if modelID != tt.wantModelID { + t.Errorf("modelID = %q, want %q", modelID, tt.wantModelID) + } + if _, ok := provider.(*HTTPProvider); !ok { + t.Fatalf("expected *HTTPProvider, got %T", provider) + } + }) } }