feat: refactor prepare service to use toolCatalog for tool resolution and add unit tests for tool catalog functionality

This commit is contained in:
mlogclub
2026-04-13 19:20:14 +08:00
parent 2effb35af3
commit 1ebbdadc1b
4 changed files with 198 additions and 120 deletions
+13 -99
View File
@@ -5,18 +5,16 @@ import (
"encoding/json"
"strings"
"cs-agent/internal/ai/runtime/registry"
"cs-agent/internal/ai/skills"
"cs-agent/internal/models"
"cs-agent/internal/pkg/toolx"
)
func newPrepareService(registry *registry.Registry) *prepareService {
return &prepareService{registry: registry}
func newPrepareService(catalog *toolCatalog) *prepareService {
return &prepareService{catalog: catalog}
}
type prepareService struct {
registry *registry.Registry
catalog *toolCatalog
}
func (s *prepareService) selectSkill(ctx context.Context, req Request) (*models.SkillDefinition, string, string, error) {
@@ -44,37 +42,30 @@ func (s *prepareService) selectSkill(ctx context.Context, req Request) (*models.
}
func (s *prepareService) prepareToolsForRun(req *Request) error {
if req == nil || req.ToolSet != nil || s.registry == nil {
if req == nil || req.ToolSet != nil || s.catalog == nil {
return nil
}
toolSet, err := s.registry.Resolve(registry.Context{
Conversation: req.Conversation,
AIAgent: req.AIAgent,
AIConfig: req.AIConfig,
UserMessage: req.UserMessage,
AllowedToolCodes: resolveAllowedToolCodes(req.AIAgent, req.SelectedSkill),
})
toolSet, err := s.catalog.resolveForRun(req)
if err != nil {
return err
}
req.ToolSet = toolSet
if toolSet != nil {
req.ToolSet = toolSet
}
return nil
}
func (s *prepareService) prepareToolsForResume(req *ResumeRequest) error {
if req == nil || req.ToolSet != nil || s.registry == nil {
if req == nil || req.ToolSet != nil || s.catalog == nil {
return nil
}
toolSet, err := s.registry.Resolve(registry.Context{
Conversation: req.Conversation,
AIAgent: req.AIAgent,
AIConfig: req.AIConfig,
AllowedToolCodes: parseAgentAllowedToolCodes(req.AIAgent),
})
toolSet, err := s.catalog.resolveForResume(req)
if err != nil {
return err
}
req.ToolSet = toolSet
if toolSet != nil {
req.ToolSet = toolSet
}
return nil
}
@@ -96,80 +87,3 @@ func cloneSkillDefinition(item *models.SkillDefinition) *models.SkillDefinition
clone := *item
return &clone
}
func 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
}
ret := make([]string, 0, len(items))
for _, item := range items {
item = strings.TrimSpace(item)
item = toolx.NormalizeToolCodeAlias(item)
if item == "" {
continue
}
ret = append(ret, item)
}
return ret
}
func parseAgentAllowedToolCodes(aiAgent *models.AIAgent) []string {
if aiAgent == nil || strings.TrimSpace(aiAgent.AllowedMCPTools) == "" {
return nil
}
items, err := toolx.ParseAgentMCPToolsJSON(aiAgent.AllowedMCPTools)
if err != nil {
return nil
}
ret := make([]string, 0, len(items))
for _, item := range items {
toolCode := strings.TrimSpace(item.ToolCode)
toolCode = toolx.NormalizeToolCodeAlias(toolCode)
if toolCode == "" {
continue
}
ret = append(ret, toolCode)
}
return ret
}
func resolveAllowedToolCodes(aiAgent *models.AIAgent, skill *models.SkillDefinition) []string {
agentAllowed := parseAgentAllowedToolCodes(aiAgent)
skillAllowed := parseSkillAllowedToolCodes(skill)
switch {
case len(agentAllowed) == 0:
return skillAllowed
case len(skillAllowed) == 0:
return agentAllowed
default:
skillSet := make(map[string]struct{}, len(skillAllowed))
for _, item := range skillAllowed {
item = strings.TrimSpace(item)
item = toolx.NormalizeToolCodeAlias(item)
if item == "" {
continue
}
skillSet[item] = struct{}{}
}
ret := make([]string, 0, len(agentAllowed))
for _, item := range agentAllowed {
item = strings.TrimSpace(item)
item = toolx.NormalizeToolCodeAlias(item)
if item == "" {
continue
}
if _, ok := skillSet[item]; ok {
ret = append(ret, item)
}
}
return ret
}
}