Files
aichat/main_test.go
T
2026-06-17 13:13:06 +08:00

399 lines
15 KiB
Go

package main
import (
"context"
"errors"
"strings"
"testing"
"time"
"aichat/agents/time"
agents "aichat/agenttool"
"aichat/config"
"aichat/llm"
"aichat/message"
"aichat/stream"
"aichat/toolmanager"
"aichat/toolrouter"
"github.com/volcengine/volcengine-go-sdk/service/arkruntime/model"
)
const testOpenAIBaseURL = "https://ark.cn-beijing.volces.com/api/v3"
func newTestToolManager(tools ...agents.LoadedTool) *toolmanager.Manager {
if len(tools) == 0 {
tools = []agents.LoadedTool{timeagent.NewLoadedTool(nil)}
}
return toolmanager.NewForTest(tools...)
}
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.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")
}
defaults := config.DefaultToolRouterConfig()
if cfg.ToolRouter.Timeout != defaults.Timeout {
t.Fatalf("timeout = %d", cfg.ToolRouter.Timeout)
}
if cfg.ToolRouter.MaxTokens != defaults.MaxTokens {
t.Fatalf("max_tokens = %d", cfg.ToolRouter.MaxTokens)
}
if strings.TrimSpace(cfg.ToolRouter.SystemPrompt) == "" {
t.Fatal("system prompt should be defaulted")
}
if len(cfg.ToolRouter.Tools) != 0 {
t.Fatalf("unexpected tools: %#v", cfg.ToolRouter.Tools)
}
}
func TestNormalizeOpenAIConfigDefaultsContextWindow(t *testing.T) {
cfg := config.Default()
cfg.OpenAI[0].ContextWindowTokens = 0
changed, err := config.NormalizeOpenAIConfigs(&cfg)
if err != nil {
t.Fatal(err)
}
if !changed {
t.Fatal("expected context window default to change config")
}
if cfg.OpenAI[0].ContextWindowTokens != 262144 {
t.Fatalf("context_window_tokens = %d", cfg.OpenAI[0].ContextWindowTokens)
}
}
func TestNormalizeToolRouterConfigKeepsConfiguredTools(t *testing.T) {
cfg := &config.Config{ToolRouter: config.ToolRouterConfig{
Enabled: true,
Timeout: 1,
MaxTokens: 1,
SystemPrompt: "tools",
Tools: []config.ToolRouteConfig{
{Name: " Search ", Enabled: true},
{Name: "sql", Enabled: true},
},
}}
changed, err := config.NormalizeToolRouterConfig(cfg)
if err != nil {
t.Fatal(err)
}
if !changed {
t.Fatal("expected configured tools to be normalized")
}
if len(cfg.ToolRouter.Tools) != 2 || cfg.ToolRouter.Tools[0].Name != "search" || cfg.ToolRouter.Tools[1].Name != "sql" {
t.Fatalf("unexpected tool order: %#v", cfg.ToolRouter.Tools)
}
}
func TestNormalizeToolRouterConfigDuplicateTools(t *testing.T) {
cfg := &config.Config{ToolRouter: config.ToolRouterConfig{
Enabled: true,
Timeout: 1,
MaxTokens: 1,
SystemPrompt: "tools",
Tools: []config.ToolRouteConfig{
{Name: "sql", Enabled: true},
{Name: " SQL ", Enabled: true},
},
}}
_, err := config.NormalizeToolRouterConfig(cfg)
if err == nil {
t.Fatal("expected duplicate tool error")
}
}
func TestAvailableAgentToolsUsesConfigOrderAndEnabled(t *testing.T) {
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: []config.ToolRouteConfig{
{Name: "search", Enabled: true},
{Name: "calculator", Enabled: true, Description: "custom calculator"},
{Name: "time", Enabled: true, Description: "custom time"},
{Name: "sql", Enabled: false},
},
}, ai)
if err != nil {
t.Fatal(err)
}
tools := toolrouter.AvailableAgentTools(router, ai.ActiveProfile(), newTestToolManager(), nil)
if len(tools) != 1 {
t.Fatalf("tools length = %d", len(tools))
}
if tools[0].Name() != "time" {
t.Fatalf("unexpected tools: %#v", tools)
}
definition := tools[0].Definition()
if definition.Function == nil || definition.Function.Description != "custom time" {
t.Fatalf("unexpected definition: %#v", definition)
}
}
func TestRunAgentToolLoopAppendsToolMessages(t *testing.T) {
ai := newTestAI(t, []config.OpenAIConfig{{Name: "test", APIKey: "key", BaseURL: testOpenAIBaseURL, Model: "test", Timeout: 1, Active: true}})
calls := 0
router, err := toolrouter.NewState(&config.ToolRouterConfig{
Enabled: true,
Timeout: 1,
MaxTokens: 128,
SystemPrompt: "use tools",
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)
}
if len(req.Tools) != 1 || req.Tools[0].Function == nil || req.Tools[0].Function.Name != "time" {
t.Fatalf("unexpected tools: %#v", req.Tools)
}
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: message.StringContent("done")}}}}, nil
}))
if err != nil {
t.Fatal(err)
}
messages, err := toolrouter.RunAgentToolLoop(context.Background(), router, ai.ActiveProfile(), []message.ChatMessage{{Role: "user", Content: "今天几号"}}, newTestToolManager(), nil)
if err != nil {
t.Fatal(err)
}
if calls != 2 {
t.Fatalf("calls = %d", calls)
}
if len(messages) < 4 {
t.Fatalf("expected system/user/assistant/tool messages, got %d", len(messages))
}
last := messages[len(messages)-1]
if last.Role != model.ChatMessageRoleTool || last.ToolCallID != "call_1" {
t.Fatalf("unexpected last message: %#v", last)
}
if last.Content == nil || last.Content.StringValue == nil || !strings.Contains(*last.Content.StringValue, "时间工具结果") {
t.Fatalf("unexpected tool content: %#v", last.Content)
}
}
func TestExecuteAgentToolCallUnknownAndError(t *testing.T) {
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 := 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)
}
}
func TestRunAgentToolLoopMaxIterations(t *testing.T) {
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: []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 := toolrouter.RunAgentToolLoop(context.Background(), router, ai.ActiveProfile(), []message.ChatMessage{{Role: "user", Content: "今天"}}, newTestToolManager(), nil)
if err != nil {
t.Fatal(err)
}
last := messages[len(messages)-1]
if last.Role != model.ChatMessageRoleSystem || last.Content == nil || last.Content.StringValue == nil || !strings.Contains(*last.Content.StringValue, "工具调用轮数已达到上限") {
t.Fatalf("unexpected last message: %#v", last)
}
}
func TestBuildArkMessageImageTextOrder(t *testing.T) {
msg, err := message.BuildArkMessage(message.ChatMessage{Role: "user", Content: "请描述图片", ImageURL: "data:image/png;base64,aGVsbG8="})
if err != nil {
t.Fatal(err)
}
if msg.Content == nil || len(msg.Content.ListValue) != 2 {
t.Fatalf("unexpected content: %#v", msg.Content)
}
if msg.Content.ListValue[0].Type != model.ChatCompletionMessageContentPartTypeText || msg.Content.ListValue[0].Text != "请描述图片" {
t.Fatalf("first part should be text: %#v", msg.Content.ListValue[0])
}
if msg.Content.ListValue[1].Type != model.ChatCompletionMessageContentPartTypeImageURL || msg.Content.ListValue[1].ImageURL == nil {
t.Fatalf("second part should be image: %#v", msg.Content.ListValue[1])
}
}
func TestBuildArkMessageImageOnly(t *testing.T) {
msg, err := message.BuildArkMessage(message.ChatMessage{Role: "user", ImageURL: "data:image/png;base64,aGVsbG8="})
if err != nil {
t.Fatal(err)
}
if msg.Content == nil || len(msg.Content.ListValue) != 1 || msg.Content.ListValue[0].Type != model.ChatCompletionMessageContentPartTypeImageURL {
t.Fatalf("unexpected content: %#v", msg.Content)
}
}
func TestThinkTagParserSingleChunk(t *testing.T) {
parser := &stream.Parser{}
visible, reasoning := parser.Accept("hello <think>abc</think> world")
flushVisible, flushReasoning := parser.Flush()
visible += flushVisible
reasoning += flushReasoning
if visible != "hello world" || reasoning != "abc" {
t.Fatalf("visible=%q reasoning=%q", visible, reasoning)
}
}
func TestThinkTagParserAcrossChunks(t *testing.T) {
parser := &stream.Parser{}
var visible, reasoning string
for _, chunk := range []string{"hello <thi", "nk>abc</thi", "nk> world"} {
v, r := parser.Accept(chunk)
visible += v
reasoning += r
}
v, r := parser.Flush()
visible += v
reasoning += r
if visible != "hello world" || reasoning != "abc" {
t.Fatalf("visible=%q reasoning=%q", visible, reasoning)
}
}
func TestThinkTagParserUnclosedThink(t *testing.T) {
parser := &stream.Parser{}
visible, reasoning := parser.Accept("answer <think>still thinking")
v, r := parser.Flush()
visible += v
reasoning += r
if visible != "answer " || reasoning != "still thinking" {
t.Fatalf("visible=%q reasoning=%q", visible, reasoning)
}
}
func TestShouldParseThinkTags(t *testing.T) {
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 llm.ShouldParseThinkTags(&llm.Profile{Config: config.OpenAIConfig{BaseURL: testOpenAIBaseURL}}) {
t.Fatal("expected remote profile not to parse think tags by default")
}
falseValue := false
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 !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 := message.BuildToolDecisionMessages([]message.ChatMessage{{Role: "user", Content: "描述这张图", ImageURL: "data:image/png;base64,aGVsbG8="}})
if err != nil {
t.Fatal(err)
}
if len(messages) != 1 || messages[0].Content == nil || messages[0].Content.StringValue == nil {
t.Fatalf("unexpected messages: %#v", messages)
}
if !strings.Contains(*messages[0].Content.StringValue, "工具判断阶段不读取图片内容") {
t.Fatalf("missing image placeholder: %q", *messages[0].Content.StringValue)
}
if messages[0].Content.ListValue != nil {
t.Fatalf("decision message should be text-only: %#v", messages[0].Content)
}
}
func TestRunAgentToolLoopImageUsesTextOnlyDecisionMessages(t *testing.T) {
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: []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)
}
}
joined := ""
for _, msg := range req.Messages {
if msg.Content != nil && msg.Content.StringValue != nil {
joined += *msg.Content.StringValue
}
}
if !strings.Contains(joined, "工具判断阶段不读取图片内容") {
t.Fatalf("missing placeholder in decision messages: %q", joined)
}
return model.ChatCompletionResponse{Choices: []*model.ChatCompletionChoice{{Message: model.ChatCompletionMessage{Content: message.StringContent("no tool")}}}}, nil
}))
if err != nil {
t.Fatal(err)
}
messages, err := toolrouter.RunAgentToolLoop(context.Background(), router, ai.ActiveProfile(), []message.ChatMessage{{Role: "user", Content: "描述这张图", ImageURL: "data:image/png;base64,aGVsbG8="}}, newTestToolManager(), nil)
if err != nil {
t.Fatal(err)
}
foundImage := false
for _, msg := range messages {
if msg.Content != nil && len(msg.Content.ListValue) > 0 {
foundImage = true
}
}
if !foundImage {
t.Fatalf("final messages should retain image: %#v", messages)
}
}
func TestRunAgentToolLoopUsesConfiguredRouterProfile(t *testing.T) {
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},
})
router, err := toolrouter.NewState(&config.ToolRouterConfig{
Enabled: true,
OpenAIName: "router",
Timeout: 1,
MaxTokens: 128,
SystemPrompt: "use tools",
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: message.StringContent("no tool")}}}}, nil
}))
if err != nil {
t.Fatal(err)
}
_, err = toolrouter.RunAgentToolLoop(context.Background(), router, ai.ActiveProfile(), []message.ChatMessage{{Role: "user", Content: "今天"}}, newTestToolManager(), nil)
if err != nil {
t.Fatal(err)
}
}