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)
+}