diff --git a/completion/completion.go b/completion/completion.go new file mode 100644 index 0000000..07019c4 --- /dev/null +++ b/completion/completion.go @@ -0,0 +1,84 @@ +package completion + +import ( + "context" + "errors" + "io" + "strings" + "time" + + "aichat/llm" + "aichat/message" + "aichat/stream" + "aichat/utils" + + "github.com/volcengine/volcengine-go-sdk/service/arkruntime/model" +) + +type ChatCompleter func(context.Context, *llm.Profile, model.CreateChatCompletionRequest, time.Duration) (model.ChatCompletionResponse, error) + +func CompleteText(ctx context.Context, profile *llm.Profile, chatMessages []message.ChatMessage, maxTokens int) (string, error) { + return CompleteTextWithTimeout(ctx, profile, chatMessages, maxTokens, time.Duration(profile.Config.Timeout)*time.Second) +} + +func CompleteTextWithTimeout(ctx context.Context, profile *llm.Profile, chatMessages []message.ChatMessage, maxTokens int, timeout time.Duration) (string, error) { + messages, err := message.BuildArkMessages(chatMessages) + if err != nil { + return "", err + } + completionCtx, cancel := context.WithTimeout(ctx, timeout) + defer cancel() + streamResp, err := profile.Client.CreateChatCompletionStream(completionCtx, model.CreateChatCompletionRequest{ + Model: profile.Config.Model, + Messages: messages, + MaxTokens: utils.IntPtr(maxTokens), + }.WithStream(true)) + if err != nil { + return "", err + } + defer streamResp.Close() + + promptTokens := stream.EstimateChatMessagesTokens(chatMessages) + completionTokens := 0 + parseThinkTags := llm.ShouldParseThinkTags(profile) + thinkParser := &stream.Parser{} + var b strings.Builder + appendVisible := func(delta string) { + if delta == "" { + return + } + b.WriteString(delta) + completionTokens += stream.EstimateTokenCount(delta) + } + for { + resp, err := streamResp.Recv() + if errors.Is(err, io.EOF) { + if parseThinkTags { + visible, _ := thinkParser.Flush() + appendVisible(visible) + } + if tracker := stream.TrackerFromContext(ctx); tracker != nil { + tracker.AddTool(promptTokens, completionTokens) + } + return b.String(), nil + } + if err != nil { + return "", err + } + if len(resp.Choices) > 0 { + delta := resp.Choices[0].Delta.Content + if parseThinkTags { + visible, _ := thinkParser.Accept(delta) + appendVisible(visible) + } else { + appendVisible(delta) + } + } + } +} + +func CompleteChatWithTimeout(ctx context.Context, profile *llm.Profile, request model.CreateChatCompletionRequest, timeout time.Duration) (model.ChatCompletionResponse, error) { + completionCtx, cancel := context.WithTimeout(ctx, timeout) + defer cancel() + return profile.Client.CreateChatCompletion(completionCtx, request.WithStream(false)) +} diff --git a/config/config.go b/config/config.go new file mode 100644 index 0000000..6354834 --- /dev/null +++ b/config/config.go @@ -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 +} diff --git a/conversation/store.go b/conversation/store.go new file mode 100644 index 0000000..a5162a4 --- /dev/null +++ b/conversation/store.go @@ -0,0 +1,152 @@ +package conversation + +import ( + "encoding/json" + "errors" + "fmt" + "os" + "path/filepath" + "sort" + "strings" + "sync" + "time" + + "aichat/message" + "aichat/utils" +) + +type Store struct { + dir string + mu sync.Mutex +} + +func NewStore(dir string) *Store { + os.MkdirAll(dir, 0755) + return &Store{dir: dir} +} + +func (s *Store) path(id string) string { + return filepath.Join(s.dir, id+".json") +} + +func (s *Store) Create() (*message.Conversation, error) { + conv := &message.Conversation{ + ID: utils.NewUUID(), + Title: "新对话", + CreatedAt: time.Now(), + UpdatedAt: time.Now(), + } + if err := s.Save(conv); err != nil { + return nil, err + } + return conv, nil +} + +func (s *Store) Save(conv *message.Conversation) error { + s.mu.Lock() + defer s.mu.Unlock() + conv.UpdatedAt = time.Now() + return atomicWriteJSON(s.path(conv.ID), conv) +} + +func (s *Store) Get(id string) (*message.Conversation, error) { + s.mu.Lock() + defer s.mu.Unlock() + data, err := os.ReadFile(s.path(id)) + if err != nil { + if os.IsNotExist(err) { + return nil, errors.New("对话不存在") + } + return nil, fmt.Errorf("读取对话失败: %w", err) + } + var conv message.Conversation + if err := json.Unmarshal(data, &conv); err != nil { + return nil, fmt.Errorf("解析对话失败: %w", err) + } + return &conv, nil +} + +func (s *Store) List() ([]message.Conversation, error) { + s.mu.Lock() + defer s.mu.Unlock() + + entries, err := os.ReadDir(s.dir) + if err != nil { + return nil, fmt.Errorf("读取对话目录失败: %w", err) + } + + var list []message.Conversation + for _, e := range entries { + if e.IsDir() || filepath.Ext(e.Name()) != ".json" { + continue + } + data, err := os.ReadFile(filepath.Join(s.dir, e.Name())) + if err != nil { + continue + } + var conv message.Conversation + if err := json.Unmarshal(data, &conv); err != nil { + continue + } + conv.Messages = nil // 列表不返回消息体 + list = append(list, conv) + } + + sort.Slice(list, func(i, j int) bool { + return list[i].UpdatedAt.After(list[j].UpdatedAt) + }) + return list, nil +} + +func (s *Store) Delete(id string) error { + s.mu.Lock() + defer s.mu.Unlock() + if err := os.Remove(s.path(id)); err != nil && !os.IsNotExist(err) { + return fmt.Errorf("删除对话失败: %w", err) + } + return nil +} + +func atomicWriteJSON(path string, v any) error { + tmp := path + ".tmp" + data, err := json.Marshal(v) + if err != nil { + return err + } + if err := os.WriteFile(tmp, data, 0644); err != nil { + return err + } + return os.Rename(tmp, path) +} + +func SaveMessages(store *Store, id string, messages []message.ChatMessage, assistantContent string) error { + conv, err := store.Get(id) + if err != nil { + return err + } + conv.Messages = append([]message.ChatMessage(nil), messages...) + conv.Messages = append(conv.Messages, message.ChatMessage{Role: "assistant", Content: assistantContent}) + if conv.Title == "" || conv.Title == "新对话" { + conv.Title = GenTitle(conv.Messages) + } + return store.Save(conv) +} + +func GenTitle(messages []message.ChatMessage) string { + for _, m := range messages { + if m.Hidden { + continue + } + if m.Role == "user" && strings.TrimSpace(m.Content) != "" { + title := strings.TrimSpace(m.Content) + title = strings.ReplaceAll(title, "\r\n", " ") + title = strings.ReplaceAll(title, "\n", " ") + runes := []rune(title) + if len(runes) > 30 { + return string(runes[:30]) + "..." + } + return title + } + } + return "新对话" +} diff --git a/llm/ollama.go b/llm/ollama.go new file mode 100644 index 0000000..8eaf2c8 --- /dev/null +++ b/llm/ollama.go @@ -0,0 +1,46 @@ +package llm + +import ( + "errors" + "net/url" + "strings" +) + +func IsOllamaProfile(profile *Profile) bool { + if profile == nil { + return false + } + u, err := url.Parse(strings.TrimSpace(profile.Config.BaseURL)) + if err != nil { + return strings.Contains(profile.Config.BaseURL, ":11434") + } + host := strings.ToLower(u.Hostname()) + port := u.Port() + return port == "11434" && (host == "127.0.0.1" || host == "localhost" || host == "::1") +} + +func ShouldParseThinkTags(profile *Profile) bool { + if profile == nil { + return false + } + if profile.Config.ParseThinkTags != nil { + return *profile.Config.ParseThinkTags + } + return IsOllamaProfile(profile) +} + +func OllamaBaseURL(profile *Profile) (string, error) { + if profile == nil { + return "", errors.New("Ollama 配置为空") + } + u, err := url.Parse(strings.TrimSpace(profile.Config.BaseURL)) + if err != nil { + return "", err + } + if strings.TrimRight(u.Path, "/") == "/v1" { + u.Path = strings.TrimSuffix(strings.TrimRight(u.Path, "/"), "/v1") + } + u.RawQuery = "" + u.Fragment = "" + return strings.TrimRight(u.String(), "/"), nil +} diff --git a/llm/state.go b/llm/state.go new file mode 100644 index 0000000..6ca858e --- /dev/null +++ b/llm/state.go @@ -0,0 +1,131 @@ +package llm + +import ( + "errors" + "fmt" + "strings" + "sync" + "time" + + "aichat/config" + + ark "github.com/volcengine/volcengine-go-sdk/service/arkruntime" +) + +type Profile struct { + Config config.OpenAIConfig + Client *ark.Client +} + +type State struct { + mu sync.RWMutex + profiles map[string]*Profile + order []string + activeName string +} + +type ListResponse struct { + Active string `json:"active"` + Profiles []config.OpenAIConfig `json:"profiles"` +} + +func NewState(configs []config.OpenAIConfig) (*State, error) { + state := &State{ + profiles: make(map[string]*Profile, len(configs)), + order: make([]string, 0, len(configs)), + } + for _, item := range configs { + if strings.TrimSpace(item.Name) == "" { + return nil, errors.New("openai.name 不能为空") + } + if strings.TrimSpace(item.APIKey) == "" { + return nil, fmt.Errorf("openai.%s.api_key 未配置,也未设置环境变量 ARK_API_KEY", item.Name) + } + if strings.TrimSpace(item.Model) == "" { + return nil, fmt.Errorf("openai.%s.model 未配置", item.Name) + } + if strings.TrimSpace(item.BaseURL) == "" { + return nil, fmt.Errorf("openai.%s.base_url 未配置", item.Name) + } + if item.Timeout <= 0 { + return nil, fmt.Errorf("openai.%s.timeout 必须大于 0", item.Name) + } + if _, ok := state.profiles[item.Name]; ok { + return nil, fmt.Errorf("openai 配置名称重复: %s", item.Name) + } + state.profiles[item.Name] = &Profile{ + Config: item, + Client: ark.NewClientWithApiKey( + item.APIKey, + ark.WithBaseUrl(item.BaseURL), + ark.WithTimeout(time.Duration(item.Timeout)*time.Second), + ), + } + state.order = append(state.order, item.Name) + if item.Active && state.activeName == "" { + state.activeName = item.Name + } + } + if len(state.order) == 0 { + return nil, errors.New("openai 配置不能为空") + } + if state.activeName == "" { + state.activeName = state.order[0] + } + return state, nil +} + +func (s *State) ActiveProfile() *Profile { + s.mu.RLock() + defer s.mu.RUnlock() + return s.profiles[s.activeName] +} + +func (s *State) GetProfile(name string) (*Profile, error) { + s.mu.RLock() + defer s.mu.RUnlock() + if strings.TrimSpace(name) == "" { + return s.profiles[s.activeName], nil + } + profile, ok := s.profiles[strings.TrimSpace(name)] + if !ok { + return nil, fmt.Errorf("OpenAI 配置不存在: %s", name) + } + return profile, nil +} + +func (s *State) SwitchActive(name string) (*Profile, error) { + name = strings.TrimSpace(name) + if name == "" { + return nil, errors.New("OpenAI 配置名称不能为空") + } + s.mu.Lock() + defer s.mu.Unlock() + profile, ok := s.profiles[name] + if !ok { + return nil, fmt.Errorf("OpenAI 配置不存在: %s", name) + } + s.activeName = name + return profile, nil +} + +func (s *State) ListProfiles() ListResponse { + s.mu.RLock() + defer s.mu.RUnlock() + profiles := make([]config.OpenAIConfig, 0, len(s.order)) + for _, name := range s.order { + profile := s.profiles[name] + cfg := profile.Config + cfg.APIKey = "" + cfg.Active = name == s.activeName + profiles = append(profiles, cfg) + } + return ListResponse{Active: s.activeName, Profiles: profiles} +} + +func PublicConfig(profile *Profile, active bool) config.OpenAIConfig { + cfg := profile.Config + cfg.APIKey = "" + cfg.Active = active + return cfg +} diff --git a/main.go b/main.go index d8e9c09..63dd7a9 100644 --- a/main.go +++ b/main.go @@ -1,1969 +1,28 @@ package main import ( - "bufio" - "bytes" - "context" - "crypto/rand" - "encoding/base64" - "encoding/hex" - "encoding/json" - "errors" "fmt" - "io" - "net" - "net/http" - "net/url" "os" - "path/filepath" - "sort" "strings" - "sync" - "time" - "unicode" searchagent "aichat/agents/search" sqlquery "aichat/agents/sql" - timeagent "aichat/agents/time" - - "github.com/gin-gonic/gin" - ark "github.com/volcengine/volcengine-go-sdk/service/arkruntime" - "github.com/volcengine/volcengine-go-sdk/service/arkruntime/model" - "gopkg.in/yaml.v3" + "aichat/config" + "aichat/conversation" + "aichat/llm" + "aichat/server" + "aichat/toolrouter" ) -// ─── 配置 ───────────────────────────────────────────────── - -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 defaultConfig() 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 loadConfig(path string) (*Config, error) { - if err := ensureConfigFile(path); err != nil { - return nil, err - } - - data, err := os.ReadFile(path) - if err != nil { - return nil, fmt.Errorf("读取配置文件失败: %w", err) - } - var cfg Config - if err = yaml.Unmarshal(data, &cfg); err != nil { - return nil, fmt.Errorf("解析配置文件失败: %w", err) - } - if _, err := normalizeOpenAIConfigs(&cfg); err != nil { - return 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, err - } - return &cfg, nil -} - -func ensureConfigFile(path string) error { - defaults := defaultConfig() - if _, err := os.Stat(path); err != nil { - if !os.IsNotExist(err) { - return fmt.Errorf("检查配置文件失败: %w", err) - } - return writeConfig(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 writeConfig(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) { - 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 writeConfig(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 -} - -// ─── 请求结构 ───────────────────────────────────────────── - -type ChatMessage struct { - Role string `json:"role"` - Content string `json:"content"` - ImageURL string `json:"image_url,omitempty"` // base64 data URI 或 http URL - ImageURLAlias string `json:"imageURL,omitempty"` - Hidden bool `json:"hidden,omitempty"` -} - -type ChatRequest struct { - ConversationID string `json:"conversation_id,omitempty"` - Messages []ChatMessage `json:"messages"` - WebSearch bool `json:"web_search,omitempty"` - OpenAIName string `json:"openai_name,omitempty"` -} - -type Conversation struct { - ID string `json:"id"` - Title string `json:"title"` - CreatedAt time.Time `json:"created_at"` - UpdatedAt time.Time `json:"updated_at"` - Messages []ChatMessage `json:"messages,omitempty"` -} - -type ConvStore struct { - dir string - mu sync.Mutex -} - -type OpenAIProfile struct { - Config OpenAIConfig - Client *ark.Client -} - -type OpenAIState struct { - mu sync.RWMutex - profiles map[string]*OpenAIProfile - order []string - activeName string -} - -type activeProfileRequest struct { - Name string `json:"name"` -} - -type openAIListResponse struct { - Active string `json:"active"` - Profiles []OpenAIConfig `json:"profiles"` -} - -type chatCompleter func(context.Context, *OpenAIProfile, model.CreateChatCompletionRequest, time.Duration) (model.ChatCompletionResponse, error) - -type ToolRouterState struct { - cfg *ToolRouterConfig - ai *OpenAIState - complete chatCompleter -} - -func NewToolRouterState(config *ToolRouterConfig, ai *OpenAIState) (*ToolRouterState, error) { - if config == nil { - cfg := defaultToolRouterConfig() - config = &cfg - } - if ai == nil { - return nil, errors.New("工具路由需要 OpenAI 状态") - } - if config.Enabled && strings.TrimSpace(config.OpenAIName) != "" { - if _, err := ai.GetProfile(config.OpenAIName); err != nil { - return nil, fmt.Errorf("tool_router.openai_name 配置无效: %w", err) - } - } - return &ToolRouterState{cfg: config, ai: ai, complete: completeChatWithTimeout}, nil -} - -func NewOpenAIState(configs []OpenAIConfig) (*OpenAIState, error) { - state := &OpenAIState{ - profiles: make(map[string]*OpenAIProfile, len(configs)), - order: make([]string, 0, len(configs)), - } - for _, config := range configs { - if strings.TrimSpace(config.Name) == "" { - return nil, errors.New("openai.name 不能为空") - } - if strings.TrimSpace(config.APIKey) == "" { - return nil, fmt.Errorf("openai.%s.api_key 未配置,也未设置环境变量 ARK_API_KEY", config.Name) - } - if strings.TrimSpace(config.Model) == "" { - return nil, fmt.Errorf("openai.%s.model 未配置", config.Name) - } - if strings.TrimSpace(config.BaseURL) == "" { - return nil, fmt.Errorf("openai.%s.base_url 未配置", config.Name) - } - if config.Timeout <= 0 { - return nil, fmt.Errorf("openai.%s.timeout 必须大于 0", config.Name) - } - if _, ok := state.profiles[config.Name]; ok { - return nil, fmt.Errorf("openai 配置名称重复: %s", config.Name) - } - state.profiles[config.Name] = &OpenAIProfile{ - Config: config, - Client: ark.NewClientWithApiKey( - config.APIKey, - ark.WithBaseUrl(config.BaseURL), - ark.WithTimeout(time.Duration(config.Timeout)*time.Second), - ), - } - state.order = append(state.order, config.Name) - if config.Active && state.activeName == "" { - state.activeName = config.Name - } - } - if len(state.order) == 0 { - return nil, errors.New("openai 配置不能为空") - } - if state.activeName == "" { - state.activeName = state.order[0] - } - return state, nil -} - -func (s *OpenAIState) ActiveProfile() *OpenAIProfile { - s.mu.RLock() - defer s.mu.RUnlock() - return s.profiles[s.activeName] -} - -func (s *OpenAIState) GetProfile(name string) (*OpenAIProfile, error) { - s.mu.RLock() - defer s.mu.RUnlock() - if strings.TrimSpace(name) == "" { - return s.profiles[s.activeName], nil - } - profile, ok := s.profiles[strings.TrimSpace(name)] - if !ok { - return nil, fmt.Errorf("OpenAI 配置不存在: %s", name) - } - return profile, nil -} - -func (s *OpenAIState) SwitchActive(name string) (*OpenAIProfile, error) { - name = strings.TrimSpace(name) - if name == "" { - return nil, errors.New("OpenAI 配置名称不能为空") - } - s.mu.Lock() - defer s.mu.Unlock() - profile, ok := s.profiles[name] - if !ok { - return nil, fmt.Errorf("OpenAI 配置不存在: %s", name) - } - s.activeName = name - return profile, nil -} - -func (s *OpenAIState) ListProfiles() openAIListResponse { - s.mu.RLock() - defer s.mu.RUnlock() - profiles := make([]OpenAIConfig, 0, len(s.order)) - for _, name := range s.order { - profile := s.profiles[name] - config := profile.Config - config.APIKey = "" - config.Active = name == s.activeName - profiles = append(profiles, config) - } - return openAIListResponse{Active: s.activeName, Profiles: profiles} -} - -func publicOpenAIConfig(profile *OpenAIProfile, active bool) OpenAIConfig { - config := profile.Config - config.APIKey = "" - config.Active = active - return config -} - -func (s *ToolRouterState) RouterProfile(fallback *OpenAIProfile) *OpenAIProfile { - if s == nil || s.cfg == nil || s.ai == nil { - return fallback - } - name := strings.TrimSpace(s.cfg.OpenAIName) - if name == "" { - return fallback - } - profile, err := s.ai.GetProfile(name) - if err != nil { - return fallback - } - return profile -} - -func isOllamaProfile(profile *OpenAIProfile) bool { - if profile == nil { - return false - } - u, err := url.Parse(strings.TrimSpace(profile.Config.BaseURL)) - if err != nil { - return strings.Contains(profile.Config.BaseURL, ":11434") - } - host := strings.ToLower(u.Hostname()) - port := u.Port() - return port == "11434" && (host == "127.0.0.1" || host == "localhost" || host == "::1") -} - -func shouldParseThinkTags(profile *OpenAIProfile) bool { - if profile == nil { - return false - } - if profile.Config.ParseThinkTags != nil { - return *profile.Config.ParseThinkTags - } - return isOllamaProfile(profile) -} - -// ─── 全局变量 ───────────────────────────────────────────── - -var ( - cfg *Config - aiState *OpenAIState - searchState *searchagent.State - legacySearchProfiles []searchagent.ProfileConfig - toolRouterState *ToolRouterState - sqlState *sqlquery.State - store *ConvStore -) - -type chatSSEFrame struct { - Type string `json:"type"` - Text string `json:"text,omitempty"` - Message string `json:"message,omitempty"` - Tool string `json:"tool,omitempty"` - Stage string `json:"stage,omitempty"` - Status string `json:"status,omitempty"` - Data map[string]any `json:"data,omitempty"` - Stats *tokenUsageStats `json:"stats,omitempty"` - Error string `json:"error,omitempty"` -} - -type tokenUsageStats struct { - PromptTokens int `json:"prompt_tokens"` - CompletionTokens int `json:"completion_tokens"` - ToolPromptTokens int `json:"tool_prompt_tokens"` - ToolCompletionTokens int `json:"tool_completion_tokens"` - TotalTokens int `json:"total_tokens"` - CompletionTokensPerSec float64 `json:"completion_tokens_per_sec"` - PeakCompletionTokensPerSec float64 `json:"peak_completion_tokens_per_sec"` - Estimated bool `json:"estimated"` -} - -type tokenUsageTracker struct { - mu sync.Mutex - promptTokens int - completionTokens int - toolPromptTokens int - toolCompletionTokens int -} - -type tokenUsageContextKey struct{} - -func newTokenUsageTracker() *tokenUsageTracker { - return &tokenUsageTracker{} -} - -func contextWithTokenUsage(ctx context.Context, tracker *tokenUsageTracker) context.Context { - if tracker == nil { - return ctx - } - return context.WithValue(ctx, tokenUsageContextKey{}, tracker) -} - -func tokenUsageFromContext(ctx context.Context) *tokenUsageTracker { - tracker, _ := ctx.Value(tokenUsageContextKey{}).(*tokenUsageTracker) - return tracker -} - -func (t *tokenUsageTracker) addTool(promptTokens, completionTokens int) { - if t == nil { - return - } - t.mu.Lock() - defer t.mu.Unlock() - t.toolPromptTokens += promptTokens - t.toolCompletionTokens += completionTokens -} - -func (t *tokenUsageTracker) setModel(promptTokens, completionTokens int) { - if t == nil { - return - } - t.mu.Lock() - defer t.mu.Unlock() - t.promptTokens = promptTokens - t.completionTokens = completionTokens -} - -func (t *tokenUsageTracker) snapshot(tokensPerSecond, peakTokensPerSecond float64) tokenUsageStats { - if t == nil { - return tokenUsageStats{Estimated: true} - } - t.mu.Lock() - defer t.mu.Unlock() - total := t.promptTokens + t.completionTokens + t.toolPromptTokens + t.toolCompletionTokens - return tokenUsageStats{ - PromptTokens: t.promptTokens, - CompletionTokens: t.completionTokens, - ToolPromptTokens: t.toolPromptTokens, - ToolCompletionTokens: t.toolCompletionTokens, - TotalTokens: total, - CompletionTokensPerSec: tokensPerSecond, - PeakCompletionTokensPerSec: peakTokensPerSecond, - Estimated: true, - } -} - -// ─── 路由 ───────────────────────────────────────────────── - -func indexHandler(c *gin.Context) { - profile := aiState.ActiveProfile() - c.HTML(http.StatusOK, "chat.html", gin.H{ - "Title": "AI 对话", - "Model": profile.Config.Model, - "OpenAIName": profile.Config.Name, - }) -} - -func listOpenAIHandler(c *gin.Context) { - c.JSON(http.StatusOK, aiState.ListProfiles()) -} - -func switchOpenAIHandler(c *gin.Context) { - var req activeProfileRequest - if err := c.ShouldBindJSON(&req); err != nil { - c.JSON(http.StatusBadRequest, gin.H{"error": "请求格式错误: " + err.Error()}) - return - } - profile, err := aiState.SwitchActive(req.Name) - if err != nil { - c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()}) - return - } - c.JSON(http.StatusOK, gin.H{ - "active": profile.Config.Name, - "profile": publicOpenAIConfig(profile, true), - }) -} - -func listSearchHandler(c *gin.Context) { - c.JSON(http.StatusOK, searchState.ListProfiles()) -} - -func switchSearchHandler(c *gin.Context) { - var req activeProfileRequest - if err := c.ShouldBindJSON(&req); err != nil { - c.JSON(http.StatusBadRequest, gin.H{"error": "请求格式错误: " + err.Error()}) - return - } - profile, err := searchState.SwitchActive(req.Name) - if err != nil { - c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()}) - return - } - profile.APIKey = "" - profile.Active = true - c.JSON(http.StatusOK, gin.H{ - "active": profile.Name, - "profile": profile, - }) -} - -func listConversationsHandler(c *gin.Context) { - convs, err := store.List() - if err != nil { - c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()}) - return - } - c.JSON(http.StatusOK, convs) -} - -func createConversationHandler(c *gin.Context) { - conv, err := store.Create() - if err != nil { - c.JSON(http.StatusInternalServerError, gin.H{"error": "创建对话失败: " + err.Error()}) - return - } - c.JSON(http.StatusOK, conv) -} - -func getConversationHandler(c *gin.Context) { - conv, err := store.Get(c.Param("id")) - if err != nil { - status := http.StatusInternalServerError - if err.Error() == "对话不存在" { - status = http.StatusNotFound - } - c.JSON(status, gin.H{"error": err.Error()}) - return - } - c.JSON(http.StatusOK, conv) -} - -func deleteConversationHandler(c *gin.Context) { - if err := store.Delete(c.Param("id")); err != nil { - c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()}) - return - } - c.Status(http.StatusNoContent) -} - -// chatHandler 流式 SSE 对话接口 -func chatHandler(c *gin.Context) { - var req ChatRequest - if err := c.ShouldBindJSON(&req); err != nil { - c.JSON(http.StatusBadRequest, gin.H{"error": "请求格式错误: " + err.Error()}) - return - } - if len(req.Messages) == 0 { - c.JSON(http.StatusBadRequest, gin.H{"error": "消息不能为空"}) - return - } - profile, err := aiState.GetProfile(req.OpenAIName) - if err != nil { - c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()}) - return - } - - // SSE 头先写出,后续插件/模型过程都通过 trace 事件实时展示。 - c.Writer.Header().Set("Content-Type", "text/event-stream") - c.Writer.Header().Set("Cache-Control", "no-cache") - c.Writer.Header().Set("Connection", "keep-alive") - c.Writer.Header().Set("X-Accel-Buffering", "no") - c.Writer.WriteHeader(http.StatusOK) - flusher, ok := c.Writer.(http.Flusher) - if !ok { - return - } - emit := func(frame chatSSEFrame) { - writeSSEJSON(c.Writer, frame) - flusher.Flush() - } - emitTrace := func(tool, stage, status, message string, data map[string]any) { - emit(chatSSEFrame{Type: "trace", Tool: tool, Stage: stage, Status: status, Message: message, Data: data}) - } - emitError := func(err error) { - emit(chatSSEFrame{Type: "error", Error: err.Error()}) - } - - // 超时 context - timeout := time.Duration(profile.Config.Timeout) * time.Second - ctx, cancel := context.WithTimeout(c.Request.Context(), timeout) - defer cancel() - usage := newTokenUsageTracker() - ctx = contextWithTokenUsage(ctx, usage) - - // 用 Function Calling 工具循环替代旧的路由+隐藏上下文机制 - messages, err := runAgentToolLoop(ctx, profile, req.Messages, emit) - if err != nil { - fmt.Fprintln(os.Stderr, "Agent 工具循环失败:", err) - messages, err = buildArkMessages(req.Messages) - if err != nil { - emitError(err) - return - } - } - promptTokens := estimateChatMessagesTokens(req.Messages) - - if isOllamaProfile(profile) && hasImageMessage(req.Messages) { - emitTrace("model", "request", "running", "正在通过 Ollama 原生接口调用视觉模型", nil) - err = streamOllamaChat(ctx, profile, messages, promptTokens, usage, emit, func(content string) { - if req.ConversationID != "" { - if err := saveConversationMessages(req.ConversationID, req.Messages, content); err != nil { - fmt.Fprintln(os.Stderr, "保存对话失败:", err) - } - } - }) - if err != nil { - emitError(err) - } - return - } - - emitTrace("model", "request", "running", "正在调用模型生成回答", nil) - stream, err := profile.Client.CreateChatCompletionStream(ctx, model.CreateChatCompletionRequest{ - Model: profile.Config.Model, - Messages: messages, - MaxTokens: intPtr(4096), - }.WithStream(true)) - if err != nil { - emitError(err) - return - } - defer stream.Close() - emitTrace("model", "stream", "running", "模型已开始输出", nil) - - var full strings.Builder - completionTokens := 0 - streamStarted := time.Now() - windowStarted := streamStarted - windowTokens := 0 - peakTokensPerSecond := 0.0 - parseThinkTags := shouldParseThinkTags(profile) - thinkParser := &thinkTagParser{} - emitDelta := func(delta string) { - if delta == "" { - return - } - now := time.Now() - deltaTokens := estimateTokenCount(delta) - windowTokens += deltaTokens - windowElapsed := now.Sub(windowStarted).Seconds() - if windowElapsed >= 1 { - windowSpeed := float64(windowTokens) / windowElapsed - if windowSpeed > peakTokensPerSecond { - peakTokensPerSecond = windowSpeed - } - windowStarted = now - windowTokens = 0 - } else if peakTokensPerSecond == 0 && windowElapsed > 0.25 { - peakTokensPerSecond = float64(windowTokens) / windowElapsed - } - full.WriteString(delta) - completionTokens += deltaTokens - usage.setModel(promptTokens, completionTokens) - stats := usage.snapshot(tokensPerSecond(completionTokens, streamStarted), peakTokensPerSecond) - emit(chatSSEFrame{Type: "delta", Text: delta, Stats: &stats}) - } - emitModelContent := func(delta string) { - if delta == "" { - return - } - if !parseThinkTags { - emitDelta(delta) - return - } - visible, reasoning := thinkParser.Accept(delta) - if reasoning != "" { - emit(chatSSEFrame{Type: "reasoning", Text: reasoning}) - } - emitDelta(visible) - } - for { - resp, err := stream.Recv() - if errors.Is(err, io.EOF) { - if parseThinkTags { - visible, reasoning := thinkParser.Flush() - if reasoning != "" { - emit(chatSSEFrame{Type: "reasoning", Text: reasoning}) - } - emitDelta(visible) - } - usage.setModel(promptTokens, completionTokens) - if windowTokens > 0 { - windowElapsed := time.Since(windowStarted).Seconds() - if windowElapsed > 0.25 { - windowSpeed := float64(windowTokens) / windowElapsed - if windowSpeed > peakTokensPerSecond { - peakTokensPerSecond = windowSpeed - } - } - } - if peakTokensPerSecond == 0 { - peakTokensPerSecond = tokensPerSecond(completionTokens, streamStarted) - } - if req.ConversationID != "" { - if err := saveConversationMessages(req.ConversationID, req.Messages, full.String()); err != nil { - fmt.Fprintln(os.Stderr, "保存对话失败:", err) - } - } - finalStats := usage.snapshot(tokensPerSecond(completionTokens, streamStarted), peakTokensPerSecond) - emit(chatSSEFrame{Type: "stats", Stats: &finalStats}) - emitTrace("model", "stream", "success", "回答生成完成", nil) - fmt.Fprintf(c.Writer, "data: [DONE]\n\n") - flusher.Flush() - return - } - if err != nil { - emitError(err) - return - } - if len(resp.Choices) > 0 { - emitModelContent(resp.Choices[0].Delta.Content) - // 思考过程 reasoning_content 单独事件推送 - if resp.Choices[0].Delta.ReasoningContent != nil && *resp.Choices[0].Delta.ReasoningContent != "" { - emit(chatSSEFrame{Type: "reasoning", Text: *resp.Choices[0].Delta.ReasoningContent}) - } - } - } -} - -// ─── 辅助函数 ───────────────────────────────────────────── - -func estimateChatMessagesTokens(messages []ChatMessage) int { - total := 0 - for _, msg := range messages { - total += estimateTokenCount(msg.Role) + estimateTokenCount(msg.Content) + 4 - if msg.ImageURL != "" || msg.ImageURLAlias != "" { - total += 85 - } - } - return total -} - -func estimateTokenCount(text string) int { - text = strings.TrimSpace(text) - if text == "" { - return 0 - } - tokens := 0 - asciiRunes := 0 - flushASCII := func() { - if asciiRunes > 0 { - tokens += (asciiRunes + 3) / 4 - asciiRunes = 0 - } - } - for _, r := range text { - if unicode.IsSpace(r) { - flushASCII() - continue - } - if r <= unicode.MaxASCII { - asciiRunes++ - continue - } - flushASCII() - tokens++ - } - flushASCII() - if tokens == 0 { - return 1 - } - return tokens -} - -func tokensPerSecond(tokens int, start time.Time) float64 { - elapsed := time.Since(start).Seconds() - if tokens <= 0 || elapsed <= 0 { - return 0 - } - return float64(tokens) / elapsed -} - -type agentTool struct { - name string - definition *model.Tool - execute func(context.Context, string) (string, error) -} - -func (t agentTool) Name() string { return t.name } - -const maxAgentToolIterations = 6 - -func availableAgentTools(profile *OpenAIProfile, emit func(chatSSEFrame)) []agentTool { - if toolRouterState == nil || toolRouterState.cfg == nil || !toolRouterState.cfg.Enabled { - return nil - } - tools := make([]agentTool, 0, len(toolRouterState.cfg.Tools)) - for _, item := range toolRouterState.cfg.Tools { - if !item.Enabled { - continue - } - description := strings.TrimSpace(item.Description) - switch item.Name { - case timeagent.ToolName: - tools = append(tools, agentTool{ - name: timeagent.ToolName, - definition: timeagent.ToolDefinition(description), - execute: func(ctx context.Context, args string) (string, error) { - result, err := timeagent.ExecuteTool(args, time.Now()) - if err == nil && emit != nil { - emit(chatSSEFrame{Type: "trace", Tool: timeagent.ToolName, Stage: "resolve", Status: "success", Message: "已获取当前时间上下文"}) - } - return result, err - }, - }) - case searchagent.ToolName: - if searchState == nil || !searchState.Enabled() { - continue - } - tools = append(tools, agentTool{ - name: searchagent.ToolName, - definition: searchState.ToolDefinition(description), - execute: func(ctx context.Context, args string) (string, error) { - if emit != nil { - emit(chatSSEFrame{Type: "trace", Tool: searchagent.ToolName, Stage: "request", Status: "running", Message: "正在联网搜索"}) - } - result, err := searchState.ExecuteTool(ctx, args) - if emit != nil { - status := "success" - message := "联网搜索完成" - if err != nil { - status = "error" - message = "联网搜索失败" - } - emit(chatSSEFrame{Type: "trace", Tool: searchagent.ToolName, Stage: "results", Status: status, Message: message}) - } - return result, err - }, - }) - case sqlquery.ToolName: - if sqlState == nil || !sqlState.Enabled() { - continue - } - tools = append(tools, agentTool{ - name: sqlquery.ToolName, - definition: sqlState.ToolDefinition(description), - execute: func(ctx context.Context, args string) (string, error) { - if emit != nil { - emit(chatSSEFrame{Type: "trace", Tool: sqlquery.ToolName, Stage: "execute", Status: "running", Message: "正在查询数据库"}) - } - generator := func(ctx context.Context, prompt string, maxTokens int) (string, error) { - return completeText(ctx, profile, []ChatMessage{{Role: "system", Content: prompt}}, maxTokens) - } - result, err := sqlState.ExecuteTool(ctx, args, generator) - if emit != nil { - status := "success" - message := "数据库查询完成" - if err != nil { - status = "error" - message = "数据库查询失败" - } - emit(chatSSEFrame{Type: "trace", Tool: sqlquery.ToolName, Stage: "execute", Status: status, Message: message}) - } - return result, err - }, - }) - } - } - return tools -} - -func runAgentToolLoop(ctx context.Context, profile *OpenAIProfile, chatMessages []ChatMessage, emit func(chatSSEFrame)) ([]*model.ChatCompletionMessage, error) { - finalMessages, err := buildArkMessages(chatMessages) - if err != nil { - return nil, err - } - routerProfile := profile - if toolRouterState != nil { - routerProfile = toolRouterState.RouterProfile(profile) - } - tools := availableAgentTools(routerProfile, emit) - if len(tools) == 0 { - return finalMessages, nil - } - decisionMessages := append([]*model.ChatCompletionMessage(nil), finalMessages...) - if hasImageMessage(chatMessages) { - decisionMessages, err = buildToolDecisionMessages(chatMessages) - if err != nil { - return nil, err - } - if emit != nil { - emit(chatSSEFrame{Type: "trace", Tool: "agent_tools", Stage: "prepare", Status: "success", Message: "检测到图片输入,工具判断阶段将使用纯文本上下文"}) - } - } - toolByName := make(map[string]agentTool, len(tools)) - definitions := make([]*model.Tool, 0, len(tools)) - availableNames := make([]string, 0, len(tools)) - toolDescriptions := make([]string, 0, len(tools)) - for _, tool := range tools { - toolByName[tool.name] = tool - definitions = append(definitions, tool.definition) - availableNames = append(availableNames, tool.name) - if tool.definition != nil && tool.definition.Function != nil { - toolDescriptions = append(toolDescriptions, fmt.Sprintf("%s: %s", tool.name, tool.definition.Function.Description)) - } - } - if emit != nil { - emit(chatSSEFrame{Type: "trace", Tool: "agent_tools", Stage: "prepare", Status: "success", Message: "已准备可用工具", Data: map[string]any{"tools": availableNames, "tool_descriptions": toolDescriptions}}) - } - if prompt := strings.TrimSpace(toolRouterState.cfg.SystemPrompt); prompt != "" { - systemMessage := &model.ChatCompletionMessage{Role: model.ChatMessageRoleSystem, Content: stringContent(prompt)} - finalMessages = append([]*model.ChatCompletionMessage{systemMessage}, finalMessages...) - decisionMessages = append([]*model.ChatCompletionMessage{systemMessage}, decisionMessages...) - } - for i := 0; i < maxAgentToolIterations; i++ { - if emit != nil { - emit(chatSSEFrame{Type: "trace", Tool: "agent_tools", Stage: "request", Status: "running", Message: fmt.Sprintf("正在进行第 %d 轮工具判断", i+1), Data: map[string]any{"iteration": i + 1, "max_iterations": maxAgentToolIterations, "tools": availableNames}}) - } - resp, err := toolRouterState.complete(ctx, routerProfile, model.CreateChatCompletionRequest{ - Model: routerProfile.Config.Model, - Messages: decisionMessages, - MaxTokens: intPtr(toolRouterState.cfg.MaxTokens), - Tools: definitions, - ToolChoice: model.ToolChoiceStringTypeAuto, - ParallelToolCalls: boolPtr(false), - }, time.Duration(toolRouterState.cfg.Timeout)*time.Second) - if err != nil { - return finalMessages, err - } - if tracker := tokenUsageFromContext(ctx); tracker != nil { - tracker.addTool(resp.Usage.PromptTokens, resp.Usage.CompletionTokens) - } - if len(resp.Choices) == 0 { - return finalMessages, nil - } - choice := resp.Choices[0] - decisionPreview := chatMessageContentString(choice.Message.Content) - if emit != nil { - emit(chatSSEFrame{Type: "trace", Tool: "agent_tools", Stage: "decision", Status: "success", Message: "工具判断响应已返回", Data: map[string]any{"iteration": i + 1, "finish_reason": string(choice.FinishReason), "content_preview": truncateString(decisionPreview, 800)}}) - } - calls := choice.Message.ToolCalls - if len(calls) == 0 && choice.Message.FunctionCall != nil { - calls = []*model.ToolCall{{ID: "legacy_function_call", Type: model.ToolTypeFunction, Function: *choice.Message.FunctionCall}} - } - if len(calls) == 0 { - if emit != nil { - emit(chatSSEFrame{Type: "trace", Tool: "agent_tools", Stage: "request", Status: "success", Message: "模型未请求工具,进入回答生成"}) - } - return finalMessages, nil - } - callNames := make([]string, 0, len(calls)) - for _, call := range calls { - if call != nil { - callNames = append(callNames, call.Function.Name) - } - } - if emit != nil { - emit(chatSSEFrame{Type: "trace", Tool: "agent_tools", Stage: "tool_calls", Status: "running", Message: fmt.Sprintf("模型请求调用 %d 个工具", len(calls)), Data: map[string]any{"tools": callNames, "iteration": i + 1}}) - } - assistantMessage := &model.ChatCompletionMessage{Role: model.ChatMessageRoleAssistant, ToolCalls: calls, Content: choice.Message.Content} - finalMessages = append(finalMessages, assistantMessage) - decisionMessages = append(decisionMessages, assistantMessage) - for _, call := range calls { - result := executeAgentToolCall(ctx, call, toolByName, emit) - toolMessage := &model.ChatCompletionMessage{Role: model.ChatMessageRoleTool, ToolCallID: call.ID, Content: stringContent(result)} - finalMessages = append(finalMessages, toolMessage) - decisionMessages = append(decisionMessages, toolMessage) - } - } - limitMessage := &model.ChatCompletionMessage{Role: model.ChatMessageRoleSystem, Content: stringContent("工具调用轮数已达到上限。请基于已有工具结果回答,并说明可能未完成全部工具调用。")} - finalMessages = append(finalMessages, limitMessage) - return finalMessages, nil -} - -type thinkTagParser struct { - inThink bool - buffer string -} - -const ( - thinkOpenTag = "" - thinkCloseTag = "" -) - -func (p *thinkTagParser) Accept(delta string) (visible string, reasoning string) { - p.buffer += delta - for p.buffer != "" { - if p.inThink { - idx := strings.Index(p.buffer, thinkCloseTag) - if idx >= 0 { - reasoning += p.buffer[:idx] - p.buffer = p.buffer[idx+len(thinkCloseTag):] - p.inThink = false - continue - } - keep := tagPrefixSuffixLen(p.buffer, thinkCloseTag) - if len(p.buffer) > keep { - reasoning += p.buffer[:len(p.buffer)-keep] - p.buffer = p.buffer[len(p.buffer)-keep:] - } - return visible, reasoning - } - - idx := strings.Index(p.buffer, thinkOpenTag) - if idx >= 0 { - visible += p.buffer[:idx] - p.buffer = p.buffer[idx+len(thinkOpenTag):] - p.inThink = true - continue - } - keep := tagPrefixSuffixLen(p.buffer, thinkOpenTag) - if len(p.buffer) > keep { - visible += p.buffer[:len(p.buffer)-keep] - p.buffer = p.buffer[len(p.buffer)-keep:] - } - return visible, reasoning - } - return visible, reasoning -} - -func (p *thinkTagParser) Flush() (visible string, reasoning string) { - if p.inThink { - reasoning = p.buffer - } else { - visible = p.buffer - } - p.buffer = "" - p.inThink = false - return visible, reasoning -} - -func tagPrefixSuffixLen(text, tag string) int { - limit := len(tag) - 1 - if len(text) < limit { - limit = len(text) - } - for i := limit; i > 0; i-- { - if strings.HasPrefix(tag, text[len(text)-i:]) { - return i - } - } - return 0 -} - -func executeAgentToolCall(ctx context.Context, call *model.ToolCall, tools map[string]agentTool, emit func(chatSSEFrame)) string { - if call == nil || call.Type != model.ToolTypeFunction { - result := "工具调用无效:仅支持 function 类型工具。" - if emit != nil { - emit(chatSSEFrame{Type: "trace", Tool: "agent_tools", Stage: "execute", Status: "error", Message: result}) - } - return result - } - toolName := call.Function.Name - if emit != nil { - emit(chatSSEFrame{Type: "trace", Tool: toolName, Stage: "arguments", Status: "running", Message: "准备执行工具", Data: map[string]any{"tool_call_id": call.ID, "arguments": call.Function.Arguments}}) - } - tool, ok := tools[toolName] - if !ok { - result := fmt.Sprintf("工具调用失败:未知工具 %s。", toolName) - if emit != nil { - emit(chatSSEFrame{Type: "trace", Tool: toolName, Stage: "execute", Status: "error", Message: result}) - } - return result - } - started := time.Now() - result, err := tool.execute(ctx, call.Function.Arguments) - durationMs := time.Since(started).Milliseconds() - if err != nil { - message := fmt.Sprintf("工具 %s 执行失败:%v", tool.name, err) - if emit != nil { - emit(chatSSEFrame{Type: "trace", Tool: tool.name, Stage: "execute", Status: "error", Message: "工具执行失败", Data: map[string]any{"tool_call_id": call.ID, "duration_ms": durationMs, "error": err.Error()}}) - } - return message - } - if strings.TrimSpace(result) == "" { - result = fmt.Sprintf("工具 %s 执行完成,但没有返回内容。", tool.name) - } - if emit != nil { - emit(chatSSEFrame{Type: "trace", Tool: tool.name, Stage: "result", Status: "success", Message: "工具执行完成", Data: map[string]any{"tool_call_id": call.ID, "duration_ms": durationMs, "result_preview": truncateString(result, 1200)}}) - } - return result -} - -type ollamaChatRequest struct { - Model string `json:"model"` - Messages []ollamaChatMessage `json:"messages"` - Stream bool `json:"stream"` - Options map[string]int `json:"options,omitempty"` -} - -type ollamaChatMessage struct { - Role string `json:"role"` - Content string `json:"content"` - Images []string `json:"images,omitempty"` -} - -type ollamaChatResponse struct { - Message struct { - Role string `json:"role"` - Content string `json:"content"` - Thinking string `json:"thinking"` - } `json:"message"` - Done bool `json:"done"` - PromptEvalCount int `json:"prompt_eval_count"` - EvalCount int `json:"eval_count"` - DoneReason string `json:"done_reason"` -} - -func streamOllamaChat(ctx context.Context, profile *OpenAIProfile, messages []*model.ChatCompletionMessage, promptTokens int, usage *tokenUsageTracker, emit func(chatSSEFrame), onDone func(string)) error { - requestMessages, err := buildOllamaMessages(messages) - if err != nil { - return err - } - baseURL, err := ollamaBaseURL(profile) - if err != nil { - return err - } - body, err := json.Marshal(ollamaChatRequest{ - Model: profile.Config.Model, - Messages: requestMessages, - Stream: true, - Options: map[string]int{"num_predict": 4096}, - }) - if err != nil { - return err - } - req, err := http.NewRequestWithContext(ctx, http.MethodPost, strings.TrimRight(baseURL, "/")+"/api/chat", bytes.NewReader(body)) - if err != nil { - return err - } - req.Header.Set("Content-Type", "application/json") - resp, err := http.DefaultClient.Do(req) - if err != nil { - return err - } - defer resp.Body.Close() - if resp.StatusCode < 200 || resp.StatusCode >= 300 { - data, _ := io.ReadAll(io.LimitReader(resp.Body, 4096)) - return fmt.Errorf("Ollama 原生接口调用失败: %s %s", resp.Status, strings.TrimSpace(string(data))) - } - - emit(chatSSEFrame{Type: "trace", Tool: "model", Stage: "stream", Status: "running", Message: "Ollama 视觉模型已开始输出"}) - parseThinkTags := shouldParseThinkTags(profile) - thinkParser := &thinkTagParser{} - var full strings.Builder - completionTokens := 0 - streamStarted := time.Now() - peakTokensPerSecond := 0.0 - emitDelta := func(delta string) { - if delta == "" { - return - } - full.WriteString(delta) - completionTokens += estimateTokenCount(delta) - usage.setModel(promptTokens, completionTokens) - currentSpeed := tokensPerSecond(completionTokens, streamStarted) - if currentSpeed > peakTokensPerSecond { - peakTokensPerSecond = currentSpeed - } - stats := usage.snapshot(currentSpeed, peakTokensPerSecond) - emit(chatSSEFrame{Type: "delta", Text: delta, Stats: &stats}) - } - emitContent := func(delta string) { - if delta == "" { - return - } - if !parseThinkTags { - emitDelta(delta) - return - } - visible, reasoning := thinkParser.Accept(delta) - if reasoning != "" { - emit(chatSSEFrame{Type: "reasoning", Text: reasoning}) - } - emitDelta(visible) - } - - scanner := bufio.NewScanner(resp.Body) - scanner.Buffer(make([]byte, 0, 64*1024), 10*1024*1024) - for scanner.Scan() { - line := strings.TrimSpace(scanner.Text()) - if line == "" { - continue - } - var chunk ollamaChatResponse - if err := json.Unmarshal([]byte(line), &chunk); err != nil { - return fmt.Errorf("解析 Ollama 流失败: %w", err) - } - if chunk.Message.Thinking != "" { - emit(chatSSEFrame{Type: "reasoning", Text: chunk.Message.Thinking}) - } - emitContent(chunk.Message.Content) - if chunk.Done { - if chunk.PromptEvalCount > 0 || chunk.EvalCount > 0 { - usage.setModel(chunk.PromptEvalCount, chunk.EvalCount) - } - break - } - } - if err := scanner.Err(); err != nil { - return err - } - if parseThinkTags { - visible, reasoning := thinkParser.Flush() - if reasoning != "" { - emit(chatSSEFrame{Type: "reasoning", Text: reasoning}) - } - emitDelta(visible) - } - if onDone != nil { - onDone(full.String()) - } - finalStats := usage.snapshot(tokensPerSecond(completionTokens, streamStarted), peakTokensPerSecond) - emit(chatSSEFrame{Type: "stats", Stats: &finalStats}) - emit(chatSSEFrame{Type: "trace", Tool: "model", Stage: "stream", Status: "success", Message: "回答生成完成"}) - return nil -} - -func buildOllamaMessages(messages []*model.ChatCompletionMessage) ([]ollamaChatMessage, error) { - result := make([]ollamaChatMessage, 0, len(messages)) - for _, msg := range messages { - if msg == nil { - continue - } - role := string(msg.Role) - if msg.Role == model.ChatMessageRoleTool { - role = string(model.ChatMessageRoleUser) - } - item := ollamaChatMessage{Role: role} - if msg.Content == nil { - if len(msg.ToolCalls) > 0 { - continue - } - result = append(result, item) - continue - } - if msg.Content.StringValue != nil { - item.Content = *msg.Content.StringValue - if msg.Role == model.ChatMessageRoleTool { - item.Content = "工具结果:\n" + item.Content - } - result = append(result, item) - continue - } - for _, part := range msg.Content.ListValue { - if part == nil { - continue - } - switch part.Type { - case model.ChatCompletionMessageContentPartTypeText: - if part.Text != "" { - if item.Content != "" { - item.Content += "\n" - } - item.Content += part.Text - } - case model.ChatCompletionMessageContentPartTypeImageURL: - if part.ImageURL == nil { - continue - } - image, err := ollamaImagePayload(part.ImageURL.URL) - if err != nil { - return nil, err - } - item.Images = append(item.Images, image) - } - } - result = append(result, item) - } - return result, nil -} - -func ollamaImagePayload(raw string) (string, error) { - raw = strings.TrimSpace(raw) - if strings.HasPrefix(strings.ToLower(raw), "data:") { - comma := strings.Index(raw, ",") - if comma < 0 { - return "", errors.New("图片 base64 数据格式错误") - } - return strings.TrimSpace(raw[comma+1:]), nil - } - return raw, nil -} - -func ollamaBaseURL(profile *OpenAIProfile) (string, error) { - if profile == nil { - return "", errors.New("Ollama 配置为空") - } - u, err := url.Parse(strings.TrimSpace(profile.Config.BaseURL)) - if err != nil { - return "", err - } - if strings.TrimRight(u.Path, "/") == "/v1" { - u.Path = strings.TrimSuffix(strings.TrimRight(u.Path, "/"), "/v1") - } - u.RawQuery = "" - u.Fragment = "" - return strings.TrimRight(u.String(), "/"), nil -} - -func completeText(ctx context.Context, profile *OpenAIProfile, chatMessages []ChatMessage, maxTokens int) (string, error) { - return completeTextWithTimeout(ctx, profile, chatMessages, maxTokens, time.Duration(profile.Config.Timeout)*time.Second) -} - -func completeTextWithTimeout(ctx context.Context, profile *OpenAIProfile, chatMessages []ChatMessage, maxTokens int, timeout time.Duration) (string, error) { - messages, err := buildArkMessages(chatMessages) - if err != nil { - return "", err - } - completionCtx, cancel := context.WithTimeout(ctx, timeout) - defer cancel() - stream, err := profile.Client.CreateChatCompletionStream(completionCtx, model.CreateChatCompletionRequest{ - Model: profile.Config.Model, - Messages: messages, - MaxTokens: intPtr(maxTokens), - }.WithStream(true)) - if err != nil { - return "", err - } - defer stream.Close() - - promptTokens := estimateChatMessagesTokens(chatMessages) - completionTokens := 0 - parseThinkTags := shouldParseThinkTags(profile) - thinkParser := &thinkTagParser{} - var b strings.Builder - appendVisible := func(delta string) { - if delta == "" { - return - } - b.WriteString(delta) - completionTokens += estimateTokenCount(delta) - } - for { - resp, err := stream.Recv() - if errors.Is(err, io.EOF) { - if parseThinkTags { - visible, _ := thinkParser.Flush() - appendVisible(visible) - } - if tracker := tokenUsageFromContext(ctx); tracker != nil { - tracker.addTool(promptTokens, completionTokens) - } - return b.String(), nil - } - if err != nil { - return "", err - } - if len(resp.Choices) > 0 { - delta := resp.Choices[0].Delta.Content - if parseThinkTags { - visible, _ := thinkParser.Accept(delta) - appendVisible(visible) - } else { - appendVisible(delta) - } - } - } -} - -func completeChatWithTimeout(ctx context.Context, profile *OpenAIProfile, request model.CreateChatCompletionRequest, timeout time.Duration) (model.ChatCompletionResponse, error) { - completionCtx, cancel := context.WithTimeout(ctx, timeout) - defer cancel() - return profile.Client.CreateChatCompletion(completionCtx, request.WithStream(false)) -} - -func newUUID() string { - b := make([]byte, 16) - _, _ = rand.Read(b) - b[6] = (b[6] & 0x0f) | 0x40 - b[8] = (b[8] & 0x3f) | 0x80 - return hex.EncodeToString(b[:4]) + "-" + hex.EncodeToString(b[4:6]) + "-" + - hex.EncodeToString(b[6:8]) + "-" + hex.EncodeToString(b[8:10]) + "-" + - hex.EncodeToString(b[10:]) -} - -// ─── ConvStore ───────────────────────────────────────────── - -func NewConvStore(dir string) *ConvStore { - os.MkdirAll(dir, 0755) - return &ConvStore{dir: dir} -} - -func (s *ConvStore) path(id string) string { - return filepath.Join(s.dir, id+".json") -} - -func (s *ConvStore) Create() (*Conversation, error) { - conv := &Conversation{ - ID: newUUID(), - Title: "新对话", - CreatedAt: time.Now(), - UpdatedAt: time.Now(), - } - if err := s.Save(conv); err != nil { - return nil, err - } - return conv, nil -} - -func (s *ConvStore) Save(conv *Conversation) error { - s.mu.Lock() - defer s.mu.Unlock() - conv.UpdatedAt = time.Now() - return atomicWriteJSON(s.path(conv.ID), conv) -} - -func (s *ConvStore) Get(id string) (*Conversation, error) { - s.mu.Lock() - defer s.mu.Unlock() - data, err := os.ReadFile(s.path(id)) - if err != nil { - if os.IsNotExist(err) { - return nil, errors.New("对话不存在") - } - return nil, fmt.Errorf("读取对话失败: %w", err) - } - var conv Conversation - if err := json.Unmarshal(data, &conv); err != nil { - return nil, fmt.Errorf("解析对话失败: %w", err) - } - return &conv, nil -} - -func (s *ConvStore) List() ([]Conversation, error) { - s.mu.Lock() - defer s.mu.Unlock() - - entries, err := os.ReadDir(s.dir) - if err != nil { - return nil, fmt.Errorf("读取对话目录失败: %w", err) - } - - var list []Conversation - for _, e := range entries { - if e.IsDir() || filepath.Ext(e.Name()) != ".json" { - continue - } - data, err := os.ReadFile(filepath.Join(s.dir, e.Name())) - if err != nil { - continue - } - var conv Conversation - if err := json.Unmarshal(data, &conv); err != nil { - continue - } - conv.Messages = nil // 列表不返回消息体 - list = append(list, conv) - } - - sort.Slice(list, func(i, j int) bool { - return list[i].UpdatedAt.After(list[j].UpdatedAt) - }) - return list, nil -} - -func (s *ConvStore) Delete(id string) error { - s.mu.Lock() - defer s.mu.Unlock() - if err := os.Remove(s.path(id)); err != nil && !os.IsNotExist(err) { - return fmt.Errorf("删除对话失败: %w", err) - } - return nil -} - -func atomicWriteJSON(path string, v any) error { - tmp := path + ".tmp" - data, err := json.Marshal(v) - if err != nil { - return err - } - if err := os.WriteFile(tmp, data, 0644); err != nil { - return err - } - return os.Rename(tmp, path) -} - -func saveConversationMessages(id string, messages []ChatMessage, assistantContent string) error { - conv, err := store.Get(id) - if err != nil { - return err - } - conv.Messages = append([]ChatMessage(nil), messages...) - conv.Messages = append(conv.Messages, ChatMessage{Role: "assistant", Content: assistantContent}) - if conv.Title == "" || conv.Title == "新对话" { - conv.Title = genConvTitle(conv.Messages) - } - return store.Save(conv) -} - -func genConvTitle(messages []ChatMessage) string { - for _, m := range messages { - if m.Hidden { - continue - } - if m.Role == "user" && strings.TrimSpace(m.Content) != "" { - title := strings.TrimSpace(m.Content) - title = strings.ReplaceAll(title, "\r\n", " ") - title = strings.ReplaceAll(title, "\n", " ") - runes := []rune(title) - if len(runes) > 30 { - return string(runes[:30]) + "..." - } - return title - } - } - return "新对话" -} - -const maxImageSize = 4 * 1024 * 1024 - -var allowedImageTypes = map[string]bool{ - "image/jpeg": true, - "image/png": true, - "image/webp": true, - "image/gif": true, -} - -func buildArkMessages(chatMessages []ChatMessage) ([]*model.ChatCompletionMessage, error) { - messages := make([]*model.ChatCompletionMessage, 0, len(chatMessages)) - for _, m := range chatMessages { - msg, err := buildArkMessage(m) - if err != nil { - return nil, err - } - messages = append(messages, msg) - } - return messages, nil -} - -func hasImageMessage(messages []ChatMessage) bool { - for _, msg := range messages { - if strings.TrimSpace(msg.ImageURL) != "" || strings.TrimSpace(msg.ImageURLAlias) != "" { - return true - } - } - return false -} - -func buildToolDecisionMessages(chatMessages []ChatMessage) ([]*model.ChatCompletionMessage, error) { - messages := make([]*model.ChatCompletionMessage, 0, len(chatMessages)) - for _, m := range chatMessages { - content := m.Content - if strings.TrimSpace(m.ImageURL) != "" || strings.TrimSpace(m.ImageURLAlias) != "" { - content = strings.TrimSpace(content) - placeholder := "[用户上传了一张图片。工具判断阶段不读取图片内容;如果问题主要依赖识图,应不要调用工具,交给最终多模态模型回答。]" - if content == "" { - content = placeholder - } else { - content += "\n\n" + placeholder - } - } - messages = append(messages, &model.ChatCompletionMessage{Role: m.Role, Content: stringContent(content)}) - } - return messages, nil -} - -func buildArkMessage(m ChatMessage) (*model.ChatCompletionMessage, error) { - msg := &model.ChatCompletionMessage{Role: m.Role} - - if m.ImageURL == "" && m.ImageURLAlias != "" { - m.ImageURL = m.ImageURLAlias - } - - if m.ImageURL == "" { - msg.Content = &model.ChatCompletionMessageContent{ - StringValue: &m.Content, - } - return msg, nil - } - - imageURL, err := normalizeImageURL(m.ImageURL) - if err != nil { - return nil, err - } - - // 有图片时:文字内容可有可无(图片 caption 场景),均构造多模态消息 - // 若无文字,则只传图片 part;若同时有图片和文字,先文后图 - parts := make([]*model.ChatCompletionMessageContentPart, 0, 2) - if m.Content != "" { - parts = append(parts, textPart(m.Content)) - } - parts = append(parts, imagePart(imageURL)) - msg.Content = &model.ChatCompletionMessageContent{ListValue: parts} - return msg, nil -} - -func imagePart(url string) *model.ChatCompletionMessageContentPart { - return &model.ChatCompletionMessageContentPart{ - Type: model.ChatCompletionMessageContentPartTypeImageURL, - ImageURL: &model.ChatMessageImageURL{ - URL: url, - Detail: model.ImageURLDetailAuto, - }, - } -} - -func textPart(text string) *model.ChatCompletionMessageContentPart { - return &model.ChatCompletionMessageContentPart{ - Type: model.ChatCompletionMessageContentPartTypeText, - Text: text, - } -} - -func stringContent(text string) *model.ChatCompletionMessageContent { - return &model.ChatCompletionMessageContent{StringValue: &text} -} - -func chatMessageContentString(content *model.ChatCompletionMessageContent) string { - if content == nil || content.StringValue == nil { - return "" - } - return *content.StringValue -} - -func normalizeImageURL(raw string) (string, error) { - raw = strings.TrimSpace(raw) - if raw == "" { - return "", errors.New("图片地址不能为空") - } - - lower := strings.ToLower(raw) - if strings.HasPrefix(lower, "data:") { - return normalizeImageDataURI(raw) - } - - u, err := url.Parse(raw) - if err != nil || u.Host == "" || (u.Scheme != "http" && u.Scheme != "https") { - return "", errors.New("图片地址无效,仅支持 http/https URL 或 base64 data URI") - } - return raw, nil -} - -func normalizeImageDataURI(raw string) (string, error) { - comma := strings.Index(raw, ",") - if comma < 0 { - return "", errors.New("图片 base64 数据格式错误") - } - - meta := strings.ToLower(strings.TrimSpace(raw[5:comma])) - payload := strings.TrimSpace(raw[comma+1:]) - if payload == "" { - return "", errors.New("图片 base64 数据不能为空") - } - parts := strings.Split(meta, ";") - if len(parts) < 2 || !contains(parts[1:], "base64") { - return "", errors.New("图片 data URI 必须使用 base64 编码") - } - - mime := parts[0] - if !allowedImageTypes[mime] { - return "", errors.New("图片格式不支持,仅支持 jpeg/png/webp/gif") - } - - decoded, err := base64.StdEncoding.DecodeString(payload) - if err != nil { - return "", errors.New("图片 base64 数据无效") - } - if len(decoded) > maxImageSize { - return "", errors.New("图片过大,请选择小于 4MB 的图片") - } - - return "data:" + mime + ";base64," + payload, nil -} - -func contains(items []string, target string) bool { - for _, item := range items { - if strings.TrimSpace(item) == target { - return true - } - } - return false -} - -func intPtr(i int) *int { return &i } - -func boolPtr(v bool) *bool { return &v } - -func truncateString(text string, maxRunes int) string { - runes := []rune(strings.TrimSpace(text)) - if maxRunes <= 0 || len(runes) <= maxRunes { - return string(runes) - } - return string(runes[:maxRunes]) + "..." -} - -func writeSSEJSON(w io.Writer, frame chatSSEFrame) { - data, err := json.Marshal(frame) - if err != nil { - data, _ = json.Marshal(chatSSEFrame{Type: "error", Error: "序列化流事件失败"}) - } - fmt.Fprintf(w, "data: %s\n\n", data) -} - -func toJSON(s string) string { - b, _ := json.Marshal(s) - return string(b) -} - -func toSSE(s string) string { - s = strings.ReplaceAll(s, `\`, `\\`) - s = strings.ReplaceAll(s, "\n", `\n`) - s = strings.ReplaceAll(s, "\r", "") - s = strings.ReplaceAll(s, `"`, `\"`) - return fmt.Sprintf(`"%s"`, s) -} - -// ─── 入口 ───────────────────────────────────────────────── - func main() { - var err error - cfg, err = loadConfig("config.yaml") + cfg, legacySearchProfiles, err := config.Load("config.yaml") if err != nil { fmt.Fprintln(os.Stderr, "配置加载失败:", err) os.Exit(1) } // 初始化火山方舟 SDK 客户端 - aiState, err = NewOpenAIState(cfg.OpenAI) + aiState, err := llm.NewState(cfg.OpenAI) if err != nil { fmt.Fprintln(os.Stderr, "OpenAI 配置初始化失败:", err) os.Exit(1) @@ -1973,7 +32,7 @@ func main() { fmt.Fprintln(os.Stderr, "联网搜索配置加载失败:", err) os.Exit(1) } - searchState, err = searchagent.NewState(searchConfig) + searchState, err := searchagent.NewState(searchConfig) if err != nil { fmt.Fprintln(os.Stderr, "联网搜索初始化失败:", err) os.Exit(1) @@ -1983,57 +42,23 @@ func main() { fmt.Fprintln(os.Stderr, "SQL 查询插件配置加载失败:", err) os.Exit(1) } - sqlState, err = sqlquery.NewState(sqlConfig) + sqlState, err := sqlquery.NewState(sqlConfig) if err != nil { fmt.Fprintln(os.Stderr, "SQL 查询插件初始化失败:", err) os.Exit(1) } defer sqlState.Close() - toolRouterState, err = NewToolRouterState(&cfg.ToolRouter, aiState) + toolRouterState, err := toolrouter.NewState(&cfg.ToolRouter, aiState) if err != nil { fmt.Fprintln(os.Stderr, "工具路由配置初始化失败:", err) os.Exit(1) } - store = NewConvStore("conversations") + store := conversation.NewStore("conversations") - // Gin 路由 - r := gin.Default() - r.LoadHTMLGlob("templates/*") - r.Static("/static", "./static") - - r.GET("/", indexHandler) - r.POST("/api/chat", chatHandler) - r.GET("/api/openai", listOpenAIHandler) - r.POST("/api/openai/active", switchOpenAIHandler) - r.GET("/api/search", listSearchHandler) - r.POST("/api/search/active", switchSearchHandler) - r.GET("/api/conversations", listConversationsHandler) - r.POST("/api/conversations", createConversationHandler) - r.GET("/api/conversations/:id", getConversationHandler) - r.DELETE("/api/conversations/:id", deleteConversationHandler) - - // 根据配置选择监听方式 - switch strings.ToLower(cfg.Server.Mode) { - case "unix": - socketPath := cfg.Server.Address - if _, statErr := os.Stat(socketPath); statErr == nil { - os.Remove(socketPath) - } - ln, listenErr := net.Listen("unix", socketPath) - if listenErr != nil { - fmt.Fprintln(os.Stderr, "监听 Unix socket 失败:", listenErr) - os.Exit(1) - } - fmt.Println("服务已启动,监听 Unix socket:", socketPath) - if serveErr := http.Serve(ln, r); serveErr != nil { - fmt.Fprintln(os.Stderr, "服务异常退出:", serveErr) - os.Exit(1) - } - default: - fmt.Println("服务已启动,监听 TCP:", cfg.Server.Address) - if runErr := r.Run(cfg.Server.Address); runErr != nil { - fmt.Fprintln(os.Stderr, "服务异常退出:", runErr) - os.Exit(1) - } + app := server.New(cfg, aiState, searchState, sqlState, toolRouterState, store) + cfg.Server.Mode = strings.ToLower(cfg.Server.Mode) + if err := app.Run(); err != nil { + fmt.Fprintln(os.Stderr, "服务异常退出:", err) + os.Exit(1) } } diff --git a/main_test.go b/main_test.go index 2ee35fb..695edd0 100644 --- a/main_test.go +++ b/main_test.go @@ -7,22 +7,40 @@ import ( "testing" "time" + "aichat/config" + "aichat/llm" + "aichat/message" + "aichat/stream" + "aichat/toolrouter" + "github.com/volcengine/volcengine-go-sdk/service/arkruntime/model" ) +const testOpenAIBaseURL = "https://ark.cn-beijing.volces.com/api/v3" + +func newTestAI(t *testing.T, configs []config.OpenAIConfig) *llm.State { + t.Helper() + ai, err := llm.NewState(configs) + if err != nil { + t.Fatal(err) + } + return ai +} + func TestNormalizeToolRouterConfigDefaults(t *testing.T) { - cfg := &Config{ToolRouter: ToolRouterConfig{Enabled: true}} - changed, err := normalizeToolRouterConfig(cfg) + cfg := &config.Config{ToolRouter: config.ToolRouterConfig{Enabled: true}} + changed, err := config.NormalizeToolRouterConfig(cfg) if err != nil { t.Fatal(err) } if !changed { t.Fatal("expected defaults to change config") } - if cfg.ToolRouter.Timeout != defaultToolRouterTimeout { + defaults := config.DefaultToolRouterConfig() + if cfg.ToolRouter.Timeout != defaults.Timeout { t.Fatalf("timeout = %d", cfg.ToolRouter.Timeout) } - if cfg.ToolRouter.MaxTokens != defaultToolRouterMaxTokens { + if cfg.ToolRouter.MaxTokens != defaults.MaxTokens { t.Fatalf("max_tokens = %d", cfg.ToolRouter.MaxTokens) } if strings.TrimSpace(cfg.ToolRouter.SystemPrompt) == "" { @@ -34,17 +52,17 @@ func TestNormalizeToolRouterConfigDefaults(t *testing.T) { } func TestNormalizeToolRouterConfigAddsTimeBeforeSQL(t *testing.T) { - cfg := &Config{ToolRouter: ToolRouterConfig{ + cfg := &config.Config{ToolRouter: config.ToolRouterConfig{ Enabled: true, Timeout: 1, MaxTokens: 1, SystemPrompt: "tools", - Tools: []ToolRouteConfig{ + Tools: []config.ToolRouteConfig{ {Name: "search", Enabled: true}, {Name: "sql", Enabled: true}, }, }} - changed, err := normalizeToolRouterConfig(cfg) + changed, err := config.NormalizeToolRouterConfig(cfg) if err != nil { t.Fatal(err) } @@ -57,68 +75,59 @@ func TestNormalizeToolRouterConfigAddsTimeBeforeSQL(t *testing.T) { } func TestNormalizeToolRouterConfigDuplicateTools(t *testing.T) { - cfg := &Config{ToolRouter: ToolRouterConfig{ + cfg := &config.Config{ToolRouter: config.ToolRouterConfig{ Enabled: true, Timeout: 1, MaxTokens: 1, SystemPrompt: "tools", - Tools: []ToolRouteConfig{ + Tools: []config.ToolRouteConfig{ {Name: "sql", Enabled: true}, {Name: " SQL ", Enabled: true}, }, }} - _, err := normalizeToolRouterConfig(cfg) + _, err := config.NormalizeToolRouterConfig(cfg) if err == nil { t.Fatal("expected duplicate tool error") } } func TestAvailableAgentToolsUsesConfigOrderAndEnabled(t *testing.T) { - oldRouter := toolRouterState - oldSearch := searchState - oldSQL := sqlState - defer func() { - toolRouterState = oldRouter - searchState = oldSearch - sqlState = oldSQL - }() - - toolRouterState = &ToolRouterState{cfg: &ToolRouterConfig{ + ai := newTestAI(t, []config.OpenAIConfig{{Name: "test", APIKey: "key", BaseURL: testOpenAIBaseURL, Model: "test", Timeout: 1, Active: true}}) + router, err := toolrouter.NewState(&config.ToolRouterConfig{ Enabled: true, - Tools: []ToolRouteConfig{ + Tools: []config.ToolRouteConfig{ {Name: "search", Enabled: true}, {Name: "time", Enabled: true, Description: "custom time"}, {Name: "sql", Enabled: false}, }, - }} - searchState = nil - sqlState = nil + }, ai) + if err != nil { + t.Fatal(err) + } - tools := availableAgentTools(&OpenAIProfile{}, nil) + tools := toolrouter.AvailableAgentTools(router, ai.ActiveProfile(), nil, nil, nil) if len(tools) != 1 { t.Fatalf("tools length = %d", len(tools)) } - if tools[0].name != "time" { - t.Fatalf("tool name = %s", tools[0].name) + if tools[0].Name() != "time" { + t.Fatalf("tool name = %s", tools[0].Name()) } - if tools[0].definition.Function == nil || tools[0].definition.Function.Description != "custom time" { - t.Fatalf("unexpected definition: %#v", tools[0].definition) + definition := tools[0].Definition() + if definition.Function == nil || definition.Function.Description != "custom time" { + t.Fatalf("unexpected definition: %#v", definition) } } func TestRunAgentToolLoopAppendsToolMessages(t *testing.T) { - oldRouter := toolRouterState - defer func() { toolRouterState = oldRouter }() - + ai := newTestAI(t, []config.OpenAIConfig{{Name: "test", APIKey: "key", BaseURL: testOpenAIBaseURL, Model: "test", Timeout: 1, Active: true}}) calls := 0 - toolRouterState = &ToolRouterState{cfg: &ToolRouterConfig{ + router, err := toolrouter.NewState(&config.ToolRouterConfig{ Enabled: true, Timeout: 1, MaxTokens: 128, SystemPrompt: "use tools", - Tools: []ToolRouteConfig{{Name: "time", Enabled: true}}, - }} - toolRouterState.complete = func(ctx context.Context, profile *OpenAIProfile, req model.CreateChatCompletionRequest, timeout time.Duration) (model.ChatCompletionResponse, error) { + Tools: []config.ToolRouteConfig{{Name: "time", Enabled: true}}, + }, ai, toolrouter.WithCompleter(func(ctx context.Context, profile *llm.Profile, req model.CreateChatCompletionRequest, timeout time.Duration) (model.ChatCompletionResponse, error) { calls++ if req.ToolChoice != model.ToolChoiceStringTypeAuto { t.Fatalf("tool choice = %#v", req.ToolChoice) @@ -129,10 +138,13 @@ func TestRunAgentToolLoopAppendsToolMessages(t *testing.T) { if calls == 1 { return model.ChatCompletionResponse{Choices: []*model.ChatCompletionChoice{{Message: model.ChatCompletionMessage{ToolCalls: []*model.ToolCall{{ID: "call_1", Type: model.ToolTypeFunction, Function: model.FunctionCall{Name: "time", Arguments: `{"reason":"需要当前日期"}`}}}}}}}, nil } - return model.ChatCompletionResponse{Choices: []*model.ChatCompletionChoice{{Message: model.ChatCompletionMessage{Content: stringContent("done")}}}}, nil + return model.ChatCompletionResponse{Choices: []*model.ChatCompletionChoice{{Message: model.ChatCompletionMessage{Content: message.StringContent("done")}}}}, nil + })) + if err != nil { + t.Fatal(err) } - messages, err := runAgentToolLoop(context.Background(), &OpenAIProfile{Config: OpenAIConfig{Model: "test"}}, []ChatMessage{{Role: "user", Content: "今天几号"}}, nil) + messages, err := toolrouter.RunAgentToolLoop(context.Background(), router, ai.ActiveProfile(), []message.ChatMessage{{Role: "user", Content: "今天几号"}}, nil, nil, nil) if err != nil { t.Fatal(err) } @@ -152,13 +164,13 @@ func TestRunAgentToolLoopAppendsToolMessages(t *testing.T) { } func TestExecuteAgentToolCallUnknownAndError(t *testing.T) { - unknown := executeAgentToolCall(context.Background(), &model.ToolCall{ID: "1", Type: model.ToolTypeFunction, Function: model.FunctionCall{Name: "missing"}}, map[string]agentTool{}, nil) + unknown := toolrouter.ExecuteAgentToolCall(context.Background(), &model.ToolCall{ID: "1", Type: model.ToolTypeFunction, Function: model.FunctionCall{Name: "missing"}}, map[string]toolrouter.AgentTool{}, nil) if !strings.Contains(unknown, "未知工具") { t.Fatalf("unknown result = %q", unknown) } - failed := executeAgentToolCall(context.Background(), &model.ToolCall{ID: "2", Type: model.ToolTypeFunction, Function: model.FunctionCall{Name: "boom"}}, map[string]agentTool{ - "boom": {name: "boom", execute: func(context.Context, string) (string, error) { return "", errors.New("bad args") }}, + failed := toolrouter.ExecuteAgentToolCall(context.Background(), &model.ToolCall{ID: "2", Type: model.ToolTypeFunction, Function: model.FunctionCall{Name: "boom"}}, map[string]toolrouter.AgentTool{ + "boom": toolrouter.NewAgentTool("boom", nil, func(context.Context, string) (string, error) { return "", errors.New("bad args") }), }, nil) if !strings.Contains(failed, "bad args") { t.Fatalf("failed result = %q", failed) @@ -166,21 +178,21 @@ func TestExecuteAgentToolCallUnknownAndError(t *testing.T) { } func TestRunAgentToolLoopMaxIterations(t *testing.T) { - oldRouter := toolRouterState - defer func() { toolRouterState = oldRouter }() - - toolRouterState = &ToolRouterState{cfg: &ToolRouterConfig{ + ai := newTestAI(t, []config.OpenAIConfig{{Name: "test", APIKey: "key", BaseURL: testOpenAIBaseURL, Model: "test", Timeout: 1, Active: true}}) + router, err := toolrouter.NewState(&config.ToolRouterConfig{ Enabled: true, Timeout: 1, MaxTokens: 128, SystemPrompt: "use tools", - Tools: []ToolRouteConfig{{Name: "time", Enabled: true}}, - }} - toolRouterState.complete = func(context.Context, *OpenAIProfile, model.CreateChatCompletionRequest, time.Duration) (model.ChatCompletionResponse, error) { + Tools: []config.ToolRouteConfig{{Name: "time", Enabled: true}}, + }, ai, toolrouter.WithCompleter(func(context.Context, *llm.Profile, model.CreateChatCompletionRequest, time.Duration) (model.ChatCompletionResponse, error) { return model.ChatCompletionResponse{Choices: []*model.ChatCompletionChoice{{Message: model.ChatCompletionMessage{ToolCalls: []*model.ToolCall{{ID: "loop", Type: model.ToolTypeFunction, Function: model.FunctionCall{Name: "time", Arguments: `{}`}}}}}}}, nil + })) + if err != nil { + t.Fatal(err) } - messages, err := runAgentToolLoop(context.Background(), &OpenAIProfile{Config: OpenAIConfig{Model: "test"}}, []ChatMessage{{Role: "user", Content: "今天"}}, nil) + messages, err := toolrouter.RunAgentToolLoop(context.Background(), router, ai.ActiveProfile(), []message.ChatMessage{{Role: "user", Content: "今天"}}, nil, nil, nil) if err != nil { t.Fatal(err) } @@ -191,7 +203,7 @@ func TestRunAgentToolLoopMaxIterations(t *testing.T) { } func TestBuildArkMessageImageTextOrder(t *testing.T) { - msg, err := buildArkMessage(ChatMessage{Role: "user", Content: "请描述图片", ImageURL: "data:image/png;base64,aGVsbG8="}) + msg, err := message.BuildArkMessage(message.ChatMessage{Role: "user", Content: "请描述图片", ImageURL: "data:image/png;base64,aGVsbG8="}) if err != nil { t.Fatal(err) } @@ -207,7 +219,7 @@ func TestBuildArkMessageImageTextOrder(t *testing.T) { } func TestBuildArkMessageImageOnly(t *testing.T) { - msg, err := buildArkMessage(ChatMessage{Role: "user", ImageURL: "data:image/png;base64,aGVsbG8="}) + msg, err := message.BuildArkMessage(message.ChatMessage{Role: "user", ImageURL: "data:image/png;base64,aGVsbG8="}) if err != nil { t.Fatal(err) } @@ -217,7 +229,7 @@ func TestBuildArkMessageImageOnly(t *testing.T) { } func TestThinkTagParserSingleChunk(t *testing.T) { - parser := &thinkTagParser{} + parser := &stream.Parser{} visible, reasoning := parser.Accept("hello abc world") flushVisible, flushReasoning := parser.Flush() visible += flushVisible @@ -228,7 +240,7 @@ func TestThinkTagParserSingleChunk(t *testing.T) { } func TestThinkTagParserAcrossChunks(t *testing.T) { - parser := &thinkTagParser{} + parser := &stream.Parser{} var visible, reasoning string for _, chunk := range []string{"hello abc world"} { v, r := parser.Accept(chunk) @@ -244,7 +256,7 @@ func TestThinkTagParserAcrossChunks(t *testing.T) { } func TestThinkTagParserUnclosedThink(t *testing.T) { - parser := &thinkTagParser{} + parser := &stream.Parser{} visible, reasoning := parser.Accept("answer still thinking") v, r := parser.Flush() visible += v @@ -255,24 +267,24 @@ func TestThinkTagParserUnclosedThink(t *testing.T) { } func TestShouldParseThinkTags(t *testing.T) { - if !shouldParseThinkTags(&OpenAIProfile{Config: OpenAIConfig{BaseURL: "http://127.0.0.1:11434/v1"}}) { + if !llm.ShouldParseThinkTags(&llm.Profile{Config: config.OpenAIConfig{BaseURL: "http://127.0.0.1:11434/v1"}}) { t.Fatal("expected local ollama to parse think tags") } - if shouldParseThinkTags(&OpenAIProfile{Config: OpenAIConfig{BaseURL: defaultOpenAIBaseURL}}) { + if llm.ShouldParseThinkTags(&llm.Profile{Config: config.OpenAIConfig{BaseURL: testOpenAIBaseURL}}) { t.Fatal("expected remote profile not to parse think tags by default") } falseValue := false - if shouldParseThinkTags(&OpenAIProfile{Config: OpenAIConfig{BaseURL: "http://127.0.0.1:11434/v1", ParseThinkTags: &falseValue}}) { + if llm.ShouldParseThinkTags(&llm.Profile{Config: config.OpenAIConfig{BaseURL: "http://127.0.0.1:11434/v1", ParseThinkTags: &falseValue}}) { t.Fatal("explicit false should disable think parsing") } trueValue := true - if !shouldParseThinkTags(&OpenAIProfile{Config: OpenAIConfig{BaseURL: defaultOpenAIBaseURL, ParseThinkTags: &trueValue}}) { + if !llm.ShouldParseThinkTags(&llm.Profile{Config: config.OpenAIConfig{BaseURL: testOpenAIBaseURL, ParseThinkTags: &trueValue}}) { t.Fatal("explicit true should enable think parsing") } } func TestBuildToolDecisionMessagesRemovesImages(t *testing.T) { - messages, err := buildToolDecisionMessages([]ChatMessage{{Role: "user", Content: "描述这张图", ImageURL: "data:image/png;base64,aGVsbG8="}}) + messages, err := message.BuildToolDecisionMessages([]message.ChatMessage{{Role: "user", Content: "描述这张图", ImageURL: "data:image/png;base64,aGVsbG8="}}) if err != nil { t.Fatal(err) } @@ -288,17 +300,14 @@ func TestBuildToolDecisionMessagesRemovesImages(t *testing.T) { } func TestRunAgentToolLoopImageUsesTextOnlyDecisionMessages(t *testing.T) { - oldRouter := toolRouterState - defer func() { toolRouterState = oldRouter }() - - toolRouterState = &ToolRouterState{cfg: &ToolRouterConfig{ + ai := newTestAI(t, []config.OpenAIConfig{{Name: "chat", APIKey: "key", BaseURL: testOpenAIBaseURL, Model: "chat", Timeout: 1, Active: true}}) + router, err := toolrouter.NewState(&config.ToolRouterConfig{ Enabled: true, Timeout: 1, MaxTokens: 128, SystemPrompt: "use tools", - Tools: []ToolRouteConfig{{Name: "time", Enabled: true}}, - }} - toolRouterState.complete = func(ctx context.Context, profile *OpenAIProfile, req model.CreateChatCompletionRequest, timeout time.Duration) (model.ChatCompletionResponse, error) { + Tools: []config.ToolRouteConfig{{Name: "time", Enabled: true}}, + }, ai, toolrouter.WithCompleter(func(ctx context.Context, profile *llm.Profile, req model.CreateChatCompletionRequest, timeout time.Duration) (model.ChatCompletionResponse, error) { for _, msg := range req.Messages { if msg.Content != nil && len(msg.Content.ListValue) > 0 { t.Fatalf("tool decision should not receive multimodal content: %#v", msg.Content) @@ -313,10 +322,13 @@ func TestRunAgentToolLoopImageUsesTextOnlyDecisionMessages(t *testing.T) { if !strings.Contains(joined, "工具判断阶段不读取图片内容") { t.Fatalf("missing placeholder in decision messages: %q", joined) } - return model.ChatCompletionResponse{Choices: []*model.ChatCompletionChoice{{Message: model.ChatCompletionMessage{Content: stringContent("no tool")}}}}, nil + return model.ChatCompletionResponse{Choices: []*model.ChatCompletionChoice{{Message: model.ChatCompletionMessage{Content: message.StringContent("no tool")}}}}, nil + })) + if err != nil { + t.Fatal(err) } - messages, err := runAgentToolLoop(context.Background(), &OpenAIProfile{Config: OpenAIConfig{Model: "chat"}}, []ChatMessage{{Role: "user", Content: "描述这张图", ImageURL: "data:image/png;base64,aGVsbG8="}}, nil) + messages, err := toolrouter.RunAgentToolLoop(context.Background(), router, ai.ActiveProfile(), []message.ChatMessage{{Role: "user", Content: "描述这张图", ImageURL: "data:image/png;base64,aGVsbG8="}}, nil, nil, nil) if err != nil { t.Fatal(err) } @@ -332,32 +344,28 @@ func TestRunAgentToolLoopImageUsesTextOnlyDecisionMessages(t *testing.T) { } func TestRunAgentToolLoopUsesConfiguredRouterProfile(t *testing.T) { - oldRouter := toolRouterState - defer func() { toolRouterState = oldRouter }() - - ai, err := NewOpenAIState([]OpenAIConfig{ - {Name: "chat", APIKey: "key", BaseURL: defaultOpenAIBaseURL, Model: "chat-model", Timeout: 1, Active: true}, - {Name: "router", APIKey: "key", BaseURL: defaultOpenAIBaseURL, Model: "router-model", Timeout: 1}, + ai := newTestAI(t, []config.OpenAIConfig{ + {Name: "chat", APIKey: "key", BaseURL: testOpenAIBaseURL, Model: "chat-model", Timeout: 1, Active: true}, + {Name: "router", APIKey: "key", BaseURL: testOpenAIBaseURL, Model: "router-model", Timeout: 1}, }) - if err != nil { - t.Fatal(err) - } - toolRouterState = &ToolRouterState{ai: ai, cfg: &ToolRouterConfig{ + router, err := toolrouter.NewState(&config.ToolRouterConfig{ Enabled: true, OpenAIName: "router", Timeout: 1, MaxTokens: 128, SystemPrompt: "use tools", - Tools: []ToolRouteConfig{{Name: "time", Enabled: true}}, - }} - toolRouterState.complete = func(ctx context.Context, profile *OpenAIProfile, req model.CreateChatCompletionRequest, timeout time.Duration) (model.ChatCompletionResponse, error) { + Tools: []config.ToolRouteConfig{{Name: "time", Enabled: true}}, + }, ai, toolrouter.WithCompleter(func(ctx context.Context, profile *llm.Profile, req model.CreateChatCompletionRequest, timeout time.Duration) (model.ChatCompletionResponse, error) { if profile.Config.Name != "router" || req.Model != "router-model" { t.Fatalf("router profile not used: profile=%s model=%s", profile.Config.Name, req.Model) } - return model.ChatCompletionResponse{Choices: []*model.ChatCompletionChoice{{Message: model.ChatCompletionMessage{Content: stringContent("no tool")}}}}, nil + return model.ChatCompletionResponse{Choices: []*model.ChatCompletionChoice{{Message: model.ChatCompletionMessage{Content: message.StringContent("no tool")}}}}, nil + })) + if err != nil { + t.Fatal(err) } - _, err = runAgentToolLoop(context.Background(), &OpenAIProfile{Config: OpenAIConfig{Name: "chat", Model: "chat-model"}}, []ChatMessage{{Role: "user", Content: "今天"}}, nil) + _, err = toolrouter.RunAgentToolLoop(context.Background(), router, ai.ActiveProfile(), []message.ChatMessage{{Role: "user", Content: "今天"}}, nil, nil, nil) if err != nil { t.Fatal(err) } diff --git a/message/ark.go b/message/ark.go new file mode 100644 index 0000000..88c1994 --- /dev/null +++ b/message/ark.go @@ -0,0 +1,124 @@ +package message + +import ( + "errors" + "net/url" + "strings" + + "github.com/volcengine/volcengine-go-sdk/service/arkruntime/model" +) + +func BuildArkMessages(chatMessages []ChatMessage) ([]*model.ChatCompletionMessage, error) { + messages := make([]*model.ChatCompletionMessage, 0, len(chatMessages)) + for _, m := range chatMessages { + msg, err := BuildArkMessage(m) + if err != nil { + return nil, err + } + messages = append(messages, msg) + } + return messages, nil +} + +func HasImageMessage(messages []ChatMessage) bool { + for _, msg := range messages { + if strings.TrimSpace(msg.ImageURL) != "" || strings.TrimSpace(msg.ImageURLAlias) != "" { + return true + } + } + return false +} + +func BuildToolDecisionMessages(chatMessages []ChatMessage) ([]*model.ChatCompletionMessage, error) { + messages := make([]*model.ChatCompletionMessage, 0, len(chatMessages)) + for _, m := range chatMessages { + content := m.Content + if strings.TrimSpace(m.ImageURL) != "" || strings.TrimSpace(m.ImageURLAlias) != "" { + content = strings.TrimSpace(content) + placeholder := "[用户上传了一张图片。工具判断阶段不读取图片内容;如果问题主要依赖识图,应不要调用工具,交给最终多模态模型回答。]" + if content == "" { + content = placeholder + } else { + content += "\n\n" + placeholder + } + } + messages = append(messages, &model.ChatCompletionMessage{Role: m.Role, Content: StringContent(content)}) + } + return messages, nil +} + +func BuildArkMessage(m ChatMessage) (*model.ChatCompletionMessage, error) { + msg := &model.ChatCompletionMessage{Role: m.Role} + + if m.ImageURL == "" && m.ImageURLAlias != "" { + m.ImageURL = m.ImageURLAlias + } + + if m.ImageURL == "" { + msg.Content = &model.ChatCompletionMessageContent{ + StringValue: &m.Content, + } + return msg, nil + } + + imageURL, err := NormalizeImageURL(m.ImageURL) + if err != nil { + return nil, err + } + + // 有图片时:文字内容可有可无(图片 caption 场景),均构造多模态消息 + // 若无文字,则只传图片 part;若同时有图片和文字,先文后图 + parts := make([]*model.ChatCompletionMessageContentPart, 0, 2) + if m.Content != "" { + parts = append(parts, TextPart(m.Content)) + } + parts = append(parts, ImagePart(imageURL)) + msg.Content = &model.ChatCompletionMessageContent{ListValue: parts} + return msg, nil +} + +func ImagePart(url string) *model.ChatCompletionMessageContentPart { + return &model.ChatCompletionMessageContentPart{ + Type: model.ChatCompletionMessageContentPartTypeImageURL, + ImageURL: &model.ChatMessageImageURL{ + URL: url, + Detail: model.ImageURLDetailAuto, + }, + } +} + +func TextPart(text string) *model.ChatCompletionMessageContentPart { + return &model.ChatCompletionMessageContentPart{ + Type: model.ChatCompletionMessageContentPartTypeText, + Text: text, + } +} + +func StringContent(text string) *model.ChatCompletionMessageContent { + return &model.ChatCompletionMessageContent{StringValue: &text} +} + +func ChatMessageContentString(content *model.ChatCompletionMessageContent) string { + if content == nil || content.StringValue == nil { + return "" + } + return *content.StringValue +} + +func NormalizeImageURL(raw string) (string, error) { + raw = strings.TrimSpace(raw) + if raw == "" { + return "", errors.New("图片地址不能为空") + } + + lower := strings.ToLower(raw) + if strings.HasPrefix(lower, "data:") { + return normalizeImageDataURI(raw) + } + + u, err := url.Parse(raw) + if err != nil || u.Host == "" || (u.Scheme != "http" && u.Scheme != "https") { + return "", errors.New("图片地址无效,仅支持 http/https URL 或 base64 data URI") + } + return raw, nil +} diff --git a/message/image.go b/message/image.go new file mode 100644 index 0000000..85cd887 --- /dev/null +++ b/message/image.go @@ -0,0 +1,50 @@ +package message + +import ( + "encoding/base64" + "errors" + "strings" + + "aichat/utils" +) + +const maxImageSize = 4 * 1024 * 1024 + +var allowedImageTypes = map[string]bool{ + "image/jpeg": true, + "image/png": true, + "image/webp": true, + "image/gif": true, +} + +func normalizeImageDataURI(raw string) (string, error) { + comma := strings.Index(raw, ",") + if comma < 0 { + return "", errors.New("图片 base64 数据格式错误") + } + + meta := strings.ToLower(strings.TrimSpace(raw[5:comma])) + payload := strings.TrimSpace(raw[comma+1:]) + if payload == "" { + return "", errors.New("图片 base64 数据不能为空") + } + parts := strings.Split(meta, ";") + if len(parts) < 2 || !utils.Contains(parts[1:], "base64") { + return "", errors.New("图片 data URI 必须使用 base64 编码") + } + + mime := parts[0] + if !allowedImageTypes[mime] { + return "", errors.New("图片格式不支持,仅支持 jpeg/png/webp/gif") + } + + decoded, err := base64.StdEncoding.DecodeString(payload) + if err != nil { + return "", errors.New("图片 base64 数据无效") + } + if len(decoded) > maxImageSize { + return "", errors.New("图片过大,请选择小于 4MB 的图片") + } + + return "data:" + mime + ";base64," + payload, nil +} diff --git a/message/types.go b/message/types.go new file mode 100644 index 0000000..448bf38 --- /dev/null +++ b/message/types.go @@ -0,0 +1,26 @@ +package message + +import "time" + +type ChatMessage struct { + Role string `json:"role"` + Content string `json:"content"` + ImageURL string `json:"image_url,omitempty"` // base64 data URI 或 http URL + ImageURLAlias string `json:"imageURL,omitempty"` + Hidden bool `json:"hidden,omitempty"` +} + +type ChatRequest struct { + ConversationID string `json:"conversation_id,omitempty"` + Messages []ChatMessage `json:"messages"` + WebSearch bool `json:"web_search,omitempty"` + OpenAIName string `json:"openai_name,omitempty"` +} + +type Conversation struct { + ID string `json:"id"` + Title string `json:"title"` + CreatedAt time.Time `json:"created_at"` + UpdatedAt time.Time `json:"updated_at"` + Messages []ChatMessage `json:"messages,omitempty"` +} diff --git a/server/handlers.go b/server/handlers.go new file mode 100644 index 0000000..04afe55 --- /dev/null +++ b/server/handlers.go @@ -0,0 +1,298 @@ +package server + +import ( + "context" + "errors" + "fmt" + "io" + "net/http" + "os" + "strings" + "time" + + "aichat/conversation" + "aichat/llm" + "aichat/message" + "aichat/stream" + "aichat/toolrouter" + "aichat/utils" + + "github.com/gin-gonic/gin" + "github.com/volcengine/volcengine-go-sdk/service/arkruntime/model" +) + +type activeProfileRequest struct { + Name string `json:"name"` +} + +func (s *Server) indexHandler(c *gin.Context) { + profile := s.aiState.ActiveProfile() + c.HTML(http.StatusOK, "chat.html", gin.H{ + "Title": "AI 对话", + "Model": profile.Config.Model, + "OpenAIName": profile.Config.Name, + }) +} + +func (s *Server) listOpenAIHandler(c *gin.Context) { + c.JSON(http.StatusOK, s.aiState.ListProfiles()) +} + +func (s *Server) switchOpenAIHandler(c *gin.Context) { + var req activeProfileRequest + if err := c.ShouldBindJSON(&req); err != nil { + c.JSON(http.StatusBadRequest, gin.H{"error": "请求格式错误: " + err.Error()}) + return + } + profile, err := s.aiState.SwitchActive(req.Name) + if err != nil { + c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()}) + return + } + c.JSON(http.StatusOK, gin.H{ + "active": profile.Config.Name, + "profile": llm.PublicConfig(profile, true), + }) +} + +func (s *Server) listSearchHandler(c *gin.Context) { + c.JSON(http.StatusOK, s.searchState.ListProfiles()) +} + +func (s *Server) switchSearchHandler(c *gin.Context) { + var req activeProfileRequest + if err := c.ShouldBindJSON(&req); err != nil { + c.JSON(http.StatusBadRequest, gin.H{"error": "请求格式错误: " + err.Error()}) + return + } + profile, err := s.searchState.SwitchActive(req.Name) + if err != nil { + c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()}) + return + } + profile.APIKey = "" + profile.Active = true + c.JSON(http.StatusOK, gin.H{ + "active": profile.Name, + "profile": profile, + }) +} + +func (s *Server) listConversationsHandler(c *gin.Context) { + convs, err := s.store.List() + if err != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()}) + return + } + c.JSON(http.StatusOK, convs) +} + +func (s *Server) createConversationHandler(c *gin.Context) { + conv, err := s.store.Create() + if err != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": "创建对话失败: " + err.Error()}) + return + } + c.JSON(http.StatusOK, conv) +} + +func (s *Server) getConversationHandler(c *gin.Context) { + conv, err := s.store.Get(c.Param("id")) + if err != nil { + status := http.StatusInternalServerError + if err.Error() == "对话不存在" { + status = http.StatusNotFound + } + c.JSON(status, gin.H{"error": err.Error()}) + return + } + c.JSON(http.StatusOK, conv) +} + +func (s *Server) deleteConversationHandler(c *gin.Context) { + if err := s.store.Delete(c.Param("id")); err != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()}) + return + } + c.Status(http.StatusNoContent) +} + +// chatHandler 流式 SSE 对话接口 +func (s *Server) chatHandler(c *gin.Context) { + var req message.ChatRequest + if err := c.ShouldBindJSON(&req); err != nil { + c.JSON(http.StatusBadRequest, gin.H{"error": "请求格式错误: " + err.Error()}) + return + } + if len(req.Messages) == 0 { + c.JSON(http.StatusBadRequest, gin.H{"error": "消息不能为空"}) + return + } + profile, err := s.aiState.GetProfile(req.OpenAIName) + if err != nil { + c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()}) + return + } + + // SSE 头先写出,后续插件/模型过程都通过 trace 事件实时展示。 + c.Writer.Header().Set("Content-Type", "text/event-stream") + c.Writer.Header().Set("Cache-Control", "no-cache") + c.Writer.Header().Set("Connection", "keep-alive") + c.Writer.Header().Set("X-Accel-Buffering", "no") + c.Writer.WriteHeader(http.StatusOK) + flusher, ok := c.Writer.(http.Flusher) + if !ok { + return + } + emit := func(frame stream.Frame) { + stream.WriteSSEJSON(c.Writer, frame) + flusher.Flush() + } + emitTrace := func(tool, stage, status, message string, data map[string]any) { + emit(stream.Frame{Type: "trace", Tool: tool, Stage: stage, Status: status, Message: message, Data: data}) + } + emitError := func(err error) { + emit(stream.Frame{Type: "error", Error: err.Error()}) + } + + // 超时 context + timeout := time.Duration(profile.Config.Timeout) * time.Second + ctx, cancel := context.WithTimeout(c.Request.Context(), timeout) + defer cancel() + usage := stream.NewTracker() + ctx = stream.ContextWithTracker(ctx, usage) + + // 用 Function Calling 工具循环替代旧的路由+隐藏上下文机制 + messages, err := toolrouter.RunAgentToolLoop(ctx, s.toolRouterState, profile, req.Messages, s.searchState, s.sqlState, emit) + if err != nil { + fmt.Fprintln(os.Stderr, "Agent 工具循环失败:", err) + messages, err = message.BuildArkMessages(req.Messages) + if err != nil { + emitError(err) + return + } + } + promptTokens := stream.EstimateChatMessagesTokens(req.Messages) + + if llm.IsOllamaProfile(profile) && message.HasImageMessage(req.Messages) { + emitTrace("model", "request", "running", "正在通过 Ollama 原生接口调用视觉模型", nil) + err = stream.StreamOllamaChat(ctx, profile, messages, promptTokens, usage, emit, func(content string) { + if req.ConversationID != "" { + if err := conversation.SaveMessages(s.store, req.ConversationID, req.Messages, content); err != nil { + fmt.Fprintln(os.Stderr, "保存对话失败:", err) + } + } + }) + if err != nil { + emitError(err) + } + return + } + + emitTrace("model", "request", "running", "正在调用模型生成回答", nil) + modelStream, err := profile.Client.CreateChatCompletionStream(ctx, model.CreateChatCompletionRequest{ + Model: profile.Config.Model, + Messages: messages, + MaxTokens: utils.IntPtr(4096), + }.WithStream(true)) + if err != nil { + emitError(err) + return + } + defer modelStream.Close() + emitTrace("model", "stream", "running", "模型已开始输出", nil) + + var full strings.Builder + completionTokens := 0 + streamStarted := time.Now() + windowStarted := streamStarted + windowTokens := 0 + peakTokensPerSecond := 0.0 + parseThinkTags := llm.ShouldParseThinkTags(profile) + thinkParser := &stream.Parser{} + emitDelta := func(delta string) { + if delta == "" { + return + } + now := time.Now() + deltaTokens := stream.EstimateTokenCount(delta) + windowTokens += deltaTokens + windowElapsed := now.Sub(windowStarted).Seconds() + if windowElapsed >= 1 { + windowSpeed := float64(windowTokens) / windowElapsed + if windowSpeed > peakTokensPerSecond { + peakTokensPerSecond = windowSpeed + } + windowStarted = now + windowTokens = 0 + } else if peakTokensPerSecond == 0 && windowElapsed > 0.25 { + peakTokensPerSecond = float64(windowTokens) / windowElapsed + } + full.WriteString(delta) + completionTokens += deltaTokens + usage.SetModel(promptTokens, completionTokens) + stats := usage.Snapshot(stream.TokensPerSecond(completionTokens, streamStarted), peakTokensPerSecond) + emit(stream.Frame{Type: "delta", Text: delta, Stats: &stats}) + } + emitModelContent := func(delta string) { + if delta == "" { + return + } + if !parseThinkTags { + emitDelta(delta) + return + } + visible, reasoning := thinkParser.Accept(delta) + if reasoning != "" { + emit(stream.Frame{Type: "reasoning", Text: reasoning}) + } + emitDelta(visible) + } + for { + resp, err := modelStream.Recv() + if errors.Is(err, io.EOF) { + if parseThinkTags { + visible, reasoning := thinkParser.Flush() + if reasoning != "" { + emit(stream.Frame{Type: "reasoning", Text: reasoning}) + } + emitDelta(visible) + } + usage.SetModel(promptTokens, completionTokens) + if windowTokens > 0 { + windowElapsed := time.Since(windowStarted).Seconds() + if windowElapsed > 0.25 { + windowSpeed := float64(windowTokens) / windowElapsed + if windowSpeed > peakTokensPerSecond { + peakTokensPerSecond = windowSpeed + } + } + } + if peakTokensPerSecond == 0 { + peakTokensPerSecond = stream.TokensPerSecond(completionTokens, streamStarted) + } + if req.ConversationID != "" { + if err := conversation.SaveMessages(s.store, req.ConversationID, req.Messages, full.String()); err != nil { + fmt.Fprintln(os.Stderr, "保存对话失败:", err) + } + } + finalStats := usage.Snapshot(stream.TokensPerSecond(completionTokens, streamStarted), peakTokensPerSecond) + emit(stream.Frame{Type: "stats", Stats: &finalStats}) + emitTrace("model", "stream", "success", "回答生成完成", nil) + fmt.Fprintf(c.Writer, "data: [DONE]\n\n") + flusher.Flush() + return + } + if err != nil { + emitError(err) + return + } + if len(resp.Choices) > 0 { + emitModelContent(resp.Choices[0].Delta.Content) + // 思考过程 reasoning_content 单独事件推送 + if resp.Choices[0].Delta.ReasoningContent != nil && *resp.Choices[0].Delta.ReasoningContent != "" { + emit(stream.Frame{Type: "reasoning", Text: *resp.Choices[0].Delta.ReasoningContent}) + } + } + } +} diff --git a/server/routes.go b/server/routes.go new file mode 100644 index 0000000..6986274 --- /dev/null +++ b/server/routes.go @@ -0,0 +1,14 @@ +package server + +func (s *Server) registerRoutes() { + s.router.GET("/", s.indexHandler) + s.router.POST("/api/chat", s.chatHandler) + s.router.GET("/api/openai", s.listOpenAIHandler) + s.router.POST("/api/openai/active", s.switchOpenAIHandler) + s.router.GET("/api/search", s.listSearchHandler) + s.router.POST("/api/search/active", s.switchSearchHandler) + s.router.GET("/api/conversations", s.listConversationsHandler) + s.router.POST("/api/conversations", s.createConversationHandler) + s.router.GET("/api/conversations/:id", s.getConversationHandler) + s.router.DELETE("/api/conversations/:id", s.deleteConversationHandler) +} diff --git a/server/server.go b/server/server.go new file mode 100644 index 0000000..4dacf74 --- /dev/null +++ b/server/server.go @@ -0,0 +1,65 @@ +package server + +import ( + "fmt" + "net" + "net/http" + "os" + + searchagent "aichat/agents/search" + sqlquery "aichat/agents/sql" + "aichat/config" + "aichat/conversation" + "aichat/llm" + "aichat/toolrouter" + + "github.com/gin-gonic/gin" +) + +type Server struct { + cfg *config.Config + aiState *llm.State + searchState *searchagent.State + sqlState *sqlquery.State + toolRouterState *toolrouter.State + store *conversation.Store + router *gin.Engine +} + +func New(cfg *config.Config, aiState *llm.State, searchState *searchagent.State, sqlState *sqlquery.State, toolRouterState *toolrouter.State, store *conversation.Store) *Server { + s := &Server{ + cfg: cfg, + aiState: aiState, + searchState: searchState, + sqlState: sqlState, + toolRouterState: toolRouterState, + store: store, + router: gin.Default(), + } + s.router.LoadHTMLGlob("templates/*") + s.router.Static("/static", "./static") + s.registerRoutes() + return s +} + +func (s *Server) Run() error { + switch s.cfg.Server.Mode { + case "unix": + return s.runUnix(s.cfg.Server.Address) + default: + fmt.Println("服务已启动,监听 TCP:", s.cfg.Server.Address) + return s.router.Run(s.cfg.Server.Address) + } +} + +func (s *Server) runUnix(socketPath string) error { + if _, statErr := os.Stat(socketPath); statErr == nil { + os.Remove(socketPath) + } + ln, err := net.Listen("unix", socketPath) + if err != nil { + return fmt.Errorf("监听 Unix socket 失败: %w", err) + } + fmt.Println("服务已启动,监听 Unix socket:", socketPath) + return http.Serve(ln, s.router) +} diff --git a/stream/ollama.go b/stream/ollama.go new file mode 100644 index 0000000..7ba4fdb --- /dev/null +++ b/stream/ollama.go @@ -0,0 +1,219 @@ +package stream + +import ( + "bufio" + "bytes" + "context" + "encoding/json" + "errors" + "fmt" + "io" + "net/http" + "strings" + "time" + + "aichat/llm" + + "github.com/volcengine/volcengine-go-sdk/service/arkruntime/model" +) + +type ollamaChatRequest struct { + Model string `json:"model"` + Messages []ollamaChatMessage `json:"messages"` + Stream bool `json:"stream"` + Options map[string]int `json:"options,omitempty"` +} + +type ollamaChatMessage struct { + Role string `json:"role"` + Content string `json:"content"` + Images []string `json:"images,omitempty"` +} + +type ollamaChatResponse struct { + Message struct { + Role string `json:"role"` + Content string `json:"content"` + Thinking string `json:"thinking"` + } `json:"message"` + Done bool `json:"done"` + PromptEvalCount int `json:"prompt_eval_count"` + EvalCount int `json:"eval_count"` + DoneReason string `json:"done_reason"` +} + +func StreamOllamaChat(ctx context.Context, profile *llm.Profile, messages []*model.ChatCompletionMessage, promptTokens int, usage *Tracker, emit EmitFunc, onDone func(string)) error { + requestMessages, err := buildOllamaMessages(messages) + if err != nil { + return err + } + baseURL, err := llm.OllamaBaseURL(profile) + if err != nil { + return err + } + body, err := json.Marshal(ollamaChatRequest{ + Model: profile.Config.Model, + Messages: requestMessages, + Stream: true, + Options: map[string]int{"num_predict": 4096}, + }) + if err != nil { + return err + } + req, err := http.NewRequestWithContext(ctx, http.MethodPost, strings.TrimRight(baseURL, "/")+"/api/chat", bytes.NewReader(body)) + if err != nil { + return err + } + req.Header.Set("Content-Type", "application/json") + resp, err := http.DefaultClient.Do(req) + if err != nil { + return err + } + defer resp.Body.Close() + if resp.StatusCode < 200 || resp.StatusCode >= 300 { + data, _ := io.ReadAll(io.LimitReader(resp.Body, 4096)) + return fmt.Errorf("Ollama 原生接口调用失败: %s %s", resp.Status, strings.TrimSpace(string(data))) + } + + emit(Frame{Type: "trace", Tool: "model", Stage: "stream", Status: "running", Message: "Ollama 视觉模型已开始输出"}) + parseThinkTags := llm.ShouldParseThinkTags(profile) + thinkParser := &Parser{} + var full strings.Builder + completionTokens := 0 + streamStarted := time.Now() + peakTokensPerSecond := 0.0 + emitDelta := func(delta string) { + if delta == "" { + return + } + full.WriteString(delta) + completionTokens += EstimateTokenCount(delta) + usage.SetModel(promptTokens, completionTokens) + currentSpeed := TokensPerSecond(completionTokens, streamStarted) + if currentSpeed > peakTokensPerSecond { + peakTokensPerSecond = currentSpeed + } + stats := usage.Snapshot(currentSpeed, peakTokensPerSecond) + emit(Frame{Type: "delta", Text: delta, Stats: &stats}) + } + emitContent := func(delta string) { + if delta == "" { + return + } + if !parseThinkTags { + emitDelta(delta) + return + } + visible, reasoning := thinkParser.Accept(delta) + if reasoning != "" { + emit(Frame{Type: "reasoning", Text: reasoning}) + } + emitDelta(visible) + } + + scanner := bufio.NewScanner(resp.Body) + scanner.Buffer(make([]byte, 0, 64*1024), 10*1024*1024) + for scanner.Scan() { + line := strings.TrimSpace(scanner.Text()) + if line == "" { + continue + } + var chunk ollamaChatResponse + if err := json.Unmarshal([]byte(line), &chunk); err != nil { + return fmt.Errorf("解析 Ollama 流失败: %w", err) + } + if chunk.Message.Thinking != "" { + emit(Frame{Type: "reasoning", Text: chunk.Message.Thinking}) + } + emitContent(chunk.Message.Content) + if chunk.Done { + if chunk.PromptEvalCount > 0 || chunk.EvalCount > 0 { + usage.SetModel(chunk.PromptEvalCount, chunk.EvalCount) + } + break + } + } + if err := scanner.Err(); err != nil { + return err + } + if parseThinkTags { + visible, reasoning := thinkParser.Flush() + if reasoning != "" { + emit(Frame{Type: "reasoning", Text: reasoning}) + } + emitDelta(visible) + } + if onDone != nil { + onDone(full.String()) + } + finalStats := usage.Snapshot(TokensPerSecond(completionTokens, streamStarted), peakTokensPerSecond) + emit(Frame{Type: "stats", Stats: &finalStats}) + emit(Frame{Type: "trace", Tool: "model", Stage: "stream", Status: "success", Message: "回答生成完成"}) + return nil +} + +func buildOllamaMessages(messages []*model.ChatCompletionMessage) ([]ollamaChatMessage, error) { + result := make([]ollamaChatMessage, 0, len(messages)) + for _, msg := range messages { + if msg == nil { + continue + } + role := string(msg.Role) + if msg.Role == model.ChatMessageRoleTool { + role = string(model.ChatMessageRoleUser) + } + item := ollamaChatMessage{Role: role} + if msg.Content == nil { + if len(msg.ToolCalls) > 0 { + continue + } + result = append(result, item) + continue + } + if msg.Content.StringValue != nil { + item.Content = *msg.Content.StringValue + if msg.Role == model.ChatMessageRoleTool { + item.Content = "工具结果:\n" + item.Content + } + result = append(result, item) + continue + } + for _, part := range msg.Content.ListValue { + if part == nil { + continue + } + switch part.Type { + case model.ChatCompletionMessageContentPartTypeText: + if part.Text != "" { + if item.Content != "" { + item.Content += "\n" + } + item.Content += part.Text + } + case model.ChatCompletionMessageContentPartTypeImageURL: + if part.ImageURL == nil { + continue + } + image, err := ollamaImagePayload(part.ImageURL.URL) + if err != nil { + return nil, err + } + item.Images = append(item.Images, image) + } + } + result = append(result, item) + } + return result, nil +} + +func ollamaImagePayload(raw string) (string, error) { + raw = strings.TrimSpace(raw) + if strings.HasPrefix(strings.ToLower(raw), "data:") { + comma := strings.Index(raw, ",") + if comma < 0 { + return "", errors.New("图片 base64 数据格式错误") + } + return strings.TrimSpace(raw[comma+1:]), nil + } + return raw, nil +} diff --git a/stream/sse.go b/stream/sse.go new file mode 100644 index 0000000..15ef853 --- /dev/null +++ b/stream/sse.go @@ -0,0 +1,29 @@ +package stream + +import ( + "encoding/json" + "fmt" + "io" +) + +type Frame struct { + Type string `json:"type"` + Text string `json:"text,omitempty"` + Message string `json:"message,omitempty"` + Tool string `json:"tool,omitempty"` + Stage string `json:"stage,omitempty"` + Status string `json:"status,omitempty"` + Data map[string]any `json:"data,omitempty"` + Stats *Stats `json:"stats,omitempty"` + Error string `json:"error,omitempty"` +} + +type EmitFunc func(Frame) + +func WriteSSEJSON(w io.Writer, frame Frame) { + data, err := json.Marshal(frame) + if err != nil { + data, _ = json.Marshal(Frame{Type: "error", Error: "序列化流事件失败"}) + } + fmt.Fprintf(w, "data: %s\n\n", data) +} diff --git a/stream/think.go b/stream/think.go new file mode 100644 index 0000000..88e3476 --- /dev/null +++ b/stream/think.go @@ -0,0 +1,73 @@ +package stream + +import "strings" + +type Parser struct { + inThink bool + buffer string +} + +const ( + thinkOpenTag = "" + thinkCloseTag = "" +) + +func (p *Parser) Accept(delta string) (visible string, reasoning string) { + p.buffer += delta + for p.buffer != "" { + if p.inThink { + idx := strings.Index(p.buffer, thinkCloseTag) + if idx >= 0 { + reasoning += p.buffer[:idx] + p.buffer = p.buffer[idx+len(thinkCloseTag):] + p.inThink = false + continue + } + keep := tagPrefixSuffixLen(p.buffer, thinkCloseTag) + if len(p.buffer) > keep { + reasoning += p.buffer[:len(p.buffer)-keep] + p.buffer = p.buffer[len(p.buffer)-keep:] + } + return visible, reasoning + } + + idx := strings.Index(p.buffer, thinkOpenTag) + if idx >= 0 { + visible += p.buffer[:idx] + p.buffer = p.buffer[idx+len(thinkOpenTag):] + p.inThink = true + continue + } + keep := tagPrefixSuffixLen(p.buffer, thinkOpenTag) + if len(p.buffer) > keep { + visible += p.buffer[:len(p.buffer)-keep] + p.buffer = p.buffer[len(p.buffer)-keep:] + } + return visible, reasoning + } + return visible, reasoning +} + +func (p *Parser) Flush() (visible string, reasoning string) { + if p.inThink { + reasoning = p.buffer + } else { + visible = p.buffer + } + p.buffer = "" + p.inThink = false + return visible, reasoning +} + +func tagPrefixSuffixLen(text, tag string) int { + limit := len(tag) - 1 + if len(text) < limit { + limit = len(text) + } + for i := limit; i > 0; i-- { + if strings.HasPrefix(tag, text[len(text)-i:]) { + return i + } + } + return 0 +} diff --git a/stream/tokens.go b/stream/tokens.go new file mode 100644 index 0000000..1f62461 --- /dev/null +++ b/stream/tokens.go @@ -0,0 +1,138 @@ +package stream + +import ( + "context" + "strings" + "sync" + "time" + "unicode" + + "aichat/message" +) + +type Stats struct { + PromptTokens int `json:"prompt_tokens"` + CompletionTokens int `json:"completion_tokens"` + ToolPromptTokens int `json:"tool_prompt_tokens"` + ToolCompletionTokens int `json:"tool_completion_tokens"` + TotalTokens int `json:"total_tokens"` + CompletionTokensPerSec float64 `json:"completion_tokens_per_sec"` + PeakCompletionTokensPerSec float64 `json:"peak_completion_tokens_per_sec"` + Estimated bool `json:"estimated"` +} + +type Tracker struct { + mu sync.Mutex + promptTokens int + completionTokens int + toolPromptTokens int + toolCompletionTokens int +} + +type tokenUsageContextKey struct{} + +func NewTracker() *Tracker { + return &Tracker{} +} + +func ContextWithTracker(ctx context.Context, tracker *Tracker) context.Context { + if tracker == nil { + return ctx + } + return context.WithValue(ctx, tokenUsageContextKey{}, tracker) +} + +func TrackerFromContext(ctx context.Context) *Tracker { + tracker, _ := ctx.Value(tokenUsageContextKey{}).(*Tracker) + return tracker +} + +func (t *Tracker) AddTool(promptTokens, completionTokens int) { + if t == nil { + return + } + t.mu.Lock() + defer t.mu.Unlock() + t.toolPromptTokens += promptTokens + t.toolCompletionTokens += completionTokens +} + +func (t *Tracker) SetModel(promptTokens, completionTokens int) { + if t == nil { + return + } + t.mu.Lock() + defer t.mu.Unlock() + t.promptTokens = promptTokens + t.completionTokens = completionTokens +} + +func (t *Tracker) Snapshot(tokensPerSecond, peakTokensPerSecond float64) Stats { + if t == nil { + return Stats{Estimated: true} + } + t.mu.Lock() + defer t.mu.Unlock() + total := t.promptTokens + t.completionTokens + t.toolPromptTokens + t.toolCompletionTokens + return Stats{ + PromptTokens: t.promptTokens, + CompletionTokens: t.completionTokens, + ToolPromptTokens: t.toolPromptTokens, + ToolCompletionTokens: t.toolCompletionTokens, + TotalTokens: total, + CompletionTokensPerSec: tokensPerSecond, + PeakCompletionTokensPerSec: peakTokensPerSecond, + Estimated: true, + } +} + +func EstimateChatMessagesTokens(messages []message.ChatMessage) int { + total := 0 + for _, msg := range messages { + total += EstimateTokenCount(msg.Role) + EstimateTokenCount(msg.Content) + 4 + if msg.ImageURL != "" || msg.ImageURLAlias != "" { + total += 85 + } + } + return total +} + +func EstimateTokenCount(text string) int { + text = strings.TrimSpace(text) + if text == "" { + return 0 + } + tokens := 0 + asciiRunes := 0 + flushASCII := func() { + if asciiRunes > 0 { + tokens += (asciiRunes + 3) / 4 + asciiRunes = 0 + } + } + for _, r := range text { + if unicode.IsSpace(r) { + flushASCII() + continue + } + if r <= unicode.MaxASCII { + asciiRunes++ + continue + } + flushASCII() + tokens++ + } + flushASCII() + if tokens == 0 { + return 1 + } + return tokens +} + +func TokensPerSecond(tokens int, start time.Time) float64 { + elapsed := time.Since(start).Seconds() + if tokens <= 0 || elapsed <= 0 { + return 0 + } + return float64(tokens) / elapsed +} diff --git a/toolrouter/loop.go b/toolrouter/loop.go new file mode 100644 index 0000000..a568e33 --- /dev/null +++ b/toolrouter/loop.go @@ -0,0 +1,125 @@ +package toolrouter + +import ( + "context" + "fmt" + "strings" + "time" + + searchagent "aichat/agents/search" + sqlquery "aichat/agents/sql" + "aichat/llm" + "aichat/message" + "aichat/stream" + "aichat/utils" + + "github.com/volcengine/volcengine-go-sdk/service/arkruntime/model" +) + +const maxAgentToolIterations = 6 + +func RunAgentToolLoop(ctx context.Context, state *State, profile *llm.Profile, chatMessages []message.ChatMessage, searchState *searchagent.State, sqlState *sqlquery.State, emit stream.EmitFunc) ([]*model.ChatCompletionMessage, error) { + finalMessages, err := message.BuildArkMessages(chatMessages) + if err != nil { + return nil, err + } + routerProfile := profile + if state != nil { + routerProfile = state.RouterProfile(profile) + } + tools := availableAgentTools(state, routerProfile, searchState, sqlState, emit) + if len(tools) == 0 { + return finalMessages, nil + } + decisionMessages := append([]*model.ChatCompletionMessage(nil), finalMessages...) + if message.HasImageMessage(chatMessages) { + decisionMessages, err = message.BuildToolDecisionMessages(chatMessages) + if err != nil { + return nil, err + } + if emit != nil { + emit(stream.Frame{Type: "trace", Tool: "agent_tools", Stage: "prepare", Status: "success", Message: "检测到图片输入,工具判断阶段将使用纯文本上下文"}) + } + } + toolByName := make(map[string]AgentTool, len(tools)) + definitions := make([]*model.Tool, 0, len(tools)) + availableNames := make([]string, 0, len(tools)) + toolDescriptions := make([]string, 0, len(tools)) + for _, tool := range tools { + toolByName[tool.name] = tool + definitions = append(definitions, tool.definition) + availableNames = append(availableNames, tool.name) + if tool.definition != nil && tool.definition.Function != nil { + toolDescriptions = append(toolDescriptions, fmt.Sprintf("%s: %s", tool.name, tool.definition.Function.Description)) + } + } + if emit != nil { + emit(stream.Frame{Type: "trace", Tool: "agent_tools", Stage: "prepare", Status: "success", Message: "已准备可用工具", Data: map[string]any{"tools": availableNames, "tool_descriptions": toolDescriptions}}) + } + if state == nil || state.cfg == nil { + return finalMessages, nil + } + if prompt := strings.TrimSpace(state.cfg.SystemPrompt); prompt != "" { + systemMessage := &model.ChatCompletionMessage{Role: model.ChatMessageRoleSystem, Content: message.StringContent(prompt)} + finalMessages = append([]*model.ChatCompletionMessage{systemMessage}, finalMessages...) + decisionMessages = append([]*model.ChatCompletionMessage{systemMessage}, decisionMessages...) + } + for i := 0; i < maxAgentToolIterations; i++ { + if emit != nil { + emit(stream.Frame{Type: "trace", Tool: "agent_tools", Stage: "request", Status: "running", Message: fmt.Sprintf("正在进行第 %d 轮工具判断", i+1), Data: map[string]any{"iteration": i + 1, "max_iterations": maxAgentToolIterations, "tools": availableNames}}) + } + resp, err := state.complete(ctx, routerProfile, model.CreateChatCompletionRequest{ + Model: routerProfile.Config.Model, + Messages: decisionMessages, + MaxTokens: utils.IntPtr(state.cfg.MaxTokens), + Tools: definitions, + ToolChoice: model.ToolChoiceStringTypeAuto, + ParallelToolCalls: utils.BoolPtr(false), + }, time.Duration(state.cfg.Timeout)*time.Second) + if err != nil { + return finalMessages, err + } + if tracker := stream.TrackerFromContext(ctx); tracker != nil { + tracker.AddTool(resp.Usage.PromptTokens, resp.Usage.CompletionTokens) + } + if len(resp.Choices) == 0 { + return finalMessages, nil + } + choice := resp.Choices[0] + decisionPreview := message.ChatMessageContentString(choice.Message.Content) + if emit != nil { + emit(stream.Frame{Type: "trace", Tool: "agent_tools", Stage: "decision", Status: "success", Message: "工具判断响应已返回", Data: map[string]any{"iteration": i + 1, "finish_reason": string(choice.FinishReason), "content_preview": utils.TruncateString(decisionPreview, 800)}}) + } + calls := choice.Message.ToolCalls + if len(calls) == 0 && choice.Message.FunctionCall != nil { + calls = []*model.ToolCall{{ID: "legacy_function_call", Type: model.ToolTypeFunction, Function: *choice.Message.FunctionCall}} + } + if len(calls) == 0 { + if emit != nil { + emit(stream.Frame{Type: "trace", Tool: "agent_tools", Stage: "request", Status: "success", Message: "模型未请求工具,进入回答生成"}) + } + return finalMessages, nil + } + callNames := make([]string, 0, len(calls)) + for _, call := range calls { + if call != nil { + callNames = append(callNames, call.Function.Name) + } + } + if emit != nil { + emit(stream.Frame{Type: "trace", Tool: "agent_tools", Stage: "tool_calls", Status: "running", Message: fmt.Sprintf("模型请求调用 %d 个工具", len(calls)), Data: map[string]any{"tools": callNames, "iteration": i + 1}}) + } + assistantMessage := &model.ChatCompletionMessage{Role: model.ChatMessageRoleAssistant, ToolCalls: calls, Content: choice.Message.Content} + finalMessages = append(finalMessages, assistantMessage) + decisionMessages = append(decisionMessages, assistantMessage) + for _, call := range calls { + result := ExecuteAgentToolCall(ctx, call, toolByName, emit) + toolMessage := &model.ChatCompletionMessage{Role: model.ChatMessageRoleTool, ToolCallID: call.ID, Content: message.StringContent(result)} + finalMessages = append(finalMessages, toolMessage) + decisionMessages = append(decisionMessages, toolMessage) + } + } + limitMessage := &model.ChatCompletionMessage{Role: model.ChatMessageRoleSystem, Content: message.StringContent("工具调用轮数已达到上限。请基于已有工具结果回答,并说明可能未完成全部工具调用。")} + finalMessages = append(finalMessages, limitMessage) + return finalMessages, nil +} diff --git a/toolrouter/state.go b/toolrouter/state.go new file mode 100644 index 0000000..38106a3 --- /dev/null +++ b/toolrouter/state.go @@ -0,0 +1,62 @@ +package toolrouter + +import ( + "errors" + "fmt" + "strings" + + "aichat/completion" + "aichat/config" + "aichat/llm" +) + +type State struct { + cfg *config.ToolRouterConfig + ai *llm.State + complete completion.ChatCompleter +} + +type Option func(*State) + +func WithCompleter(completer completion.ChatCompleter) Option { + return func(s *State) { + if completer != nil { + s.complete = completer + } + } +} + +func NewState(cfg *config.ToolRouterConfig, ai *llm.State, options ...Option) (*State, error) { + if cfg == nil { + defaultConfig := config.DefaultToolRouterConfig() + cfg = &defaultConfig + } + if ai == nil { + return nil, errors.New("工具路由需要 OpenAI 状态") + } + if cfg.Enabled && strings.TrimSpace(cfg.OpenAIName) != "" { + if _, err := ai.GetProfile(cfg.OpenAIName); err != nil { + return nil, fmt.Errorf("tool_router.openai_name 配置无效: %w", err) + } + } + state := &State{cfg: cfg, ai: ai, complete: completion.CompleteChatWithTimeout} + for _, option := range options { + option(state) + } + return state, nil +} + +func (s *State) RouterProfile(fallback *llm.Profile) *llm.Profile { + if s == nil || s.cfg == nil || s.ai == nil { + return fallback + } + name := strings.TrimSpace(s.cfg.OpenAIName) + if name == "" { + return fallback + } + profile, err := s.ai.GetProfile(name) + if err != nil { + return fallback + } + return profile +} diff --git a/toolrouter/tools.go b/toolrouter/tools.go new file mode 100644 index 0000000..2784dff --- /dev/null +++ b/toolrouter/tools.go @@ -0,0 +1,155 @@ +package toolrouter + +import ( + "context" + "fmt" + "strings" + "time" + + searchagent "aichat/agents/search" + sqlquery "aichat/agents/sql" + timeagent "aichat/agents/time" + "aichat/completion" + "aichat/llm" + "aichat/message" + "aichat/stream" + "aichat/utils" + + "github.com/volcengine/volcengine-go-sdk/service/arkruntime/model" +) + +type AgentTool struct { + name string + definition *model.Tool + execute func(context.Context, string) (string, error) +} + +func NewAgentTool(name string, definition *model.Tool, execute func(context.Context, string) (string, error)) AgentTool { + return AgentTool{name: name, definition: definition, execute: execute} +} + +func (t AgentTool) Name() string { return t.name } + +func (t AgentTool) Definition() *model.Tool { return t.definition } + +func AvailableAgentTools(state *State, profile *llm.Profile, searchState *searchagent.State, sqlState *sqlquery.State, emit stream.EmitFunc) []AgentTool { + return availableAgentTools(state, profile, searchState, sqlState, emit) +} + +func availableAgentTools(state *State, profile *llm.Profile, searchState *searchagent.State, sqlState *sqlquery.State, emit stream.EmitFunc) []AgentTool { + if state == nil || state.cfg == nil || !state.cfg.Enabled { + return nil + } + tools := make([]AgentTool, 0, len(state.cfg.Tools)) + for _, item := range state.cfg.Tools { + if !item.Enabled { + continue + } + description := strings.TrimSpace(item.Description) + switch item.Name { + case timeagent.ToolName: + tools = append(tools, AgentTool{ + name: timeagent.ToolName, + definition: timeagent.ToolDefinition(description), + execute: func(ctx context.Context, args string) (string, error) { + result, err := timeagent.ExecuteTool(args, time.Now()) + if err == nil && emit != nil { + emit(stream.Frame{Type: "trace", Tool: timeagent.ToolName, Stage: "resolve", Status: "success", Message: "已获取当前时间上下文"}) + } + return result, err + }, + }) + case searchagent.ToolName: + if searchState == nil || !searchState.Enabled() { + continue + } + tools = append(tools, AgentTool{ + name: searchagent.ToolName, + definition: searchState.ToolDefinition(description), + execute: func(ctx context.Context, args string) (string, error) { + if emit != nil { + emit(stream.Frame{Type: "trace", Tool: searchagent.ToolName, Stage: "request", Status: "running", Message: "正在联网搜索"}) + } + result, err := searchState.ExecuteTool(ctx, args) + if emit != nil { + status := "success" + messageText := "联网搜索完成" + if err != nil { + status = "error" + messageText = "联网搜索失败" + } + emit(stream.Frame{Type: "trace", Tool: searchagent.ToolName, Stage: "results", Status: status, Message: messageText}) + } + return result, err + }, + }) + case sqlquery.ToolName: + if sqlState == nil || !sqlState.Enabled() { + continue + } + tools = append(tools, AgentTool{ + name: sqlquery.ToolName, + definition: sqlState.ToolDefinition(description), + execute: func(ctx context.Context, args string) (string, error) { + if emit != nil { + emit(stream.Frame{Type: "trace", Tool: sqlquery.ToolName, Stage: "execute", Status: "running", Message: "正在查询数据库"}) + } + generator := func(ctx context.Context, prompt string, maxTokens int) (string, error) { + return completion.CompleteText(ctx, profile, []message.ChatMessage{{Role: "system", Content: prompt}}, maxTokens) + } + result, err := sqlState.ExecuteTool(ctx, args, generator) + if emit != nil { + status := "success" + messageText := "数据库查询完成" + if err != nil { + status = "error" + messageText = "数据库查询失败" + } + emit(stream.Frame{Type: "trace", Tool: sqlquery.ToolName, Stage: "execute", Status: status, Message: messageText}) + } + return result, err + }, + }) + } + } + return tools +} + +func ExecuteAgentToolCall(ctx context.Context, call *model.ToolCall, tools map[string]AgentTool, emit stream.EmitFunc) string { + if call == nil || call.Type != model.ToolTypeFunction { + result := "工具调用无效:仅支持 function 类型工具。" + if emit != nil { + emit(stream.Frame{Type: "trace", Tool: "agent_tools", Stage: "execute", Status: "error", Message: result}) + } + return result + } + toolName := call.Function.Name + if emit != nil { + emit(stream.Frame{Type: "trace", Tool: toolName, Stage: "arguments", Status: "running", Message: "准备执行工具", Data: map[string]any{"tool_call_id": call.ID, "arguments": call.Function.Arguments}}) + } + tool, ok := tools[toolName] + if !ok { + result := fmt.Sprintf("工具调用失败:未知工具 %s。", toolName) + if emit != nil { + emit(stream.Frame{Type: "trace", Tool: toolName, Stage: "execute", Status: "error", Message: result}) + } + return result + } + started := time.Now() + result, err := tool.execute(ctx, call.Function.Arguments) + durationMs := time.Since(started).Milliseconds() + if err != nil { + messageText := fmt.Sprintf("工具 %s 执行失败:%v", tool.name, err) + if emit != nil { + emit(stream.Frame{Type: "trace", Tool: tool.name, Stage: "execute", Status: "error", Message: "工具执行失败", Data: map[string]any{"tool_call_id": call.ID, "duration_ms": durationMs, "error": err.Error()}}) + } + return messageText + } + if strings.TrimSpace(result) == "" { + result = fmt.Sprintf("工具 %s 执行完成,但没有返回内容。", tool.name) + } + if emit != nil { + emit(stream.Frame{Type: "trace", Tool: tool.name, Stage: "result", Status: "success", Message: "工具执行完成", Data: map[string]any{"tool_call_id": call.ID, "duration_ms": durationMs, "result_preview": utils.TruncateString(result, 1200)}}) + } + return result +} diff --git a/utils/helpers.go b/utils/helpers.go new file mode 100644 index 0000000..82f4094 --- /dev/null +++ b/utils/helpers.go @@ -0,0 +1,53 @@ +package utils + +import ( + "crypto/rand" + "encoding/hex" + "encoding/json" + "fmt" + "strings" +) + +func IntPtr(i int) *int { return &i } + +func BoolPtr(v bool) *bool { return &v } + +func TruncateString(text string, maxRunes int) string { + runes := []rune(strings.TrimSpace(text)) + if maxRunes <= 0 || len(runes) <= maxRunes { + return string(runes) + } + return string(runes[:maxRunes]) + "..." +} + +func Contains(items []string, target string) bool { + for _, item := range items { + if strings.TrimSpace(item) == target { + return true + } + } + return false +} + +func NewUUID() string { + b := make([]byte, 16) + _, _ = rand.Read(b) + b[6] = (b[6] & 0x0f) | 0x40 + b[8] = (b[8] & 0x3f) | 0x80 + return hex.EncodeToString(b[:4]) + "-" + hex.EncodeToString(b[4:6]) + "-" + + hex.EncodeToString(b[6:8]) + "-" + hex.EncodeToString(b[8:10]) + "-" + + hex.EncodeToString(b[10:]) +} + +func ToJSON(s string) string { + b, _ := json.Marshal(s) + return string(b) +} + +func ToSSE(s string) string { + s = strings.ReplaceAll(s, `\`, `\\`) + s = strings.ReplaceAll(s, "\n", `\n`) + s = strings.ReplaceAll(s, "\r", "") + s = strings.ReplaceAll(s, `"`, `\"`) + return fmt.Sprintf(`"%s"`, s) +}