Files
ai-agent/internal/ai/application/runtime/tool_catalog_test.go
T
2026-06-22 00:26:16 +08:00

160 lines
4.9 KiB
Go

package runtime
import (
"encoding/json"
"strings"
"testing"
"agent-desk/internal/ai/workflow/dsl"
workflowregistry "agent-desk/internal/ai/workflow/registry"
"agent-desk/internal/models"
"agent-desk/internal/pkg/enums"
"agent-desk/internal/pkg/toolx"
"github.com/glebarez/sqlite"
"github.com/mlogclub/simple/sqls"
"gorm.io/gorm"
)
func TestNormalizeAllowedToolCodes(t *testing.T) {
ret := toolx.NormalizeToolCodes([]string{
" ",
"graph/create_ticket_with_confirmation",
"builtin/create_ticket_with_confirmation",
"graph/handoff_to_human",
"graph/handoff_to_human",
})
if len(ret) != 2 {
t.Fatalf("expected 2 tool codes, got %d: %#v", len(ret), ret)
}
if ret[0] != "graph/create_ticket_with_confirmation" {
t.Fatalf("unexpected first tool code: %s", ret[0])
}
if ret[1] != "graph/handoff_to_human" {
t.Fatalf("unexpected second tool code: %s", ret[1])
}
}
func TestToolCatalogResolveAllowedToolCodes(t *testing.T) {
catalog := newToolCatalog()
agent := models.AIAgent{
AllowedMCPTools: `[{"toolCode":"graph/create_ticket_with_confirmation"},{"toolCode":"graph/handoff_to_human"}]`,
}
skill := &models.SkillDefinition{
ToolWhitelist: `["builtin/create_ticket_with_confirmation","graph/prepare_ticket_draft"]`,
}
ret := catalog.resolveAllowedToolCodes(agent, skill)
if len(ret) != 1 {
t.Fatalf("expected 1 tool code, got %d: %#v", len(ret), ret)
}
if ret[0] != "graph/create_ticket_with_confirmation" {
t.Fatalf("unexpected tool code: %s", ret[0])
}
}
func TestToolCatalogResolveAllowedToolCodesFallsBackWhenSkillEmpty(t *testing.T) {
catalog := newToolCatalog()
agent := models.AIAgent{
AllowedMCPTools: `[{"toolCode":"graph/create_ticket_with_confirmation"},{"toolCode":"graph/handoff_to_human"}]`,
}
ret := catalog.resolveAllowedToolCodes(agent, nil)
if len(ret) != 2 {
t.Fatalf("expected 2 tool codes, got %d: %#v", len(ret), ret)
}
}
func TestBuildRuntimeStaticTools(t *testing.T) {
ret := buildRuntimeStaticTools()
if len(ret) != len(toolx.ListRuntimeStaticToolSpecs()) {
t.Fatalf("expected %d runtime static tools, got %d", len(toolx.ListRuntimeStaticToolSpecs()), len(ret))
}
}
func TestToolCatalogIncludesPublishedWorkflowGraphTools(t *testing.T) {
setupWorkflowRuntimeTestDB(t)
version := createWorkflowRuntimeTestVersion(t, dsl.Definition{
SchemaVersion: 1,
EntryNodeID: "start",
Nodes: []dsl.Node{
{ID: "start", Type: workflowregistry.NodeTypeStart, Name: "Start"},
{ID: "draft", Type: workflowregistry.NodeTypePrepareTicketDraft, Name: "Draft Ticket"},
{ID: "create", Type: workflowregistry.NodeTypeCreateTicket, Name: "Create Ticket"},
{ID: "handoff", Type: workflowregistry.NodeTypeHandoffToHuman, Name: "Handoff"},
},
})
catalog := newToolCatalog()
ret := catalog.parseAgentAllowedToolCodes(models.AIAgent{
RuntimeMode: enums.AIAgentRuntimeModeWorkflow,
WorkflowVersionID: version.ID,
})
assertContainsToolCode(t, ret, toolx.GraphPrepareTicketDraft.Code)
assertContainsToolCode(t, ret, toolx.GraphCreateTicketConfirm.Code)
assertContainsToolCode(t, ret, toolx.GraphHandoffConversation.Code)
}
func TestApplyWorkflowInstructionAppendsPublishedWorkflow(t *testing.T) {
setupWorkflowRuntimeTestDB(t)
version := createWorkflowRuntimeTestVersion(t, dsl.Definition{
SchemaVersion: 1,
EntryNodeID: "start",
Nodes: []dsl.Node{
{ID: "start", Type: workflowregistry.NodeTypeStart, Name: "Start"},
{ID: "handoff", Type: workflowregistry.NodeTypeHandoffToHuman, Name: "Handoff"},
},
})
agent := applyWorkflowInstruction(models.AIAgent{
SystemPrompt: "Base prompt.",
RuntimeMode: enums.AIAgentRuntimeModeWorkflow,
WorkflowVersionID: version.ID,
})
if agent.SystemPrompt == "Base prompt." {
t.Fatalf("expected workflow appendix to be appended")
}
if !strings.Contains(agent.SystemPrompt, "Published customer-service workflow") {
t.Fatalf("missing workflow appendix: %s", agent.SystemPrompt)
}
}
func setupWorkflowRuntimeTestDB(t *testing.T) {
t.Helper()
db, err := gorm.Open(sqlite.Open("file:"+t.Name()+"?mode=memory&cache=shared"), &gorm.Config{})
if err != nil {
t.Fatalf("open sqlite db: %v", err)
}
if err := db.AutoMigrate(&models.AIWorkflowVersion{}); err != nil {
t.Fatalf("auto migrate: %v", err)
}
sqls.SetDB(db)
}
func createWorkflowRuntimeTestVersion(t *testing.T, def dsl.Definition) *models.AIWorkflowVersion {
t.Helper()
definition, err := json.Marshal(def)
if err != nil {
t.Fatalf("marshal definition: %v", err)
}
version := &models.AIWorkflowVersion{
WorkflowID: 1,
Version: 1,
Status: enums.StatusOk,
Definition: string(definition),
}
if err := sqls.DB().Create(version).Error; err != nil {
t.Fatalf("create workflow version: %v", err)
}
return version
}
func assertContainsToolCode(t *testing.T, items []string, want string) {
t.Helper()
for _, item := range items {
if item == want {
return
}
}
t.Fatalf("expected tool code %s in %#v", want, items)
}