From 301d2f8bf9347d8ff28b2bd8f81cb327e0073249 Mon Sep 17 00:00:00 2001 From: XYSK-lilong007 <267018309+XYSK-lilong007@users.noreply.github.com> Date: Thu, 12 Mar 2026 12:56:38 +0800 Subject: [PATCH] fix(provider): preserve @cf model identifiers --- pkg/providers/factory_provider.go | 3 +++ pkg/providers/factory_provider_test.go | 26 ++++++++++++++++++++++++++ pkg/providers/model_ref.go | 7 +++++++ pkg/providers/model_ref_test.go | 13 +++++++++++++ 4 files changed, 49 insertions(+) diff --git a/pkg/providers/factory_provider.go b/pkg/providers/factory_provider.go index a798154cb..c269ae664 100644 --- a/pkg/providers/factory_provider.go +++ b/pkg/providers/factory_provider.go @@ -44,6 +44,9 @@ func createCodexAuthProvider() (LLMProvider, error) { // - "gpt-4o" -> ("openai", "gpt-4o") // default protocol func ExtractProtocol(model string) (protocol, modelID string) { model = strings.TrimSpace(model) + if strings.HasPrefix(model, "@") { + return "openai", model + } protocol, modelID, found := strings.Cut(model, "/") if !found { return "openai", model diff --git a/pkg/providers/factory_provider_test.go b/pkg/providers/factory_provider_test.go index 17bc55d25..ac250ec4d 100644 --- a/pkg/providers/factory_provider_test.go +++ b/pkg/providers/factory_provider_test.go @@ -64,6 +64,12 @@ func TestExtractProtocol(t *testing.T) { wantProtocol: "nvidia", wantModelID: "meta/llama-3.1-8b", }, + { + name: "cloudflare model id keeps full path", + model: "@cf/qwen/qwen1.5-0.5b-chat", + wantProtocol: "openai", + wantModelID: "@cf/qwen/qwen1.5-0.5b-chat", + }, } for _, tt := range tests { @@ -99,6 +105,26 @@ func TestCreateProviderFromConfig_OpenAI(t *testing.T) { } } +func TestCreateProviderFromConfig_CloudflareModelID(t *testing.T) { + cfg := &config.ModelConfig{ + ModelName: "cf-qwen", + Model: "@cf/qwen/qwen1.5-0.5b-chat", + APIKey: "test-key", + APIBase: "https://api.cloudflare.com/client/v4/accounts/test/ai/v1", + } + + 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 != "@cf/qwen/qwen1.5-0.5b-chat" { + t.Fatalf("modelID = %q, want %q", modelID, "@cf/qwen/qwen1.5-0.5b-chat") + } +} + 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..69320a8c7 100644 --- a/pkg/providers/model_ref.go +++ b/pkg/providers/model_ref.go @@ -17,6 +17,13 @@ func ParseModelRef(raw string, defaultProvider string) *ModelRef { return nil } + if strings.HasPrefix(raw, "@") { + return &ModelRef{ + Provider: NormalizeProvider(defaultProvider), + Model: raw, + } + } + if idx := strings.Index(raw, "/"); idx > 0 { provider := NormalizeProvider(raw[:idx]) model := strings.TrimSpace(raw[idx+1:]) diff --git a/pkg/providers/model_ref_test.go b/pkg/providers/model_ref_test.go index 6dd25167f..c846df20f 100644 --- a/pkg/providers/model_ref_test.go +++ b/pkg/providers/model_ref_test.go @@ -123,3 +123,16 @@ func TestParseModelRef_DefaultProviderNormalization(t *testing.T) { t.Errorf("provider = %q, want openai (normalized from GPT)", ref.Provider) } } + +func TestParseModelRef_CloudflareModelPathUsesDefaultProvider(t *testing.T) { + ref := ParseModelRef("@cf/qwen/qwen1.5-0.5b-chat", "openai") + if ref == nil { + t.Fatal("expected non-nil ref") + } + if ref.Provider != "openai" { + t.Fatalf("provider = %q, want openai", ref.Provider) + } + if ref.Model != "@cf/qwen/qwen1.5-0.5b-chat" { + t.Fatalf("model = %q, want %q", ref.Model, "@cf/qwen/qwen1.5-0.5b-chat") + } +}