system_prompt 提取到独立文件,移除无用的 port 配置
This commit is contained in:
+36
-17
@@ -15,9 +15,11 @@ import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
const (
|
const (
|
||||||
configDir = "data"
|
configDir = "data"
|
||||||
configFile = "config.yaml"
|
configFile = "config.yaml"
|
||||||
toolConfigDir = "data/tools"
|
toolConfigDir = "data/tools"
|
||||||
|
systemPromptFile = "data/system_prompt.md"
|
||||||
|
defaultSystemPrompt = "你是一个乐于助人的 AI 助手。"
|
||||||
)
|
)
|
||||||
|
|
||||||
type ModelConfig struct {
|
type ModelConfig struct {
|
||||||
@@ -61,9 +63,8 @@ type Provider struct {
|
|||||||
|
|
||||||
type Config struct {
|
type Config struct {
|
||||||
BotName string `yaml:"bot_name"`
|
BotName string `yaml:"bot_name"`
|
||||||
Port int `yaml:"port"`
|
|
||||||
LogLevel string `yaml:"log_level"`
|
LogLevel string `yaml:"log_level"`
|
||||||
SystemPrompt string `yaml:"system_prompt"`
|
SystemPrompt string `yaml:"-"` // 来源:data/system_prompt.md,不写入 yaml
|
||||||
Providers []Provider `yaml:"providers"`
|
Providers []Provider `yaml:"providers"`
|
||||||
DefaultProvider string `yaml:"default_provider"`
|
DefaultProvider string `yaml:"default_provider"`
|
||||||
DefaultModel string `yaml:"default_model"`
|
DefaultModel string `yaml:"default_model"`
|
||||||
@@ -126,9 +127,37 @@ func Load() (*Config, error) {
|
|||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
prompt, err := LoadSystemPrompt()
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
cfg.SystemPrompt = prompt
|
||||||
return cfg, nil
|
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 {
|
func GetConfig() *Config {
|
||||||
return cfg
|
return cfg
|
||||||
}
|
}
|
||||||
@@ -164,18 +193,10 @@ func applyDefaults(c *Config) (changed bool) {
|
|||||||
c.BotName = "ai-bot"
|
c.BotName = "ai-bot"
|
||||||
changed = true
|
changed = true
|
||||||
}
|
}
|
||||||
if c.Port == 0 {
|
|
||||||
c.Port = 8080
|
|
||||||
changed = true
|
|
||||||
}
|
|
||||||
if c.LogLevel == "" {
|
if c.LogLevel == "" {
|
||||||
c.LogLevel = "info"
|
c.LogLevel = "info"
|
||||||
changed = true
|
changed = true
|
||||||
}
|
}
|
||||||
if c.SystemPrompt == "" {
|
|
||||||
c.SystemPrompt = "你是一个乐于助人的 AI 助手。"
|
|
||||||
changed = true
|
|
||||||
}
|
|
||||||
if c.DefaultProvider == "" && len(c.Providers) > 0 {
|
if c.DefaultProvider == "" && len(c.Providers) > 0 {
|
||||||
c.DefaultProvider = c.Providers[0].Name
|
c.DefaultProvider = c.Providers[0].Name
|
||||||
changed = true
|
changed = true
|
||||||
@@ -360,10 +381,8 @@ func contains(list []string, s string) bool {
|
|||||||
|
|
||||||
func writeDefault(path string) error {
|
func writeDefault(path string) error {
|
||||||
cfg = &Config{
|
cfg = &Config{
|
||||||
BotName: "ai-bot",
|
BotName: "ai-bot",
|
||||||
Port: 8080,
|
LogLevel: "info",
|
||||||
LogLevel: "info",
|
|
||||||
SystemPrompt: "你是一个乐于助人的 AI 助手。",
|
|
||||||
Providers: []Provider{
|
Providers: []Provider{
|
||||||
{
|
{
|
||||||
Name: "openai",
|
Name: "openai",
|
||||||
|
|||||||
@@ -24,7 +24,6 @@ func TestApplyDatabaseDefaults(t *testing.T) {
|
|||||||
func TestApplyDefaultsNoChange(t *testing.T) {
|
func TestApplyDefaultsNoChange(t *testing.T) {
|
||||||
c := &Config{
|
c := &Config{
|
||||||
BotName: "x",
|
BotName: "x",
|
||||||
Port: 8081,
|
|
||||||
LogLevel: "debug",
|
LogLevel: "debug",
|
||||||
SystemPrompt: "sp",
|
SystemPrompt: "sp",
|
||||||
DefaultProvider: "p",
|
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) {
|
func TestValidateDatabase(t *testing.T) {
|
||||||
c := &Config{
|
c := &Config{
|
||||||
DefaultProvider: "p",
|
DefaultProvider: "p",
|
||||||
|
|||||||
Reference in New Issue
Block a user