diff --git a/src/Libraries/Microsoft.Extensions.AI/ChatCompletion/FunctionInvokingChatClient.cs b/src/Libraries/Microsoft.Extensions.AI/ChatCompletion/FunctionInvokingChatClient.cs index 9fd95df0424..f1858a295cf 100644 --- a/src/Libraries/Microsoft.Extensions.AI/ChatCompletion/FunctionInvokingChatClient.cs +++ b/src/Libraries/Microsoft.Extensions.AI/ChatCompletion/FunctionInvokingChatClient.cs @@ -460,7 +460,10 @@ public override async IAsyncEnumerable GetStreamingResponseA foreach (var message in preDownstreamCallHistory) { yield return ConvertToolResultMessageToUpdate(message, options?.ConversationId, message.MessageId); - Activity.Current = activity; // workaround for https://github.com/dotnet/runtime/issues/47802 + if (activity is not null) + { + Activity.Current = activity; // workaround for https://github.com/dotnet/runtime/issues/47802 + } } } @@ -474,7 +477,10 @@ public override async IAsyncEnumerable GetStreamingResponseA { message.MessageId = toolMessageId; yield return ConvertToolResultMessageToUpdate(message, options?.ConversationId, message.MessageId); - Activity.Current = activity; // workaround for https://github.com/dotnet/runtime/issues/47802 + if (activity is not null) + { + Activity.Current = activity; // workaround for https://github.com/dotnet/runtime/issues/47802 + } } if (shouldTerminate) @@ -557,8 +563,10 @@ public override async IAsyncEnumerable GetStreamingResponseA // we can yield the update as-is. lastYieldedUpdateIndex++; yield return update; - Activity.Current = activity; // workaround for https://github.com/dotnet/runtime/issues/47802 - + if (activity is not null) + { + Activity.Current = activity; // workaround for https://github.com/dotnet/runtime/issues/47802 + } continue; } @@ -584,7 +592,10 @@ public override async IAsyncEnumerable GetStreamingResponseA } yield return updateToYield; - Activity.Current = activity; // workaround for https://github.com/dotnet/runtime/issues/47802 + if (activity is not null) + { + Activity.Current = activity; // workaround for https://github.com/dotnet/runtime/issues/47802 + } } continue; @@ -601,7 +612,10 @@ public override async IAsyncEnumerable GetStreamingResponseA { var updateToYield = updates[lastYieldedUpdateIndex]; yield return updateToYield; - Activity.Current = activity; // workaround for https://github.com/dotnet/runtime/issues/47802 + if (activity is not null) + { + Activity.Current = activity; // workaround for https://github.com/dotnet/runtime/issues/47802 + } } // If there's nothing more to do, break out of the loop and allow the handling at the @@ -632,7 +646,10 @@ public override async IAsyncEnumerable GetStreamingResponseA foreach (var message in modeAndMessages.MessagesAdded) { yield return ConvertToolResultMessageToUpdate(message, response.ConversationId, toolMessageId); - Activity.Current = activity; // workaround for https://github.com/dotnet/runtime/issues/47802 + if (activity is not null) + { + Activity.Current = activity; // workaround for https://github.com/dotnet/runtime/issues/47802 + } } if (modeAndMessages.ShouldTerminate) diff --git a/test/Libraries/Microsoft.Extensions.AI.Tests/ChatCompletion/FunctionInvokingChatClientTests.cs b/test/Libraries/Microsoft.Extensions.AI.Tests/ChatCompletion/FunctionInvokingChatClientTests.cs index 2ddf757d185..dd2a139b4f6 100644 --- a/test/Libraries/Microsoft.Extensions.AI.Tests/ChatCompletion/FunctionInvokingChatClientTests.cs +++ b/test/Libraries/Microsoft.Extensions.AI.Tests/ChatCompletion/FunctionInvokingChatClientTests.cs @@ -1713,6 +1713,63 @@ public async Task DoesNotCreateOrchestrateToolsSpanWhenInvokeAgentIsParent(strin Assert.All(childActivities, activity => Assert.Same(invokeAgent, activity.Parent)); } + [Fact] + public async Task StreamingPreservesTraceContextWhenInvokeAgentWithNameIsParent() + { + string agentSourceName = Guid.NewGuid().ToString(); + string clientSourceName = Guid.NewGuid().ToString(); + var activities = new List(); + + using TracerProvider tracerProvider = OpenTelemetry.Sdk.CreateTracerProviderBuilder() + .AddSource(agentSourceName) + .AddSource(clientSourceName) + .AddInMemoryExporter(activities) + .Build(); + + int callCount = 0; + using var innerClient = new TestChatClient + { + GetStreamingResponseAsyncCallback = (messages, options, ct) => + { + callCount++; + ChatMessage message = callCount == 1 + ? new(ChatRole.Assistant, [new FunctionCallContent("call1", "Func1")]) + : new(ChatRole.Assistant, "Done"); + return YieldAsync(new ChatResponse(message).ToChatResponseUpdates()); + } + }; + + var client = innerClient.AsBuilder() + .Use(c => new FunctionInvokingChatClient( + new OpenTelemetryChatClient(c, sourceName: clientSourceName))) + .Build(); + + var options = new ChatOptions + { + Tools = [AIFunctionFactory.Create(() => "Result 1", "Func1")] + }; + + using var agentSource = new ActivitySource(agentSourceName); + using var invokeAgentActivity = agentSource.StartActivity("invoke_agent MyAgent(agent-123)"); + Assert.NotNull(invokeAgentActivity); + + await foreach (var update in client.GetStreamingResponseAsync( + [new ChatMessage(ChatRole.User, "hello")], options)) + { + // consume all updates + } + + Assert.Equal(2, callCount); + + var chatActivities = activities.Where(a => a.DisplayName.StartsWith("chat", StringComparison.Ordinal)).ToList(); + Assert.Equal(2, chatActivities.Count); + + // All child activities must share the same trace as invoke_agent + var nonAgentActivities = activities.Where(a => a != invokeAgentActivity).ToList(); + Assert.All(nonAgentActivities, a => + Assert.Equal(invokeAgentActivity.TraceId, a.TraceId)); + } + [Theory] [InlineData("invoke_agen")] [InlineData("invoke_agent_extra")]