feat(tool-filter): implement runtime tool filtering middleware and associated tests
This commit is contained in:
@@ -214,4 +214,7 @@ func syncSkillSummaryFromCollector(summary *RunResult, collector *callbacks.Runt
|
|||||||
summary.SkillRouteReason = strings.TrimSpace(trace.RouteReason)
|
summary.SkillRouteReason = strings.TrimSpace(trace.RouteReason)
|
||||||
summary.SkillRouteTrace = strings.TrimSpace(trace.RouteTrace)
|
summary.SkillRouteTrace = strings.TrimSpace(trace.RouteTrace)
|
||||||
summary.SkillAllowedToolCodes = append([]string(nil), trace.AllowedToolCodes...)
|
summary.SkillAllowedToolCodes = append([]string(nil), trace.AllowedToolCodes...)
|
||||||
|
if len(trace.FilteredToolCodes) > 0 {
|
||||||
|
summary.ToolCodes = append([]string(nil), trace.FilteredToolCodes...)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -100,6 +100,15 @@ func (c *RuntimeTraceCollector) ActivateSkill(skill SkillMetadata, routeReason s
|
|||||||
c.Data.Skill.RouteTrace = routeTrace
|
c.Data.Skill.RouteTrace = routeTrace
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (c *RuntimeTraceCollector) SetFilteredToolCodes(toolCodes []string) {
|
||||||
|
if c == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
c.mu.Lock()
|
||||||
|
defer c.mu.Unlock()
|
||||||
|
c.Data.Skill.FilteredToolCodes = append([]string(nil), toolCodes...)
|
||||||
|
}
|
||||||
|
|
||||||
func (c *RuntimeTraceCollector) SetRetrieverSummary(summary RetrieverTraceSummary) {
|
func (c *RuntimeTraceCollector) SetRetrieverSummary(summary RetrieverTraceSummary) {
|
||||||
if c == nil {
|
if c == nil {
|
||||||
return
|
return
|
||||||
|
|||||||
@@ -12,6 +12,8 @@ type ToolTraceItem struct {
|
|||||||
LatencyMs int64 `json:"latencyMs,omitempty"`
|
LatencyMs int64 `json:"latencyMs,omitempty"`
|
||||||
Status string `json:"status,omitempty"`
|
Status string `json:"status,omitempty"`
|
||||||
ErrorMessage string `json:"errorMessage,omitempty"`
|
ErrorMessage string `json:"errorMessage,omitempty"`
|
||||||
|
Blocked bool `json:"blocked,omitempty"`
|
||||||
|
BlockedReason string `json:"blockedReason,omitempty"`
|
||||||
}
|
}
|
||||||
|
|
||||||
type ToolSearchTraceItem struct {
|
type ToolSearchTraceItem struct {
|
||||||
@@ -146,6 +148,7 @@ type SkillTraceData struct {
|
|||||||
RouteReason string `json:"routeReason,omitempty"`
|
RouteReason string `json:"routeReason,omitempty"`
|
||||||
RouteTrace string `json:"routeTrace,omitempty"`
|
RouteTrace string `json:"routeTrace,omitempty"`
|
||||||
AllowedToolCodes []string `json:"allowedToolCodes,omitempty"`
|
AllowedToolCodes []string `json:"allowedToolCodes,omitempty"`
|
||||||
|
FilteredToolCodes []string `json:"filteredToolCodes,omitempty"`
|
||||||
MiddlewareEnabled bool `json:"middlewareEnabled,omitempty"`
|
MiddlewareEnabled bool `json:"middlewareEnabled,omitempty"`
|
||||||
MiddlewareToolName string `json:"middlewareToolName,omitempty"`
|
MiddlewareToolName string `json:"middlewareToolName,omitempty"`
|
||||||
VisibleCodes []string `json:"visibleCodes,omitempty"`
|
VisibleCodes []string `json:"visibleCodes,omitempty"`
|
||||||
|
|||||||
@@ -2,6 +2,7 @@ package factory
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"strings"
|
||||||
|
|
||||||
einocallbacks "cs-agent/internal/ai/runtime/internal/impl/callbacks"
|
einocallbacks "cs-agent/internal/ai/runtime/internal/impl/callbacks"
|
||||||
"cs-agent/internal/ai/runtime/registry"
|
"cs-agent/internal/ai/runtime/registry"
|
||||||
@@ -36,7 +37,26 @@ func NewAgentHandlerService(skillMiddleware *SkillMiddlewareService) *AgentHandl
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (s *AgentHandlerService) Build(ctx context.Context, input BuildAgentHandlersInput) ([]adk.ChatModelAgentMiddleware, error) {
|
func (s *AgentHandlerService) Build(ctx context.Context, input BuildAgentHandlersInput) ([]adk.ChatModelAgentMiddleware, error) {
|
||||||
handlers := make([]adk.ChatModelAgentMiddleware, 0, 3)
|
handlers := make([]adk.ChatModelAgentMiddleware, 0, 4)
|
||||||
|
skillMetadataByCode := buildRuntimeSkillMetadataMap(input.AIAgent)
|
||||||
|
toolMetadataBy := buildRuntimeTraceToolMetadata(input.DynamicToolDefinitions, input.StaticToolMetadata, len(skillMetadataByCode) > 0)
|
||||||
|
traceSkillMetadata := make(map[string]einocallbacks.SkillMetadata, len(skillMetadataByCode))
|
||||||
|
for code, item := range skillMetadataByCode {
|
||||||
|
traceSkillMetadata[code] = einocallbacks.SkillMetadata{
|
||||||
|
Code: item.Code,
|
||||||
|
Name: item.Name,
|
||||||
|
Description: item.Description,
|
||||||
|
AllowedToolCodes: append([]string(nil), item.AllowedToolCodes...),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if input.Collector != nil {
|
||||||
|
if len(skillMetadataByCode) > 0 {
|
||||||
|
input.Collector.SetSkillMiddleware(true, toolx.BuiltinSkill.Name)
|
||||||
|
}
|
||||||
|
input.Collector.SetVisibleSkills(traceSkillMetadata)
|
||||||
|
input.Collector.SetInstructionSummary(input.InstructionSummary)
|
||||||
|
handlers = append(handlers, einocallbacks.NewRuntimeTraceHandler(input.Collector, toolMetadataBy, traceSkillMetadata))
|
||||||
|
}
|
||||||
if len(input.DynamicTools) > 0 {
|
if len(input.DynamicTools) > 0 {
|
||||||
toolSearchHandler, err := einotoolsearch.New(ctx, &einotoolsearch.Config{
|
toolSearchHandler, err := einotoolsearch.New(ctx, &einotoolsearch.Config{
|
||||||
DynamicTools: input.DynamicTools,
|
DynamicTools: input.DynamicTools,
|
||||||
@@ -46,31 +66,33 @@ func (s *AgentHandlerService) Build(ctx context.Context, input BuildAgentHandler
|
|||||||
}
|
}
|
||||||
handlers = append(handlers, toolSearchHandler)
|
handlers = append(handlers, toolSearchHandler)
|
||||||
}
|
}
|
||||||
skillMetadataByCode := buildRuntimeSkillMetadataMap(input.AIAgent)
|
|
||||||
if len(skillMetadataByCode) > 0 {
|
if len(skillMetadataByCode) > 0 {
|
||||||
skillHandler, err := s.skillMiddleware.Build(ctx, input.AIAgent, input.InstructionToolDefinitions)
|
skillHandler, err := s.skillMiddleware.Build(ctx, input.AIAgent, input.InstructionToolDefinitions)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
handlers = append(handlers, skillHandler)
|
handlers = append(handlers, skillHandler)
|
||||||
}
|
handlers = append(handlers, NewRuntimeToolFilterMiddleware(
|
||||||
if input.Collector != nil {
|
input.Collector,
|
||||||
toolMetadataBy := buildRuntimeTraceToolMetadata(input.DynamicToolDefinitions, input.StaticToolMetadata, len(skillMetadataByCode) > 0)
|
toolMetadataBy,
|
||||||
traceSkillMetadata := make(map[string]einocallbacks.SkillMetadata, len(skillMetadataByCode))
|
traceSkillMetadata,
|
||||||
for code, item := range skillMetadataByCode {
|
dynamicToolModelNames(input.DynamicToolDefinitions),
|
||||||
traceSkillMetadata[code] = einocallbacks.SkillMetadata{
|
))
|
||||||
Code: item.Code,
|
|
||||||
Name: item.Name,
|
|
||||||
Description: item.Description,
|
|
||||||
AllowedToolCodes: append([]string(nil), item.AllowedToolCodes...),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if len(skillMetadataByCode) > 0 {
|
|
||||||
input.Collector.SetSkillMiddleware(true, toolx.BuiltinSkill.Name)
|
|
||||||
}
|
|
||||||
input.Collector.SetVisibleSkills(traceSkillMetadata)
|
|
||||||
input.Collector.SetInstructionSummary(input.InstructionSummary)
|
|
||||||
handlers = append(handlers, einocallbacks.NewRuntimeTraceHandler(input.Collector, toolMetadataBy, traceSkillMetadata))
|
|
||||||
}
|
}
|
||||||
return handlers, nil
|
return handlers, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func dynamicToolModelNames(definitions []runtimetooling.MCPToolDefinition) []string {
|
||||||
|
ret := make([]string, 0, len(definitions))
|
||||||
|
for _, item := range definitions {
|
||||||
|
modelName := strings.TrimSpace(item.ModelName)
|
||||||
|
if modelName == "" {
|
||||||
|
modelName = strings.TrimSpace(runtimetooling.BuildModelToolName(item))
|
||||||
|
}
|
||||||
|
if modelName == "" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
ret = append(ret, modelName)
|
||||||
|
}
|
||||||
|
return ret
|
||||||
|
}
|
||||||
|
|||||||
@@ -0,0 +1,402 @@
|
|||||||
|
package factory
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
einocallbacks "cs-agent/internal/ai/runtime/internal/impl/callbacks"
|
||||||
|
"cs-agent/internal/pkg/toolx"
|
||||||
|
|
||||||
|
"github.com/cloudwego/eino/adk"
|
||||||
|
"github.com/cloudwego/eino/components/model"
|
||||||
|
einotool "github.com/cloudwego/eino/components/tool"
|
||||||
|
"github.com/cloudwego/eino/schema"
|
||||||
|
)
|
||||||
|
|
||||||
|
const activeSkillRunLocalKey = "runtime_active_skill_code"
|
||||||
|
|
||||||
|
type RuntimeToolFilterMiddleware struct {
|
||||||
|
*adk.BaseChatModelAgentMiddleware
|
||||||
|
collector *einocallbacks.RuntimeTraceCollector
|
||||||
|
toolMetadataByName map[string]einocallbacks.ToolMetadata
|
||||||
|
skillMetadataBy map[string]einocallbacks.SkillMetadata
|
||||||
|
dynamicToolNames []string
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewRuntimeToolFilterMiddleware(
|
||||||
|
collector *einocallbacks.RuntimeTraceCollector,
|
||||||
|
toolMetadataByName map[string]einocallbacks.ToolMetadata,
|
||||||
|
skillMetadataBy map[string]einocallbacks.SkillMetadata,
|
||||||
|
dynamicToolNames []string,
|
||||||
|
) *RuntimeToolFilterMiddleware {
|
||||||
|
return &RuntimeToolFilterMiddleware{
|
||||||
|
BaseChatModelAgentMiddleware: &adk.BaseChatModelAgentMiddleware{},
|
||||||
|
collector: collector,
|
||||||
|
toolMetadataByName: cloneToolMetadataMap(toolMetadataByName),
|
||||||
|
skillMetadataBy: cloneSkillMetadataMap(skillMetadataBy),
|
||||||
|
dynamicToolNames: append([]string(nil), dynamicToolNames...),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *RuntimeToolFilterMiddleware) WrapModel(_ context.Context, cm model.BaseChatModel, mc *adk.ModelContext) (model.BaseChatModel, error) {
|
||||||
|
if mc == nil {
|
||||||
|
return cm, nil
|
||||||
|
}
|
||||||
|
return &runtimeToolFilterModelWrapper{
|
||||||
|
cm: cm,
|
||||||
|
allTools: append([]*schema.ToolInfo(nil), mc.Tools...),
|
||||||
|
collector: m.collector,
|
||||||
|
toolMetadataByName: m.toolMetadataByName,
|
||||||
|
skillMetadataBy: m.skillMetadataBy,
|
||||||
|
dynamicToolNames: append([]string(nil), m.dynamicToolNames...),
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *RuntimeToolFilterMiddleware) WrapInvokableToolCall(_ context.Context, endpoint adk.InvokableToolCallEndpoint, tCtx *adk.ToolContext) (adk.InvokableToolCallEndpoint, error) {
|
||||||
|
return func(ctx context.Context, argumentsInJSON string, opts ...einotool.Option) (string, error) {
|
||||||
|
toolName := ""
|
||||||
|
if tCtx != nil {
|
||||||
|
toolName = strings.TrimSpace(tCtx.Name)
|
||||||
|
}
|
||||||
|
metadata, _ := resolveRuntimeToolMetadata(toolName, m.toolMetadataByName)
|
||||||
|
if !isRuntimeBuiltinAlwaysAllowed(metadata.ToolCode) {
|
||||||
|
activeSkill, restricted := m.resolveActiveSkill(ctx)
|
||||||
|
if restricted && !isToolCodeAllowedForSkill(metadata.ToolCode, activeSkill.AllowedToolCodes) {
|
||||||
|
return "", m.blockToolCall(metadata, argumentsInJSON, activeSkill)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
result, err := endpoint(ctx, argumentsInJSON, opts...)
|
||||||
|
if err != nil {
|
||||||
|
return result, err
|
||||||
|
}
|
||||||
|
if strings.TrimSpace(metadata.ToolCode) == toolx.BuiltinSkill.Code {
|
||||||
|
_ = m.setActiveSkill(ctx, skillCodeFromArguments(argumentsInJSON))
|
||||||
|
return result, nil
|
||||||
|
}
|
||||||
|
if strings.TrimSpace(metadata.ToolCode) == toolx.BuiltinToolSearch.Code {
|
||||||
|
activeSkill, restricted := m.resolveActiveSkill(ctx)
|
||||||
|
if !restricted {
|
||||||
|
return result, nil
|
||||||
|
}
|
||||||
|
filtered, filterErr := filterToolSearchResult(result, activeSkill.AllowedToolCodes, m.toolMetadataByName)
|
||||||
|
if filterErr == nil {
|
||||||
|
return filtered, nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return result, nil
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *RuntimeToolFilterMiddleware) blockToolCall(metadata einocallbacks.ToolMetadata, argumentsInJSON string, activeSkill einocallbacks.SkillMetadata) error {
|
||||||
|
err := fmt.Errorf("tool %s is not allowed for active skill %s", strings.TrimSpace(metadata.ToolCode), strings.TrimSpace(activeSkill.Code))
|
||||||
|
if m.collector != nil {
|
||||||
|
m.collector.AddToolItem(einocallbacks.ToolTraceItem{
|
||||||
|
ToolCode: strings.TrimSpace(metadata.ToolCode),
|
||||||
|
ServerCode: strings.TrimSpace(metadata.ServerCode),
|
||||||
|
ToolName: strings.TrimSpace(metadata.ToolName),
|
||||||
|
Arguments: parseRuntimeToolArguments(argumentsInJSON),
|
||||||
|
Status: "error",
|
||||||
|
ErrorMessage: err.Error(),
|
||||||
|
Blocked: true,
|
||||||
|
BlockedReason: "skill_tool_not_allowed",
|
||||||
|
})
|
||||||
|
}
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *RuntimeToolFilterMiddleware) setActiveSkill(ctx context.Context, skillCode string) error {
|
||||||
|
skillCode = strings.TrimSpace(skillCode)
|
||||||
|
if skillCode == "" {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return adk.SetRunLocalValue(ctx, activeSkillRunLocalKey, skillCode)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *RuntimeToolFilterMiddleware) resolveActiveSkill(ctx context.Context) (einocallbacks.SkillMetadata, bool) {
|
||||||
|
if len(m.skillMetadataBy) == 0 {
|
||||||
|
return einocallbacks.SkillMetadata{}, false
|
||||||
|
}
|
||||||
|
value, found, err := adk.GetRunLocalValue(ctx, activeSkillRunLocalKey)
|
||||||
|
if err != nil || !found {
|
||||||
|
return einocallbacks.SkillMetadata{}, false
|
||||||
|
}
|
||||||
|
code, ok := value.(string)
|
||||||
|
if !ok {
|
||||||
|
return einocallbacks.SkillMetadata{}, false
|
||||||
|
}
|
||||||
|
code = strings.TrimSpace(code)
|
||||||
|
if code == "" {
|
||||||
|
return einocallbacks.SkillMetadata{}, false
|
||||||
|
}
|
||||||
|
skill, ok := m.skillMetadataBy[code]
|
||||||
|
if !ok {
|
||||||
|
return einocallbacks.SkillMetadata{}, false
|
||||||
|
}
|
||||||
|
if len(skill.AllowedToolCodes) == 0 {
|
||||||
|
return skill, false
|
||||||
|
}
|
||||||
|
return skill, true
|
||||||
|
}
|
||||||
|
|
||||||
|
type runtimeToolFilterModelWrapper struct {
|
||||||
|
cm model.BaseChatModel
|
||||||
|
allTools []*schema.ToolInfo
|
||||||
|
collector *einocallbacks.RuntimeTraceCollector
|
||||||
|
toolMetadataByName map[string]einocallbacks.ToolMetadata
|
||||||
|
skillMetadataBy map[string]einocallbacks.SkillMetadata
|
||||||
|
dynamicToolNames []string
|
||||||
|
}
|
||||||
|
|
||||||
|
func (w *runtimeToolFilterModelWrapper) Generate(ctx context.Context, input []*schema.Message, opts ...model.Option) (*schema.Message, error) {
|
||||||
|
tools := w.filteredTools(ctx, input)
|
||||||
|
return w.cm.Generate(ctx, input, append(opts, model.WithTools(tools))...)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (w *runtimeToolFilterModelWrapper) Stream(ctx context.Context, input []*schema.Message, opts ...model.Option) (*schema.StreamReader[*schema.Message], error) {
|
||||||
|
tools := w.filteredTools(ctx, input)
|
||||||
|
return w.cm.Stream(ctx, input, append(opts, model.WithTools(tools))...)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (w *runtimeToolFilterModelWrapper) filteredTools(ctx context.Context, input []*schema.Message) []*schema.ToolInfo {
|
||||||
|
tools := filterDynamicToolInfos(w.allTools, w.dynamicToolNames, input)
|
||||||
|
activeSkill, restricted := resolveActiveSkillMetadata(ctx, w.skillMetadataBy)
|
||||||
|
if restricted {
|
||||||
|
tools = filterToolInfosBySkill(tools, w.toolMetadataByName, activeSkill.AllowedToolCodes)
|
||||||
|
}
|
||||||
|
if w.collector != nil {
|
||||||
|
w.collector.SetFilteredToolCodes(extractToolCodesFromInfos(tools, w.toolMetadataByName))
|
||||||
|
}
|
||||||
|
return tools
|
||||||
|
}
|
||||||
|
|
||||||
|
func resolveActiveSkillMetadata(ctx context.Context, skills map[string]einocallbacks.SkillMetadata) (einocallbacks.SkillMetadata, bool) {
|
||||||
|
if len(skills) == 0 {
|
||||||
|
return einocallbacks.SkillMetadata{}, false
|
||||||
|
}
|
||||||
|
value, found, err := adk.GetRunLocalValue(ctx, activeSkillRunLocalKey)
|
||||||
|
if err != nil || !found {
|
||||||
|
return einocallbacks.SkillMetadata{}, false
|
||||||
|
}
|
||||||
|
code, ok := value.(string)
|
||||||
|
if !ok {
|
||||||
|
return einocallbacks.SkillMetadata{}, false
|
||||||
|
}
|
||||||
|
code = strings.TrimSpace(code)
|
||||||
|
skill, ok := skills[code]
|
||||||
|
if !ok || len(skill.AllowedToolCodes) == 0 {
|
||||||
|
return skill, false
|
||||||
|
}
|
||||||
|
return skill, true
|
||||||
|
}
|
||||||
|
|
||||||
|
func filterDynamicToolInfos(allTools []*schema.ToolInfo, dynamicToolNames []string, messages []*schema.Message) []*schema.ToolInfo {
|
||||||
|
if len(allTools) == 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
selectedToolNames := extractSelectedDynamicToolNames(messages)
|
||||||
|
if len(dynamicToolNames) == 0 {
|
||||||
|
return append([]*schema.ToolInfo(nil), allTools...)
|
||||||
|
}
|
||||||
|
removeMap := invertStringSelection(dynamicToolNames, selectedToolNames)
|
||||||
|
ret := make([]*schema.ToolInfo, 0, len(allTools))
|
||||||
|
for _, info := range allTools {
|
||||||
|
if info == nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if _, ok := removeMap[strings.TrimSpace(info.Name)]; ok {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
ret = append(ret, info)
|
||||||
|
}
|
||||||
|
return ret
|
||||||
|
}
|
||||||
|
|
||||||
|
func extractSelectedDynamicToolNames(messages []*schema.Message) []string {
|
||||||
|
if len(messages) == 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
selected := make([]string, 0)
|
||||||
|
for _, message := range messages {
|
||||||
|
if message == nil || message.Role != schema.Tool || strings.TrimSpace(message.ToolName) != toolx.BuiltinToolSearch.Name {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
var payload struct {
|
||||||
|
SelectedTools []string `json:"selectedTools"`
|
||||||
|
}
|
||||||
|
if err := json.Unmarshal([]byte(strings.TrimSpace(message.Content)), &payload); err != nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
for _, item := range payload.SelectedTools {
|
||||||
|
item = strings.TrimSpace(item)
|
||||||
|
if item == "" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
selected = append(selected, item)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return selected
|
||||||
|
}
|
||||||
|
|
||||||
|
func invertStringSelection(all []string, selected []string) map[string]struct{} {
|
||||||
|
selectedSet := make(map[string]struct{}, len(selected))
|
||||||
|
for _, item := range selected {
|
||||||
|
item = strings.TrimSpace(item)
|
||||||
|
if item == "" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
selectedSet[item] = struct{}{}
|
||||||
|
}
|
||||||
|
ret := make(map[string]struct{})
|
||||||
|
for _, item := range all {
|
||||||
|
item = strings.TrimSpace(item)
|
||||||
|
if item == "" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if _, ok := selectedSet[item]; ok {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
ret[item] = struct{}{}
|
||||||
|
}
|
||||||
|
return ret
|
||||||
|
}
|
||||||
|
|
||||||
|
func filterToolInfosBySkill(allTools []*schema.ToolInfo, toolMetadataByName map[string]einocallbacks.ToolMetadata, allowedToolCodes []string) []*schema.ToolInfo {
|
||||||
|
if len(allTools) == 0 || len(allowedToolCodes) == 0 {
|
||||||
|
return append([]*schema.ToolInfo(nil), allTools...)
|
||||||
|
}
|
||||||
|
ret := make([]*schema.ToolInfo, 0, len(allTools))
|
||||||
|
for _, info := range allTools {
|
||||||
|
if info == nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
metadata, _ := resolveRuntimeToolMetadata(strings.TrimSpace(info.Name), toolMetadataByName)
|
||||||
|
if isRuntimeBuiltinAlwaysAllowed(metadata.ToolCode) || isToolCodeAllowedForSkill(metadata.ToolCode, allowedToolCodes) {
|
||||||
|
ret = append(ret, info)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return ret
|
||||||
|
}
|
||||||
|
|
||||||
|
func extractToolCodesFromInfos(infos []*schema.ToolInfo, toolMetadataByName map[string]einocallbacks.ToolMetadata) []string {
|
||||||
|
ret := make([]string, 0, len(infos))
|
||||||
|
for _, info := range infos {
|
||||||
|
if info == nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
metadata, ok := resolveRuntimeToolMetadata(strings.TrimSpace(info.Name), toolMetadataByName)
|
||||||
|
if !ok || strings.TrimSpace(metadata.ToolCode) == "" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
ret = append(ret, metadata.ToolCode)
|
||||||
|
}
|
||||||
|
return toolx.NormalizeToolCodes(ret)
|
||||||
|
}
|
||||||
|
|
||||||
|
func isToolCodeAllowedForSkill(toolCode string, allowedToolCodes []string) bool {
|
||||||
|
toolCode = toolx.NormalizeToolCodeAlias(strings.TrimSpace(toolCode))
|
||||||
|
if toolCode == "" {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
for _, item := range toolx.NormalizeToolCodes(allowedToolCodes) {
|
||||||
|
if item == toolCode {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
func isRuntimeBuiltinAlwaysAllowed(toolCode string) bool {
|
||||||
|
toolCode = toolx.NormalizeToolCodeAlias(strings.TrimSpace(toolCode))
|
||||||
|
return toolCode == toolx.BuiltinSkill.Code || toolCode == toolx.BuiltinToolSearch.Code
|
||||||
|
}
|
||||||
|
|
||||||
|
func resolveRuntimeToolMetadata(toolName string, toolMetadataByName map[string]einocallbacks.ToolMetadata) (einocallbacks.ToolMetadata, bool) {
|
||||||
|
toolName = strings.TrimSpace(toolName)
|
||||||
|
if toolName == "" {
|
||||||
|
return einocallbacks.ToolMetadata{}, false
|
||||||
|
}
|
||||||
|
if spec, ok := toolx.GetRegisteredToolSpecByName(toolName); ok {
|
||||||
|
resolved := toolx.ResolveToolMetadata(spec.Code, spec.Name)
|
||||||
|
return einocallbacks.ToolMetadata{
|
||||||
|
ToolCode: resolved.ToolCode,
|
||||||
|
ServerCode: resolved.ServerCode,
|
||||||
|
ToolName: resolved.ToolName,
|
||||||
|
SourceType: resolved.SourceType,
|
||||||
|
}, true
|
||||||
|
}
|
||||||
|
metadata, ok := toolMetadataByName[toolName]
|
||||||
|
return metadata, ok
|
||||||
|
}
|
||||||
|
|
||||||
|
func skillCodeFromArguments(argumentsInJSON string) string {
|
||||||
|
var args struct {
|
||||||
|
Skill string `json:"skill"`
|
||||||
|
}
|
||||||
|
if err := json.Unmarshal([]byte(strings.TrimSpace(argumentsInJSON)), &args); err != nil {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
return strings.TrimSpace(args.Skill)
|
||||||
|
}
|
||||||
|
|
||||||
|
func parseRuntimeToolArguments(argumentsInJSON string) map[string]any {
|
||||||
|
argumentsInJSON = strings.TrimSpace(argumentsInJSON)
|
||||||
|
if argumentsInJSON == "" {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
ret := make(map[string]any)
|
||||||
|
if err := json.Unmarshal([]byte(argumentsInJSON), &ret); err != nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return ret
|
||||||
|
}
|
||||||
|
|
||||||
|
func filterToolSearchResult(result string, allowedToolCodes []string, toolMetadataByName map[string]einocallbacks.ToolMetadata) (string, error) {
|
||||||
|
result = strings.TrimSpace(result)
|
||||||
|
if result == "" {
|
||||||
|
return result, nil
|
||||||
|
}
|
||||||
|
var payload struct {
|
||||||
|
SelectedTools []string `json:"selectedTools"`
|
||||||
|
}
|
||||||
|
if err := json.Unmarshal([]byte(result), &payload); err != nil {
|
||||||
|
return result, err
|
||||||
|
}
|
||||||
|
filtered := make([]string, 0, len(payload.SelectedTools))
|
||||||
|
for _, toolName := range payload.SelectedTools {
|
||||||
|
metadata, _ := resolveRuntimeToolMetadata(toolName, toolMetadataByName)
|
||||||
|
if isToolCodeAllowedForSkill(metadata.ToolCode, allowedToolCodes) {
|
||||||
|
filtered = append(filtered, strings.TrimSpace(toolName))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
payload.SelectedTools = filtered
|
||||||
|
buf, err := json.Marshal(payload)
|
||||||
|
if err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
return string(buf), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func cloneToolMetadataMap(input map[string]einocallbacks.ToolMetadata) map[string]einocallbacks.ToolMetadata {
|
||||||
|
if len(input) == 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
ret := make(map[string]einocallbacks.ToolMetadata, len(input))
|
||||||
|
for key, value := range input {
|
||||||
|
ret[key] = value
|
||||||
|
}
|
||||||
|
return ret
|
||||||
|
}
|
||||||
|
|
||||||
|
func cloneSkillMetadataMap(input map[string]einocallbacks.SkillMetadata) map[string]einocallbacks.SkillMetadata {
|
||||||
|
if len(input) == 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
ret := make(map[string]einocallbacks.SkillMetadata, len(input))
|
||||||
|
for key, value := range input {
|
||||||
|
value.AllowedToolCodes = append([]string(nil), value.AllowedToolCodes...)
|
||||||
|
ret[key] = value
|
||||||
|
}
|
||||||
|
return ret
|
||||||
|
}
|
||||||
@@ -0,0 +1,84 @@
|
|||||||
|
package factory
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
einocallbacks "cs-agent/internal/ai/runtime/internal/impl/callbacks"
|
||||||
|
"cs-agent/internal/pkg/toolx"
|
||||||
|
|
||||||
|
"github.com/cloudwego/eino/schema"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestFilterDynamicToolInfos(t *testing.T) {
|
||||||
|
allTools := []*schema.ToolInfo{
|
||||||
|
{Name: toolx.BuiltinToolSearch.Name},
|
||||||
|
{Name: "mcp_server_a"},
|
||||||
|
{Name: "mcp_server_b"},
|
||||||
|
}
|
||||||
|
messages := []*schema.Message{
|
||||||
|
{Role: schema.Tool, ToolName: toolx.BuiltinToolSearch.Name, Content: `{"selectedTools":["mcp_server_b"]}`},
|
||||||
|
}
|
||||||
|
|
||||||
|
filtered := filterDynamicToolInfos(allTools, []string{"mcp_server_a", "mcp_server_b"}, messages)
|
||||||
|
if len(filtered) != 2 {
|
||||||
|
t.Fatalf("unexpected filtered tool count: %d", len(filtered))
|
||||||
|
}
|
||||||
|
if filtered[0].Name != toolx.BuiltinToolSearch.Name || filtered[1].Name != "mcp_server_b" {
|
||||||
|
t.Fatalf("unexpected filtered tools: %#v", filtered)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFilterToolInfosBySkill(t *testing.T) {
|
||||||
|
allTools := []*schema.ToolInfo{
|
||||||
|
{Name: toolx.BuiltinSkill.Name},
|
||||||
|
{Name: toolx.BuiltinToolSearch.Name},
|
||||||
|
{Name: toolx.GraphHandoffConversation.Name},
|
||||||
|
{Name: "mcp_server_refund"},
|
||||||
|
}
|
||||||
|
toolMetadataByName := map[string]einocallbacks.ToolMetadata{
|
||||||
|
toolx.GraphHandoffConversation.Name: {ToolCode: toolx.GraphHandoffConversation.Code, ToolName: toolx.GraphHandoffConversation.Name},
|
||||||
|
"mcp_server_refund": {ToolCode: "mcp/refund", ToolName: "mcp_server_refund"},
|
||||||
|
}
|
||||||
|
|
||||||
|
filtered := filterToolInfosBySkill(allTools, toolMetadataByName, []string{toolx.GraphHandoffConversation.Code})
|
||||||
|
if len(filtered) != 3 {
|
||||||
|
t.Fatalf("unexpected filtered tool count: %d", len(filtered))
|
||||||
|
}
|
||||||
|
if filtered[0].Name != toolx.BuiltinSkill.Name || filtered[1].Name != toolx.BuiltinToolSearch.Name || filtered[2].Name != toolx.GraphHandoffConversation.Name {
|
||||||
|
t.Fatalf("unexpected filtered tools: %#v", filtered)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFilterToolSearchResult(t *testing.T) {
|
||||||
|
toolMetadataByName := map[string]einocallbacks.ToolMetadata{
|
||||||
|
"mcp_server_refund": {ToolCode: "mcp/refund", ToolName: "mcp_server_refund"},
|
||||||
|
"mcp_server_order": {ToolCode: "mcp/order", ToolName: "mcp_server_order"},
|
||||||
|
}
|
||||||
|
|
||||||
|
got, err := filterToolSearchResult(`{"selectedTools":["mcp_server_refund","mcp_server_order"]}`, []string{"mcp/order"}, toolMetadataByName)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("filterToolSearchResult returned error: %v", err)
|
||||||
|
}
|
||||||
|
if got != `{"selectedTools":["mcp_server_order"]}` {
|
||||||
|
t.Fatalf("unexpected filtered result: %s", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestExtractToolCodesFromInfos(t *testing.T) {
|
||||||
|
infos := []*schema.ToolInfo{
|
||||||
|
{Name: toolx.BuiltinSkill.Name},
|
||||||
|
{Name: toolx.GraphPrepareTicketDraft.Name},
|
||||||
|
{Name: "mcp_server_refund"},
|
||||||
|
}
|
||||||
|
toolMetadataByName := map[string]einocallbacks.ToolMetadata{
|
||||||
|
toolx.GraphPrepareTicketDraft.Name: {ToolCode: toolx.GraphPrepareTicketDraft.Code, ToolName: toolx.GraphPrepareTicketDraft.Name},
|
||||||
|
"mcp_server_refund": {ToolCode: "mcp/refund", ToolName: "mcp_server_refund"},
|
||||||
|
}
|
||||||
|
got := extractToolCodesFromInfos(infos, toolMetadataByName)
|
||||||
|
if len(got) != 3 {
|
||||||
|
t.Fatalf("unexpected tool codes: %#v", got)
|
||||||
|
}
|
||||||
|
if got[0] != toolx.BuiltinSkill.Code || got[1] != toolx.GraphPrepareTicketDraft.Code || got[2] != "mcp/refund" {
|
||||||
|
t.Fatalf("unexpected tool codes order: %#v", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user