From e0cb4072f5dd30d23317e218ac2915a8158216af Mon Sep 17 00:00:00 2001 From: jackkav Date: Mon, 23 Feb 2026 21:08:29 +0100 Subject: [PATCH] fix model parse --- .gitignore | 2 ++ pkg/providers/factory_provider_test.go | 28 ++++++++++++++++++++++++++ pkg/providers/model_ref.go | 26 +++++++++++++++++++----- pkg/providers/model_ref_test.go | 26 ++++++++++++++++++++++++ 4 files changed, 77 insertions(+), 5 deletions(-) diff --git a/.gitignore b/.gitignore index ce30d749e..c584c5d45 100644 --- a/.gitignore +++ b/.gitignore @@ -44,3 +44,5 @@ tasks/ # Added by goreleaser init: dist/ +.gocache +.gomodcache diff --git a/pkg/providers/factory_provider_test.go b/pkg/providers/factory_provider_test.go index 6b133101a..a75253e83 100644 --- a/pkg/providers/factory_provider_test.go +++ b/pkg/providers/factory_provider_test.go @@ -60,6 +60,12 @@ func TestExtractProtocol(t *testing.T) { wantProtocol: "nvidia", wantModelID: "meta/llama-3.1-8b", }, + { + name: "openrouter nested model path with suffix", + model: "openrouter/stepfun/step-3.5-flash:free", + wantProtocol: "openrouter", + wantModelID: "stepfun/step-3.5-flash:free", + }, } for _, tt := range tests { @@ -95,6 +101,28 @@ func TestCreateProviderFromConfig_OpenAI(t *testing.T) { } } +func TestCreateProviderFromConfig_OpenRouterNestedModelPath(t *testing.T) { + cfg := &config.ModelConfig{ + ModelName: "test-openrouter-step", + Model: "openrouter/stepfun/step-3.5-flash:free", + APIKey: "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 _, ok := provider.(*HTTPProvider); !ok { + t.Fatalf("expected *HTTPProvider, got %T", provider) + } + if modelID != "stepfun/step-3.5-flash:free" { + t.Errorf("modelID = %q, want %q", modelID, "stepfun/step-3.5-flash:free") + } +} + func TestCreateProviderFromConfig_DefaultAPIBase(t *testing.T) { tests := []struct { name string diff --git a/pkg/providers/model_ref.go b/pkg/providers/model_ref.go index 0d1b02d16..d0e6e51a0 100644 --- a/pkg/providers/model_ref.go +++ b/pkg/providers/model_ref.go @@ -18,12 +18,15 @@ func ParseModelRef(raw string, defaultProvider string) *ModelRef { } if idx := strings.Index(raw, "/"); idx > 0 { - provider := NormalizeProvider(raw[:idx]) - model := strings.TrimSpace(raw[idx+1:]) - if model == "" { - return nil + prefix := strings.TrimSpace(raw[:idx]) + if isKnownProviderPrefix(prefix) { + provider := NormalizeProvider(prefix) + model := strings.TrimSpace(raw[idx+1:]) + if model == "" { + return nil + } + return &ModelRef{Provider: provider, Model: model} } - return &ModelRef{Provider: provider, Model: model} } return &ModelRef{ @@ -32,6 +35,19 @@ func ParseModelRef(raw string, defaultProvider string) *ModelRef { } } +func isKnownProviderPrefix(prefix string) bool { + switch NormalizeProvider(prefix) { + case "openai", "anthropic", "openrouter", "groq", "zhipu", "gemini", + "nvidia", "ollama", "moonshot", "shengsuanyun", "deepseek", "cerebras", + "volcengine", "vllm", "qwen-portal", "mistral", "antigravity", + "claude-cli", "claudecli", "codex-cli", "codexcli", "github-copilot", + "github_copilot", "copilot", "zai", "opencode", "kimi-coding": + return true + default: + return false + } +} + // NormalizeProvider normalizes provider identifiers to canonical form. func NormalizeProvider(provider string) string { p := strings.ToLower(strings.TrimSpace(provider)) diff --git a/pkg/providers/model_ref_test.go b/pkg/providers/model_ref_test.go index 6dd25167f..df00e2975 100644 --- a/pkg/providers/model_ref_test.go +++ b/pkg/providers/model_ref_test.go @@ -123,3 +123,29 @@ func TestParseModelRef_DefaultProviderNormalization(t *testing.T) { t.Errorf("provider = %q, want openai (normalized from GPT)", ref.Provider) } } + +func TestParseModelRef_OpenRouterNestedModelWithExplicitProvider(t *testing.T) { + ref := ParseModelRef("openrouter/stepfun/step-3.5-flash:free", "openai") + if ref == nil { + t.Fatal("expected non-nil ref") + } + if ref.Provider != "openrouter" { + t.Errorf("provider = %q, want openrouter", ref.Provider) + } + if ref.Model != "stepfun/step-3.5-flash:free" { + t.Errorf("model = %q, want stepfun/step-3.5-flash:free", ref.Model) + } +} + +func TestParseModelRef_UnknownPrefixUsesDefaultProvider(t *testing.T) { + ref := ParseModelRef("stepfun/step-3.5-flash:free", "openrouter") + if ref == nil { + t.Fatal("expected non-nil ref") + } + if ref.Provider != "openrouter" { + t.Errorf("provider = %q, want openrouter", ref.Provider) + } + if ref.Model != "stepfun/step-3.5-flash:free" { + t.Errorf("model = %q, want stepfun/step-3.5-flash:free", ref.Model) + } +}