Files
2026-06-17 13:13:06 +08:00

334 lines
8.8 KiB
Go

package config
import (
"fmt"
"os"
"strings"
"gopkg.in/yaml.v3"
)
const (
defaultOpenAIBaseURL = "https://ark.cn-beijing.volces.com/api/v3"
defaultOpenAITimeout = 120
defaultContextWindowTokens = 262144
defaultToolRouterTimeout = 30
defaultToolRouterMaxTokens = 512
defaultToolRouterSystemText = `你可以按需直接调用可用工具来回答用户问题。
每个工具的 description 描述了它的适用场景和调用条件。
工具结果优先于模型内置知识;工具失败时必须如实说明,不要编造结果。
只调用确实必要的工具。`
)
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"`
ContextWindowTokens int `yaml:"context_window_tokens" json:"context_window_tokens"`
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,omitempty" json:"tools,omitempty"`
}
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,
ContextWindowTokens: defaultContextWindowTokens,
}
}
func DefaultToolRouterConfig() ToolRouterConfig {
return ToolRouterConfig{
Enabled: true,
OpenAIName: "",
Timeout: defaultToolRouterTimeout,
MaxTokens: defaultToolRouterMaxTokens,
SystemPrompt: defaultToolRouterSystemText,
}
}
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, []map[string]any, 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) {
return normalizeOpenAIConfigs(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.ContextWindowTokens <= 0 {
profile.ContextWindowTokens = defaultContextWindowTokens
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
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
}
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
}
return changed, nil
}
func readLegacySearchProfiles(data []byte) []map[string]any {
var legacy struct {
Search []map[string]any `yaml:"search"`
}
if err := yaml.Unmarshal(data, &legacy); err != nil {
return nil
}
return 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
}