refactor: enhance MCP tool management with risk policies and confirmation handling

This commit is contained in:
mlogclub
2026-07-28 11:33:15 +08:00
parent ed109047d3
commit 07370fe9a2
17 changed files with 448 additions and 117 deletions
@@ -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 {
+2
View File
@@ -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"`
}