@@ -0,0 +1,367 @@
|
||||
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
|
||||
}
|
||||
Reference in New Issue
Block a user