拆分 main 功能模块

Co-Authored-By: Claude <noreply@anthropic.com>
This commit is contained in:
2026-06-17 11:29:51 +08:00
co-authored by Claude
parent 132ab2a1cb
commit ccc1260fe0
21 changed files with 2318 additions and 2074 deletions
+84
View File
@@ -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))
}
+367
View File
@@ -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
}
+152
View File
@@ -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 "新对话"
}
+46
View File
@@ -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
}
+131
View File
@@ -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
}
+16 -1991
View File
File diff suppressed because it is too large Load Diff
+91 -83
View File
@@ -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 <think>abc</think> 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 <thi", "nk>abc</thi", "nk> 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 <think>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)
}
+124
View File
@@ -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
}
+50
View File
@@ -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
}
+26
View File
@@ -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"`
}
+298
View File
@@ -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})
}
}
}
}
+14
View File
@@ -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)
}
+65
View File
@@ -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)
}
+219
View File
@@ -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
}
+29
View File
@@ -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)
}
+73
View File
@@ -0,0 +1,73 @@
package stream
import "strings"
type Parser struct {
inThink bool
buffer string
}
const (
thinkOpenTag = "<think>"
thinkCloseTag = "</think>"
)
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
}
+138
View File
@@ -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
}
+125
View File
@@ -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
}
+62
View File
@@ -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
}
+155
View File
@@ -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
}
+53
View File
@@ -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)
}