From aeb2e124c7730dd635d28cea63794fc9f63b0f5a Mon Sep 17 00:00:00 2001 From: Jacob Alber Date: Fri, 8 May 2026 14:33:50 -0400 Subject: [PATCH] test: Add unit tests for Handoff FunctionCall/Result matching fix --- .../StreamingToolCallResultPairMatcher.cs | 11 +- .../HandoffOrchestrationTests.cs | 251 ++++++++++++++++++ 2 files changed, 260 insertions(+), 2 deletions(-) diff --git a/dotnet/src/Microsoft.Agents.AI.Workflows/Specialized/Magentic/StreamingToolCallResultPairMatcher.cs b/dotnet/src/Microsoft.Agents.AI.Workflows/Specialized/Magentic/StreamingToolCallResultPairMatcher.cs index bb72b8390a..80fa89b0a2 100644 --- a/dotnet/src/Microsoft.Agents.AI.Workflows/Specialized/Magentic/StreamingToolCallResultPairMatcher.cs +++ b/dotnet/src/Microsoft.Agents.AI.Workflows/Specialized/Magentic/StreamingToolCallResultPairMatcher.cs @@ -3,13 +3,14 @@ using System; using System.Collections.Generic; using System.Diagnostics.CodeAnalysis; +using System.Linq; using Microsoft.Extensions.AI; namespace Microsoft.Agents.AI.Workflows.Specialized.Magentic; internal sealed class StreamingToolCallResultPairMatcher { - private enum CallType + internal enum CallType { Function, McpServerTool @@ -17,7 +18,7 @@ internal sealed class StreamingToolCallResultPairMatcher private record CallSummaryKey(CallType Type, string CallId); - private struct ToolCallSummary(CallType callType, string callId, string name) + internal struct ToolCallSummary(CallType callType, string callId, string name) { public CallType CallType => callType; @@ -28,6 +29,12 @@ internal sealed class StreamingToolCallResultPairMatcher private readonly Dictionary _callSummaries = new(); + public bool HasUnmatchedCalls => this._callSummaries.Count > 0; + + public IEnumerable UnmatchedCalls => this.HasUnmatchedCalls + ? this._callSummaries.Values.ToList() + : []; + private void Collect(CallType callType, string callId, string name, string callContentTypeName, string resultContentTypeName) { CallSummaryKey key = new(callType, callId); diff --git a/dotnet/tests/Microsoft.Agents.AI.Workflows.UnitTests/HandoffOrchestrationTests.cs b/dotnet/tests/Microsoft.Agents.AI.Workflows.UnitTests/HandoffOrchestrationTests.cs index 89cb52a89a..73c34cae53 100644 --- a/dotnet/tests/Microsoft.Agents.AI.Workflows.UnitTests/HandoffOrchestrationTests.cs +++ b/dotnet/tests/Microsoft.Agents.AI.Workflows.UnitTests/HandoffOrchestrationTests.cs @@ -12,6 +12,7 @@ using System.Threading.Tasks; using FluentAssertions; using Microsoft.Agents.AI.Workflows.InProc; using Microsoft.Agents.AI.Workflows.Specialized; +using Microsoft.Agents.AI.Workflows.Specialized.Magentic; using Microsoft.Extensions.AI; using Microsoft.Extensions.Logging; @@ -859,6 +860,256 @@ public class HandoffOrchestrationTests Assert.Equal(1, specialistCallCount); // specialist NOT called } + private static MockChatClient CreateFunctionCallResultValidatingClient(Func, ChatOptions?, ChatResponse> innerResponseFactory) + { + return new MockChatClient(InvokeResponseFactory); + + ChatResponse InvokeResponseFactory(IEnumerable chatMessages, ChatOptions? options) + { + // We do not need to keep the callResolver around because ChatClientAgent owns making sure that the function call is properly + // resent to the underlying agent. + StreamingToolCallResultPairMatcher callResolver = new(); + List incomingMessages = chatMessages.ToList(); + foreach (ChatMessage message in incomingMessages) + { + foreach (AIContent content in message.Contents) + { + switch (content) + { + case FunctionCallContent functionCallContent: + { + callResolver.CollectFunctionCall(functionCallContent); + break; + } + + case FunctionResultContent functionResultContent: + { + if (!callResolver.TryResolveFunctionCall(functionResultContent, out _)) + { + throw new InvalidOperationException($"Received unexpected function result: {functionResultContent.CallId}"); + } + break; + } + + case McpServerToolCallContent mcpServerToolCallContent: + { + callResolver.CollectMcpServerToolCall(mcpServerToolCallContent); + break; + } + + case McpServerToolResultContent mcpServerToolResultContent: + { + if (!callResolver.TryResolveMcpServerToolCall(mcpServerToolResultContent, out _)) + { + throw new InvalidOperationException($"Received unexpected tool result: {mcpServerToolResultContent.CallId}"); + } + break; + } + } + } + } + + // If there are still unmatched calls, we have an error + callResolver.UnmatchedCalls.Should().BeEmpty(); + + // Now we can invoke the inner response factory to generate the response + ChatResponse response = innerResponseFactory(incomingMessages, options); + + foreach (ChatMessage message in response.Messages) + { + foreach (AIContent content in message.Contents) + { + switch (content) + { + case FunctionCallContent functionCallContent: + callResolver.CollectFunctionCall(functionCallContent); + break; + case McpServerToolCallContent mcpServerToolCallContent: + callResolver.CollectMcpServerToolCall(mcpServerToolCallContent); + break; + case FunctionResultContent functionResultContent: + { + if (!callResolver.TryResolveFunctionCall(functionResultContent, out string? name)) + { + throw new InvalidOperationException($"Produced unexpected function result: {functionResultContent.CallId}"); + } + break; + } + case McpServerToolResultContent mcpServerToolResultContent: + { + if (!callResolver.TryResolveMcpServerToolCall(mcpServerToolResultContent, out string? name)) + { + throw new InvalidOperationException($"Produced unexpected tool result: {mcpServerToolResultContent.CallId}"); + } + break; + } + } + } + } + + return response; + } + } + + [Fact] + public async Task Handoffs_ReentrantHandoff_FunctionResultSentToAgentOnSubsequentInvocationAsync() + { + // Regression test: When an agent requests a handoff, the synthesized FunctionResult for the handoff + // must be sent back to the agent on subsequent invocations. If this doesn't happen, the agent's + // conversation state will be broken because the LLM will receive a FunctionCall without a + // corresponding FunctionResult. + + List>? specialistInvocations = []; + int coordinatorCallCount = 0; + int specialistCallCount = 0; + + var coordinator = new ChatClientAgent(CreateFunctionCallResultValidatingClient((messages, options) => + { + coordinatorCallCount++; + // Always hand off to specialist + string? transferFuncName = options?.Tools?.FirstOrDefault(t => t.Name.StartsWith("handoff_to_", StringComparison.Ordinal))?.Name; + Assert.NotNull(transferFuncName); + return new(new ChatMessage(ChatRole.Assistant, [new FunctionCallContent($"coordinator_handoff_call_{coordinatorCallCount}", transferFuncName)])); + }), name: "coordinator"); + + var specialist = new ChatClientAgent(CreateFunctionCallResultValidatingClient((messages, options) => + { + specialistCallCount++; + specialistInvocations.Add(messages.ToList()); + + if (specialistCallCount == 1) + { + // First call: hand back to coordinator + string? transferFuncName = options?.Tools?.FirstOrDefault(t => t.Name.StartsWith("handoff_to_", StringComparison.Ordinal))?.Name; + Assert.NotNull(transferFuncName); + return new(new ChatMessage(ChatRole.Assistant, [new FunctionCallContent("specialist_handoff_call", transferFuncName)])); + } + + // Subsequent calls: respond normally + return new(new ChatMessage(ChatRole.Assistant, "specialist final response")); + }), name: "specialist", description: "The specialist agent"); + + var workflow = AgentWorkflowBuilder.CreateHandoffBuilderWith(coordinator) + .WithHandoff(coordinator, specialist) + .WithHandoff(specialist, coordinator) + .EnableReturnToPrevious() + .Build(); + + CheckpointManager checkpointManager = CheckpointManager.CreateInMemory(); + const ExecutionEnvironment Environment = ExecutionEnvironment.InProcess_Lockstep; + + // Turn 1: coordinator -> specialist -> coordinator -> specialist (specialist responds on 2nd call) + // Flow: coordinator(1) hands off -> specialist(1) hands off -> coordinator(2) hands off -> specialist(2) responds + WorkflowRunResult result = await RunWorkflowCheckpointedAsync(workflow, [new ChatMessage(ChatRole.User, "start")], Environment, checkpointManager); + Assert.Equal(2, coordinatorCallCount); // initial + receiving handback from specialist + Assert.Equal(2, specialistCallCount); // specialist invoked twice (once handed off, once responded) + + Assert.Equal(2, specialistInvocations.Count); + } + + [Fact] + public async Task Handoffs_MultiTurnWithHandoffAndReturn_AllFunctionCallsHaveMatchingResultsAsync() + { + // This test verifies that across multiple turns with handoffs going back and forth, + // the FunctionCall/FunctionResult pairing rule is always maintained for any agent + // that is re-invoked after previously requesting a handoff. + + List> coordinatorInvocations = []; + List> specialistInvocations = []; + + var coordinator = new ChatClientAgent(CreateFunctionCallResultValidatingClient((messages, options) => + { + coordinatorInvocations.Add(messages.ToList()); + int callCount = coordinatorInvocations.Count; + + // Coordinator always hands off to specialist + string? transferFuncName = options?.Tools?.FirstOrDefault(t => t.Name.StartsWith("handoff_to_", StringComparison.Ordinal))?.Name; + Assert.NotNull(transferFuncName); + return new(new ChatMessage(ChatRole.Assistant, [new FunctionCallContent($"coord_call_{callCount}", transferFuncName)])); + }), name: "coordinator"); + + var specialist = new ChatClientAgent(CreateFunctionCallResultValidatingClient((messages, options) => + { + specialistInvocations.Add(messages.ToList()); + int callCount = specialistInvocations.Count; + + if (callCount % 2 == 1) + { + // Odd invocations: hand back to coordinator + string? transferFuncName = options?.Tools?.FirstOrDefault(t => t.Name.StartsWith("handoff_to_", StringComparison.Ordinal))?.Name; + Assert.NotNull(transferFuncName); + return new(new ChatMessage(ChatRole.Assistant, [new FunctionCallContent($"spec_call_{callCount}", transferFuncName)])); + } + + // Even invocations: respond normally + return new(new ChatMessage(ChatRole.Assistant, $"specialist response {callCount}")); + }), name: "specialist", description: "The specialist agent"); + + var workflow = AgentWorkflowBuilder.CreateHandoffBuilderWith(coordinator) + .WithHandoff(coordinator, specialist) + .WithHandoff(specialist, coordinator) + .EnableReturnToPrevious() + .Build(); + + CheckpointManager checkpointManager = CheckpointManager.CreateInMemory(); + const ExecutionEnvironment Environment = ExecutionEnvironment.InProcess_Lockstep; + + // Turn 1: coordinator -> specialist -> coordinator -> specialist (ends with response) + WorkflowRunResult result = await RunWorkflowCheckpointedAsync( + workflow, + [new ChatMessage(ChatRole.User, "turn 1")], + Environment, + checkpointManager); + + // Verify FunctionCall/FunctionResult pairing for all invocations + VerifyFunctionCallResultPairing(coordinatorInvocations, "coordinator"); + VerifyFunctionCallResultPairing(specialistInvocations, "specialist"); + + // Turn 2: conversation continues + _ = await RunWorkflowCheckpointedAsync( + workflow, + [new ChatMessage(ChatRole.User, "turn 2")], + Environment, + checkpointManager, + result.LastCheckpoint); + + // Verify pairing again after second turn + VerifyFunctionCallResultPairing(coordinatorInvocations, "coordinator"); + VerifyFunctionCallResultPairing(specialistInvocations, "specialist"); + } + + /// + /// Verifies that for each invocation of an agent, all FunctionCallContent items + /// that appear in the message history have corresponding FunctionResultContent items. + /// + private static void VerifyFunctionCallResultPairing(List> invocations, string agentName) + { + for (int i = 0; i < invocations.Count; i++) + { + List messages = invocations[i]; + + // Get all FunctionCallContent and FunctionResultContent items from the messages + var functionCalls = messages + .SelectMany(m => m.Contents.OfType()) + .ToList(); + + var functionResults = messages + .SelectMany(m => m.Contents.OfType()) + .ToList(); + + // Create lookup of call IDs that have results + var resultCallIds = new HashSet(functionResults.Select(r => r.CallId)); + + // Verify each function call has a matching result + foreach (var call in functionCalls) + { + Assert.True(resultCallIds.Contains(call.CallId), + $"Agent '{agentName}' invocation {i + 1}: FunctionCallContent with CallId '{call.CallId}' (Name: '{call.Name}') " + + "has no matching FunctionResultContent. This violates the LLM's requirement that all FunctionCalls have results."); + } + } + } + #region Helper Types and Methods private sealed record WorkflowRunResult(string UpdateText, List? Result, CheckpointInfo? LastCheckpoint, List PendingRequests);