fix model parse

This commit is contained in:
jackkav 2026-02-23 21:08:29 +01:00
parent ae74fa3812
commit e0cb4072f5
4 changed files with 77 additions and 5 deletions

2
.gitignore vendored
View file

@ -44,3 +44,5 @@ tasks/
# Added by goreleaser init: # Added by goreleaser init:
dist/ dist/
.gocache
.gomodcache

View file

@ -60,6 +60,12 @@ func TestExtractProtocol(t *testing.T) {
wantProtocol: "nvidia", wantProtocol: "nvidia",
wantModelID: "meta/llama-3.1-8b", 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 { 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) { func TestCreateProviderFromConfig_DefaultAPIBase(t *testing.T) {
tests := []struct { tests := []struct {
name string name string

View file

@ -18,13 +18,16 @@ func ParseModelRef(raw string, defaultProvider string) *ModelRef {
} }
if idx := strings.Index(raw, "/"); idx > 0 { if idx := strings.Index(raw, "/"); idx > 0 {
provider := NormalizeProvider(raw[:idx]) prefix := strings.TrimSpace(raw[:idx])
if isKnownProviderPrefix(prefix) {
provider := NormalizeProvider(prefix)
model := strings.TrimSpace(raw[idx+1:]) model := strings.TrimSpace(raw[idx+1:])
if model == "" { if model == "" {
return nil return nil
} }
return &ModelRef{Provider: provider, Model: model} return &ModelRef{Provider: provider, Model: model}
} }
}
return &ModelRef{ return &ModelRef{
Provider: NormalizeProvider(defaultProvider), Provider: NormalizeProvider(defaultProvider),
@ -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. // NormalizeProvider normalizes provider identifiers to canonical form.
func NormalizeProvider(provider string) string { func NormalizeProvider(provider string) string {
p := strings.ToLower(strings.TrimSpace(provider)) p := strings.ToLower(strings.TrimSpace(provider))

View file

@ -123,3 +123,29 @@ func TestParseModelRef_DefaultProviderNormalization(t *testing.T) {
t.Errorf("provider = %q, want openai (normalized from GPT)", ref.Provider) 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)
}
}