368 lines
9.8 KiB
Go
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
|
|
}
|