refactor: remove unused tool filter middleware and helpers
- Deleted tool_filter_middleware_test.go and tool_helpers.go as they are no longer needed. - Removed associated test cases in tool_helpers_test.go. - Refactored knowledge retriever logic by moving it to a new file and updating imports. - Introduced tooling package for tool result reduction logic. - Updated traces package to include new trace types and structures. - Adjusted executor and tool search tool to reflect new package structure.
This commit is contained in:
@@ -1,25 +0,0 @@
|
||||
package runtime
|
||||
|
||||
import "agent-desk/internal/ai/runtime/registry"
|
||||
|
||||
func newPrepareService(catalog *toolCatalog) *prepareService {
|
||||
return &prepareService{catalog: catalog}
|
||||
}
|
||||
|
||||
type prepareService struct {
|
||||
catalog *toolCatalog
|
||||
}
|
||||
|
||||
func (s *prepareService) prepareToolsForRun(req Request) (*registry.ToolSet, error) {
|
||||
if req.ToolSet != nil {
|
||||
return req.ToolSet, nil
|
||||
}
|
||||
return s.catalog.resolveForRun(req)
|
||||
}
|
||||
|
||||
func (s *prepareService) prepareToolsForResume(req ResumeRequest) (*registry.ToolSet, error) {
|
||||
if req.ToolSet != nil {
|
||||
return req.ToolSet, nil
|
||||
}
|
||||
return s.catalog.resolveForResume(req)
|
||||
}
|
||||
@@ -6,7 +6,6 @@ import (
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"agent-desk/internal/ai/runtime/executor"
|
||||
workflowexecutor "agent-desk/internal/ai/runtime/workflow"
|
||||
"agent-desk/internal/models"
|
||||
"agent-desk/internal/pkg/errorsx"
|
||||
@@ -17,9 +16,6 @@ import (
|
||||
)
|
||||
|
||||
type Service struct {
|
||||
runtime *executor.Service
|
||||
catalog *toolCatalog
|
||||
prepare *prepareService
|
||||
}
|
||||
|
||||
const (
|
||||
@@ -29,12 +25,7 @@ const (
|
||||
)
|
||||
|
||||
func NewService() *Service {
|
||||
catalog := newToolCatalog()
|
||||
return &Service{
|
||||
runtime: executor.NewService(),
|
||||
catalog: catalog,
|
||||
prepare: newPrepareService(catalog),
|
||||
}
|
||||
return &Service{}
|
||||
}
|
||||
|
||||
func (s *Service) Run(ctx context.Context, req Request) (*Summary, error) {
|
||||
@@ -106,23 +97,7 @@ func (s *Service) Resume(ctx context.Context, req ResumeRequest) (*Summary, erro
|
||||
return toWorkflowSummary(workflowResult, req.AIConfig.ModelName, workflow, workflowRunID), nil
|
||||
}
|
||||
}
|
||||
toolSet, err := s.prepare.prepareToolsForResume(req)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
req.ToolSet = toolSet
|
||||
summary, err := s.runtime.ExecuteResume(ctx, executor.ResumeInput{
|
||||
Conversation: req.Conversation,
|
||||
AIAgent: req.AIAgent,
|
||||
AIConfig: req.AIConfig,
|
||||
CheckPointID: req.CheckPointID,
|
||||
ResumeData: req.ResumeData,
|
||||
ToolSet: req.ToolSet,
|
||||
})
|
||||
if err != nil {
|
||||
return toSummary(summary), err
|
||||
}
|
||||
return toSummary(summary), nil
|
||||
return nil, errorsx.InvalidParam("legacy checkpoint is not supported; please start a new workflow reply")
|
||||
}
|
||||
|
||||
func firstWorkflowResumeText(data map[string]string) string {
|
||||
|
||||
@@ -1,45 +0,0 @@
|
||||
package runtime
|
||||
|
||||
import (
|
||||
"agent-desk/internal/ai/runtime/executor"
|
||||
"strings"
|
||||
)
|
||||
|
||||
func toSummary(summary *executor.RunResult) *Summary {
|
||||
if summary == nil {
|
||||
return nil
|
||||
}
|
||||
ret := &Summary{
|
||||
RunID: summary.RunID,
|
||||
Status: summary.Status,
|
||||
ReplyText: summary.ReplyText,
|
||||
PlannedSkillID: summary.SelectedSkillID,
|
||||
PlannedSkillName: strings.TrimSpace(summary.SelectedSkillName),
|
||||
PlanReason: strings.TrimSpace(summary.SkillRouteReason),
|
||||
SkillRouteTrace: strings.TrimSpace(summary.SkillRouteTrace),
|
||||
SkillAllowedToolCodes: append([]string(nil), summary.SkillAllowedToolCodes...),
|
||||
ModelName: summary.ModelName,
|
||||
PromptTokens: summary.PromptTokens,
|
||||
CompletionTokens: summary.CompletionTokens,
|
||||
HistoryMessageCount: summary.HistoryMessageCount,
|
||||
RetrieverCount: summary.RetrieverCount,
|
||||
ToolCallCount: summary.ToolCallCount,
|
||||
ToolCodes: append([]string(nil), summary.ToolCodes...),
|
||||
InvokedToolCodes: append([]string(nil), summary.InvokedToolCodes...),
|
||||
CheckPointID: summary.CheckPointID,
|
||||
Interrupted: summary.Interrupted,
|
||||
TraceData: summary.TraceData,
|
||||
ErrorMessage: summary.ErrorMessage,
|
||||
}
|
||||
if len(summary.Interrupts) > 0 {
|
||||
ret.Interrupts = make([]InterruptContextSummary, 0, len(summary.Interrupts))
|
||||
for _, item := range summary.Interrupts {
|
||||
ret.Interrupts = append(ret.Interrupts, InterruptContextSummary{
|
||||
Type: item.Type,
|
||||
ID: item.ID,
|
||||
InfoPreview: item.InfoPreview,
|
||||
})
|
||||
}
|
||||
}
|
||||
return ret
|
||||
}
|
||||
@@ -1,93 +0,0 @@
|
||||
package runtime
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"strings"
|
||||
|
||||
"agent-desk/internal/ai/runtime/registry"
|
||||
"agent-desk/internal/ai/runtime/tools"
|
||||
"agent-desk/internal/models"
|
||||
"agent-desk/internal/pkg/toolx"
|
||||
)
|
||||
|
||||
type toolCatalog struct {
|
||||
registry *registry.Registry
|
||||
}
|
||||
|
||||
func newToolCatalog() *toolCatalog {
|
||||
return &toolCatalog{
|
||||
registry: registry.NewRegistry(buildRuntimeStaticTools()...),
|
||||
}
|
||||
}
|
||||
|
||||
func buildRuntimeStaticTools() []registry.Tool {
|
||||
ret := make([]registry.Tool, 0, len(toolx.ListRuntimeStaticToolSpecs()))
|
||||
for _, spec := range toolx.ListRuntimeStaticToolSpecs() {
|
||||
tool := tools.NewRuntimeStaticTool(spec.Code)
|
||||
if tool == nil {
|
||||
continue
|
||||
}
|
||||
ret = append(ret, tool)
|
||||
}
|
||||
return ret
|
||||
}
|
||||
|
||||
func (c *toolCatalog) resolveForRun(req Request) (*registry.ToolSet, error) {
|
||||
return c.registry.Resolve(registry.Context{
|
||||
Conversation: req.Conversation,
|
||||
AIAgent: req.AIAgent,
|
||||
AIConfig: req.AIConfig,
|
||||
UserMessage: req.UserMessage,
|
||||
AllowedToolCodes: c.parseAgentAllowedToolCodes(req.AIAgent),
|
||||
})
|
||||
}
|
||||
|
||||
func (c *toolCatalog) resolveForResume(req ResumeRequest) (*registry.ToolSet, error) {
|
||||
return c.registry.Resolve(registry.Context{
|
||||
Conversation: req.Conversation,
|
||||
AIAgent: req.AIAgent,
|
||||
AIConfig: req.AIConfig,
|
||||
AllowedToolCodes: c.parseAgentAllowedToolCodes(req.AIAgent),
|
||||
})
|
||||
}
|
||||
|
||||
func (c *toolCatalog) parseSkillAllowedToolCodes(skill *models.SkillDefinition) []string {
|
||||
if skill == nil {
|
||||
return nil
|
||||
}
|
||||
raw := strings.TrimSpace(skill.ToolWhitelist)
|
||||
if raw == "" {
|
||||
return nil
|
||||
}
|
||||
var items []string
|
||||
if err := json.Unmarshal([]byte(raw), &items); err != nil {
|
||||
return nil
|
||||
}
|
||||
return toolx.NormalizeToolCodes(items)
|
||||
}
|
||||
|
||||
func (c *toolCatalog) parseAgentAllowedToolCodes(aiAgent models.AIAgent) []string {
|
||||
ret := make([]string, 0)
|
||||
if raw := strings.TrimSpace(aiAgent.AllowedMCPTools); raw != "" {
|
||||
items, err := toolx.ParseAgentMCPToolsJSON(raw)
|
||||
if err == nil {
|
||||
for _, item := range items {
|
||||
ret = append(ret, item.ToolCode)
|
||||
}
|
||||
}
|
||||
}
|
||||
if raw := strings.TrimSpace(aiAgent.AllowedGraphTools); raw != "" {
|
||||
var graphTools []string
|
||||
if err := json.Unmarshal([]byte(raw), &graphTools); err == nil {
|
||||
ret = append(ret, graphTools...)
|
||||
}
|
||||
}
|
||||
if workflow, err := resolveAgentWorkflow(aiAgent); err == nil {
|
||||
ret = append(ret, workflow.Compiled.ToolCodes...)
|
||||
}
|
||||
return toolx.NormalizeToolCodes(ret)
|
||||
}
|
||||
|
||||
func (c *toolCatalog) resolveAllowedToolCodes(aiAgent models.AIAgent, skill *models.SkillDefinition) []string {
|
||||
return toolx.IntersectToolCodes(c.parseAgentAllowedToolCodes(aiAgent), c.parseSkillAllowedToolCodes(skill))
|
||||
}
|
||||
@@ -1,255 +0,0 @@
|
||||
package runtime
|
||||
|
||||
import (
|
||||
"context"
|
||||
"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{
|
||||
WorkflowVersionID: version.ID,
|
||||
})
|
||||
|
||||
assertContainsToolCode(t, ret, toolx.GraphPrepareTicketDraft.Code)
|
||||
assertContainsToolCode(t, ret, toolx.GraphCreateTicketConfirm.Code)
|
||||
assertContainsToolCode(t, ret, toolx.GraphHandoffConversation.Code)
|
||||
}
|
||||
|
||||
func TestPrepareWorkflowAgentAppendsPublishedWorkflow(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, _, err := prepareWorkflowAgent(models.AIAgent{
|
||||
SystemPrompt: "Base prompt.",
|
||||
WorkflowVersionID: version.ID,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("prepare workflow agent: %v", err)
|
||||
}
|
||||
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 TestPrepareWorkflowAgentRejectsMissingPublishedWorkflow(t *testing.T) {
|
||||
_, _, err := prepareWorkflowAgent(models.AIAgent{})
|
||||
if err == nil {
|
||||
t.Fatalf("expected missing workflow version error")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "AI Agent workflow is not published") {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPrepareWorkflowAgentRejectsDeletedPublishedWorkflow(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: "end", Type: workflowregistry.NodeTypeEnd, Name: "End"},
|
||||
},
|
||||
})
|
||||
if err := sqls.DB().Model(&models.AIWorkflowVersion{}).Where("id = ?", version.ID).Update("status", enums.StatusDeleted).Error; err != nil {
|
||||
t.Fatalf("delete workflow version: %v", err)
|
||||
}
|
||||
|
||||
_, _, err := prepareWorkflowAgent(models.AIAgent{WorkflowVersionID: version.ID})
|
||||
if err == nil {
|
||||
t.Fatalf("expected invalid workflow version error")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "workflow version does not exist") {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestServiceRunExecutesPublishedWorkflow(t *testing.T) {
|
||||
setupWorkflowRuntimeTestDB(t)
|
||||
version := createWorkflowRuntimeTestVersion(t, dsl.Definition{
|
||||
SchemaVersion: 1,
|
||||
EntryNodeID: "start_1",
|
||||
Nodes: []dsl.Node{
|
||||
{ID: "start_1", Type: workflowregistry.NodeTypeStart, Name: "Start"},
|
||||
{ID: "reply_1", Type: workflowregistry.NodeTypeLLMReply, Name: "Reply", Config: []byte(`{"staticReply":"workflow reply"}`)},
|
||||
{ID: "send_1", Type: workflowregistry.NodeTypeSendReply, Name: "Send", Inputs: map[string]dsl.VariableSelector{
|
||||
"replyText": {NodeID: "reply_1", Field: "replyText"},
|
||||
}},
|
||||
{ID: "end_1", Type: workflowregistry.NodeTypeEnd, Name: "End"},
|
||||
},
|
||||
Edges: []dsl.Edge{
|
||||
{ID: "edge_start_reply", Source: "start_1", Target: "reply_1"},
|
||||
{ID: "edge_reply_send", Source: "reply_1", Target: "send_1"},
|
||||
{ID: "edge_send_end", Source: "send_1", Target: "end_1"},
|
||||
},
|
||||
})
|
||||
|
||||
summary, err := NewService().Run(context.Background(), Request{
|
||||
UserMessage: models.Message{Content: "hello"},
|
||||
AIAgent: models.AIAgent{
|
||||
WorkflowVersionID: version.ID,
|
||||
},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("run workflow: %v", err)
|
||||
}
|
||||
if summary.ReplyText != "workflow reply" {
|
||||
t.Fatalf("unexpected workflow reply: %q", summary.ReplyText)
|
||||
}
|
||||
var runCount int64
|
||||
if err := sqls.DB().Model(&models.AIWorkflowRun{}).Count(&runCount).Error; err != nil {
|
||||
t.Fatalf("count workflow runs: %v", err)
|
||||
}
|
||||
if runCount != 1 {
|
||||
t.Fatalf("expected one workflow run, got %d", runCount)
|
||||
}
|
||||
var nodeRunCount int64
|
||||
if err := sqls.DB().Model(&models.AIWorkflowNodeRun{}).Count(&nodeRunCount).Error; err != nil {
|
||||
t.Fatalf("count workflow node runs: %v", err)
|
||||
}
|
||||
if nodeRunCount != 4 {
|
||||
t.Fatalf("expected four workflow node runs, got %d", nodeRunCount)
|
||||
}
|
||||
var replyNodeRun models.AIWorkflowNodeRun
|
||||
if err := sqls.DB().First(&replyNodeRun, "node_id = ?", "reply_1").Error; err != nil {
|
||||
t.Fatalf("find reply node run: %v", err)
|
||||
}
|
||||
if replyNodeRun.InputPreview == "" || replyNodeRun.OutputPreview == "" {
|
||||
t.Fatalf("expected node input/output previews, got input=%q output=%q", replyNodeRun.InputPreview, replyNodeRun.OutputPreview)
|
||||
}
|
||||
if !strings.Contains(replyNodeRun.OutputPreview, "workflow reply") {
|
||||
t.Fatalf("expected reply output preview, got %q", replyNodeRun.OutputPreview)
|
||||
}
|
||||
if replyNodeRun.DurationMS < 0 {
|
||||
t.Fatalf("unexpected negative duration: %d", replyNodeRun.DurationMS)
|
||||
}
|
||||
}
|
||||
|
||||
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{}, &models.AIWorkflowRun{}, &models.AIWorkflowNodeRun{}); 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)
|
||||
}
|
||||
@@ -1,7 +1,6 @@
|
||||
package runtime
|
||||
|
||||
import (
|
||||
"agent-desk/internal/ai/runtime/registry"
|
||||
"agent-desk/internal/models"
|
||||
)
|
||||
|
||||
@@ -11,7 +10,6 @@ type Request struct {
|
||||
AIAgent models.AIAgent
|
||||
AIConfig models.AIConfig
|
||||
CheckPointID string
|
||||
ToolSet *registry.ToolSet
|
||||
}
|
||||
|
||||
type ResumeRequest struct {
|
||||
@@ -21,7 +19,6 @@ type ResumeRequest struct {
|
||||
AIConfig models.AIConfig
|
||||
CheckPointID string
|
||||
ResumeData map[string]string
|
||||
ToolSet *registry.ToolSet
|
||||
}
|
||||
|
||||
type InterruptContextSummary struct {
|
||||
|
||||
Reference in New Issue
Block a user