Files
aichat/config/config.go
T
2026-06-17 11:29:51 +08:00

368 lines
9.8 KiB
Go

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
}