+23
-13
@@ -7,10 +7,13 @@ import (
|
||||
"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"
|
||||
@@ -18,6 +21,13 @@ import (
|
||||
|
||||
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)
|
||||
@@ -46,7 +56,7 @@ func TestNormalizeToolRouterConfigDefaults(t *testing.T) {
|
||||
if strings.TrimSpace(cfg.ToolRouter.SystemPrompt) == "" {
|
||||
t.Fatal("system prompt should be defaulted")
|
||||
}
|
||||
if len(cfg.ToolRouter.Tools) != 4 || cfg.ToolRouter.Tools[0].Name != "calculator" || cfg.ToolRouter.Tools[1].Name != "time" || cfg.ToolRouter.Tools[2].Name != "search" || cfg.ToolRouter.Tools[3].Name != "sql" || !cfg.ToolRouter.Tools[0].Enabled || !cfg.ToolRouter.Tools[1].Enabled || !cfg.ToolRouter.Tools[2].Enabled || !cfg.ToolRouter.Tools[3].Enabled {
|
||||
if len(cfg.ToolRouter.Tools) != 0 {
|
||||
t.Fatalf("unexpected tools: %#v", cfg.ToolRouter.Tools)
|
||||
}
|
||||
}
|
||||
@@ -66,14 +76,14 @@ func TestNormalizeOpenAIConfigDefaultsContextWindow(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalizeToolRouterConfigAddsCalculatorAndTimeBeforeSQL(t *testing.T) {
|
||||
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: " Search ", Enabled: true},
|
||||
{Name: "sql", Enabled: true},
|
||||
},
|
||||
}}
|
||||
@@ -82,9 +92,9 @@ func TestNormalizeToolRouterConfigAddsCalculatorAndTimeBeforeSQL(t *testing.T) {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !changed {
|
||||
t.Fatal("expected calculator and time tools to be added")
|
||||
t.Fatal("expected configured tools to be normalized")
|
||||
}
|
||||
if len(cfg.ToolRouter.Tools) < 4 || cfg.ToolRouter.Tools[0].Name != "calculator" || cfg.ToolRouter.Tools[1].Name != "time" || cfg.ToolRouter.Tools[3].Name != "sql" {
|
||||
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)
|
||||
}
|
||||
}
|
||||
@@ -121,15 +131,15 @@ func TestAvailableAgentToolsUsesConfigOrderAndEnabled(t *testing.T) {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
tools := toolrouter.AvailableAgentTools(router, ai.ActiveProfile(), nil, nil, nil)
|
||||
if len(tools) != 2 {
|
||||
tools := toolrouter.AvailableAgentTools(router, ai.ActiveProfile(), newTestToolManager(), nil)
|
||||
if len(tools) != 1 {
|
||||
t.Fatalf("tools length = %d", len(tools))
|
||||
}
|
||||
if tools[0].Name() != "calculator" || tools[1].Name() != "time" {
|
||||
if tools[0].Name() != "time" {
|
||||
t.Fatalf("unexpected tools: %#v", tools)
|
||||
}
|
||||
definition := tools[0].Definition()
|
||||
if definition.Function == nil || definition.Function.Description != "custom calculator" {
|
||||
if definition.Function == nil || definition.Function.Description != "custom time" {
|
||||
t.Fatalf("unexpected definition: %#v", definition)
|
||||
}
|
||||
}
|
||||
@@ -160,7 +170,7 @@ func TestRunAgentToolLoopAppendsToolMessages(t *testing.T) {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
messages, err := toolrouter.RunAgentToolLoop(context.Background(), router, ai.ActiveProfile(), []message.ChatMessage{{Role: "user", Content: "今天几号"}}, nil, nil, nil)
|
||||
messages, err := toolrouter.RunAgentToolLoop(context.Background(), router, ai.ActiveProfile(), []message.ChatMessage{{Role: "user", Content: "今天几号"}}, newTestToolManager(), nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
@@ -208,7 +218,7 @@ func TestRunAgentToolLoopMaxIterations(t *testing.T) {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
messages, err := toolrouter.RunAgentToolLoop(context.Background(), router, ai.ActiveProfile(), []message.ChatMessage{{Role: "user", Content: "今天"}}, nil, nil, nil)
|
||||
messages, err := toolrouter.RunAgentToolLoop(context.Background(), router, ai.ActiveProfile(), []message.ChatMessage{{Role: "user", Content: "今天"}}, newTestToolManager(), nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
@@ -344,7 +354,7 @@ func TestRunAgentToolLoopImageUsesTextOnlyDecisionMessages(t *testing.T) {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
messages, err := toolrouter.RunAgentToolLoop(context.Background(), router, ai.ActiveProfile(), []message.ChatMessage{{Role: "user", Content: "描述这张图", ImageURL: "data:image/png;base64,aGVsbG8="}}, nil, nil, nil)
|
||||
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)
|
||||
}
|
||||
@@ -381,7 +391,7 @@ func TestRunAgentToolLoopUsesConfiguredRouterProfile(t *testing.T) {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
_, err = toolrouter.RunAgentToolLoop(context.Background(), router, ai.ActiveProfile(), []message.ChatMessage{{Role: "user", Content: "今天"}}, nil, nil, nil)
|
||||
_, err = toolrouter.RunAgentToolLoop(context.Background(), router, ai.ActiveProfile(), []message.ChatMessage{{Role: "user", Content: "今天"}}, newTestToolManager(), nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user