334 lines
8.8 KiB
Go
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
|
|
}
|