fix model parse
This commit is contained in:
parent
ae74fa3812
commit
e0cb4072f5
4 changed files with 77 additions and 5 deletions
2
.gitignore
vendored
2
.gitignore
vendored
|
|
@ -44,3 +44,5 @@ tasks/
|
||||||
|
|
||||||
# Added by goreleaser init:
|
# Added by goreleaser init:
|
||||||
dist/
|
dist/
|
||||||
|
.gocache
|
||||||
|
.gomodcache
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -18,12 +18,15 @@ 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])
|
||||||
model := strings.TrimSpace(raw[idx+1:])
|
if isKnownProviderPrefix(prefix) {
|
||||||
if model == "" {
|
provider := NormalizeProvider(prefix)
|
||||||
return nil
|
model := strings.TrimSpace(raw[idx+1:])
|
||||||
|
if model == "" {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return &ModelRef{Provider: provider, Model: model}
|
||||||
}
|
}
|
||||||
return &ModelRef{Provider: provider, Model: model}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
return &ModelRef{
|
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.
|
// 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))
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue