package config import ( "fmt" "os" "strings" searchagent "aichat/agents/search" "gopkg.in/yaml.v3" ) const ( defaultOpenAIBaseURL = "https://ark.cn-beijing.volces.com/api/v3" defaultOpenAITimeout = 120 defaultToolRouterTimeout = 30 defaultToolRouterMaxTokens = 512 defaultToolRouterSystemText = `你可以按需直接调用可用工具来回答用户问题。 如果用户问题包含今天、今日、明天、昨天、本周、本月、本年、最近等相对时间,且后续需要搜索或查询数据库,应先调用 time 获取绝对日期范围。 需要实时网页资料、新闻、当前版本、近期事件、网页核验或用户明确要求联网时,调用 search。 需要查询本地业务数据、日程、会议、待办、记录、统计或时间范围内数据时,调用 sql。 工具结果优先于模型内置知识;工具失败时必须如实说明,不要编造结果。 只调用确实必要的工具。` ) type OpenAIConfig struct { Name string `yaml:"name" json:"name"` Active bool `yaml:"active,omitempty" json:"active"` APIKey string `yaml:"api_key" json:"-"` BaseURL string `yaml:"base_url" json:"base_url"` Model string `yaml:"model" json:"model"` Timeout int `yaml:"timeout" json:"timeout"` ParseThinkTags *bool `yaml:"parse_think_tags,omitempty" json:"parse_think_tags,omitempty"` } type OpenAIConfigs []OpenAIConfig type ToolRouterConfig struct { Enabled bool `yaml:"enabled" json:"enabled"` OpenAIName string `yaml:"openai_name" json:"openai_name"` Timeout int `yaml:"timeout" json:"timeout"` MaxTokens int `yaml:"max_tokens" json:"max_tokens"` SystemPrompt string `yaml:"system_prompt" json:"system_prompt"` Tools []ToolRouteConfig `yaml:"tools" json:"tools"` } type ToolRouteConfig struct { Name string `yaml:"name" json:"name"` Enabled bool `yaml:"enabled" json:"enabled"` Description string `yaml:"description" json:"description"` } func (configs *OpenAIConfigs) UnmarshalYAML(value *yaml.Node) error { switch value.Kind { case yaml.SequenceNode: var items []OpenAIConfig if err := value.Decode(&items); err != nil { return err } *configs = items case yaml.MappingNode: var item OpenAIConfig if err := value.Decode(&item); err != nil { return err } *configs = []OpenAIConfig{item} case yaml.ScalarNode: if value.Tag == "!!null" { *configs = nil return nil } return fmt.Errorf("openai 配置格式无效") default: return fmt.Errorf("openai 配置格式无效") } return nil } type Config struct { Server struct { Mode string `yaml:"mode"` Address string `yaml:"address"` } `yaml:"server"` OpenAI OpenAIConfigs `yaml:"openai"` ToolRouter ToolRouterConfig `yaml:"tool_router"` } func defaultOpenAIConfig() OpenAIConfig { return OpenAIConfig{ Name: "default", Active: true, BaseURL: defaultOpenAIBaseURL, Timeout: defaultOpenAITimeout, } } func DefaultToolRouterConfig() ToolRouterConfig { return ToolRouterConfig{ Enabled: true, OpenAIName: "", Timeout: defaultToolRouterTimeout, MaxTokens: defaultToolRouterMaxTokens, SystemPrompt: defaultToolRouterSystemText, Tools: []ToolRouteConfig{ {Name: "time", Enabled: true, Description: ""}, {Name: "search", Enabled: true, Description: ""}, {Name: "sql", Enabled: true, Description: ""}, }, } } func Default() Config { var cfg Config cfg.Server.Mode = "tcp" cfg.Server.Address = "0.0.0.0:8080" cfg.OpenAI = OpenAIConfigs{defaultOpenAIConfig()} cfg.ToolRouter = DefaultToolRouterConfig() return cfg } func Load(path string) (*Config, []searchagent.ProfileConfig, error) { if err := ensureFile(path); err != nil { return nil, nil, err } data, err := os.ReadFile(path) if err != nil { return nil, nil, fmt.Errorf("读取配置文件失败: %w", err) } var cfg Config if err = yaml.Unmarshal(data, &cfg); err != nil { return nil, nil, fmt.Errorf("解析配置文件失败: %w", err) } if _, err := normalizeOpenAIConfigs(&cfg); err != nil { return nil, nil, err } // 环境变量优先 if key := os.Getenv("ARK_API_KEY"); key != "" { for i := range cfg.OpenAI { cfg.OpenAI[i].APIKey = key } } legacySearchProfiles := readLegacySearchProfiles(data) if _, err := normalizeToolRouterConfig(&cfg); err != nil { return nil, nil, err } return &cfg, legacySearchProfiles, nil } func ensureFile(path string) error { defaults := Default() if _, err := os.Stat(path); err != nil { if !os.IsNotExist(err) { return fmt.Errorf("检查配置文件失败: %w", err) } return Write(path, defaults) } data, err := os.ReadFile(path) if err != nil { return fmt.Errorf("读取配置文件失败: %w", err) } var cfg Config if err = yaml.Unmarshal(data, &cfg); err != nil { return fmt.Errorf("解析配置文件失败: %w", err) } var raw map[string]any if err = yaml.Unmarshal(data, &raw); err != nil { return fmt.Errorf("解析配置文件失败: %w", err) } changed := false server, _ := raw["server"].(map[string]any) if server == nil { cfg.Server = defaults.Server changed = true } else { if _, ok := server["mode"]; !ok { cfg.Server.Mode = defaults.Server.Mode changed = true } if _, ok := server["address"]; !ok { cfg.Server.Address = defaults.Server.Address changed = true } } if _, ok := raw["openai"].([]any); !ok { changed = true } if normalized, err := normalizeOpenAIConfigs(&cfg); err != nil { return err } else if normalized { changed = true } if _, ok := raw["tool_router"]; !ok { cfg.ToolRouter = defaults.ToolRouter changed = true } else if normalized, err := normalizeToolRouterConfig(&cfg); err != nil { return err } else if normalized { changed = true } if !changed { return nil } return Write(path, cfg) } func normalizeOpenAIConfigs(cfg *Config) (bool, error) { changed := false if len(cfg.OpenAI) == 0 { cfg.OpenAI = OpenAIConfigs{defaultOpenAIConfig()} changed = true } activeIndex := -1 seen := map[string]bool{} for i := range cfg.OpenAI { profile := &cfg.OpenAI[i] name := strings.TrimSpace(profile.Name) if name == "" { name = strings.TrimSpace(profile.Model) if name == "" { name = fmt.Sprintf("openai-%d", i+1) } profile.Name = name changed = true } else if name != profile.Name { profile.Name = name changed = true } if seen[name] { return changed, fmt.Errorf("openai 配置名称重复: %s", name) } seen[name] = true if strings.TrimSpace(profile.BaseURL) == "" { profile.BaseURL = defaultOpenAIBaseURL changed = true } if profile.Timeout <= 0 { profile.Timeout = defaultOpenAITimeout changed = true } if profile.Active { if activeIndex == -1 { activeIndex = i } else { profile.Active = false changed = true } } } if activeIndex == -1 { cfg.OpenAI[0].Active = true changed = true } return changed, nil } func isLegacyToolRouterPrompt(prompt string) bool { prompt = strings.TrimSpace(prompt) return strings.Contains(prompt, "工具路由器") || strings.Contains(prompt, "route_tools") || strings.Contains(prompt, `"tools":[`) } func NormalizeToolRouterConfig(cfg *Config) (bool, error) { return normalizeToolRouterConfig(cfg) } func normalizeToolRouterConfig(cfg *Config) (bool, error) { changed := false defaults := DefaultToolRouterConfig() cfg.ToolRouter.OpenAIName = strings.TrimSpace(cfg.ToolRouter.OpenAIName) if cfg.ToolRouter.Timeout <= 0 { cfg.ToolRouter.Timeout = defaultToolRouterTimeout changed = true } if cfg.ToolRouter.MaxTokens <= 0 { cfg.ToolRouter.MaxTokens = defaultToolRouterMaxTokens changed = true } systemPrompt := strings.TrimSpace(cfg.ToolRouter.SystemPrompt) if systemPrompt == "" || isLegacyToolRouterPrompt(systemPrompt) { cfg.ToolRouter.SystemPrompt = defaultToolRouterSystemText changed = true } else if systemPrompt != cfg.ToolRouter.SystemPrompt { cfg.ToolRouter.SystemPrompt = systemPrompt changed = true } if len(cfg.ToolRouter.Tools) == 0 { cfg.ToolRouter.Tools = defaults.Tools changed = true } seen := map[string]bool{} for i := range cfg.ToolRouter.Tools { tool := &cfg.ToolRouter.Tools[i] name := strings.ToLower(strings.TrimSpace(tool.Name)) if name == "" { name = fmt.Sprintf("tool-%d", i+1) } if name != tool.Name { tool.Name = name changed = true } tool.Description = strings.TrimSpace(tool.Description) if seen[name] { return changed, fmt.Errorf("tool_router.tools 配置名称重复: %s", name) } seen[name] = true } byName := map[string]ToolRouteConfig{} for _, tool := range cfg.ToolRouter.Tools { byName[tool.Name] = tool } merged := make([]ToolRouteConfig, 0, len(cfg.ToolRouter.Tools)+len(defaults.Tools)) used := map[string]bool{} for _, tool := range defaults.Tools { if existing, ok := byName[tool.Name]; ok { merged = append(merged, existing) } else { merged = append(merged, tool) changed = true } used[tool.Name] = true } for _, tool := range cfg.ToolRouter.Tools { if !used[tool.Name] { merged = append(merged, tool) } } if len(merged) != len(cfg.ToolRouter.Tools) { changed = true } else { for i := range merged { if merged[i].Name != cfg.ToolRouter.Tools[i].Name { changed = true break } } } cfg.ToolRouter.Tools = merged return changed, nil } func readLegacySearchProfiles(data []byte) []searchagent.ProfileConfig { var legacy struct { Search searchagent.ProfileConfigs `yaml:"search"` } if err := yaml.Unmarshal(data, &legacy); err != nil { return nil } return []searchagent.ProfileConfig(legacy.Search) } func Write(path string, cfg Config) error { data, err := yaml.Marshal(&cfg) if err != nil { return fmt.Errorf("生成配置文件失败: %w", err) } if err := os.WriteFile(path, data, 0644); err != nil { return fmt.Errorf("写入配置文件失败: %w", err) } return nil }