refactor: enhance MCP tool management with risk policies and confirmation handling
This commit is contained in:
@@ -206,6 +206,7 @@ func (e *AgentLoopEngine) Resume(ctx context.Context, req ResumeInput) (*RunResu
|
||||
return nil, err
|
||||
}
|
||||
req.AIAgent, req.AIConfig = snapshot.Agent, snapshot.AIConfig
|
||||
req.ResumeData = normalizeAgentLoopResumeData(req.UserMessage.MessageType, req.ResumeData)
|
||||
if interrupt.WorkflowRunID > 0 {
|
||||
workflowRun, _ := svc.AIWorkflowService.GetRunDetail(interrupt.WorkflowRunID)
|
||||
if workflowRun == nil {
|
||||
@@ -240,16 +241,35 @@ func (e *AgentLoopEngine) Resume(ctx context.Context, req ResumeInput) (*RunResu
|
||||
if err := json.Unmarshal([]byte(interrupt.RequestData), &checkpoint); err != nil {
|
||||
return nil, errorsx.InvalidParam("invalid MCP checkpoint data")
|
||||
}
|
||||
if !isAgentLoopConfirmation(firstAgentLoopResumeText(req.ResumeData)) {
|
||||
tool, err := configuredMCPTool(req.AIAgent.AllowedMCPTools, checkpoint.ToolCode)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
switch parseAgentLoopConfirmation(firstAgentLoopResumeText(req.ResumeData)) {
|
||||
case agentLoopConfirmationCancelled:
|
||||
ret := &RunResult{
|
||||
Status: "completed", ReplyText: "操作已取消。", ModelName: req.AIConfig.ModelName,
|
||||
AgentRunID: interrupt.AgentRunID,
|
||||
}
|
||||
return ret, recordAgentLoopResume(interrupt.AgentRunID, 0, ret.Status, ret.ReplyText, nil)
|
||||
}
|
||||
tool, err := configuredMCPTool(req.AIAgent.AllowedMCPTools, checkpoint.ToolCode)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
case agentLoopConfirmationUnknown:
|
||||
prompt := buildAgentLoopMCPConfirmationRetryPrompt(tool.Title)
|
||||
ret := &RunResult{
|
||||
Status: "interrupted",
|
||||
ReplyText: prompt,
|
||||
ModelName: req.AIConfig.ModelName,
|
||||
AgentRunID: interrupt.AgentRunID,
|
||||
CheckPointID: interrupt.CheckPointID,
|
||||
CheckPointData: interrupt.RequestData,
|
||||
Interrupted: true,
|
||||
Interrupts: []InterruptContextSummary{{
|
||||
Type: "tool_confirmation",
|
||||
ID: checkpoint.ToolCode,
|
||||
DisplayName: tool.Title,
|
||||
PromptText: prompt,
|
||||
}},
|
||||
}
|
||||
return ret, recordAgentLoopResume(interrupt.AgentRunID, 0, ret.Status, ret.ReplyText, nil)
|
||||
}
|
||||
policy := parseAgentLoopToolPolicy(req.AIAgent.ToolPolicy)
|
||||
executionPolicy := aitooling.Policy{
|
||||
@@ -299,12 +319,32 @@ func firstAgentLoopResumeText(data map[string]string) string {
|
||||
return ""
|
||||
}
|
||||
|
||||
func isAgentLoopConfirmation(value string) bool {
|
||||
switch strings.ToLower(strings.TrimSpace(value)) {
|
||||
type agentLoopConfirmationDecision int
|
||||
|
||||
const (
|
||||
agentLoopConfirmationUnknown agentLoopConfirmationDecision = iota
|
||||
agentLoopConfirmationConfirmed
|
||||
agentLoopConfirmationCancelled
|
||||
)
|
||||
|
||||
func normalizeAgentLoopResumeData(messageType enums.IMMessageType, data map[string]string) map[string]string {
|
||||
ret := make(map[string]string, len(data))
|
||||
for key, value := range data {
|
||||
ret[key] = strings.TrimSpace(utils.BuildRuntimeMessageText(messageType, value))
|
||||
}
|
||||
return ret
|
||||
}
|
||||
|
||||
func parseAgentLoopConfirmation(value string) agentLoopConfirmationDecision {
|
||||
normalized := strings.ToLower(strings.TrimSpace(value))
|
||||
normalized = strings.TrimSpace(strings.Trim(normalized, "。.!!??"))
|
||||
switch normalized {
|
||||
case "确认", "确认执行", "同意", "继续", "是", "yes", "y", "confirm", "approve", "approved":
|
||||
return true
|
||||
return agentLoopConfirmationConfirmed
|
||||
case "取消", "取消执行", "不同意", "拒绝", "否", "不要", "停止", "no", "n", "cancel", "reject", "rejected":
|
||||
return agentLoopConfirmationCancelled
|
||||
default:
|
||||
return false
|
||||
return agentLoopConfirmationUnknown
|
||||
}
|
||||
}
|
||||
|
||||
@@ -610,9 +650,15 @@ func executeAgentLoopMCP(ctx context.Context, runInput RunInput, toolCode string
|
||||
checkpoint := agentLoopMCPCheckpoint{ToolCode: toolCode, Arguments: arguments}
|
||||
data, _ := json.Marshal(checkpoint)
|
||||
checkPointID := fmt.Sprintf("tool:%d:%d", runInput.Conversation.ID, time.Now().UnixNano())
|
||||
prompt := buildAgentLoopMCPConfirmationPrompt(configured.Title)
|
||||
state.Interrupted = &RunResult{
|
||||
Status: "interrupted", ReplyText: "请确认是否执行该操作。", CheckPointID: checkPointID, CheckPointData: string(data), Interrupted: true,
|
||||
Interrupts: []InterruptContextSummary{{Type: "tool_confirmation", ID: toolCode, InfoPreview: configured.Title}},
|
||||
Status: "interrupted", ReplyText: prompt, CheckPointID: checkPointID, CheckPointData: string(data), Interrupted: true,
|
||||
Interrupts: []InterruptContextSummary{{
|
||||
Type: "tool_confirmation",
|
||||
ID: toolCode,
|
||||
DisplayName: configured.Title,
|
||||
PromptText: prompt,
|
||||
}},
|
||||
}
|
||||
return definition, "", &agentLoopInterruptError{reason: "Agent Loop interrupted for MCP confirmation"}
|
||||
}
|
||||
@@ -633,12 +679,28 @@ func configuredMCPTool(raw, toolCode string) (request.AIAgentMCPToolRequest, err
|
||||
}
|
||||
for _, item := range items {
|
||||
if item.ToolCode == toolCode {
|
||||
return item, nil
|
||||
return toolx.ApplyTrustedMCPToolPolicy(item), nil
|
||||
}
|
||||
}
|
||||
return request.AIAgentMCPToolRequest{}, errorsx.InvalidParam("MCP tool is not configured for this Agent")
|
||||
}
|
||||
|
||||
func buildAgentLoopMCPConfirmationPrompt(title string) string {
|
||||
title = strings.TrimSpace(title)
|
||||
if title == "" {
|
||||
return "即将执行一项操作,是否确认继续?"
|
||||
}
|
||||
return fmt.Sprintf("即将执行“%s”,是否确认继续?", title)
|
||||
}
|
||||
|
||||
func buildAgentLoopMCPConfirmationRetryPrompt(title string) string {
|
||||
title = strings.TrimSpace(title)
|
||||
if title == "" {
|
||||
return "未识别您的选择,请回复“确认”继续执行,或回复“取消”终止操作。"
|
||||
}
|
||||
return fmt.Sprintf("未识别您的选择。若要继续执行“%s”,请回复“确认”;若要终止,请回复“取消”。", title)
|
||||
}
|
||||
|
||||
func executeAgentLoopReadTool(ctx context.Context, conversation models.Conversation, agent models.AIAgent, toolCode string, arguments map[string]any, policy aitooling.Policy) (aitooling.Definition, string, error) {
|
||||
toolCode = toolx.NormalizeToolCodeAlias(strings.TrimSpace(toolCode))
|
||||
if toolCode != toolx.BuiltinConversationContext.Code && toolCode != toolx.BuiltinKnowledgeRetrieve.Code && toolCode != toolx.GraphTriageServiceRequest.Code && toolCode != toolx.GraphAnalyzeConversation.Code && toolCode != toolx.GraphPrepareTicketDraft.Code {
|
||||
@@ -728,6 +790,7 @@ func buildAgentLoopSystemPrompt(agent models.AIAgent, hasKnowledgeBase bool, kno
|
||||
if prompt == "" {
|
||||
prompt = "You are a customer service assistant. Answer accurately, ask for clarification when evidence is insufficient, and do not invent facts."
|
||||
}
|
||||
prompt += "\n\nMaintain conversational continuity. If the immediately preceding assistant message already welcomed the customer and the current customer message is only a greeting, reply briefly without repeating the welcome wording, service capabilities, or service scope."
|
||||
if retrieveErr != nil {
|
||||
prompt += "\n\nKnowledge retrieval is temporarily unavailable for this message. You may answer greetings, acknowledgements, gratitude, farewells, and requests for clarification naturally. For product facts, policies, pricing, functions, procedures, timing, refunds, accounts, permissions, or after-sales questions, do not claim that any detail is verified. Explain that you cannot verify it now, ask one focused question when useful, or offer human handoff."
|
||||
} else if hasKnowledgeBase && strings.TrimSpace(knowledgeContext) == "" {
|
||||
|
||||
@@ -74,6 +74,11 @@ func TestAgentLoopInterruptsBeforeWriteMCPTool(t *testing.T) {
|
||||
if state.Interrupted == nil || !state.Interrupted.Interrupted || !strings.HasPrefix(state.Interrupted.CheckPointID, "tool:9:") {
|
||||
t.Fatalf("missing MCP confirmation checkpoint: %#v", state.Interrupted)
|
||||
}
|
||||
if state.Interrupted.ReplyText != "即将执行“更新客户”,是否确认继续?" ||
|
||||
len(state.Interrupted.Interrupts) != 1 ||
|
||||
state.Interrupted.Interrupts[0].PromptText != state.Interrupted.ReplyText {
|
||||
t.Fatalf("unexpected customer confirmation prompt: %#v", state.Interrupted)
|
||||
}
|
||||
if len(calls) != 1 || calls[0].RiskLevel != "write" || !calls[0].RequireConfirm || calls[0].Status != "interrupted" {
|
||||
t.Fatalf("unexpected MCP safety audit: %#v", calls)
|
||||
}
|
||||
@@ -142,6 +147,46 @@ func TestAgentLoopKnowledgeFallbackCanRequestHandoff(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestAgentLoopConfirmationNormalizesHTMLAndKeepsUnknownPending(t *testing.T) {
|
||||
data := normalizeAgentLoopResumeData(enums.IMMessageTypeHTML, map[string]string{
|
||||
"message": "<p>确认。</p>",
|
||||
})
|
||||
if got := parseAgentLoopConfirmation(firstAgentLoopResumeText(data)); got != agentLoopConfirmationConfirmed {
|
||||
t.Fatalf("expected HTML confirmation, got %v from %#v", got, data)
|
||||
}
|
||||
if got := parseAgentLoopConfirmation("取消!"); got != agentLoopConfirmationCancelled {
|
||||
t.Fatalf("expected cancellation, got %v", got)
|
||||
}
|
||||
if got := parseAgentLoopConfirmation("稍后再说"); got != agentLoopConfirmationUnknown {
|
||||
t.Fatalf("ambiguous input must stay pending, got %v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestConfiguredMCPToolAppliesTrustedSystemPolicy(t *testing.T) {
|
||||
configured, _ := json.Marshal([]request.AIAgentMCPToolRequest{{
|
||||
ToolCode: "system/server_time",
|
||||
ServerCode: "system",
|
||||
ToolName: "server_time",
|
||||
Title: "server_time",
|
||||
RiskLevel: "write",
|
||||
RequireConfirmation: true,
|
||||
}})
|
||||
tool, err := configuredMCPTool(string(configured), "system/server_time")
|
||||
if err != nil {
|
||||
t.Fatalf("resolve configured system tool: %v", err)
|
||||
}
|
||||
if tool.Title != "获取当前时间" || tool.RiskLevel != "read" || tool.RequireConfirmation {
|
||||
t.Fatalf("trusted policy was not applied at runtime: %#v", tool)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAgentLoopPromptAvoidsRepeatingWelcomeMessage(t *testing.T) {
|
||||
prompt := buildAgentLoopSystemPrompt(models.AIAgent{}, false, "", nil)
|
||||
if !strings.Contains(prompt, "without repeating the welcome wording") {
|
||||
t.Fatalf("conversation continuity instruction missing: %q", prompt)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAgentTurnPublishesAllConfiguredCapabilityKinds(t *testing.T) {
|
||||
db, err := gorm.Open(sqlite.Open("file:"+strings.ReplaceAll(t.Name(), "/", "_")+"?mode=memory&cache=shared"), &gorm.Config{})
|
||||
if err != nil {
|
||||
|
||||
@@ -29,6 +29,8 @@ type ResumeInput struct {
|
||||
type InterruptContextSummary struct {
|
||||
Type string `json:"type,omitempty"`
|
||||
ID string `json:"id"`
|
||||
DisplayName string `json:"displayName,omitempty"`
|
||||
PromptText string `json:"promptText,omitempty"`
|
||||
InfoPreview string `json:"infoPreview,omitempty"`
|
||||
}
|
||||
|
||||
|
||||
@@ -56,12 +56,14 @@ func (c *Client) ListTools(ctx context.Context, cfg ServerConfig) ([]ToolInfo, e
|
||||
}
|
||||
ret := make([]ToolInfo, 0, len(result.Tools))
|
||||
for _, tool := range result.Tools {
|
||||
readOnlyHint := tool.Annotations != nil && tool.Annotations.ReadOnlyHint
|
||||
ret = append(ret, ToolInfo{
|
||||
Name: tool.Name,
|
||||
Title: tool.Title,
|
||||
Description: tool.Description,
|
||||
InputSchema: tool.InputSchema,
|
||||
OutputSchema: tool.OutputSchema,
|
||||
ReadOnlyHint: readOnlyHint,
|
||||
})
|
||||
}
|
||||
return ret, nil
|
||||
|
||||
@@ -23,7 +23,11 @@ func (p *systemToolProvider) Register(server *mcp.Server) error {
|
||||
server,
|
||||
&mcp.Tool{
|
||||
Name: "server_time",
|
||||
Title: "获取当前时间",
|
||||
Description: "获取当前服务端时间,可选传入时区。",
|
||||
Annotations: &mcp.ToolAnnotations{
|
||||
ReadOnlyHint: true,
|
||||
},
|
||||
},
|
||||
func(_ context.Context, _ *mcp.CallToolRequest, args serverTimeArgs) (*mcp.CallToolResult, map[string]any, error) {
|
||||
loc := time.Local
|
||||
@@ -46,7 +50,11 @@ func (p *systemToolProvider) Register(server *mcp.Server) error {
|
||||
server,
|
||||
&mcp.Tool{
|
||||
Name: "service_info",
|
||||
Title: "查看服务信息",
|
||||
Description: "查看当前 agent-desk 服务的基础运行信息。",
|
||||
Annotations: &mcp.ToolAnnotations{
|
||||
ReadOnlyHint: true,
|
||||
},
|
||||
},
|
||||
func(ctx context.Context, req *mcp.CallToolRequest, _ struct{}) (*mcp.CallToolResult, map[string]any, error) {
|
||||
cfg := config.Current()
|
||||
|
||||
@@ -28,6 +28,7 @@ type ToolInfo struct {
|
||||
Description string `json:"description"`
|
||||
InputSchema any `json:"inputSchema"`
|
||||
OutputSchema any `json:"outputSchema,omitempty"`
|
||||
ReadOnlyHint bool `json:"readOnlyHint"`
|
||||
}
|
||||
|
||||
type ToolResultContent struct {
|
||||
|
||||
@@ -46,12 +46,24 @@ func resolveInterruptPrompt(summary *applicationruntime.RunResult) string {
|
||||
if summary == nil || len(summary.Interrupts) == 0 {
|
||||
return i18nx.Get("conversation.interrupt.defaultPrompt")
|
||||
}
|
||||
if prompt := extractInterruptMessage(summary.Interrupts[0].InfoPreview); prompt != "" {
|
||||
interrupt := summary.Interrupts[0]
|
||||
if prompt := strings.TrimSpace(interrupt.PromptText); prompt != "" {
|
||||
return prompt
|
||||
}
|
||||
if prompt := strings.TrimSpace(summary.Interrupts[0].InfoPreview); prompt != "" {
|
||||
if prompt := extractInterruptMessage(interrupt.InfoPreview); prompt != "" {
|
||||
return prompt
|
||||
}
|
||||
if prompt := strings.TrimSpace(summary.ReplyText); prompt != "" {
|
||||
return prompt
|
||||
}
|
||||
if interrupt.Type != "tool_confirmation" {
|
||||
if prompt := strings.TrimSpace(interrupt.InfoPreview); prompt != "" {
|
||||
return prompt
|
||||
}
|
||||
}
|
||||
if displayName := strings.TrimSpace(interrupt.DisplayName); displayName != "" {
|
||||
return "即将执行“" + displayName + "”,是否确认继续?"
|
||||
}
|
||||
return i18nx.Get("conversation.interrupt.defaultPrompt")
|
||||
}
|
||||
|
||||
|
||||
@@ -103,6 +103,22 @@ func TestResolveInterruptPrompt(t *testing.T) {
|
||||
if got := resolveInterruptPrompt(summary); got != "直接补充手机号" {
|
||||
t.Fatalf("unexpected raw interrupt prompt: %q", got)
|
||||
}
|
||||
|
||||
summary.ReplyText = ""
|
||||
summary.Interrupts[0] = applicationruntime.InterruptContextSummary{
|
||||
ID: "system/server_time",
|
||||
Type: "tool_confirmation",
|
||||
DisplayName: "获取当前时间",
|
||||
InfoPreview: "system/server_time",
|
||||
}
|
||||
if got := resolveInterruptPrompt(summary); got != "即将执行“获取当前时间”,是否确认继续?" {
|
||||
t.Fatalf("tool code leaked into customer prompt: %q", got)
|
||||
}
|
||||
|
||||
summary.Interrupts[0].PromptText = "请确认是否更新客户资料。"
|
||||
if got := resolveInterruptPrompt(summary); got != "请确认是否更新客户资料。" {
|
||||
t.Fatalf("explicit customer prompt was not preferred: %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func newConversationFixture() models.Conversation {
|
||||
|
||||
@@ -36,15 +36,18 @@ func MCPAnyCatalog(ctx *gin.Context) {
|
||||
ret := make([]response.MCPToolCatalogResponse, 0, len(items))
|
||||
for _, item := range items {
|
||||
ret = append(ret, response.MCPToolCatalogResponse{
|
||||
ToolCode: item.ToolCode,
|
||||
ServerCode: item.ServerCode,
|
||||
ToolName: item.ToolName,
|
||||
SourceType: item.SourceType,
|
||||
AutoInjected: item.AutoInjected,
|
||||
Title: item.Title,
|
||||
Description: item.Description,
|
||||
InputSchema: item.InputSchema,
|
||||
OutputSchema: item.OutputSchema,
|
||||
ToolCode: item.ToolCode,
|
||||
ServerCode: item.ServerCode,
|
||||
ToolName: item.ToolName,
|
||||
SourceType: item.SourceType,
|
||||
AutoInjected: item.AutoInjected,
|
||||
Title: item.Title,
|
||||
Description: item.Description,
|
||||
InputSchema: item.InputSchema,
|
||||
OutputSchema: item.OutputSchema,
|
||||
RiskLevel: item.RiskLevel,
|
||||
RequireConfirmation: item.RequireConfirmation,
|
||||
RiskEditable: item.RiskEditable,
|
||||
})
|
||||
}
|
||||
httpx.WriteJSON(ctx, ret)
|
||||
|
||||
@@ -52,6 +52,7 @@ type MCPToolInfoResponse struct {
|
||||
Description string `json:"description"`
|
||||
InputSchema any `json:"inputSchema"`
|
||||
OutputSchema any `json:"outputSchema,omitempty"`
|
||||
ReadOnlyHint bool `json:"readOnlyHint"`
|
||||
}
|
||||
|
||||
func BuildMCPToolInfoResponses(items []mcps.ToolInfo) []MCPToolInfoResponse {
|
||||
@@ -63,21 +64,25 @@ func BuildMCPToolInfoResponses(items []mcps.ToolInfo) []MCPToolInfoResponse {
|
||||
Description: item.Description,
|
||||
InputSchema: item.InputSchema,
|
||||
OutputSchema: item.OutputSchema,
|
||||
ReadOnlyHint: item.ReadOnlyHint,
|
||||
})
|
||||
}
|
||||
return ret
|
||||
}
|
||||
|
||||
type MCPToolCatalogResponse struct {
|
||||
ToolCode string `json:"toolCode"`
|
||||
ServerCode string `json:"serverCode"`
|
||||
ToolName string `json:"toolName"`
|
||||
SourceType enums.ToolSourceType `json:"sourceType"`
|
||||
AutoInjected bool `json:"autoInjected"`
|
||||
Title string `json:"title"`
|
||||
Description string `json:"description"`
|
||||
InputSchema any `json:"inputSchema"`
|
||||
OutputSchema any `json:"outputSchema,omitempty"`
|
||||
ToolCode string `json:"toolCode"`
|
||||
ServerCode string `json:"serverCode"`
|
||||
ToolName string `json:"toolName"`
|
||||
SourceType enums.ToolSourceType `json:"sourceType"`
|
||||
AutoInjected bool `json:"autoInjected"`
|
||||
Title string `json:"title"`
|
||||
Description string `json:"description"`
|
||||
InputSchema any `json:"inputSchema"`
|
||||
OutputSchema any `json:"outputSchema,omitempty"`
|
||||
RiskLevel string `json:"riskLevel"`
|
||||
RequireConfirmation bool `json:"requireConfirmation"`
|
||||
RiskEditable bool `json:"riskEditable"`
|
||||
}
|
||||
|
||||
type MCPToolResultContentResponse struct {
|
||||
|
||||
@@ -0,0 +1,51 @@
|
||||
package toolx
|
||||
|
||||
import (
|
||||
"strings"
|
||||
|
||||
"agent-desk/internal/pkg/dto/request"
|
||||
)
|
||||
|
||||
const (
|
||||
MCPRiskLevelRead = "read"
|
||||
MCPRiskLevelWrite = "write"
|
||||
)
|
||||
|
||||
type TrustedMCPToolPolicy struct {
|
||||
ToolCode string
|
||||
Title string
|
||||
RiskLevel string
|
||||
RequireConfirmation bool
|
||||
}
|
||||
|
||||
var trustedMCPToolPolicies = map[string]TrustedMCPToolPolicy{
|
||||
"system/server_time": {
|
||||
ToolCode: "system/server_time",
|
||||
Title: "获取当前时间",
|
||||
RiskLevel: MCPRiskLevelRead,
|
||||
RequireConfirmation: false,
|
||||
},
|
||||
"system/service_info": {
|
||||
ToolCode: "system/service_info",
|
||||
Title: "查看服务信息",
|
||||
RiskLevel: MCPRiskLevelRead,
|
||||
RequireConfirmation: false,
|
||||
},
|
||||
}
|
||||
|
||||
func GetTrustedMCPToolPolicy(toolCode string) (TrustedMCPToolPolicy, bool) {
|
||||
policy, ok := trustedMCPToolPolicies[NormalizeToolCodeAlias(strings.TrimSpace(toolCode))]
|
||||
return policy, ok
|
||||
}
|
||||
|
||||
func ApplyTrustedMCPToolPolicy(item request.AIAgentMCPToolRequest) request.AIAgentMCPToolRequest {
|
||||
policy, ok := GetTrustedMCPToolPolicy(item.ToolCode)
|
||||
if !ok {
|
||||
return item
|
||||
}
|
||||
item.ToolCode = policy.ToolCode
|
||||
item.Title = policy.Title
|
||||
item.RiskLevel = policy.RiskLevel
|
||||
item.RequireConfirmation = policy.RequireConfirmation
|
||||
return item
|
||||
}
|
||||
@@ -0,0 +1,39 @@
|
||||
package services
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"agent-desk/internal/pkg/dto/request"
|
||||
)
|
||||
|
||||
func TestValidateMCPToolRiskPolicyRejectsTrustedToolOverride(t *testing.T) {
|
||||
_, err := validateMCPToolRiskPolicy(request.AIAgentMCPToolRequest{
|
||||
ToolCode: "system/server_time",
|
||||
RiskLevel: "write",
|
||||
RequireConfirmation: true,
|
||||
})
|
||||
if err == nil {
|
||||
t.Fatal("expected trusted system tool policy override to be rejected")
|
||||
}
|
||||
|
||||
item, err := validateMCPToolRiskPolicy(request.AIAgentMCPToolRequest{
|
||||
ToolCode: "system/server_time",
|
||||
RiskLevel: "read",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("validate trusted system tool policy: %v", err)
|
||||
}
|
||||
if item.Title != "获取当前时间" || item.RiskLevel != "read" || item.RequireConfirmation {
|
||||
t.Fatalf("unexpected normalized trusted policy: %#v", item)
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateMCPToolRiskPolicyRequiresWriteConfirmation(t *testing.T) {
|
||||
_, err := validateMCPToolRiskPolicy(request.AIAgentMCPToolRequest{
|
||||
ToolCode: "crm/update_customer",
|
||||
RiskLevel: "write",
|
||||
})
|
||||
if err == nil {
|
||||
t.Fatal("expected write tool without confirmation to be rejected")
|
||||
}
|
||||
}
|
||||
@@ -216,11 +216,8 @@ func (s *aIAgentService) validatePublishableAgent(db *gorm.DB, agent *models.AIA
|
||||
if err != nil || definition.InputSchema == nil {
|
||||
return errorsx.InvalidParam("ai agent MCP tool definition is unavailable")
|
||||
}
|
||||
if item.RiskLevel != aitooling.RiskLevelRead && item.RiskLevel != aitooling.RiskLevelWrite {
|
||||
return errorsx.InvalidParam("ai agent MCP tool risk level is invalid")
|
||||
}
|
||||
if item.RiskLevel == aitooling.RiskLevelWrite && !item.RequireConfirmation {
|
||||
return errorsx.InvalidParam("write MCP tools must require confirmation")
|
||||
if _, err := validateMCPToolRiskPolicy(item); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
for _, binding := range s.ListEnabledWorkflowBindings(db, agent.ID) {
|
||||
@@ -522,12 +519,10 @@ func (s *aIAgentService) normalizeMCPTools(input []request.AIAgentMCPToolRequest
|
||||
return nil, err
|
||||
}
|
||||
normalized.RiskLevel = strings.ToLower(strings.TrimSpace(item.RiskLevel))
|
||||
if normalized.RiskLevel != aitooling.RiskLevelRead && normalized.RiskLevel != aitooling.RiskLevelWrite {
|
||||
return nil, errorsx.InvalidParam("MCP tool risk level must be read or write")
|
||||
}
|
||||
normalized.RequireConfirmation = item.RequireConfirmation
|
||||
if normalized.RiskLevel == aitooling.RiskLevelWrite && !normalized.RequireConfirmation {
|
||||
return nil, errorsx.InvalidParam("write MCP tools must require confirmation")
|
||||
normalized, err = validateMCPToolRiskPolicy(normalized)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
key := strings.TrimSpace(normalized.ToolCode)
|
||||
if _, exists := seen[key]; exists {
|
||||
@@ -539,6 +534,22 @@ func (s *aIAgentService) normalizeMCPTools(input []request.AIAgentMCPToolRequest
|
||||
return ret, nil
|
||||
}
|
||||
|
||||
func validateMCPToolRiskPolicy(item request.AIAgentMCPToolRequest) (request.AIAgentMCPToolRequest, error) {
|
||||
if policy, ok := toolx.GetTrustedMCPToolPolicy(item.ToolCode); ok {
|
||||
if item.RiskLevel != policy.RiskLevel || item.RequireConfirmation != policy.RequireConfirmation {
|
||||
return request.AIAgentMCPToolRequest{}, errorsx.InvalidParam("system MCP tool risk policy cannot be changed")
|
||||
}
|
||||
return toolx.ApplyTrustedMCPToolPolicy(item), nil
|
||||
}
|
||||
if item.RiskLevel != aitooling.RiskLevelRead && item.RiskLevel != aitooling.RiskLevelWrite {
|
||||
return request.AIAgentMCPToolRequest{}, errorsx.InvalidParam("MCP tool risk level must be read or write")
|
||||
}
|
||||
if item.RiskLevel == aitooling.RiskLevelWrite && !item.RequireConfirmation {
|
||||
return request.AIAgentMCPToolRequest{}, errorsx.InvalidParam("write MCP tools must require confirmation")
|
||||
}
|
||||
return item, nil
|
||||
}
|
||||
|
||||
func (s *aIAgentService) UpdateSort(ids []int64) error {
|
||||
return sqls.WithTransaction(func(ctx *sqls.TxContext) error {
|
||||
for i, id := range ids {
|
||||
|
||||
@@ -22,15 +22,18 @@ func newToolCatalogService() *toolCatalogService {
|
||||
type toolCatalogService struct{}
|
||||
|
||||
type MCPToolCatalogItem struct {
|
||||
ToolCode string
|
||||
ServerCode string
|
||||
ToolName string
|
||||
SourceType enums.ToolSourceType
|
||||
AutoInjected bool
|
||||
Title string
|
||||
Description string
|
||||
InputSchema any
|
||||
OutputSchema any
|
||||
ToolCode string
|
||||
ServerCode string
|
||||
ToolName string
|
||||
SourceType enums.ToolSourceType
|
||||
AutoInjected bool
|
||||
Title string
|
||||
Description string
|
||||
InputSchema any
|
||||
OutputSchema any
|
||||
RiskLevel string
|
||||
RequireConfirmation bool
|
||||
RiskEditable bool
|
||||
}
|
||||
|
||||
func (s *toolCatalogService) ListMCPTools(ctx context.Context) ([]MCPToolCatalogItem, error) {
|
||||
@@ -71,16 +74,37 @@ func (s *toolCatalogService) ListMCPToolsWithLocale(ctx context.Context, locale
|
||||
return nil, err
|
||||
}
|
||||
for _, item := range tools {
|
||||
toolCode := toolx.BuildMCPToolCode(serverCode, item.Name)
|
||||
title := strings.TrimSpace(item.Title)
|
||||
riskLevel := toolx.MCPRiskLevelWrite
|
||||
requireConfirmation := true
|
||||
riskEditable := true
|
||||
if item.ReadOnlyHint {
|
||||
riskLevel = toolx.MCPRiskLevelRead
|
||||
requireConfirmation = false
|
||||
}
|
||||
if policy, ok := toolx.GetTrustedMCPToolPolicy(toolCode); ok {
|
||||
title = policy.Title
|
||||
riskLevel = policy.RiskLevel
|
||||
requireConfirmation = policy.RequireConfirmation
|
||||
riskEditable = false
|
||||
}
|
||||
if title == "" {
|
||||
title = strings.TrimSpace(item.Name)
|
||||
}
|
||||
ret = append(ret, MCPToolCatalogItem{
|
||||
ToolCode: toolx.BuildMCPToolCode(serverCode, item.Name),
|
||||
ServerCode: serverCode,
|
||||
ToolName: strings.TrimSpace(item.Name),
|
||||
SourceType: enums.ToolSourceTypeMCP,
|
||||
AutoInjected: false,
|
||||
Title: strings.TrimSpace(item.Title),
|
||||
Description: strings.TrimSpace(item.Description),
|
||||
InputSchema: item.InputSchema,
|
||||
OutputSchema: item.OutputSchema,
|
||||
ToolCode: toolCode,
|
||||
ServerCode: serverCode,
|
||||
ToolName: strings.TrimSpace(item.Name),
|
||||
SourceType: enums.ToolSourceTypeMCP,
|
||||
AutoInjected: false,
|
||||
Title: title,
|
||||
Description: strings.TrimSpace(item.Description),
|
||||
InputSchema: item.InputSchema,
|
||||
OutputSchema: item.OutputSchema,
|
||||
RiskLevel: riskLevel,
|
||||
RequireConfirmation: requireConfirmation,
|
||||
RiskEditable: riskEditable,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user