refactor: resolve approval tools from SDK call ids

- correlate permission requests to tool names via ToolCallId when the
  workflow request stream exposes matching CallId values
- keep the url-to-web_fetch alias only as a narrow fallback when a
  matching tool call cannot be recovered
- cover generic permission-category lookup for url, shell, and read in
  sidecar tests

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
This commit is contained in:
David Kaya
2026-03-24 21:03:16 +01:00
co-authored by Copilot
parent 7622d047b1
commit 44d9c43b09
2 changed files with 129 additions and 20 deletions
@@ -43,6 +43,7 @@ public sealed class CopilotWorkflowRunner : ITurnWorkflowRunner
throw new InvalidOperationException(validationError.Message); throw new InvalidOperationException(validationError.Message);
} }
ConcurrentDictionary<string, string> toolNamesByCallId = new(StringComparer.Ordinal);
await using AgentBundle bundle = await AgentBundle.CreateAsync( await using AgentBundle bundle = await AgentBundle.CreateAsync(
command, command,
(agent, request, invocation) => RequestApprovalAsync( (agent, request, invocation) => RequestApprovalAsync(
@@ -50,6 +51,7 @@ public sealed class CopilotWorkflowRunner : ITurnWorkflowRunner
agent, agent,
request, request,
invocation, invocation,
toolNamesByCallId,
onApproval, onApproval,
cancellationToken), cancellationToken),
cancellationToken); cancellationToken);
@@ -85,7 +87,8 @@ public sealed class CopilotWorkflowRunner : ITurnWorkflowRunner
AgentActivityEventDto? activity = TryCreateActivityFromRequest( AgentActivityEventDto? activity = TryCreateActivityFromRequest(
command, command,
requestInfo, requestInfo,
activeAgent); activeAgent,
toolNamesByCallId);
if (activity is not null) if (activity is not null)
{ {
@@ -239,10 +242,11 @@ public sealed class CopilotWorkflowRunner : ITurnWorkflowRunner
PatternAgentDefinitionDto agent, PatternAgentDefinitionDto agent,
PermissionRequest request, PermissionRequest request,
PermissionInvocation invocation, PermissionInvocation invocation,
IReadOnlyDictionary<string, string> toolNamesByCallId,
Func<ApprovalRequestedEventDto, Task> onApproval, Func<ApprovalRequestedEventDto, Task> onApproval,
CancellationToken cancellationToken) CancellationToken cancellationToken)
{ {
TryGetApprovalToolName(request, out string? toolName); TryGetApprovalToolName(request, toolNamesByCallId, out string? toolName);
if (!RequiresToolCallApproval(command.Pattern.ApprovalPolicy, agent.Id, toolName)) if (!RequiresToolCallApproval(command.Pattern.ApprovalPolicy, agent.Id, toolName))
{ {
@@ -294,7 +298,8 @@ public sealed class CopilotWorkflowRunner : ITurnWorkflowRunner
private static AgentActivityEventDto? TryCreateActivityFromRequest( private static AgentActivityEventDto? TryCreateActivityFromRequest(
RunTurnCommandDto command, RunTurnCommandDto command,
RequestInfoEvent requestInfo, RequestInfoEvent requestInfo,
AgentIdentity? activeAgent) AgentIdentity? activeAgent,
ConcurrentDictionary<string, string> toolNamesByCallId)
{ {
if (TryGetHandoffTarget(command.Pattern, requestInfo, out AgentIdentity handoffAgent)) if (TryGetHandoffTarget(command.Pattern, requestInfo, out AgentIdentity handoffAgent))
{ {
@@ -305,11 +310,17 @@ public sealed class CopilotWorkflowRunner : ITurnWorkflowRunner
sourceAgent: activeAgent); sourceAgent: activeAgent);
} }
if (!activeAgent.HasValue || !TryGetToolName(requestInfo, out string toolName)) if (!activeAgent.HasValue
|| !TryGetToolRequestInfo(requestInfo, out string toolName, out string? toolCallId))
{ {
return null; return null;
} }
if (!string.IsNullOrWhiteSpace(toolCallId))
{
toolNamesByCallId[toolCallId] = toolName;
}
return CreateActivityEvent( return CreateActivityEvent(
command, command,
activityType: "tool-calling", activityType: "tool-calling",
@@ -353,11 +364,15 @@ public sealed class CopilotWorkflowRunner : ITurnWorkflowRunner
return !string.IsNullOrWhiteSpace(agent.AgentName); return !string.IsNullOrWhiteSpace(agent.AgentName);
} }
private static bool TryGetToolName(RequestInfoEvent requestInfo, out string toolName) private static bool TryGetToolRequestInfo(
RequestInfoEvent requestInfo,
out string toolName,
out string? toolCallId)
{ {
if (TryReadPortableValue(requestInfo.Request.Data, FunctionCallContentType, out object? functionCall)) if (TryReadPortableValue(requestInfo.Request.Data, FunctionCallContentType, out object? functionCall))
{ {
toolName = GetStringProperty(functionCall, "Name") ?? "function"; toolName = GetStringProperty(functionCall, "Name") ?? "function";
toolCallId = NormalizeOptionalString(GetStringProperty(functionCall, "CallId"));
return true; return true;
} }
@@ -366,22 +381,26 @@ public sealed class CopilotWorkflowRunner : ITurnWorkflowRunner
toolName = GetStringProperty(mcpToolCall, "ToolName") toolName = GetStringProperty(mcpToolCall, "ToolName")
?? GetStringProperty(mcpToolCall, "ServerName") ?? GetStringProperty(mcpToolCall, "ServerName")
?? string.Empty; ?? string.Empty;
toolCallId = NormalizeOptionalString(GetStringProperty(mcpToolCall, "CallId"));
return !string.IsNullOrWhiteSpace(toolName); return !string.IsNullOrWhiteSpace(toolName);
} }
if (TryReadPortableValue(requestInfo.Request.Data, CodeInterpreterToolCallContentType, out _)) if (TryReadPortableValue(requestInfo.Request.Data, CodeInterpreterToolCallContentType, out object? codeInterpreterToolCall))
{ {
toolName = "code interpreter"; toolName = "code interpreter";
toolCallId = NormalizeOptionalString(GetStringProperty(codeInterpreterToolCall, "CallId"));
return true; return true;
} }
if (TryReadPortableValue(requestInfo.Request.Data, ImageGenerationToolCallContentType, out _)) if (TryReadPortableValue(requestInfo.Request.Data, ImageGenerationToolCallContentType, out _))
{ {
toolName = "image generation"; toolName = "image generation";
toolCallId = null;
return true; return true;
} }
toolName = string.Empty; toolName = string.Empty;
toolCallId = null;
return false; return false;
} }
@@ -489,13 +508,36 @@ public sealed class CopilotWorkflowRunner : ITurnWorkflowRunner
string.Equals(candidate, toolName, StringComparison.OrdinalIgnoreCase)); string.Equals(candidate, toolName, StringComparison.OrdinalIgnoreCase));
} }
internal static bool TryGetApprovalToolName(PermissionRequest request, out string? toolName) internal static bool TryGetApprovalToolName(
PermissionRequest request,
IReadOnlyDictionary<string, string>? toolNamesByCallId,
out string? toolName)
{ {
toolName = request switch toolName = request switch
{ {
PermissionRequestMcp mcp when !string.IsNullOrWhiteSpace(mcp.ToolName) => mcp.ToolName.Trim(), PermissionRequestMcp mcp when !string.IsNullOrWhiteSpace(mcp.ToolName) => mcp.ToolName.Trim(),
PermissionRequestCustomTool customTool when !string.IsNullOrWhiteSpace(customTool.ToolName) => customTool.ToolName.Trim(), PermissionRequestCustomTool customTool when !string.IsNullOrWhiteSpace(customTool.ToolName) => customTool.ToolName.Trim(),
PermissionRequestHook hook when !string.IsNullOrWhiteSpace(hook.ToolName) => hook.ToolName.Trim(), PermissionRequestHook hook when !string.IsNullOrWhiteSpace(hook.ToolName) => hook.ToolName.Trim(),
_ => null,
};
if (!string.IsNullOrWhiteSpace(toolName))
{
return true;
}
string? toolCallId = NormalizeOptionalString(GetStringProperty(request, "ToolCallId"));
if (toolCallId is not null
&& toolNamesByCallId is not null
&& toolNamesByCallId.TryGetValue(toolCallId, out string? resolvedToolName)
&& !string.IsNullOrWhiteSpace(resolvedToolName))
{
toolName = resolvedToolName.Trim();
return true;
}
toolName = request switch
{
PermissionRequestUrl => "web_fetch", PermissionRequestUrl => "web_fetch",
_ => null, _ => null,
}; };
@@ -503,6 +545,9 @@ public sealed class CopilotWorkflowRunner : ITurnWorkflowRunner
return !string.IsNullOrWhiteSpace(toolName); return !string.IsNullOrWhiteSpace(toolName);
} }
internal static bool TryGetApprovalToolName(PermissionRequest request, out string? toolName)
=> TryGetApprovalToolName(request, toolNamesByCallId: null, out toolName);
private static string CreateApprovalRequestId() private static string CreateApprovalRequestId()
{ {
return $"approval-{Guid.NewGuid():N}"; return $"approval-{Guid.NewGuid():N}";
@@ -530,6 +575,11 @@ public sealed class CopilotWorkflowRunner : ITurnWorkflowRunner
return instance?.GetType().GetProperty(propertyName)?.GetValue(instance) as string; return instance?.GetType().GetProperty(propertyName)?.GetValue(instance) as string;
} }
private static string? NormalizeOptionalString(string? value)
{
return string.IsNullOrWhiteSpace(value) ? null : value.Trim();
}
private static StreamingSegment GetOrCreateSegment(List<StreamingSegment> segments, string messageId, string authorName) private static StreamingSegment GetOrCreateSegment(List<StreamingSegment> segments, string messageId, string authorName)
{ {
StreamingSegment? existing = segments.LastOrDefault(segment => segment.MessageId == messageId); StreamingSegment? existing = segments.LastOrDefault(segment => segment.MessageId == messageId);
@@ -258,7 +258,7 @@ public sealed class CopilotWorkflowRunnerTests
} }
[Fact] [Fact]
public void TryGetApprovalToolName_ReadsMcpCustomHookAndUrlRequests() public void TryGetApprovalToolName_ReadsMcpCustomAndHookRequests()
{ {
Assert.True( Assert.True(
CopilotWorkflowRunner.TryGetApprovalToolName( CopilotWorkflowRunner.TryGetApprovalToolName(
@@ -296,18 +296,6 @@ public sealed class CopilotWorkflowRunnerTests
out string? hookToolName)); out string? hookToolName));
Assert.Equal("web_fetch", hookToolName); Assert.Equal("web_fetch", hookToolName);
Assert.True(
CopilotWorkflowRunner.TryGetApprovalToolName(
new PermissionRequestUrl
{
Kind = "url",
ToolCallId = "tool-call-1",
Intention = "Fetch the requested page",
Url = "https://example.com/docs",
},
out string? urlToolName));
Assert.Equal("web_fetch", urlToolName);
Assert.False( Assert.False(
CopilotWorkflowRunner.TryGetApprovalToolName( CopilotWorkflowRunner.TryGetApprovalToolName(
new PermissionRequestShell new PermissionRequestShell
@@ -325,6 +313,77 @@ public sealed class CopilotWorkflowRunnerTests
Assert.Null(shellToolName); Assert.Null(shellToolName);
} }
[Fact]
public void TryGetApprovalToolName_UsesToolCallLookupForPermissionCategoriesWithoutDirectToolNames()
{
Dictionary<string, string> toolNamesByCallId = new(StringComparer.Ordinal)
{
["tool-call-url"] = "web_fetch",
["tool-call-shell"] = "shell",
["tool-call-read"] = "view",
};
Assert.True(
CopilotWorkflowRunner.TryGetApprovalToolName(
new PermissionRequestUrl
{
Kind = "url",
ToolCallId = "tool-call-url",
Intention = "Fetch the requested page",
Url = "https://example.com/docs",
},
toolNamesByCallId,
out string? urlToolName));
Assert.Equal("web_fetch", urlToolName);
Assert.True(
CopilotWorkflowRunner.TryGetApprovalToolName(
new PermissionRequestShell
{
Kind = "shell",
ToolCallId = "tool-call-shell",
FullCommandText = "curl https://example.com/docs",
Intention = "Fetch documentation with curl",
Commands = [],
PossiblePaths = [],
PossibleUrls = [],
HasWriteFileRedirection = false,
CanOfferSessionApproval = false,
},
toolNamesByCallId,
out string? shellToolName));
Assert.Equal("shell", shellToolName);
Assert.True(
CopilotWorkflowRunner.TryGetApprovalToolName(
new PermissionRequestRead
{
Kind = "read",
ToolCallId = "tool-call-read",
Intention = "Inspect a file",
Path = "README.md",
},
toolNamesByCallId,
out string? readToolName));
Assert.Equal("view", readToolName);
}
[Fact]
public void TryGetApprovalToolName_FallsBackToWebFetchForUncorrelatedUrlRequests()
{
Assert.True(
CopilotWorkflowRunner.TryGetApprovalToolName(
new PermissionRequestUrl
{
Kind = "url",
ToolCallId = "tool-call-1",
Intention = "Fetch the requested page",
Url = "https://example.com/docs",
},
out string? urlToolName));
Assert.Equal("web_fetch", urlToolName);
}
[Fact] [Fact]
public void BuildPermissionApprovalEvent_IncludesToolContextWhenKnown() public void BuildPermissionApprovalEvent_IncludesToolContextWhenKnown()
{ {