diff --git a/internal/config/config.go b/internal/config/config.go index 20b611c..7a5831e 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -82,10 +82,15 @@ func Load() (*Config, error) { return nil, err } } - applyDefaults(cfg) + changed := applyDefaults(cfg) if err := validate(cfg); err != nil { return nil, err } + if changed { + if err := writeFile(path, cfg); err != nil { + return nil, err + } + } return cfg, nil } @@ -119,36 +124,46 @@ func migrateLegacy(path string, data []byte) error { return writeFile(path, cfg) } -func applyDefaults(c *Config) { +func applyDefaults(c *Config) (changed bool) { if c.BotName == "" { c.BotName = "ai-bot" + changed = true } if c.Port == 0 { c.Port = 8080 + changed = true } if c.LogLevel == "" { c.LogLevel = "info" + changed = true } if c.SystemPrompt == "" { c.SystemPrompt = "你是一个乐于助人的 AI 助手。" + changed = true } if c.DefaultProvider == "" && len(c.Providers) > 0 { c.DefaultProvider = c.Providers[0].Name + changed = true } if c.Database.Driver == "" { c.Database.Driver = "sqlite3" + changed = true } if c.Database.File == "" { c.Database.File = "data/memory.db" + changed = true } if c.Database.Driver == "mysql" { if c.Database.Host == "" { c.Database.Host = "127.0.0.1" + changed = true } if c.Database.Port == 0 { c.Database.Port = 3306 + changed = true } } + return changed } func validate(c *Config) error { diff --git a/internal/config/config_test.go b/internal/config/config_test.go index 307a4ac..a7b88bf 100644 --- a/internal/config/config_test.go +++ b/internal/config/config_test.go @@ -1,10 +1,18 @@ package config -import "testing" +import ( + "os" + "path/filepath" + "strings" + "testing" +) func TestApplyDatabaseDefaults(t *testing.T) { c := &Config{Providers: []Provider{{Name: "p", BaseURL: "x", Models: []string{"m"}}}} - applyDefaults(c) + changed := applyDefaults(c) + if !changed { + t.Error("缺失字段应返回 changed=true") + } if c.Database.Driver != "sqlite3" { t.Errorf("默认 driver = %q, want sqlite3", c.Database.Driver) } @@ -13,12 +21,30 @@ func TestApplyDatabaseDefaults(t *testing.T) { } } +func TestApplyDefaultsNoChange(t *testing.T) { + c := &Config{ + BotName: "x", + Port: 8081, + LogLevel: "debug", + SystemPrompt: "sp", + DefaultProvider: "p", + Providers: []Provider{{Name: "p", BaseURL: "x", Models: []string{"m"}}}, + Database: DatabaseConfig{Driver: "mysql", File: "f", Host: "h", Port: 3307, Name: "n"}, + } + if applyDefaults(c) { + t.Error("完整配置不应返回 changed") + } +} + func TestApplyMySQLDefaults(t *testing.T) { c := &Config{ Providers: []Provider{{Name: "p", BaseURL: "x", Models: []string{"m"}}}, Database: DatabaseConfig{Driver: "mysql", Name: "memory"}, } - applyDefaults(c) + changed := applyDefaults(c) + if !changed { + t.Error("mysql 缺 host/port 应返回 changed") + } if c.Database.Host != "127.0.0.1" { t.Errorf("默认 host = %q, want 127.0.0.1", c.Database.Host) } @@ -27,6 +53,66 @@ func TestApplyMySQLDefaults(t *testing.T) { } } +func TestLoadWritesBackDatabase(t *testing.T) { + t.Chdir(t.TempDir()) + cfg = nil + path := filepath.Join("data", "config.yaml") + old := "bot_name: test-bot\ndefault_provider: deepseek\ndefault_model: deepseek-v4-flash\nproviders:\n - name: deepseek\n base_url: https://api.deepseek.com\n models:\n - deepseek-v4-flash\n" + if err := os.MkdirAll("data", 0o755); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(path, []byte(old), 0o644); err != nil { + t.Fatal(err) + } + c, err := Load() + if err != nil { + t.Fatalf("Load 出错: %v", err) + } + if c.Database.Driver != "sqlite3" { + t.Errorf("默认 driver = %q, want sqlite3", c.Database.Driver) + } + data, err := os.ReadFile(path) + if err != nil { + t.Fatalf("读取写回后的配置失败: %v", err) + } + if !strings.Contains(string(data), "database:") || !strings.Contains(string(data), "sqlite3") { + t.Errorf("写回的文件缺少 database 段:\n%s", data) + } + if !strings.Contains(string(data), "test-bot") { + t.Errorf("写回不应覆盖已有字段:\n%s", data) + } +} + +func TestLoadNoRewriteWhenComplete(t *testing.T) { + t.Chdir(t.TempDir()) + cfg = nil + path := filepath.Join("data", "config.yaml") + full := "bot_name: test-bot\ndefault_provider: deepseek\ndefault_model: deepseek-v4-flash\nproviders:\n - name: deepseek\n base_url: https://api.deepseek.com\n models:\n - deepseek-v4-flash\n" + if err := os.MkdirAll("data", 0o755); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(path, []byte(full), 0o644); err != nil { + t.Fatal(err) + } + // 完整配置(含 database 段) + complete := full + "database:\n driver: sqlite3\n file: data/memory.db\n" + if err := os.WriteFile(path, []byte(complete), 0o644); err != nil { + t.Fatal(err) + } + if _, err := Load(); err != nil { + t.Fatalf("Load 出错: %v", err) + } + info, _ := os.Stat(path) + before := info.ModTime() + if _, err := Load(); err != nil { + t.Fatalf("第二次 Load 出错: %v", err) + } + info, _ = os.Stat(path) + if !info.ModTime().Equal(before) { + t.Error("完整配置不应重复写回文件") + } +} + func TestValidateDatabase(t *testing.T) { c := &Config{ DefaultProvider: "p",