From dfc649f08e74e38de3cb79286f9b328fd0881526 Mon Sep 17 00:00:00 2001 From: kevin Date: Fri, 14 Aug 2026 20:30:48 +0800 Subject: [PATCH] =?UTF-8?q?system=5Fprompt=20=E6=8F=90=E5=8F=96=E5=88=B0?= =?UTF-8?q?=E7=8B=AC=E7=AB=8B=E6=96=87=E4=BB=B6=EF=BC=8C=E7=A7=BB=E9=99=A4?= =?UTF-8?q?=E6=97=A0=E7=94=A8=E7=9A=84=20port=20=E9=85=8D=E7=BD=AE?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- internal/config/config.go | 53 ++++++++++++++++++++---------- internal/config/config_test.go | 59 +++++++++++++++++++++++++++++++++- 2 files changed, 94 insertions(+), 18 deletions(-) diff --git a/internal/config/config.go b/internal/config/config.go index 69925bf..146da7e 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -15,9 +15,11 @@ import ( ) const ( - configDir = "data" - configFile = "config.yaml" - toolConfigDir = "data/tools" + configDir = "data" + configFile = "config.yaml" + toolConfigDir = "data/tools" + systemPromptFile = "data/system_prompt.md" + defaultSystemPrompt = "你是一个乐于助人的 AI 助手。" ) type ModelConfig struct { @@ -61,9 +63,8 @@ type Provider struct { type Config struct { BotName string `yaml:"bot_name"` - Port int `yaml:"port"` LogLevel string `yaml:"log_level"` - SystemPrompt string `yaml:"system_prompt"` + SystemPrompt string `yaml:"-"` // 来源:data/system_prompt.md,不写入 yaml Providers []Provider `yaml:"providers"` DefaultProvider string `yaml:"default_provider"` DefaultModel string `yaml:"default_model"` @@ -126,9 +127,37 @@ func Load() (*Config, error) { return nil, err } } + prompt, err := LoadSystemPrompt() + if err != nil { + return nil, err + } + cfg.SystemPrompt = prompt return cfg, nil } +// LoadSystemPrompt 读取 data/system_prompt.md;文件不存在时创建默认内容。 +func LoadSystemPrompt() (string, error) { + data, err := os.ReadFile(systemPromptFile) + if err == nil { + return string(data), nil + } + if !os.IsNotExist(err) { + return "", fmt.Errorf("读取系统提示词失败: %w", err) + } + if err := os.MkdirAll(configDir, 0o755); err != nil { + return "", err + } + if err := os.WriteFile(systemPromptFile, []byte(defaultSystemPrompt), 0o644); err != nil { + return "", fmt.Errorf("创建系统提示词文件失败: %w", err) + } + return defaultSystemPrompt, nil +} + +// SystemPromptPath 返回系统提示词文件的路径。 +func SystemPromptPath() string { + return systemPromptFile +} + func GetConfig() *Config { return cfg } @@ -164,18 +193,10 @@ func applyDefaults(c *Config) (changed bool) { 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 @@ -360,10 +381,8 @@ func contains(list []string, s string) bool { func writeDefault(path string) error { cfg = &Config{ - BotName: "ai-bot", - Port: 8080, - LogLevel: "info", - SystemPrompt: "你是一个乐于助人的 AI 助手。", + BotName: "ai-bot", + LogLevel: "info", Providers: []Provider{ { Name: "openai", diff --git a/internal/config/config_test.go b/internal/config/config_test.go index 24b8977..5bd2331 100644 --- a/internal/config/config_test.go +++ b/internal/config/config_test.go @@ -24,7 +24,6 @@ func TestApplyDatabaseDefaults(t *testing.T) { func TestApplyDefaultsNoChange(t *testing.T) { c := &Config{ BotName: "x", - Port: 8081, LogLevel: "debug", SystemPrompt: "sp", DefaultProvider: "p", @@ -113,6 +112,64 @@ func TestLoadNoRewriteWhenComplete(t *testing.T) { } } +func TestLoadSystemPromptCreatesDefault(t *testing.T) { + t.Chdir(t.TempDir()) + prompt, err := LoadSystemPrompt() + if err != nil { + t.Fatalf("LoadSystemPrompt 出错: %v", err) + } + if prompt != defaultSystemPrompt { + t.Errorf("默认提示词 = %q, want %q", prompt, defaultSystemPrompt) + } + data, err := os.ReadFile(systemPromptFile) + if err != nil { + t.Fatalf("默认文件未创建: %v", err) + } + if string(data) != defaultSystemPrompt { + t.Errorf("文件内容 = %q", string(data)) + } +} + +func TestLoadSystemPromptReadsExisting(t *testing.T) { + t.Chdir(t.TempDir()) + if err := os.MkdirAll("data", 0o755); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(systemPromptFile, []byte("自定义提示词\n多行"), 0o644); err != nil { + t.Fatal(err) + } + prompt, err := LoadSystemPrompt() + if err != nil { + t.Fatalf("LoadSystemPrompt 出错: %v", err) + } + if prompt != "自定义提示词\n多行" { + t.Errorf("应返回文件内容: %q", prompt) + } +} + +func TestLoadWritesBackExcludesSystemPrompt(t *testing.T) { + t.Chdir(t.TempDir()) + cfg = nil + path := filepath.Join("data", "config.yaml") + old := "bot_name: test-bot\nsystem_prompt: 旧值\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) + } + if _, err := Load(); err != nil { + t.Fatalf("Load 出错: %v", err) + } + data, err := os.ReadFile(path) + if err != nil { + t.Fatal(err) + } + if strings.Contains(string(data), "system_prompt") { + t.Errorf("写回的配置不应包含 system_prompt:\n%s", data) + } +} + func TestValidateDatabase(t *testing.T) { c := &Config{ DefaultProvider: "p",