启动时检查配置缺失字段自动补全并写回
This commit is contained in:
@@ -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 {
|
||||
|
||||
@@ -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",
|
||||
|
||||
Reference in New Issue
Block a user