From a9738e6ee94808f5cada4cb41fadfc5a816e78b7 Mon Sep 17 00:00:00 2001 From: duomi Date: Fri, 13 Mar 2026 23:52:48 +0800 Subject: [PATCH] fix(config): migrate legacy provider-only configs --- pkg/config/config.go | 14 +++++++ pkg/config/config_test.go | 78 +++++++++++++++++++++++++++++++++++ pkg/providers/factory_test.go | 33 +++++++++++++++ 3 files changed, 125 insertions(+) diff --git a/pkg/config/config.go b/pkg/config/config.go index 190341224..1083fe93b 100644 --- a/pkg/config/config.go +++ b/pkg/config/config.go @@ -832,6 +832,16 @@ func LoadConfig(path string) (*Config, error) { if len(tmp.ModelList) > 0 { cfg.ModelList = nil } + if len(tmp.ModelList) == 0 && tmp.HasProvidersConfig() { + // Legacy provider-only configs should not inherit the template + // model_list or template default model selection. Clear both so the + // old provider block can be migrated into an explicit model_list below. + cfg.ModelList = nil + if tmp.Agents.Defaults.ModelName == "" && tmp.Agents.Defaults.Model == "" { + cfg.Agents.Defaults.ModelName = "" + cfg.Agents.Defaults.Model = "" + } + } if err := json.Unmarshal(data, cfg); err != nil { return nil, err @@ -847,6 +857,10 @@ func LoadConfig(path string) (*Config, error) { // Auto-migrate: if only legacy providers config exists, convert to model_list if len(cfg.ModelList) == 0 && cfg.HasProvidersConfig() { cfg.ModelList = ConvertProvidersToModelList(cfg) + if cfg.Agents.Defaults.GetModelName() == "" && len(cfg.ModelList) > 0 { + cfg.Agents.Defaults.ModelName = cfg.ModelList[0].ModelName + cfg.Agents.Defaults.Model = "" + } } // Validate model_list for uniqueness and required fields diff --git a/pkg/config/config_test.go b/pkg/config/config_test.go index c5bdbf3c3..d6650d7c3 100644 --- a/pkg/config/config_test.go +++ b/pkg/config/config_test.go @@ -460,6 +460,84 @@ func TestLoadConfig_WebToolsProxy(t *testing.T) { } } +func TestLoadConfig_LegacyProvidersConfigMigratesWithoutTemplateModelLeak(t *testing.T) { + tmpDir := t.TempDir() + configPath := filepath.Join(tmpDir, "config.json") + configJSON := `{ + "providers": { + "gemini": { + "api_key": "gemini-test-key" + } + } +}` + if err := os.WriteFile(configPath, []byte(configJSON), 0o600); err != nil { + t.Fatalf("os.WriteFile() error: %v", err) + } + + cfg, err := LoadConfig(configPath) + if err != nil { + t.Fatalf("LoadConfig() error: %v", err) + } + + if got := cfg.Agents.Defaults.GetModelName(); got != "gemini" { + t.Fatalf("GetModelName() = %q, want %q", got, "gemini") + } + if len(cfg.ModelList) != 1 { + t.Fatalf("len(ModelList) = %d, want 1", len(cfg.ModelList)) + } + if cfg.ModelList[0].ModelName != "gemini" { + t.Fatalf("ModelList[0].ModelName = %q, want %q", cfg.ModelList[0].ModelName, "gemini") + } + if cfg.ModelList[0].Model != "gemini/gemini-pro" { + t.Fatalf("ModelList[0].Model = %q, want %q", cfg.ModelList[0].Model, "gemini/gemini-pro") + } + + modelCfg, err := cfg.GetModelConfig(cfg.Agents.Defaults.GetModelName()) + if err != nil { + t.Fatalf("GetModelConfig() error = %v", err) + } + if modelCfg.Model != "gemini/gemini-pro" { + t.Fatalf("modelCfg.Model = %q, want %q", modelCfg.Model, "gemini/gemini-pro") + } +} + +func TestLoadConfig_LegacyProvidersConfigKeepsExplicitLegacyModel(t *testing.T) { + tmpDir := t.TempDir() + configPath := filepath.Join(tmpDir, "config.json") + configJSON := `{ + "agents": { + "defaults": { + "model": "gemini-pro" + } + }, + "providers": { + "gemini": { + "api_key": "gemini-test-key" + } + } +}` + if err := os.WriteFile(configPath, []byte(configJSON), 0o600); err != nil { + t.Fatalf("os.WriteFile() error: %v", err) + } + + cfg, err := LoadConfig(configPath) + if err != nil { + t.Fatalf("LoadConfig() error: %v", err) + } + + if got := cfg.Agents.Defaults.GetModelName(); got != "gemini-pro" { + t.Fatalf("GetModelName() = %q, want %q", got, "gemini-pro") + } + + modelCfg, err := cfg.GetModelConfig(cfg.Agents.Defaults.GetModelName()) + if err != nil { + t.Fatalf("GetModelConfig() error = %v", err) + } + if modelCfg.Model != "gemini/gemini-pro" { + t.Fatalf("modelCfg.Model = %q, want %q", modelCfg.Model, "gemini/gemini-pro") + } +} + // TestDefaultConfig_DMScope verifies the default dm_scope value // TestDefaultConfig_SummarizationThresholds verifies summarization defaults func TestDefaultConfig_SummarizationThresholds(t *testing.T) { diff --git a/pkg/providers/factory_test.go b/pkg/providers/factory_test.go index 91469f25b..89b7c4edc 100644 --- a/pkg/providers/factory_test.go +++ b/pkg/providers/factory_test.go @@ -1,6 +1,8 @@ package providers import ( + "os" + "path/filepath" "strings" "testing" @@ -268,6 +270,37 @@ func TestCreateProviderReturnsHTTPProviderForOpenRouter(t *testing.T) { } } +func TestCreateProviderWithLegacyGeminiProvidersConfig(t *testing.T) { + tmpDir := t.TempDir() + configPath := filepath.Join(tmpDir, "config.json") + configJSON := `{ + "providers": { + "gemini": { + "api_key": "gemini-test-key" + } + } +}` + if err := os.WriteFile(configPath, []byte(configJSON), 0o600); err != nil { + t.Fatalf("WriteFile() error = %v", err) + } + + cfg, err := config.LoadConfig(configPath) + if err != nil { + t.Fatalf("LoadConfig() error = %v", err) + } + + provider, modelID, err := CreateProvider(cfg) + if err != nil { + t.Fatalf("CreateProvider() error = %v", err) + } + if _, ok := provider.(*HTTPProvider); !ok { + t.Fatalf("provider type = %T, want *HTTPProvider", provider) + } + if modelID != "gemini-pro" { + t.Fatalf("modelID = %q, want %q", modelID, "gemini-pro") + } +} + func TestCreateProviderReturnsCodexCliProviderForCodexCode(t *testing.T) { cfg := config.DefaultConfig() cfg.Agents.Defaults.Model = "test-codex"