diff --git a/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative.Mcp/DefaultMcpToolHandler.cs b/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative.Mcp/DefaultMcpToolHandler.cs index 61aedec347f..976d2b99815 100644 --- a/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative.Mcp/DefaultMcpToolHandler.cs +++ b/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative.Mcp/DefaultMcpToolHandler.cs @@ -6,6 +6,7 @@ using System.IO; using System.Linq; using System.Net.Http; +using System.Runtime.ExceptionServices; using System.Security.Cryptography; using System.Text; using System.Text.Json; @@ -28,12 +29,15 @@ namespace Microsoft.Agents.AI.Workflows.Declarative.Mcp; /// a pre-configured for each server. /// Provider-backed invocations create and dispose a separate MCP session for every call, including /// tools/list, because provider authentication is not represented in the session cache key. -/// Without a provider, sessions are cached by server URL, label, connection name, and explicit headers. +/// Without a provider, workflow invocations are cached by workflow session, server URL, label, +/// connection name, and explicit headers in a bounded least-recently-used cache. Evicted sessions +/// are disposed after their active invocations finish. /// Non-cancellation cleanup failures are reported through warnings without replacing /// the invocation result or error. /// -public sealed class DefaultMcpToolHandler : IMcpToolHandler, IAsyncDisposable +public sealed class DefaultMcpToolHandler : IWorkflowScopedMcpToolHandler, IAsyncDisposable { + private const int DefaultClientCacheMaxSize = 32; private const string FilenameAdditionalPropertyName = "filename"; /// @@ -46,9 +50,15 @@ public sealed class DefaultMcpToolHandler : IMcpToolHandler, IAsyncDisposable private readonly Func>? _httpClientProvider; private readonly Func _httpMessageHandlerFactory; - private readonly Dictionary<(string Url, string Label, string Connection, string HeadersHash), ClientConnection> _clients = []; - private readonly Dictionary _ownedHttpClients = []; + private readonly Dictionary<(string WorkflowSession, string Url, string Label, string Connection, string HeadersHash), CachedClient> _clients = []; + private readonly Dictionary<(string WorkflowSession, string Url, string Label, string Connection, string HeadersHash), TaskCompletionSource> _clientCreations = []; + private readonly HashSet> _clientCreationLifetimes = []; + private readonly LinkedList<(string WorkflowSession, string Url, string Label, string Connection, string HeadersHash)> _clientLru = []; + private readonly HashSet _retiredClients = []; + private readonly Dictionary _ownedHttpClients = []; private readonly SemaphoreSlim _clientLock = new(1, 1); + private readonly SemaphoreSlim _clientCreationSemaphore; + private readonly int _clientCacheMaxSize; private readonly AsyncLocal _providerInvocationContext = new(); private TaskCompletionSource? _providerInvocationsDrained; private int _activeProviderInvocations; @@ -83,10 +93,18 @@ public DefaultMcpToolHandler(Func>? internal DefaultMcpToolHandler( Func>? httpClientProvider, - Func httpMessageHandlerFactory) + Func httpMessageHandlerFactory, + int clientCacheMaxSize = DefaultClientCacheMaxSize) { + if (clientCacheMaxSize <= 0) + { + throw new ArgumentOutOfRangeException(nameof(clientCacheMaxSize), "The MCP client cache size must be positive."); + } + this._httpClientProvider = httpClientProvider; this._httpMessageHandlerFactory = Throw.IfNull(httpMessageHandlerFactory); + this._clientCacheMaxSize = clientCacheMaxSize; + this._clientCreationSemaphore = new(clientCacheMaxSize, clientCacheMaxSize); } /// @@ -98,6 +116,26 @@ public async Task InvokeToolAsync( IDictionary? headers, string? connectionName, CancellationToken cancellationToken = default) + => await this.InvokeToolInWorkflowSessionAsync( + serverUrl, + serverLabel, + toolName, + arguments, + headers, + connectionName, + workflowSessionId: string.Empty, + cancellationToken).ConfigureAwait(false); + + /// + public async Task InvokeToolInWorkflowSessionAsync( + string serverUrl, + string? serverLabel, + string toolName, + IDictionary? arguments, + IDictionary? headers, + string? connectionName, + string workflowSessionId, + CancellationToken cancellationToken = default) { if (IsListToolsToolName(toolName)) { @@ -126,7 +164,7 @@ public async Task InvokeToolAsync( try { ClientConnection invocationClient = await this.CreateClientAsync( - serverUrl.Trim(), serverLabel, headers, httpClientCacheKey: null, cancellationToken).ConfigureAwait(false); + serverUrl.Trim(), serverLabel, headers, cancellationToken).ConfigureAwait(false); await using (invocationClient.ConfigureAwait(false)) { return await InvokeClientAsync(invocationClient.Client, toolName, arguments, cancellationToken).ConfigureAwait(false); @@ -151,8 +189,16 @@ public async Task InvokeToolAsync( } } - McpClient client = await this.GetOrCreateClientAsync(serverUrl, serverLabel, headers, connectionName, cancellationToken).ConfigureAwait(false); - return await InvokeClientAsync(client, toolName, arguments, cancellationToken).ConfigureAwait(false); + CachedClient cachedClient = await this.AcquireClientAsync( + serverUrl, serverLabel, headers, connectionName, workflowSessionId, cancellationToken).ConfigureAwait(false); + try + { + return await InvokeClientAsync(cachedClient.Connection.Client, toolName, arguments, cancellationToken).ConfigureAwait(false); + } + finally + { + await this.ReleaseClientAsync(cachedClient).ConfigureAwait(false); + } } private static async Task InvokeClientAsync( @@ -218,12 +264,14 @@ public async ValueTask DisposeAsync() } Task? providerInvocations; + List> clientCreationLifetimes; await this._clientLock.WaitAsync().ConfigureAwait(false); try { this.ThrowIfDisposing(); this._disposing = true; providerInvocations = this._providerInvocationsDrained?.Task; + clientCreationLifetimes = this._clientCreationLifetimes.Select(source => source.Task).ToList(); } finally { @@ -235,32 +283,79 @@ public async ValueTask DisposeAsync() await providerInvocations.ConfigureAwait(false); } + foreach (Task clientCreationLifetime in clientCreationLifetimes) + { + await clientCreationLifetime.ConfigureAwait(false); + } + + List cachedClients; + List clientsToDispose = []; await this._clientLock.WaitAsync().ConfigureAwait(false); try { - foreach (ClientConnection client in this._clients.Values) + cachedClients = [.. this._clients.Values, .. this._retiredClients]; + foreach (CachedClient client in cachedClients) { - await client.DisposeAsync().ConfigureAwait(false); + client.Evicted = true; + this._retiredClients.Add(client); + if (client.ActiveInvocations == 0 && !client.DisposalClaimed) + { + client.DisposalClaimed = true; + clientsToDispose.Add(client); + } } this._clients.Clear(); - - // Dispose only HttpClients that the handler created (not user-provided ones) - foreach (HttpClient httpClient in this._ownedHttpClients.Values) - { - httpClient.Dispose(); - } - - this._ownedHttpClients.Clear(); + this._clientLru.Clear(); } finally { this._clientLock.Release(); } - this._clientLock.Dispose(); + try + { + await DrainCleanupAsync( + clientsToDispose.Select(client => this.DisposeCachedClientAsync(client)), + cachedClients.Select(client => client.Disposed.Task)).ConfigureAwait(false); + } + finally + { + this._clientLock.Dispose(); + this._clientCreationSemaphore.Dispose(); + } } + internal static async Task DrainCleanupAsync(IEnumerable cleanupTasks, IEnumerable completionTasks) + { + Exception? cleanupException = null; + try + { + await Task.WhenAll(cleanupTasks).ConfigureAwait(false); + } + catch (Exception exception) when (!IsFatalException(exception)) + { + cleanupException = exception; + } + + await Task.WhenAll(completionTasks).ConfigureAwait(false); + + if (cleanupException is not null) + { + ExceptionDispatchInfo.Capture(cleanupException).Throw(); + } + } + + private static bool IsFatalException(Exception exception) => + exception is OutOfMemoryException and not InsufficientMemoryException + or StackOverflowException + or AccessViolationException + or AppDomainUnloadedException + or BadImageFormatException + or CannotUnloadAppDomainException + or InvalidProgramException + or ThreadAbortException; + [System.Diagnostics.CodeAnalysis.SuppressMessage("Maintainability", "CA1513:Use ObjectDisposedException throw helper", Justification = "The helper is not available on .NET Framework or .NET Standard 2.0.")] private void ThrowIfDisposing() @@ -271,89 +366,336 @@ private void ThrowIfDisposing() } } - private async Task GetOrCreateClientAsync( + private async Task AcquireClientAsync( string serverUrl, string? serverLabel, IDictionary? headers, string? connectionName, + string workflowSessionId, CancellationToken cancellationToken) { string trimmedUrl = serverUrl.Trim(); - var clientCacheKey = BuildCacheKey(trimmedUrl, serverLabel, connectionName, headers); + var clientCacheKey = BuildCacheKey(workflowSessionId, trimmedUrl, serverLabel, connectionName, headers); + CachedClient? clientToDispose = null; + TaskCompletionSource? clientCreation; + TaskCompletionSource? clientCreationLifetime = null; + bool ownsClientCreation = false; await this._clientLock.WaitAsync(cancellationToken).ConfigureAwait(false); try { this.ThrowIfDisposing(); - if (this._clients.TryGetValue(clientCacheKey, out ClientConnection? existingClient)) + if (this._clients.TryGetValue(clientCacheKey, out CachedClient? existingClient)) + { + existingClient.ActiveInvocations++; + this._clientLru.Remove(existingClient.LruNode); + this._clientLru.AddLast(existingClient.LruNode); + return existingClient; + } + + if (!this._clientCreations.TryGetValue(clientCacheKey, out clientCreation)) + { + clientCreation = new(TaskCreationOptions.RunContinuationsAsynchronously); + this._clientCreations[clientCacheKey] = clientCreation; + clientCreationLifetime = new(TaskCreationOptions.RunContinuationsAsynchronously); + this._clientCreationLifetimes.Add(clientCreationLifetime); + ownsClientCreation = true; + } + } + finally + { + this._clientLock.Release(); + } + + if (!ownsClientCreation) + { + try + { + await WaitForClientCreationAsync(clientCreation.Task, cancellationToken).ConfigureAwait(false); + } + catch (OperationCanceledException) when (!cancellationToken.IsCancellationRequested) + { + return await this.AcquireClientAsync( + serverUrl, serverLabel, headers, connectionName, workflowSessionId, cancellationToken).ConfigureAwait(false); + } + + return await this.AcquireClientAsync( + serverUrl, serverLabel, headers, connectionName, workflowSessionId, cancellationToken).ConfigureAwait(false); + } + + TaskCompletionSource ownedClientCreationLifetime = clientCreationLifetime ?? + throw new InvalidOperationException("Missing MCP client creation lifetime."); + ClientConnection? connection = null; + bool creationSemaphoreEntered = false; + try + { + await this._clientCreationSemaphore.WaitAsync(cancellationToken).ConfigureAwait(false); + creationSemaphoreEntered = true; + await this._clientLock.WaitAsync(cancellationToken).ConfigureAwait(false); + try + { + this.ThrowIfDisposing(); + } + finally + { + this._clientLock.Release(); + } + + connection = await this.CreateClientAsync( + trimmedUrl, serverLabel, headers, cancellationToken).ConfigureAwait(false); + } + catch (Exception exception) + { + try + { + await this.CompleteClientCreationFailureAsync(clientCacheKey, clientCreation, exception).ConfigureAwait(false); + } + finally + { + if (creationSemaphoreEntered) + { + this._clientCreationSemaphore.Release(); + creationSemaphoreEntered = false; + } + + await this.CompleteClientCreationLifetimeAsync(ownedClientCreationLifetime).ConfigureAwait(false); + } + + throw; + } + + ObjectDisposedException? disposedException = null; + CachedClient? result = null; + await this._clientLock.WaitAsync(CancellationToken.None).ConfigureAwait(false); + try + { + this._clientCreations.Remove(clientCacheKey); + if (this._disposing) + { + disposedException = new ObjectDisposedException(nameof(DefaultMcpToolHandler)); + } + else + { + LinkedListNode<(string WorkflowSession, string Url, string Label, string Connection, string HeadersHash)> node = + this._clientLru.AddLast(clientCacheKey); + CachedClient newClient = new(connection, node) { ActiveInvocations = 1 }; + connection = null; + this._clients[clientCacheKey] = newClient; + + if (this._clients.Count > this._clientCacheMaxSize) + { + LinkedListNode<(string WorkflowSession, string Url, string Label, string Connection, string HeadersHash)> evictedNode = + this._clientLru.First!; + this._clientLru.RemoveFirst(); + CachedClient evictedClient = this._clients[evictedNode.Value]; + this._clients.Remove(evictedNode.Value); + evictedClient.Evicted = true; + this._retiredClients.Add(evictedClient); + if (evictedClient.ActiveInvocations == 0) + { + evictedClient.DisposalClaimed = true; + clientToDispose = evictedClient; + } + } + + result = newClient; + clientCreation.TrySetResult(newClient); + } + } + finally + { + this._clientLock.Release(); + } + + try + { + if (connection is not null) + { + await connection.DisposeAsync().ConfigureAwait(false); + } + + if (disposedException is not null) + { + clientCreation.TrySetException(disposedException); + throw disposedException; + } + + if (clientToDispose is not null) + { + await this.DisposeCachedClientAsync(clientToDispose).ConfigureAwait(false); + } + } + catch (Exception exception) + { + if (disposedException is not null) + { + clientCreation.TrySetException(exception); + } + + if (result is not null) + { + await this.ReleaseClientAsync(result).ConfigureAwait(false); + } + + throw; + } + finally + { + if (creationSemaphoreEntered) + { + this._clientCreationSemaphore.Release(); + } + + await this.CompleteClientCreationLifetimeAsync(ownedClientCreationLifetime).ConfigureAwait(false); + } + + return result ?? throw new InvalidOperationException("Failed to acquire MCP client."); + } + + private static async Task WaitForClientCreationAsync(Task clientCreation, CancellationToken cancellationToken) + { + if (!cancellationToken.CanBeCanceled || clientCreation.IsCompleted) + { + await clientCreation.ConfigureAwait(false); + return; + } + + TaskCompletionSource cancellation = new(TaskCreationOptions.RunContinuationsAsynchronously); + using CancellationTokenRegistration registration = cancellationToken.Register( + static state => ((TaskCompletionSource)state!).TrySetResult(true), + cancellation); + + if (await Task.WhenAny(clientCreation, cancellation.Task).ConfigureAwait(false) == cancellation.Task) + { + cancellationToken.ThrowIfCancellationRequested(); + } + + await clientCreation.ConfigureAwait(false); + } + + private async Task CompleteClientCreationFailureAsync( + (string WorkflowSession, string Url, string Label, string Connection, string HeadersHash) clientCacheKey, + TaskCompletionSource clientCreation, + Exception exception) + { + await this._clientLock.WaitAsync(CancellationToken.None).ConfigureAwait(false); + try + { + if (this._clientCreations.TryGetValue(clientCacheKey, out TaskCompletionSource? existingCreation) && + ReferenceEquals(existingCreation, clientCreation)) { - return existingClient.Client; + this._clientCreations.Remove(clientCacheKey); } - ClientConnection newClient = await this.CreateClientAsync(trimmedUrl, serverLabel, headers, trimmedUrl, cancellationToken).ConfigureAwait(false); - this._clients[clientCacheKey] = newClient; - return newClient.Client; + clientCreation.TrySetException(exception); + _ = clientCreation.Task.Exception; + } + finally + { + this._clientLock.Release(); + } + } + + private async Task CompleteClientCreationLifetimeAsync(TaskCompletionSource clientCreationLifetime) + { + await this._clientLock.WaitAsync(CancellationToken.None).ConfigureAwait(false); + try + { + this._clientCreationLifetimes.Remove(clientCreationLifetime); } finally { this._clientLock.Release(); } + + clientCreationLifetime.TrySetResult(true); + } + + private async ValueTask ReleaseClientAsync(CachedClient client) + { + bool dispose = false; + await this._clientLock.WaitAsync(CancellationToken.None).ConfigureAwait(false); + try + { + client.ActiveInvocations--; + if (client.ActiveInvocations == 0 && client.Evicted && !client.DisposalClaimed) + { + client.DisposalClaimed = true; + dispose = true; + } + } + finally + { + this._clientLock.Release(); + } + + if (dispose) + { + await this.DisposeCachedClientAsync(client).ConfigureAwait(false); + } + } + + private async Task DisposeCachedClientAsync(CachedClient client) + { + try + { + await client.Connection.DisposeAsync().ConfigureAwait(false); + } + finally + { + await this._clientLock.WaitAsync(CancellationToken.None).ConfigureAwait(false); + try + { + this._retiredClients.Remove(client); + } + finally + { + this._clientLock.Release(); + } + + client.Disposed.TrySetResult(true); + } } /// - /// Builds the per-client cache key as a 4-tuple of - /// (trimmed serverUrl, serverLabel, connectionName, headers hash). All four components - /// participate so that callers using different labels/connections/headers receive - /// distinct instances even when targeting the same URL. + /// Builds the per-client cache key as a 5-tuple of + /// (workflowSessionId, trimmed serverUrl, serverLabel, connectionName, headers hash). + /// All five components participate so that separate workflow sessions and callers using + /// different labels/connections/headers receive distinct instances + /// even when targeting the same URL. /// - internal static (string Url, string Label, string Connection, string HeadersHash) BuildCacheKey( + internal static (string WorkflowSession, string Url, string Label, string Connection, string HeadersHash) BuildCacheKey( + string workflowSessionId, string trimmedUrl, string? serverLabel, string? connectionName, IDictionary? headers) => - (trimmedUrl, serverLabel ?? string.Empty, connectionName ?? string.Empty, ComputeHeadersHash(headers)); + (workflowSessionId, trimmedUrl, serverLabel ?? string.Empty, connectionName ?? string.Empty, ComputeHeadersHash(headers)); private async Task CreateClientAsync( string serverUrl, string? serverLabel, IDictionary? headers, - string? httpClientCacheKey, CancellationToken cancellationToken) { - // Only the no-provider path shares handler-owned HTTP clients. HttpClient? httpClient = null; bool ownsHttpClient = false; + OwnedHttpClientLease? ownedHttpClientLease = null; if (this._httpClientProvider is not null) { httpClient = await this._httpClientProvider(serverUrl, cancellationToken).ConfigureAwait(false); } - if (httpClient is null && - (httpClientCacheKey is null || !this._ownedHttpClients.TryGetValue(httpClientCacheKey, out httpClient))) + if (httpClient is null && this._httpClientProvider is not null) { - // Pin credential headers to the configured server origin as defense-in-depth. Forcing - // StreamableHttp (below) already removes the primary vector (a server-advertised cross-origin - // SSE message endpoint), and AllowAutoRedirect=false blocks auto-redirects. This handler is the - // backstop: it guarantees the Authorization token and other credentials never leave the pinned - // origin even if a future change re-enables AutoDetect or redirects, or the SDK constructs a - // request to a new URI (AdditionalHeaders are re-stamped by the transport, so HttpClient's own - // redirect header-stripping does not cover them). - // Caveat for future maintainers: the handler strips credentials on every cross-origin request. - // That is safe today because this handler carries auth via static request headers, not the - // SDK's built-in OAuth flow. If OAuth support is ever added here, requests to the authorization - // server (a different origin than the MCP resource server) legitimately carry credentials, so - // those auth-server origins would need to be allow-listed to avoid breaking the token exchange. - OriginPinningHandler pinningHandler = new(new Uri(serverUrl)) { InnerHandler = this._httpMessageHandlerFactory() }; - httpClient = new HttpClient(pinningHandler); - if (httpClientCacheKey is null) - { - ownsHttpClient = true; - } - else - { - this._ownedHttpClients[httpClientCacheKey] = httpClient; - } + httpClient = this.CreatePinnedHttpClient(serverUrl); + ownsHttpClient = true; + } + else if (httpClient is null) + { + ownedHttpClientLease = await this.AcquireOwnedHttpClientAsync(serverUrl, cancellationToken).ConfigureAwait(false); + httpClient = ownedHttpClientLease.Client; } HttpClientTransportOptions transportOptions = new() @@ -368,20 +710,93 @@ private async Task CreateClientAsync( TransportMode = HttpTransportMode.StreamableHttp }; - HttpClientTransport transport = new(transportOptions, httpClient, ownsHttpClient: ownsHttpClient); - + HttpClientTransport? transport = null; try { + HttpClient resolvedHttpClient = httpClient ?? throw new InvalidOperationException("Failed to resolve MCP HTTP client."); + transport = new(transportOptions, resolvedHttpClient, ownsHttpClient: ownsHttpClient); McpClient client = await McpClient.CreateAsync(transport, cancellationToken: cancellationToken).ConfigureAwait(false); - return new ClientConnection(client, transport); + ClientConnection connection = new(client, transport, ownedHttpClientLease); + ownedHttpClientLease = null; + return connection; } catch { - await DisposeResourceAsync(transport, "transport").ConfigureAwait(false); + try + { + if (transport is not null) + { + await DisposeResourceAsync(transport, "transport").ConfigureAwait(false); + } + } + finally + { + if (ownedHttpClientLease is not null) + { + await ownedHttpClientLease.DisposeAsync().ConfigureAwait(false); + } + } + throw; } } + private async Task AcquireOwnedHttpClientAsync( + string serverUrl, + CancellationToken cancellationToken) + { + await this._clientLock.WaitAsync(cancellationToken).ConfigureAwait(false); + try + { + if (!this._ownedHttpClients.TryGetValue(serverUrl, out OwnedHttpClient? entry)) + { + entry = new(this.CreatePinnedHttpClient(serverUrl)); + this._ownedHttpClients[serverUrl] = entry; + } + + entry.ReferenceCount++; + return new(this, serverUrl, entry.Client); + } + finally + { + this._clientLock.Release(); + } + } + + private async ValueTask ReleaseOwnedHttpClientAsync(string serverUrl) + { + HttpClient? clientToDispose = null; + await this._clientLock.WaitAsync(CancellationToken.None).ConfigureAwait(false); + try + { + OwnedHttpClient entry = this._ownedHttpClients[serverUrl]; + if (--entry.ReferenceCount == 0) + { + this._ownedHttpClients.Remove(serverUrl); + clientToDispose = entry.Client; + } + } + finally + { + this._clientLock.Release(); + } + + clientToDispose?.Dispose(); + } + + private HttpClient CreatePinnedHttpClient(string serverUrl) + { + // Pin credential headers to the configured server origin as defense-in-depth. Forcing + // StreamableHttp (below) already removes the primary vector (a server-advertised cross-origin + // SSE message endpoint), and AllowAutoRedirect=false blocks auto-redirects. This handler is the + // backstop: it guarantees the Authorization token and other credentials never leave the pinned + // origin even if a future change re-enables AutoDetect or redirects, or the SDK constructs a + // request to a new URI (AdditionalHeaders are re-stamped by the transport, so HttpClient's own + // redirect header-stripping does not cover them). + OriginPinningHandler pinningHandler = new(new Uri(serverUrl)) { InnerHandler = this._httpMessageHandlerFactory() }; + return new HttpClient(pinningHandler); + } + private static HttpMessageHandler CreateHttpMessageHandler() => new HttpClientHandler { @@ -398,7 +813,53 @@ private sealed class ProviderInvocationContext(Task completion, ProviderInvocati public ProviderInvocationContext? Parent { get; } = parent; } - internal sealed class ClientConnection(McpClient client, IAsyncDisposable transport) : IAsyncDisposable + private sealed class CachedClient( + ClientConnection connection, + LinkedListNode<(string WorkflowSession, string Url, string Label, string Connection, string HeadersHash)> lruNode) + { + public ClientConnection Connection { get; } = connection; + + public LinkedListNode<(string WorkflowSession, string Url, string Label, string Connection, string HeadersHash)> LruNode { get; } = lruNode; + + public TaskCompletionSource Disposed { get; } = new(TaskCreationOptions.RunContinuationsAsynchronously); + + public int ActiveInvocations { get; set; } + + public bool Evicted { get; set; } + + public bool DisposalClaimed { get; set; } + } + + private sealed class OwnedHttpClient(HttpClient client) + { + public HttpClient Client { get; } = client; + + public int ReferenceCount { get; set; } + } + + private sealed class OwnedHttpClientLease( + DefaultMcpToolHandler owner, + string serverUrl, + HttpClient client) : IAsyncDisposable + { + private bool _disposed; + + public HttpClient Client { get; } = client; + + public async ValueTask DisposeAsync() + { + if (!this._disposed) + { + this._disposed = true; + await owner.ReleaseOwnedHttpClientAsync(serverUrl).ConfigureAwait(false); + } + } + } + + internal sealed class ClientConnection( + McpClient client, + IAsyncDisposable transport, + IAsyncDisposable? ownedHttpClientLease = null) : IAsyncDisposable { public McpClient Client { get; } = client; @@ -410,8 +871,18 @@ public async ValueTask DisposeAsync() } finally { - // McpClient owns the connected session, not the reusable transport factory. - await DisposeResourceAsync(transport, "transport").ConfigureAwait(false); + try + { + // McpClient owns the connected session, not the reusable transport factory. + await DisposeResourceAsync(transport, "transport").ConfigureAwait(false); + } + finally + { + if (ownedHttpClientLease is not null) + { + await ownedHttpClientLease.DisposeAsync().ConfigureAwait(false); + } + } } } } diff --git a/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative.Mcp/PublicAPI/net10.0/PublicAPI.Unshipped.txt b/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative.Mcp/PublicAPI/net10.0/PublicAPI.Unshipped.txt index ab058de62d4..3cdebd5e1a0 100644 --- a/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative.Mcp/PublicAPI/net10.0/PublicAPI.Unshipped.txt +++ b/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative.Mcp/PublicAPI/net10.0/PublicAPI.Unshipped.txt @@ -1 +1,2 @@ #nullable enable +Microsoft.Agents.AI.Workflows.Declarative.Mcp.DefaultMcpToolHandler.InvokeToolInWorkflowSessionAsync(string! serverUrl, string? serverLabel, string! toolName, System.Collections.Generic.IDictionary? arguments, System.Collections.Generic.IDictionary? headers, string? connectionName, string! workflowSessionId, System.Threading.CancellationToken cancellationToken = default(System.Threading.CancellationToken)) -> System.Threading.Tasks.Task! diff --git a/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative.Mcp/PublicAPI/net472/PublicAPI.Unshipped.txt b/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative.Mcp/PublicAPI/net472/PublicAPI.Unshipped.txt index ab058de62d4..3cdebd5e1a0 100644 --- a/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative.Mcp/PublicAPI/net472/PublicAPI.Unshipped.txt +++ b/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative.Mcp/PublicAPI/net472/PublicAPI.Unshipped.txt @@ -1 +1,2 @@ #nullable enable +Microsoft.Agents.AI.Workflows.Declarative.Mcp.DefaultMcpToolHandler.InvokeToolInWorkflowSessionAsync(string! serverUrl, string? serverLabel, string! toolName, System.Collections.Generic.IDictionary? arguments, System.Collections.Generic.IDictionary? headers, string? connectionName, string! workflowSessionId, System.Threading.CancellationToken cancellationToken = default(System.Threading.CancellationToken)) -> System.Threading.Tasks.Task! diff --git a/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative.Mcp/PublicAPI/net8.0/PublicAPI.Unshipped.txt b/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative.Mcp/PublicAPI/net8.0/PublicAPI.Unshipped.txt index ab058de62d4..3cdebd5e1a0 100644 --- a/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative.Mcp/PublicAPI/net8.0/PublicAPI.Unshipped.txt +++ b/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative.Mcp/PublicAPI/net8.0/PublicAPI.Unshipped.txt @@ -1 +1,2 @@ #nullable enable +Microsoft.Agents.AI.Workflows.Declarative.Mcp.DefaultMcpToolHandler.InvokeToolInWorkflowSessionAsync(string! serverUrl, string? serverLabel, string! toolName, System.Collections.Generic.IDictionary? arguments, System.Collections.Generic.IDictionary? headers, string? connectionName, string! workflowSessionId, System.Threading.CancellationToken cancellationToken = default(System.Threading.CancellationToken)) -> System.Threading.Tasks.Task! diff --git a/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative.Mcp/PublicAPI/net9.0/PublicAPI.Unshipped.txt b/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative.Mcp/PublicAPI/net9.0/PublicAPI.Unshipped.txt index ab058de62d4..3cdebd5e1a0 100644 --- a/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative.Mcp/PublicAPI/net9.0/PublicAPI.Unshipped.txt +++ b/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative.Mcp/PublicAPI/net9.0/PublicAPI.Unshipped.txt @@ -1 +1,2 @@ #nullable enable +Microsoft.Agents.AI.Workflows.Declarative.Mcp.DefaultMcpToolHandler.InvokeToolInWorkflowSessionAsync(string! serverUrl, string? serverLabel, string! toolName, System.Collections.Generic.IDictionary? arguments, System.Collections.Generic.IDictionary? headers, string? connectionName, string! workflowSessionId, System.Threading.CancellationToken cancellationToken = default(System.Threading.CancellationToken)) -> System.Threading.Tasks.Task! diff --git a/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative.Mcp/PublicAPI/netstandard2.0/PublicAPI.Unshipped.txt b/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative.Mcp/PublicAPI/netstandard2.0/PublicAPI.Unshipped.txt index ab058de62d4..3cdebd5e1a0 100644 --- a/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative.Mcp/PublicAPI/netstandard2.0/PublicAPI.Unshipped.txt +++ b/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative.Mcp/PublicAPI/netstandard2.0/PublicAPI.Unshipped.txt @@ -1 +1,2 @@ #nullable enable +Microsoft.Agents.AI.Workflows.Declarative.Mcp.DefaultMcpToolHandler.InvokeToolInWorkflowSessionAsync(string! serverUrl, string? serverLabel, string! toolName, System.Collections.Generic.IDictionary? arguments, System.Collections.Generic.IDictionary? headers, string? connectionName, string! workflowSessionId, System.Threading.CancellationToken cancellationToken = default(System.Threading.CancellationToken)) -> System.Threading.Tasks.Task! diff --git a/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative/IMcpToolHandler.cs b/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative/IMcpToolHandler.cs index 56b1c3deb4c..fbea21f35b4 100644 --- a/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative/IMcpToolHandler.cs +++ b/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative/IMcpToolHandler.cs @@ -39,3 +39,34 @@ Task InvokeToolAsync( string? connectionName, CancellationToken cancellationToken = default); } + +/// +/// Defines the contract for MCP handlers that isolate stateful protocol sessions by workflow session. +/// +public interface IWorkflowScopedMcpToolHandler : IMcpToolHandler +{ + /// + /// Invokes an MCP tool within the framework-owned workflow session scope. + /// + /// The URL of the MCP server. + /// An optional label identifying the server connection. + /// The name of the tool to invoke. + /// Optional arguments to pass to the tool. + /// Optional headers to include in the request. + /// An optional connection name for managed connections. + /// The framework-owned identifier for the current workflow session. + /// A token to observe cancellation. + /// + /// A task representing the asynchronous operation. The result contains a + /// with the tool invocation output. + /// + Task InvokeToolInWorkflowSessionAsync( + string serverUrl, + string? serverLabel, + string toolName, + IDictionary? arguments, + IDictionary? headers, + string? connectionName, + string workflowSessionId, + CancellationToken cancellationToken = default); +} diff --git a/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative/Interpreter/DeclarativeActionExecutor.cs b/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative/Interpreter/DeclarativeActionExecutor.cs index 597ea831296..fbbb739a55d 100644 --- a/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative/Interpreter/DeclarativeActionExecutor.cs +++ b/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative/Interpreter/DeclarativeActionExecutor.cs @@ -89,7 +89,9 @@ public override async ValueTask HandleAsync(ActionExecutorResult message, IWorkf try { - object? result = await this.ExecuteAsync(new DeclarativeWorkflowContext(context, this.State), cancellationToken).ConfigureAwait(false); + DeclarativeWorkflowContext declarativeContext = + await DeclarativeWorkflowContext.CreateAsync(context, this.State, cancellationToken).ConfigureAwait(false); + object? result = await this.ExecuteAsync(declarativeContext, cancellationToken).ConfigureAwait(false); Debug.WriteLine($"RESULT #{this.Id} - {result ?? "(null)"}"); if (this.EmitResultEvent) diff --git a/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative/Interpreter/DeclarativeWorkflowContext.cs b/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative/Interpreter/DeclarativeWorkflowContext.cs index 183b2e785e7..f88035abd12 100644 --- a/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative/Interpreter/DeclarativeWorkflowContext.cs +++ b/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative/Interpreter/DeclarativeWorkflowContext.cs @@ -15,8 +15,10 @@ namespace Microsoft.Agents.AI.Workflows.Declarative.Interpreter; -internal sealed class DeclarativeWorkflowContext : IWorkflowContext +internal sealed class DeclarativeWorkflowContext : IWorkflowContext, IWorkflowSessionContext { + internal const string WorkflowSessionIdStateKey = "__declarative_mcp_workflow_session_id"; + public static readonly FrozenSet ManagedScopes = [ VariableScopeNames.Local, @@ -24,16 +26,36 @@ internal sealed class DeclarativeWorkflowContext : IWorkflowContext VariableScopeNames.Global, ]; - public DeclarativeWorkflowContext(IWorkflowContext source, WorkflowFormulaState state) + private DeclarativeWorkflowContext(IWorkflowContext source, WorkflowFormulaState state, string sessionId) { this.Source = source; this.State = state; + this.SessionId = sessionId; + } + + public static async ValueTask CreateAsync( + IWorkflowContext source, + WorkflowFormulaState state, + CancellationToken cancellationToken = default) + { + string generatedSessionId = Guid.NewGuid().ToString("N"); + string sessionId = source is IWorkflowSessionContext sessionContext + ? sessionContext.SessionId + : await source.ReadOrInitStateAsync( + WorkflowSessionIdStateKey, + () => generatedSessionId, + VariableScopeNames.System, + cancellationToken: cancellationToken).ConfigureAwait(false); + + return new(source, state, sessionId); } private IWorkflowContext Source { get; } public WorkflowFormulaState State { get; } public IReadOnlyDictionary? TraceContext => this.Source.TraceContext; + public string SessionId { get; } + /// public bool ConcurrentRunsEnabled => this.Source.ConcurrentRunsEnabled; diff --git a/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative/Interpreter/DeclarativeWorkflowExecutor.cs b/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative/Interpreter/DeclarativeWorkflowExecutor.cs index 1126b43d8ff..2fcdea4ecce 100644 --- a/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative/Interpreter/DeclarativeWorkflowExecutor.cs +++ b/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative/Interpreter/DeclarativeWorkflowExecutor.cs @@ -154,7 +154,8 @@ private async ValueTask AdvanceAsync(ChatMessage input, IWorkflowContext context // No state to restore if we're starting from the beginning. state.SetInitialized(); - DeclarativeWorkflowContext declarativeContext = new(context, state); + DeclarativeWorkflowContext declarativeContext = + await DeclarativeWorkflowContext.CreateAsync(context, state, cancellationToken).ConfigureAwait(false); // Conversation id resolution prefers state already persisted by a prior turn, // so multi-turn invocations reuse the same backend conversation rather than diff --git a/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative/Interpreter/DelegateActionExecutor.cs b/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative/Interpreter/DelegateActionExecutor.cs index 716a3656dac..ade8dd1034b 100644 --- a/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative/Interpreter/DelegateActionExecutor.cs +++ b/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative/Interpreter/DelegateActionExecutor.cs @@ -61,7 +61,9 @@ public override async ValueTask HandleAsync(TMessage message, IWorkflowContext c { if (this._action is not null) { - await this._action.Invoke(new DeclarativeWorkflowContext(context, this._state), message, cancellationToken).ConfigureAwait(false); + DeclarativeWorkflowContext declarativeContext = + await DeclarativeWorkflowContext.CreateAsync(context, this._state, cancellationToken).ConfigureAwait(false); + await this._action.Invoke(declarativeContext, message, cancellationToken).ConfigureAwait(false); } if (this._emitResult) diff --git a/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative/Kit/ActionExecutor.cs b/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative/Kit/ActionExecutor.cs index cf636effaf8..563ff86f147 100644 --- a/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative/Kit/ActionExecutor.cs +++ b/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative/Kit/ActionExecutor.cs @@ -76,7 +76,9 @@ public ValueTask ResetAsync() [SendsMessage(typeof(ActionExecutorResult))] public override async ValueTask HandleAsync(TMessage message, IWorkflowContext context, CancellationToken cancellationToken) { - object? result = await this.ExecuteAsync(new DeclarativeWorkflowContext(context, this._session.State), message, cancellationToken).ConfigureAwait(false); + DeclarativeWorkflowContext declarativeContext = + await DeclarativeWorkflowContext.CreateAsync(context, this._session.State, cancellationToken).ConfigureAwait(false); + object? result = await this.ExecuteAsync(declarativeContext, message, cancellationToken).ConfigureAwait(false); Debug.WriteLine($"RESULT #{this.Id} - {result ?? "(null)"}"); await context.SendResultMessageAsync(this.Id, result, cancellationToken).ConfigureAwait(false); diff --git a/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative/Kit/RootExecutor.cs b/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative/Kit/RootExecutor.cs index 3a18a06c948..9e008ed391b 100644 --- a/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative/Kit/RootExecutor.cs +++ b/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative/Kit/RootExecutor.cs @@ -63,7 +63,8 @@ public ValueTask ResetAsync() [SendsMessage(typeof(ActionExecutorResult))] public override async ValueTask HandleAsync(TInput message, IWorkflowContext context, CancellationToken cancellationToken) { - DeclarativeWorkflowContext declarativeContext = new(context, this._state); + DeclarativeWorkflowContext declarativeContext = + await DeclarativeWorkflowContext.CreateAsync(context, this._state, cancellationToken).ConfigureAwait(false); ChatMessage input = (this._inputTransform ?? DefaultInputTransform).Invoke(message); diff --git a/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative/ObjectModel/InvokeMcpToolExecutor.cs b/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative/ObjectModel/InvokeMcpToolExecutor.cs index 46a5cae5fdd..ed74d30d39d 100644 --- a/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative/ObjectModel/InvokeMcpToolExecutor.cs +++ b/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative/ObjectModel/InvokeMcpToolExecutor.cs @@ -31,7 +31,6 @@ internal sealed class InvokeMcpToolExecutor( { private const string ApprovalSnapshotStateKey = nameof(_approvalSnapshots); private const string LegacyApprovalSnapshotStateKey = "_approvalSnapshot"; - /// /// Snapshots of evaluated parameters captured at approval-request time, keyed by /// per-invocation request id. Each pending approval lives here until the matching @@ -115,7 +114,8 @@ public static bool RequiresNothing(object? message) => } // No approval required - invoke the tool directly - McpServerToolResultContent resultContent = await mcpToolHandler.InvokeToolAsync( + McpServerToolResultContent resultContent = await this.InvokeToolAsync( + context, serverUrl, serverLabel, toolName, @@ -172,7 +172,8 @@ public async ValueTask CaptureResponseAsync( Dictionary? headers = this.GetHeaders(); - McpServerToolResultContent resultContent = await mcpToolHandler.InvokeToolAsync( + McpServerToolResultContent resultContent = await this.InvokeToolAsync( + context, snapshot.ServerUrl, snapshot.ServerLabel, snapshot.ToolName, @@ -184,6 +185,47 @@ public async ValueTask CaptureResponseAsync( await this.ProcessResultAsync(context, resultContent, cancellationToken).ConfigureAwait(false); } + private async Task InvokeToolAsync( + IWorkflowContext context, + string serverUrl, + string? serverLabel, + string toolName, + IDictionary? arguments, + IDictionary? headers, + string? connectionName, + CancellationToken cancellationToken) + { + if (mcpToolHandler is IWorkflowScopedMcpToolHandler scopedHandler) + { + string generatedWorkflowSessionId = Guid.NewGuid().ToString("N"); + string workflowSessionId = context is IWorkflowSessionContext sessionContext + ? sessionContext.SessionId + : await context.ReadOrInitStateAsync( + DeclarativeWorkflowContext.WorkflowSessionIdStateKey, + () => generatedWorkflowSessionId, + VariableScopeNames.System, + cancellationToken).ConfigureAwait(false); + return await scopedHandler.InvokeToolInWorkflowSessionAsync( + serverUrl, + serverLabel, + toolName, + arguments, + headers, + connectionName, + workflowSessionId, + cancellationToken).ConfigureAwait(false); + } + + return await mcpToolHandler.InvokeToolAsync( + serverUrl, + serverLabel, + toolName, + arguments, + headers, + connectionName, + cancellationToken).ConfigureAwait(false); + } + /// /// Completes the MCP tool invocation by raising the completion event. /// diff --git a/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative/PublicAPI/net10.0/PublicAPI.Unshipped.txt b/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative/PublicAPI/net10.0/PublicAPI.Unshipped.txt index 8af74cb2369..5bee0f03310 100644 --- a/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative/PublicAPI/net10.0/PublicAPI.Unshipped.txt +++ b/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative/PublicAPI/net10.0/PublicAPI.Unshipped.txt @@ -4,6 +4,8 @@ Microsoft.Agents.AI.Workflows.Declarative.DeclarativeWorkflowOptions.AllowedEnvi Microsoft.Agents.AI.Workflows.Declarative.DeclarativeWorkflowOptions.AllowProcessEnvironmentVariableFallback.get -> bool Microsoft.Agents.AI.Workflows.Declarative.DeclarativeWorkflowOptions.AllowProcessEnvironmentVariableFallback.init -> void Microsoft.Agents.AI.Workflows.Declarative.Events.ExternalInputResponse.RequestId.get -> string? +Microsoft.Agents.AI.Workflows.Declarative.IWorkflowScopedMcpToolHandler +Microsoft.Agents.AI.Workflows.Declarative.IWorkflowScopedMcpToolHandler.InvokeToolInWorkflowSessionAsync(string! serverUrl, string? serverLabel, string! toolName, System.Collections.Generic.IDictionary? arguments, System.Collections.Generic.IDictionary? headers, string? connectionName, string! workflowSessionId, System.Threading.CancellationToken cancellationToken = default(System.Threading.CancellationToken)) -> System.Threading.Tasks.Task! static Microsoft.Agents.AI.Workflows.Declarative.Kit.IWorkflowContextExtensions.ConvertValueWithSensitivityAsync(this Microsoft.Agents.AI.Workflows.IWorkflowContext! context, Microsoft.Agents.AI.Workflows.Declarative.Kit.VariableType! targetType, string! key, string? scopeName = null, System.Threading.CancellationToken cancellationToken = default(System.Threading.CancellationToken)) -> System.Threading.Tasks.ValueTask!> static Microsoft.Agents.AI.Workflows.Declarative.Kit.IWorkflowContextExtensions.EvaluateValueWithSensitivityAsync(this Microsoft.Agents.AI.Workflows.IWorkflowContext! context, string! expression, System.Threading.CancellationToken cancellationToken = default(System.Threading.CancellationToken)) -> System.Threading.Tasks.ValueTask!> static Microsoft.Agents.AI.Workflows.Declarative.Kit.IWorkflowContextExtensions.FormatTemplateWithSensitivityAsync(this Microsoft.Agents.AI.Workflows.IWorkflowContext! context, string! line) -> System.Threading.Tasks.ValueTask!> diff --git a/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative/PublicAPI/net472/PublicAPI.Unshipped.txt b/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative/PublicAPI/net472/PublicAPI.Unshipped.txt index 8af74cb2369..5bee0f03310 100644 --- a/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative/PublicAPI/net472/PublicAPI.Unshipped.txt +++ b/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative/PublicAPI/net472/PublicAPI.Unshipped.txt @@ -4,6 +4,8 @@ Microsoft.Agents.AI.Workflows.Declarative.DeclarativeWorkflowOptions.AllowedEnvi Microsoft.Agents.AI.Workflows.Declarative.DeclarativeWorkflowOptions.AllowProcessEnvironmentVariableFallback.get -> bool Microsoft.Agents.AI.Workflows.Declarative.DeclarativeWorkflowOptions.AllowProcessEnvironmentVariableFallback.init -> void Microsoft.Agents.AI.Workflows.Declarative.Events.ExternalInputResponse.RequestId.get -> string? +Microsoft.Agents.AI.Workflows.Declarative.IWorkflowScopedMcpToolHandler +Microsoft.Agents.AI.Workflows.Declarative.IWorkflowScopedMcpToolHandler.InvokeToolInWorkflowSessionAsync(string! serverUrl, string? serverLabel, string! toolName, System.Collections.Generic.IDictionary? arguments, System.Collections.Generic.IDictionary? headers, string? connectionName, string! workflowSessionId, System.Threading.CancellationToken cancellationToken = default(System.Threading.CancellationToken)) -> System.Threading.Tasks.Task! static Microsoft.Agents.AI.Workflows.Declarative.Kit.IWorkflowContextExtensions.ConvertValueWithSensitivityAsync(this Microsoft.Agents.AI.Workflows.IWorkflowContext! context, Microsoft.Agents.AI.Workflows.Declarative.Kit.VariableType! targetType, string! key, string? scopeName = null, System.Threading.CancellationToken cancellationToken = default(System.Threading.CancellationToken)) -> System.Threading.Tasks.ValueTask!> static Microsoft.Agents.AI.Workflows.Declarative.Kit.IWorkflowContextExtensions.EvaluateValueWithSensitivityAsync(this Microsoft.Agents.AI.Workflows.IWorkflowContext! context, string! expression, System.Threading.CancellationToken cancellationToken = default(System.Threading.CancellationToken)) -> System.Threading.Tasks.ValueTask!> static Microsoft.Agents.AI.Workflows.Declarative.Kit.IWorkflowContextExtensions.FormatTemplateWithSensitivityAsync(this Microsoft.Agents.AI.Workflows.IWorkflowContext! context, string! line) -> System.Threading.Tasks.ValueTask!> diff --git a/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative/PublicAPI/net8.0/PublicAPI.Unshipped.txt b/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative/PublicAPI/net8.0/PublicAPI.Unshipped.txt index 8af74cb2369..5bee0f03310 100644 --- a/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative/PublicAPI/net8.0/PublicAPI.Unshipped.txt +++ b/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative/PublicAPI/net8.0/PublicAPI.Unshipped.txt @@ -4,6 +4,8 @@ Microsoft.Agents.AI.Workflows.Declarative.DeclarativeWorkflowOptions.AllowedEnvi Microsoft.Agents.AI.Workflows.Declarative.DeclarativeWorkflowOptions.AllowProcessEnvironmentVariableFallback.get -> bool Microsoft.Agents.AI.Workflows.Declarative.DeclarativeWorkflowOptions.AllowProcessEnvironmentVariableFallback.init -> void Microsoft.Agents.AI.Workflows.Declarative.Events.ExternalInputResponse.RequestId.get -> string? +Microsoft.Agents.AI.Workflows.Declarative.IWorkflowScopedMcpToolHandler +Microsoft.Agents.AI.Workflows.Declarative.IWorkflowScopedMcpToolHandler.InvokeToolInWorkflowSessionAsync(string! serverUrl, string? serverLabel, string! toolName, System.Collections.Generic.IDictionary? arguments, System.Collections.Generic.IDictionary? headers, string? connectionName, string! workflowSessionId, System.Threading.CancellationToken cancellationToken = default(System.Threading.CancellationToken)) -> System.Threading.Tasks.Task! static Microsoft.Agents.AI.Workflows.Declarative.Kit.IWorkflowContextExtensions.ConvertValueWithSensitivityAsync(this Microsoft.Agents.AI.Workflows.IWorkflowContext! context, Microsoft.Agents.AI.Workflows.Declarative.Kit.VariableType! targetType, string! key, string? scopeName = null, System.Threading.CancellationToken cancellationToken = default(System.Threading.CancellationToken)) -> System.Threading.Tasks.ValueTask!> static Microsoft.Agents.AI.Workflows.Declarative.Kit.IWorkflowContextExtensions.EvaluateValueWithSensitivityAsync(this Microsoft.Agents.AI.Workflows.IWorkflowContext! context, string! expression, System.Threading.CancellationToken cancellationToken = default(System.Threading.CancellationToken)) -> System.Threading.Tasks.ValueTask!> static Microsoft.Agents.AI.Workflows.Declarative.Kit.IWorkflowContextExtensions.FormatTemplateWithSensitivityAsync(this Microsoft.Agents.AI.Workflows.IWorkflowContext! context, string! line) -> System.Threading.Tasks.ValueTask!> diff --git a/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative/PublicAPI/net9.0/PublicAPI.Unshipped.txt b/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative/PublicAPI/net9.0/PublicAPI.Unshipped.txt index 8af74cb2369..5bee0f03310 100644 --- a/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative/PublicAPI/net9.0/PublicAPI.Unshipped.txt +++ b/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative/PublicAPI/net9.0/PublicAPI.Unshipped.txt @@ -4,6 +4,8 @@ Microsoft.Agents.AI.Workflows.Declarative.DeclarativeWorkflowOptions.AllowedEnvi Microsoft.Agents.AI.Workflows.Declarative.DeclarativeWorkflowOptions.AllowProcessEnvironmentVariableFallback.get -> bool Microsoft.Agents.AI.Workflows.Declarative.DeclarativeWorkflowOptions.AllowProcessEnvironmentVariableFallback.init -> void Microsoft.Agents.AI.Workflows.Declarative.Events.ExternalInputResponse.RequestId.get -> string? +Microsoft.Agents.AI.Workflows.Declarative.IWorkflowScopedMcpToolHandler +Microsoft.Agents.AI.Workflows.Declarative.IWorkflowScopedMcpToolHandler.InvokeToolInWorkflowSessionAsync(string! serverUrl, string? serverLabel, string! toolName, System.Collections.Generic.IDictionary? arguments, System.Collections.Generic.IDictionary? headers, string? connectionName, string! workflowSessionId, System.Threading.CancellationToken cancellationToken = default(System.Threading.CancellationToken)) -> System.Threading.Tasks.Task! static Microsoft.Agents.AI.Workflows.Declarative.Kit.IWorkflowContextExtensions.ConvertValueWithSensitivityAsync(this Microsoft.Agents.AI.Workflows.IWorkflowContext! context, Microsoft.Agents.AI.Workflows.Declarative.Kit.VariableType! targetType, string! key, string? scopeName = null, System.Threading.CancellationToken cancellationToken = default(System.Threading.CancellationToken)) -> System.Threading.Tasks.ValueTask!> static Microsoft.Agents.AI.Workflows.Declarative.Kit.IWorkflowContextExtensions.EvaluateValueWithSensitivityAsync(this Microsoft.Agents.AI.Workflows.IWorkflowContext! context, string! expression, System.Threading.CancellationToken cancellationToken = default(System.Threading.CancellationToken)) -> System.Threading.Tasks.ValueTask!> static Microsoft.Agents.AI.Workflows.Declarative.Kit.IWorkflowContextExtensions.FormatTemplateWithSensitivityAsync(this Microsoft.Agents.AI.Workflows.IWorkflowContext! context, string! line) -> System.Threading.Tasks.ValueTask!> diff --git a/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative/PublicAPI/netstandard2.0/PublicAPI.Unshipped.txt b/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative/PublicAPI/netstandard2.0/PublicAPI.Unshipped.txt index 8af74cb2369..5bee0f03310 100644 --- a/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative/PublicAPI/netstandard2.0/PublicAPI.Unshipped.txt +++ b/dotnet/src/Microsoft.Agents.AI.Workflows.Declarative/PublicAPI/netstandard2.0/PublicAPI.Unshipped.txt @@ -4,6 +4,8 @@ Microsoft.Agents.AI.Workflows.Declarative.DeclarativeWorkflowOptions.AllowedEnvi Microsoft.Agents.AI.Workflows.Declarative.DeclarativeWorkflowOptions.AllowProcessEnvironmentVariableFallback.get -> bool Microsoft.Agents.AI.Workflows.Declarative.DeclarativeWorkflowOptions.AllowProcessEnvironmentVariableFallback.init -> void Microsoft.Agents.AI.Workflows.Declarative.Events.ExternalInputResponse.RequestId.get -> string? +Microsoft.Agents.AI.Workflows.Declarative.IWorkflowScopedMcpToolHandler +Microsoft.Agents.AI.Workflows.Declarative.IWorkflowScopedMcpToolHandler.InvokeToolInWorkflowSessionAsync(string! serverUrl, string? serverLabel, string! toolName, System.Collections.Generic.IDictionary? arguments, System.Collections.Generic.IDictionary? headers, string? connectionName, string! workflowSessionId, System.Threading.CancellationToken cancellationToken = default(System.Threading.CancellationToken)) -> System.Threading.Tasks.Task! static Microsoft.Agents.AI.Workflows.Declarative.Kit.IWorkflowContextExtensions.ConvertValueWithSensitivityAsync(this Microsoft.Agents.AI.Workflows.IWorkflowContext! context, Microsoft.Agents.AI.Workflows.Declarative.Kit.VariableType! targetType, string! key, string? scopeName = null, System.Threading.CancellationToken cancellationToken = default(System.Threading.CancellationToken)) -> System.Threading.Tasks.ValueTask!> static Microsoft.Agents.AI.Workflows.Declarative.Kit.IWorkflowContextExtensions.EvaluateValueWithSensitivityAsync(this Microsoft.Agents.AI.Workflows.IWorkflowContext! context, string! expression, System.Threading.CancellationToken cancellationToken = default(System.Threading.CancellationToken)) -> System.Threading.Tasks.ValueTask!> static Microsoft.Agents.AI.Workflows.Declarative.Kit.IWorkflowContextExtensions.FormatTemplateWithSensitivityAsync(this Microsoft.Agents.AI.Workflows.IWorkflowContext! context, string! line) -> System.Threading.Tasks.ValueTask!> diff --git a/dotnet/src/Microsoft.Agents.AI.Workflows/GroupChatManager.cs b/dotnet/src/Microsoft.Agents.AI.Workflows/GroupChatManager.cs index ab94a9fa8aa..fbda33c132a 100644 --- a/dotnet/src/Microsoft.Agents.AI.Workflows/GroupChatManager.cs +++ b/dotnet/src/Microsoft.Agents.AI.Workflows/GroupChatManager.cs @@ -152,7 +152,7 @@ protected virtual ValueTask OnCheckpointRestoredAsync(IWorkflowContext context, internal async ValueTask CheckpointAsync(IWorkflowContext context, CancellationToken cancellationToken = default) { await context.QueueStateUpdateAsync(BaseStateKey, new GroupChatManagerState(this.IterationCount), cancellationToken: cancellationToken).ConfigureAwait(false); - await this.OnCheckpointingAsync(new PrefixingWorkflowContext(context, SubclassStateKeyPrefix), cancellationToken).ConfigureAwait(false); + await this.OnCheckpointingAsync(PrefixingWorkflowContext.Create(context, SubclassStateKeyPrefix), cancellationToken).ConfigureAwait(false); } // Root restore entry point invoked by the hosting GroupChatHost. Symmetric to CheckpointAsync. @@ -160,7 +160,7 @@ internal async ValueTask RestoreCheckpointAsync(IWorkflowContext context, Cancel { GroupChatManagerState? state = await context.ReadStateAsync(BaseStateKey, cancellationToken: cancellationToken).ConfigureAwait(false); this.IterationCount = state?.IterationCount ?? 0; - await this.OnCheckpointRestoredAsync(new PrefixingWorkflowContext(context, SubclassStateKeyPrefix), cancellationToken).ConfigureAwait(false); + await this.OnCheckpointRestoredAsync(PrefixingWorkflowContext.Create(context, SubclassStateKeyPrefix), cancellationToken).ConfigureAwait(false); } } @@ -169,11 +169,16 @@ internal sealed record GroupChatManagerState(int IterationCount); // IWorkflowContext decorator that prepends a fixed prefix to every state key passed through it. // All non-state members (events, message sending, output yielding, halt requests, trace context, // and runtime characteristics) delegate directly to the wrapped context. -internal sealed class PrefixingWorkflowContext(IWorkflowContext inner, string prefix) : IWorkflowContext +internal class PrefixingWorkflowContext(IWorkflowContext inner, string prefix) : IWorkflowContext { private readonly IWorkflowContext _inner = Throw.IfNull(inner); private readonly string _prefix = Throw.IfNullOrEmpty(prefix); + public static IWorkflowContext Create(IWorkflowContext inner, string prefix) => + inner is IWorkflowSessionContext sessionContext + ? new PrefixingWorkflowSessionContext(inner, prefix, sessionContext) + : new PrefixingWorkflowContext(inner, prefix); + public IReadOnlyDictionary? TraceContext => this._inner.TraceContext; public bool ConcurrentRunsEnabled => this._inner.ConcurrentRunsEnabled; @@ -222,3 +227,11 @@ public async ValueTask QueueClearScopeAsync(string? scopeName = null, Cancellati private string Wrap(string key) => this._prefix + Throw.IfNullOrEmpty(key); } + +internal sealed class PrefixingWorkflowSessionContext( + IWorkflowContext inner, + string prefix, + IWorkflowSessionContext sessionContext) : PrefixingWorkflowContext(inner, prefix), IWorkflowSessionContext +{ + public string SessionId => sessionContext.SessionId; +} diff --git a/dotnet/src/Microsoft.Agents.AI.Workflows/IWorkflowContext.cs b/dotnet/src/Microsoft.Agents.AI.Workflows/IWorkflowContext.cs index b8b35fffd64..68ff864085f 100644 --- a/dotnet/src/Microsoft.Agents.AI.Workflows/IWorkflowContext.cs +++ b/dotnet/src/Microsoft.Agents.AI.Workflows/IWorkflowContext.cs @@ -195,3 +195,18 @@ ValueTask ReadOrInitStateAsync(string key, Func initialStateFactory, Ca /// bool ConcurrentRunsEnabled { get; } } + +/// +/// Exposes the framework-owned identifier for a workflow session. +/// +/// +/// Runtime services can use this identifier to scope stateful resources to one +/// workflow session without accepting caller-provided cache partition keys. +/// +public interface IWorkflowSessionContext +{ + /// + /// Gets the framework-owned workflow session identifier. + /// + string SessionId { get; } +} diff --git a/dotnet/src/Microsoft.Agents.AI.Workflows/InProc/InProcessRunner.cs b/dotnet/src/Microsoft.Agents.AI.Workflows/InProc/InProcessRunner.cs index 69b6c1e9bcc..c0a2e075ae1 100644 --- a/dotnet/src/Microsoft.Agents.AI.Workflows/InProc/InProcessRunner.cs +++ b/dotnet/src/Microsoft.Agents.AI.Workflows/InProc/InProcessRunner.cs @@ -31,7 +31,14 @@ public static InProcessRunner CreateTopLevelRunner(Workflow workflow, ICheckpoin knownValidInputTypes: knownValidInputTypes); } - public static InProcessRunner CreateSubworkflowRunner(Workflow workflow, ICheckpointManager? checkpointManager, string? sessionId = null, object? existingOwnerSignoff = null, bool enableConcurrentRuns = false, IEnumerable? knownValidInputTypes = null) + public static InProcessRunner CreateSubworkflowRunner( + Workflow workflow, + ICheckpointManager? checkpointManager, + string? sessionId = null, + object? existingOwnerSignoff = null, + bool enableConcurrentRuns = false, + IEnumerable? knownValidInputTypes = null, + string? workflowSessionId = null) { return new InProcessRunner(workflow, checkpointManager, @@ -39,10 +46,19 @@ public static InProcessRunner CreateSubworkflowRunner(Workflow workflow, ICheckp existingOwnerSignoff: existingOwnerSignoff, enableConcurrentRuns: enableConcurrentRuns, knownValidInputTypes: knownValidInputTypes, - subworkflow: true); + subworkflow: true, + workflowSessionId: workflowSessionId); } - private InProcessRunner(Workflow workflow, ICheckpointManager? checkpointManager, string? sessionId = null, object? existingOwnerSignoff = null, bool subworkflow = false, bool enableConcurrentRuns = false, IEnumerable? knownValidInputTypes = null) + private InProcessRunner( + Workflow workflow, + ICheckpointManager? checkpointManager, + string? sessionId = null, + object? existingOwnerSignoff = null, + bool subworkflow = false, + bool enableConcurrentRuns = false, + IEnumerable? knownValidInputTypes = null, + string? workflowSessionId = null) { if (enableConcurrentRuns && !workflow.AllowConcurrent) { @@ -54,7 +70,16 @@ private InProcessRunner(Workflow workflow, ICheckpointManager? checkpointManager this.StartExecutorId = workflow.StartExecutorId; this.Workflow = Throw.IfNull(workflow); - this.RunContext = new InProcessRunnerContext(workflow, this.SessionId, checkpointingEnabled: checkpointManager != null, this.OutgoingEvents, this.StepTracer, existingOwnerSignoff, subworkflow, enableConcurrentRuns); + this.RunContext = new InProcessRunnerContext( + workflow, + this.SessionId, + checkpointingEnabled: checkpointManager != null, + this.OutgoingEvents, + this.StepTracer, + existingOwnerSignoff, + subworkflow, + enableConcurrentRuns, + workflowSessionId: workflowSessionId); this.CheckpointManager = checkpointManager; this._knownValidInputTypes = knownValidInputTypes != null @@ -65,6 +90,8 @@ private InProcessRunner(Workflow workflow, ICheckpointManager? checkpointManager /// public string SessionId { get; } + internal string WorkflowSessionId => this.RunContext.WorkflowSessionId; + /// public string StartExecutorId { get; } diff --git a/dotnet/src/Microsoft.Agents.AI.Workflows/InProc/InProcessRunnerContext.cs b/dotnet/src/Microsoft.Agents.AI.Workflows/InProc/InProcessRunnerContext.cs index 353c38f25c9..aa8bf398592 100644 --- a/dotnet/src/Microsoft.Agents.AI.Workflows/InProc/InProcessRunnerContext.cs +++ b/dotnet/src/Microsoft.Agents.AI.Workflows/InProc/InProcessRunnerContext.cs @@ -47,7 +47,8 @@ public InProcessRunnerContext( object? existingOwnershipSignoff = null, bool subworkflow = false, bool enableConcurrentRuns = false, - ILogger? logger = null) + ILogger? logger = null, + string? workflowSessionId = null) { if (enableConcurrentRuns) { @@ -62,6 +63,7 @@ public InProcessRunnerContext( this._workflow = workflow; this._sessionId = sessionId; + this.WorkflowSessionId = workflowSessionId ?? sessionId; this._edgeMap = new(this, this._workflow, stepTracer); this._outputFilter = new(workflow); @@ -72,6 +74,8 @@ public InProcessRunnerContext( } public WorkflowTelemetryContext TelemetryContext => this._workflow.TelemetryContext; + internal string WorkflowSessionId { get; } + public IExternalRequestSink RegisterPort(string executorId, RequestPort port) { if (!this._edgeMap.TryRegisterPort(this, executorId, port)) @@ -349,8 +353,10 @@ public IExternalRequestSink RegisterPort(RequestPort port) private sealed class BoundWorkflowContext( InProcessRunnerContext RunnerContext, string ExecutorId, - Dictionary? traceContext) : IWorkflowContext + Dictionary? traceContext) : IWorkflowContext, IWorkflowSessionContext { + public string SessionId => RunnerContext.WorkflowSessionId; + public ValueTask AddEventAsync(WorkflowEvent workflowEvent, CancellationToken cancellationToken = default) => RunnerContext.AddEventAsync(workflowEvent, cancellationToken); public ValueTask SendMessageAsync(object message, string? targetId = null, CancellationToken cancellationToken = default) diff --git a/dotnet/src/Microsoft.Agents.AI.Workflows/PublicAPI/net10.0/PublicAPI.Unshipped.txt b/dotnet/src/Microsoft.Agents.AI.Workflows/PublicAPI/net10.0/PublicAPI.Unshipped.txt index ab058de62d4..724a27a74a4 100644 --- a/dotnet/src/Microsoft.Agents.AI.Workflows/PublicAPI/net10.0/PublicAPI.Unshipped.txt +++ b/dotnet/src/Microsoft.Agents.AI.Workflows/PublicAPI/net10.0/PublicAPI.Unshipped.txt @@ -1 +1,3 @@ #nullable enable +Microsoft.Agents.AI.Workflows.IWorkflowSessionContext +Microsoft.Agents.AI.Workflows.IWorkflowSessionContext.SessionId.get -> string! diff --git a/dotnet/src/Microsoft.Agents.AI.Workflows/PublicAPI/net472/PublicAPI.Unshipped.txt b/dotnet/src/Microsoft.Agents.AI.Workflows/PublicAPI/net472/PublicAPI.Unshipped.txt index ab058de62d4..724a27a74a4 100644 --- a/dotnet/src/Microsoft.Agents.AI.Workflows/PublicAPI/net472/PublicAPI.Unshipped.txt +++ b/dotnet/src/Microsoft.Agents.AI.Workflows/PublicAPI/net472/PublicAPI.Unshipped.txt @@ -1 +1,3 @@ #nullable enable +Microsoft.Agents.AI.Workflows.IWorkflowSessionContext +Microsoft.Agents.AI.Workflows.IWorkflowSessionContext.SessionId.get -> string! diff --git a/dotnet/src/Microsoft.Agents.AI.Workflows/PublicAPI/net8.0/PublicAPI.Unshipped.txt b/dotnet/src/Microsoft.Agents.AI.Workflows/PublicAPI/net8.0/PublicAPI.Unshipped.txt index ab058de62d4..724a27a74a4 100644 --- a/dotnet/src/Microsoft.Agents.AI.Workflows/PublicAPI/net8.0/PublicAPI.Unshipped.txt +++ b/dotnet/src/Microsoft.Agents.AI.Workflows/PublicAPI/net8.0/PublicAPI.Unshipped.txt @@ -1 +1,3 @@ #nullable enable +Microsoft.Agents.AI.Workflows.IWorkflowSessionContext +Microsoft.Agents.AI.Workflows.IWorkflowSessionContext.SessionId.get -> string! diff --git a/dotnet/src/Microsoft.Agents.AI.Workflows/PublicAPI/net9.0/PublicAPI.Unshipped.txt b/dotnet/src/Microsoft.Agents.AI.Workflows/PublicAPI/net9.0/PublicAPI.Unshipped.txt index ab058de62d4..724a27a74a4 100644 --- a/dotnet/src/Microsoft.Agents.AI.Workflows/PublicAPI/net9.0/PublicAPI.Unshipped.txt +++ b/dotnet/src/Microsoft.Agents.AI.Workflows/PublicAPI/net9.0/PublicAPI.Unshipped.txt @@ -1 +1,3 @@ #nullable enable +Microsoft.Agents.AI.Workflows.IWorkflowSessionContext +Microsoft.Agents.AI.Workflows.IWorkflowSessionContext.SessionId.get -> string! diff --git a/dotnet/src/Microsoft.Agents.AI.Workflows/PublicAPI/netstandard2.0/PublicAPI.Unshipped.txt b/dotnet/src/Microsoft.Agents.AI.Workflows/PublicAPI/netstandard2.0/PublicAPI.Unshipped.txt index ab058de62d4..724a27a74a4 100644 --- a/dotnet/src/Microsoft.Agents.AI.Workflows/PublicAPI/netstandard2.0/PublicAPI.Unshipped.txt +++ b/dotnet/src/Microsoft.Agents.AI.Workflows/PublicAPI/netstandard2.0/PublicAPI.Unshipped.txt @@ -1 +1,3 @@ #nullable enable +Microsoft.Agents.AI.Workflows.IWorkflowSessionContext +Microsoft.Agents.AI.Workflows.IWorkflowSessionContext.SessionId.get -> string! diff --git a/dotnet/src/Microsoft.Agents.AI.Workflows/Specialized/WorkflowHostExecutor.cs b/dotnet/src/Microsoft.Agents.AI.Workflows/Specialized/WorkflowHostExecutor.cs index 58e3a9e5239..a166d8b0e6b 100644 --- a/dotnet/src/Microsoft.Agents.AI.Workflows/Specialized/WorkflowHostExecutor.cs +++ b/dotnet/src/Microsoft.Agents.AI.Workflows/Specialized/WorkflowHostExecutor.cs @@ -20,6 +20,7 @@ internal class WorkflowHostExecutor : Executor, IAsyncDisposable private readonly Workflow _workflow; private readonly ProtocolDescriptor _workflowProtocol; private readonly object _ownershipToken; + private string? _workflowSessionId; private InProcessRunner? _activeRunner; private InMemoryCheckpointManager? _checkpointManager; @@ -56,6 +57,7 @@ protected override ProtocolBuilder ConfigureProtocol(ProtocolBuilder protocolBui private async ValueTask QueueExternalMessageAsync(PortableValue portableValue, IWorkflowContext context, CancellationToken cancellationToken) { + this.SetWorkflowSessionId(context); if (portableValue.Is(out ExternalResponse? response)) { response = this.CheckAndUnqualifyResponse(response); @@ -95,7 +97,8 @@ internal async ValueTask EnsureRunnerAsync() this._checkpointManager, this._sessionId, this._ownershipToken, - this.JoinContext.ConcurrentRunsEnabled); + this.JoinContext.ConcurrentRunsEnabled, + workflowSessionId: this._workflowSessionId); } return this._activeRunner; @@ -268,6 +271,7 @@ await context.QueueStateUpdateAsync(PendingResponsePortsStateKey, protected internal override async ValueTask OnCheckpointRestoredAsync(IWorkflowContext context, CancellationToken cancellationToken = default) { await base.OnCheckpointRestoredAsync(context, cancellationToken).ConfigureAwait(false); + this.SetWorkflowSessionId(context); InMemoryCheckpointManager manager = await context.ReadStateAsync(CheckpointManagerStateKey, cancellationToken: cancellationToken).ConfigureAwait(false) ?? new(); if (this._checkpointManager == manager) @@ -293,6 +297,14 @@ await context.ReadStateAsync>(PendingRespons await this.EnsureRunSendMessageAsync(resume: true, cancellationToken: cancellationToken).ConfigureAwait(false); } + private void SetWorkflowSessionId(IWorkflowContext context) + { + string parentSessionId = context is IWorkflowSessionContext sessionContext + ? sessionContext.SessionId + : this._sessionId; + this._workflowSessionId ??= SubworkflowBinding.CreateSubworkflowSessionId(parentSessionId, this.Id); + } + private async ValueTask ResetAsync() { if (this._run != null) diff --git a/dotnet/src/Microsoft.Agents.AI.Workflows/SubworkflowBinding.cs b/dotnet/src/Microsoft.Agents.AI.Workflows/SubworkflowBinding.cs index 11f7cf493cf..9519af12414 100644 --- a/dotnet/src/Microsoft.Agents.AI.Workflows/SubworkflowBinding.cs +++ b/dotnet/src/Microsoft.Agents.AI.Workflows/SubworkflowBinding.cs @@ -1,6 +1,7 @@ // Copyright (c) Microsoft. All rights reserved. using System; +using System.Globalization; using System.Threading.Tasks; using Microsoft.Agents.AI.Workflows.Specialized; using Microsoft.Shared.Diagnostics; @@ -35,6 +36,13 @@ async ValueTask InitHostExecutorAsync(string sessionId) } } + internal static string CreateSubworkflowSessionId(string parentSessionId, string subworkflowId) => + string.Concat( + parentSessionId.Length.ToString(CultureInfo.InvariantCulture), + ":", + parentSessionId, + subworkflowId); + /// public override bool IsSharedInstance => false; diff --git a/dotnet/tests/Microsoft.Agents.AI.Workflows.Declarative.Mcp.UnitTests/DefaultMcpToolHandlerLifetimeTests.cs b/dotnet/tests/Microsoft.Agents.AI.Workflows.Declarative.Mcp.UnitTests/DefaultMcpToolHandlerLifetimeTests.cs index 3a481fcd0e6..1e313753575 100644 --- a/dotnet/tests/Microsoft.Agents.AI.Workflows.Declarative.Mcp.UnitTests/DefaultMcpToolHandlerLifetimeTests.cs +++ b/dotnet/tests/Microsoft.Agents.AI.Workflows.Declarative.Mcp.UnitTests/DefaultMcpToolHandlerLifetimeTests.cs @@ -104,6 +104,201 @@ public async Task NoProvider_ReusesSessionUntilHandlerDisposalAsync() transport.Protected().Verify("Dispose", Times.AtLeastOnce(), ItExpr.Is(disposing => disposing)); } + [Fact] + public async Task NoProvider_SeparateWorkflowSessions_UseSeparateCachedSessionsAsync() + { + // Arrange + ProtocolStub stub = new(); + await using DefaultMcpToolHandler handler = new(null, stub.CreateMessageHandler); + using CancellationTokenSource timeout = new(TimeSpan.FromSeconds(10)); + + // Act + await InvokeScopedAsync(handler, "workflow-a", "ping", timeout.Token); + await InvokeScopedAsync(handler, "workflow-b", "ping", timeout.Token); + await InvokeScopedAsync(handler, "workflow-a", "ping", timeout.Token); + + // Assert + Assert.Equal(2, stub.Initializations); + Assert.Equal(0, stub.Terminations); + } + + [Fact] + public async Task DrainCleanupAsync_Cancellation_DrainsAllCleanupAndCompletionTasksAsync() + { + // Arrange + TaskCompletionSource firstDisposed = new(TaskCreationOptions.RunContinuationsAsynchronously); + TaskCompletionSource secondDisposed = new(TaskCreationOptions.RunContinuationsAsynchronously); + + async Task CancelledCleanupAsync() + { + try + { + await Task.Yield(); + throw new OperationCanceledException("session cleanup cancelled"); + } + finally + { + firstDisposed.TrySetResult(true); + } + } + + async Task SuccessfulCleanupAsync() + { + await Task.Yield(); + secondDisposed.TrySetResult(true); + } + + // Act + await Assert.ThrowsAnyAsync(() => DefaultMcpToolHandler.DrainCleanupAsync( + [CancelledCleanupAsync(), SuccessfulCleanupAsync()], + [firstDisposed.Task, secondDisposed.Task])); + + // Assert + Assert.True(await firstDisposed.Task); + Assert.True(await secondDisposed.Task); + } + + [Fact] + public async Task NoProvider_CacheEviction_DisposesLeastRecentlyUsedSessionAsync() + { + // Arrange + ProtocolStub stub = new(); + await using DefaultMcpToolHandler handler = new(null, stub.CreateMessageHandler, clientCacheMaxSize: 2); + using CancellationTokenSource timeout = new(TimeSpan.FromSeconds(10)); + + // Act + await InvokeScopedAsync(handler, "workflow-a", "ping", timeout.Token); + await InvokeScopedAsync(handler, "workflow-b", "ping", timeout.Token); + await InvokeScopedAsync(handler, "workflow-c", "ping", timeout.Token); + + // Assert + Assert.Equal(3, stub.Initializations); + Assert.Equal(1, stub.Terminations); + } + + [Fact] + public async Task NoProvider_CacheEviction_DisposesUnusedOwnedHttpClientAsync() + { + // Arrange + ProtocolStub stub = new(); + await using DefaultMcpToolHandler handler = new(null, stub.CreateMessageHandler, clientCacheMaxSize: 1); + using CancellationTokenSource timeout = new(TimeSpan.FromSeconds(10)); + + // Act + await handler.InvokeToolInWorkflowSessionAsync( + "https://first.example/api", null, "ping", null, null, null, "workflow-a", timeout.Token); + await handler.InvokeToolInWorkflowSessionAsync( + "https://second.example/api", null, "ping", null, null, null, "workflow-b", timeout.Token); + + // Assert + Assert.Equal(2, stub.Initializations); + Assert.Equal(1, stub.Terminations); + Assert.Equal(2, stub.Handlers.Count); + stub.Handlers[0].Protected().Verify( + "Dispose", Times.AtLeastOnce(), ItExpr.Is(disposing => disposing)); + stub.Handlers[1].Protected().Verify( + "Dispose", Times.Never(), ItExpr.Is(disposing => disposing)); + } + + [Fact] + public async Task NoProvider_CacheEviction_DefersDisposalUntilActiveInvocationCompletesAsync() + { + // Arrange + ProtocolStub stub = new(); + using SemaphoreSlim firstStarted = new(0); + using SemaphoreSlim releaseFirst = new(0); + int operations = 0; + stub.BeforeOperationAsync = async token => + { + if (Interlocked.Increment(ref operations) == 1) + { + firstStarted.Release(); + await releaseFirst.WaitAsync(token); + } + }; + DefaultMcpToolHandler handler = new(null, stub.CreateMessageHandler, clientCacheMaxSize: 1); + using CancellationTokenSource timeout = new(TimeSpan.FromSeconds(10)); + + // Act + Task first = InvokeScopedAsync(handler, "workflow-a", "ping", timeout.Token); + await firstStarted.WaitAsync(timeout.Token); + await InvokeScopedAsync(handler, "workflow-b", "ping", timeout.Token); + + // Assert + Assert.Equal(0, stub.Terminations); + Task disposal = handler.DisposeAsync().AsTask(); + Assert.False(disposal.IsCompleted); + releaseFirst.Release(); + await first; + await disposal; + Assert.Equal(2, stub.Terminations); + } + + [Fact] + public async Task NoProvider_ConcurrentWorkflowSessionCreations_AreBoundedByCacheSizeAsync() + { + // Arrange + ProtocolStub stub = new(); + using SemaphoreSlim initializationsStarted = new(0); + using SemaphoreSlim releaseInitializations = new(0); + stub.BeforeInitializationAsync = async token => + { + initializationsStarted.Release(); + await releaseInitializations.WaitAsync(token); + }; + await using DefaultMcpToolHandler handler = new(null, stub.CreateMessageHandler, clientCacheMaxSize: 2); + using CancellationTokenSource timeout = new(TimeSpan.FromSeconds(10)); + using CancellationTokenSource concurrencyTimeout = new(TimeSpan.FromSeconds(2)); + + // Act + Task first = InvokeScopedAsync(handler, "workflow-a", "ping", timeout.Token); + await initializationsStarted.WaitAsync(timeout.Token); + Task second = InvokeScopedAsync(handler, "workflow-b", "ping", timeout.Token); + await initializationsStarted.WaitAsync(concurrencyTimeout.Token); + Task third = InvokeScopedAsync(handler, "workflow-c", "ping", timeout.Token); + using CancellationTokenSource gateTimeout = new(TimeSpan.FromMilliseconds(250)); + await Assert.ThrowsAnyAsync(() => initializationsStarted.WaitAsync(gateTimeout.Token)); + releaseInitializations.Release(2); + await Task.WhenAll(first, second); + await initializationsStarted.WaitAsync(timeout.Token); + releaseInitializations.Release(); + await third; + + // Assert + Assert.Equal(3, stub.Initializations); + Assert.Equal(1, stub.Terminations); + } + + [Fact] + public async Task NoProvider_DisposalDuringFailedCreation_PreservesCreationFailureAsync() + { + // Arrange + ProtocolStub stub = new() { FailInitialization = true }; + using SemaphoreSlim initializationStarted = new(0); + using SemaphoreSlim releaseInitialization = new(0); + stub.BeforeInitializationAsync = async token => + { + initializationStarted.Release(); + await releaseInitialization.WaitAsync(token); + }; + DefaultMcpToolHandler handler = new(null, stub.CreateMessageHandler); + using CancellationTokenSource timeout = new(TimeSpan.FromSeconds(10)); + + // Act + Task invocation = + InvokeScopedAsync(handler, "workflow-a", "ping", timeout.Token); + await initializationStarted.WaitAsync(timeout.Token); + Task disposal = handler.DisposeAsync().AsTask(); + releaseInitialization.Release(); + + // Assert + await Assert.ThrowsAsync(() => invocation); + await disposal; + Assert.Single(stub.Handlers); + stub.Handlers[0].Protected().Verify( + "Dispose", Times.AtLeastOnce(), ItExpr.Is(disposing => disposing)); + } + [Fact] public async Task NoProvider_DifferentConnectionNames_UseSeparateCachedSessionsAsync() { @@ -185,6 +380,91 @@ public async Task Provider_ConcurrentInvocations_DoNotCoalesceOrSerializeSession Assert.Equal(2, stub.Terminations); } + [Fact] + public async Task NoProvider_ConcurrentSameWorkflowInvocations_CoalesceSessionCreationAsync() + { + // Arrange + ProtocolStub stub = new(); + await using DefaultMcpToolHandler handler = new(null, stub.CreateMessageHandler); + using CancellationTokenSource timeout = new(TimeSpan.FromSeconds(10)); + + // Act + await Task.WhenAll( + InvokeScopedAsync(handler, "workflow-a", "ping", timeout.Token), + InvokeScopedAsync(handler, "workflow-a", "ping", timeout.Token)); + + // Assert + Assert.Equal(1, stub.Initializations); + } + + [Fact] + public async Task NoProvider_CancelledWaiter_DoesNotWaitForSharedCreationAsync() + { + // Arrange + ProtocolStub stub = new(); + using SemaphoreSlim initializationStarted = new(0); + using SemaphoreSlim releaseInitialization = new(0); + stub.BeforeInitializationAsync = async token => + { + initializationStarted.Release(); + await releaseInitialization.WaitAsync(token); + }; + await using DefaultMcpToolHandler handler = new(null, stub.CreateMessageHandler); + using CancellationTokenSource timeout = new(TimeSpan.FromSeconds(10)); + using CancellationTokenSource waiterCancellation = new(); + + // Act + Task creator = InvokeScopedAsync(handler, "workflow-a", "ping", timeout.Token); + await initializationStarted.WaitAsync(timeout.Token); + Task waiter = InvokeScopedAsync(handler, "workflow-a", "ping", waiterCancellation.Token); + await Task.Yield(); + waiterCancellation.Cancel(); + Task completed = await Task.WhenAny(waiter, Task.Delay(TimeSpan.FromSeconds(1), timeout.Token)); + releaseInitialization.Release(); + await creator; + + // Assert + Assert.Same(waiter, completed); + await Assert.ThrowsAnyAsync(() => waiter); + Assert.Equal(1, stub.Initializations); + } + + [Fact] + public async Task NoProvider_CancelledCreator_DoesNotCancelSharedWaiterAsync() + { + // Arrange + ProtocolStub stub = new(); + using SemaphoreSlim initializationStarted = new(0); + int initializationAttempts = 0; + stub.BeforeInitializationAsync = async token => + { + initializationStarted.Release(); + if (Interlocked.Increment(ref initializationAttempts) == 1) + { + await Task.Delay(Timeout.Infinite, token); + } + }; + await using DefaultMcpToolHandler handler = new(null, stub.CreateMessageHandler); + using CancellationTokenSource creatorCancellation = new(); + using CancellationTokenSource timeout = new(TimeSpan.FromSeconds(10)); + + // Act + Task creator = + InvokeScopedAsync(handler, "workflow-a", "ping", creatorCancellation.Token); + await initializationStarted.WaitAsync(timeout.Token); + Task waiter = InvokeScopedAsync(handler, "workflow-a", "ping", timeout.Token); + await Task.Yield(); + creatorCancellation.Cancel(); + await Assert.ThrowsAnyAsync(() => creator); + await initializationStarted.WaitAsync(timeout.Token); + McpServerToolResultContent result = await waiter; + + // Assert + Assert.NotNull(result); + Assert.Equal(2, initializationAttempts); + Assert.Equal(1, stub.Initializations); + } + [Fact] public async Task Provider_OperationFailure_DisposesSessionAndPreservesCallerClientAsync() { @@ -560,6 +840,11 @@ private static Task InvokeAsync( DefaultMcpToolHandler handler, string toolName, CancellationToken cancellationToken) => handler.InvokeToolAsync("https://mcp.example/api", null, toolName, null, null, null, cancellationToken); + private static Task InvokeScopedAsync( + DefaultMcpToolHandler handler, string workflowSessionId, string toolName, CancellationToken cancellationToken) => + handler.InvokeToolInWorkflowSessionAsync( + "https://mcp.example/api", null, toolName, null, null, null, workflowSessionId, cancellationToken); + private sealed class ProtocolStub { private int _initializations; @@ -569,6 +854,7 @@ private sealed class ProtocolStub public int Terminations => this._terminations; public List> Handlers { get; } = []; public Func? BeforeOperationAsync { get; set; } + public Func? BeforeInitializationAsync { get; set; } public bool FailInitialization { get; set; } public bool FailOperation { get; set; } public bool FailTransportDisposal { get; set; } @@ -629,6 +915,11 @@ private async Task SendAsync(HttpRequestMessage request, Ca string? sessionId = null; if (method == "initialize") { + if (this.BeforeInitializationAsync is not null) + { + await this.BeforeInitializationAsync(cancellationToken); + } + if (this.FailInitialization) { return EmptyResponse(HttpStatusCode.BadRequest, request); diff --git a/dotnet/tests/Microsoft.Agents.AI.Workflows.Declarative.Mcp.UnitTests/DefaultMcpToolHandlerTests.cs b/dotnet/tests/Microsoft.Agents.AI.Workflows.Declarative.Mcp.UnitTests/DefaultMcpToolHandlerTests.cs index 610f7393eda..a6e57b63193 100644 --- a/dotnet/tests/Microsoft.Agents.AI.Workflows.Declarative.Mcp.UnitTests/DefaultMcpToolHandlerTests.cs +++ b/dotnet/tests/Microsoft.Agents.AI.Workflows.Declarative.Mcp.UnitTests/DefaultMcpToolHandlerTests.cs @@ -418,7 +418,7 @@ public void ComputeHeadersHash_DifferentHeaders_ReturnsDifferentHash() // (InvokeToolAsync against a fake server) doesn't surface cache-hit behavior // without standing up a real MCP server — McpClient.CreateAsync fails before // _clients[key] = newClient runs, so nothing ever gets cached. - // Tuple equality on the returned 4-tuple verifies that the dimensions + // Tuple equality on the returned 5-tuple verifies that the dimensions // collectively discriminate cache entries. [Fact] @@ -428,19 +428,32 @@ public void BuildCacheKey_SameInputs_ReturnsEqualKeys() Dictionary headers = new() { ["Authorization"] = "Bearer token" }; // Act - var key1 = DefaultMcpToolHandler.BuildCacheKey("http://localhost/mcp", "label", "conn", headers); - var key2 = DefaultMcpToolHandler.BuildCacheKey("http://localhost/mcp", "label", "conn", headers); + var key1 = DefaultMcpToolHandler.BuildCacheKey("workflow", "http://localhost/mcp", "label", "conn", headers); + var key2 = DefaultMcpToolHandler.BuildCacheKey("workflow", "http://localhost/mcp", "label", "conn", headers); // Assert Assert.Equal(key2, key1); } + [Fact] + public void BuildCacheKey_DifferentWorkflowSession_ReturnsDifferentKeys() + { + // Act + var key1 = DefaultMcpToolHandler.BuildCacheKey("workflow-a", "http://localhost/mcp", "label", "conn", null); + var key2 = DefaultMcpToolHandler.BuildCacheKey("workflow-b", "http://localhost/mcp", "label", "conn", null); + + // Assert + Assert.NotEqual(key2, key1); + Assert.Equal("workflow-a", key1.WorkflowSession); + Assert.Equal("workflow-b", key2.WorkflowSession); + } + [Fact] public void BuildCacheKey_DifferentConnectionName_ReturnsDifferentKeys() { // Act - var key1 = DefaultMcpToolHandler.BuildCacheKey("http://localhost/mcp", "label", "connection-a", null); - var key2 = DefaultMcpToolHandler.BuildCacheKey("http://localhost/mcp", "label", "connection-b", null); + var key1 = DefaultMcpToolHandler.BuildCacheKey("workflow", "http://localhost/mcp", "label", "connection-a", null); + var key2 = DefaultMcpToolHandler.BuildCacheKey("workflow", "http://localhost/mcp", "label", "connection-b", null); // Assert Assert.NotEqual(key2, key1); @@ -452,8 +465,8 @@ public void BuildCacheKey_DifferentConnectionName_ReturnsDifferentKeys() public void BuildCacheKey_DifferentServerLabel_ReturnsDifferentKeys() { // Act - var key1 = DefaultMcpToolHandler.BuildCacheKey("http://localhost/mcp", "label-a", null, null); - var key2 = DefaultMcpToolHandler.BuildCacheKey("http://localhost/mcp", "label-b", null, null); + var key1 = DefaultMcpToolHandler.BuildCacheKey("workflow", "http://localhost/mcp", "label-a", null, null); + var key2 = DefaultMcpToolHandler.BuildCacheKey("workflow", "http://localhost/mcp", "label-b", null, null); // Assert Assert.NotEqual(key2, key1); @@ -466,8 +479,8 @@ public void BuildCacheKey_CaseSensitiveUrlPath_ReturnsDifferentKeys() { // Arrange — RFC 3986: URL path is case-sensitive // Act - var key1 = DefaultMcpToolHandler.BuildCacheKey("http://localhost/Tools", null, null, null); - var key2 = DefaultMcpToolHandler.BuildCacheKey("http://localhost/tools", null, null, null); + var key1 = DefaultMcpToolHandler.BuildCacheKey("workflow", "http://localhost/Tools", null, null, null); + var key2 = DefaultMcpToolHandler.BuildCacheKey("workflow", "http://localhost/tools", null, null, null); // Assert Assert.NotEqual(key2, key1); @@ -481,8 +494,8 @@ public void BuildCacheKey_HeaderValuesCaseSensitive_ReturnsDifferentKeys() Dictionary headers2 = new() { ["Authorization"] = "Bearer abc" }; // Act - var key1 = DefaultMcpToolHandler.BuildCacheKey("http://localhost/mcp", null, null, headers1); - var key2 = DefaultMcpToolHandler.BuildCacheKey("http://localhost/mcp", null, null, headers2); + var key1 = DefaultMcpToolHandler.BuildCacheKey("workflow", "http://localhost/mcp", null, null, headers1); + var key2 = DefaultMcpToolHandler.BuildCacheKey("workflow", "http://localhost/mcp", null, null, headers2); // Assert — header value case must propagate into the cache key Assert.NotEqual(key2, key1); @@ -493,7 +506,7 @@ public void BuildCacheKey_HeaderValuesCaseSensitive_ReturnsDifferentKeys() public void BuildCacheKey_NullLabelAndConnection_NormalizesToEmptyString() { // Act - var key = DefaultMcpToolHandler.BuildCacheKey("http://localhost/mcp", null, null, null); + var key = DefaultMcpToolHandler.BuildCacheKey("workflow", "http://localhost/mcp", null, null, null); // Assert — verifies null-safety contract callers rely on Assert.Equal(string.Empty, key.Label); diff --git a/dotnet/tests/Microsoft.Agents.AI.Workflows.Declarative.UnitTests/Kit/IWorkflowContextExtensionsTests.cs b/dotnet/tests/Microsoft.Agents.AI.Workflows.Declarative.UnitTests/Kit/IWorkflowContextExtensionsTests.cs index d26a1b6a84b..802f71cf48d 100644 --- a/dotnet/tests/Microsoft.Agents.AI.Workflows.Declarative.UnitTests/Kit/IWorkflowContextExtensionsTests.cs +++ b/dotnet/tests/Microsoft.Agents.AI.Workflows.Declarative.UnitTests/Kit/IWorkflowContextExtensionsTests.cs @@ -23,7 +23,7 @@ public async Task FormatTemplateAsync_WithSensitiveValue_ThrowsAsync() WorkflowFormulaState state = new(RecalcEngineFactory.Create()); state.Set("SOME_SECRET", FormulaValue.New("secret-value"), VariableScopeNames.Environment, SensitivityLevel.Sensitive); state.Bind(); - DeclarativeWorkflowContext context = new(new Mock().Object, state); + DeclarativeWorkflowContext context = await CreateContextAsync(state); // Act ValueTask FormatAsync() => context.FormatTemplateAsync("={Env.SOME_SECRET}"); @@ -40,7 +40,7 @@ public async Task FormatTemplateWithSensitivityAsync_WithSensitiveValue_ReturnsS WorkflowFormulaState state = new(RecalcEngineFactory.Create()); state.Set("SOME_SECRET", FormulaValue.New("secret-value"), VariableScopeNames.Environment, SensitivityLevel.Sensitive); state.Bind(); - DeclarativeWorkflowContext context = new(new Mock().Object, state); + DeclarativeWorkflowContext context = await CreateContextAsync(state); // Act EvaluationResult result = await context.FormatTemplateWithSensitivityAsync("={Env.SOME_SECRET}"); @@ -57,7 +57,7 @@ public async Task EvaluateValueAsync_WithSensitiveValue_ThrowsAsync() WorkflowFormulaState state = new(RecalcEngineFactory.Create()); state.Set(SystemScope.Names.LastMessageText, FormulaValue.New("secret-value"), VariableScopeNames.System, SensitivityLevel.Sensitive); state.Bind(); - DeclarativeWorkflowContext context = new(new Mock().Object, state); + DeclarativeWorkflowContext context = await CreateContextAsync(state); // Act ValueTask EvaluateAsync() => context.EvaluateValueAsync("System.LastMessageText"); @@ -74,7 +74,7 @@ public async Task EvaluateValueWithSensitivityAsync_WithSensitiveValue_ReturnsSe WorkflowFormulaState state = new(RecalcEngineFactory.Create()); state.Set(SystemScope.Names.LastMessageText, FormulaValue.New("secret-value"), VariableScopeNames.System, SensitivityLevel.Sensitive); state.Bind(); - DeclarativeWorkflowContext context = new(new Mock().Object, state); + DeclarativeWorkflowContext context = await CreateContextAsync(state); // Act EvaluationResult result = await context.EvaluateValueWithSensitivityAsync("System.LastMessageText"); @@ -91,7 +91,7 @@ public async Task QueueStateUpdateAsync_WithSensitivity_RebindsStateAsync() WorkflowFormulaState state = new(RecalcEngineFactory.Create()); state.Set("TestValue", FormulaValue.New("old-value")); state.Bind(); - DeclarativeWorkflowContext context = new(new Mock().Object, state); + DeclarativeWorkflowContext context = await CreateContextAsync(state); // Act await context.QueueStateUpdateAsync(PropertyPath.Create("Local.TestValue"), FormulaValue.New("new-value"), SensitivityLevel.Sensitive); @@ -113,7 +113,8 @@ public async Task ReadStateWithSensitivityAsync_QueuesSensitiveAssignmentAsync() source .Setup(c => c.ReadStateAsync(SystemScope.Names.LastMessageText, VariableScopeNames.System, default)) .Returns(new ValueTask("secret-value")); - DeclarativeWorkflowContext context = new(source.Object, state); + source.As().SetupGet(c => c.SessionId).Returns("test-session"); + DeclarativeWorkflowContext context = await DeclarativeWorkflowContext.CreateAsync(source.Object, state); // Act var evaluatedValue = await context.ReadStateWithSensitivityAsync(SystemScope.Names.LastMessageText, VariableScopeNames.System); @@ -177,7 +178,7 @@ public async Task GeneratedForeachPattern_WithSensitiveCollection_PreservesItemS VariableScopeNames.Environment, SensitivityLevel.Sensitive); state.Bind(); - DeclarativeWorkflowContext context = new(new Mock().Object, state); + DeclarativeWorkflowContext context = await CreateContextAsync(state); // Act EvaluationResult evaluatedValue = await context.EvaluateValueWithSensitivityAsync("Env.SensitiveItems"); @@ -213,4 +214,11 @@ public async Task ConvertValueWithSensitivityAsync_WithPlainContext_PreservesSen Assert.Equal(42M, result.Value); Assert.Equal(SensitivityLevel.Sensitive, result.Sensitivity); } + + private static async ValueTask CreateContextAsync(WorkflowFormulaState state) + { + Mock source = new(); + source.As().SetupGet(c => c.SessionId).Returns("test-session"); + return await DeclarativeWorkflowContext.CreateAsync(source.Object, state); + } } diff --git a/dotnet/tests/Microsoft.Agents.AI.Workflows.Declarative.UnitTests/Kit/RootExecutorTests.cs b/dotnet/tests/Microsoft.Agents.AI.Workflows.Declarative.UnitTests/Kit/RootExecutorTests.cs index 4fedeb6d0f5..4d156d8c115 100644 --- a/dotnet/tests/Microsoft.Agents.AI.Workflows.Declarative.UnitTests/Kit/RootExecutorTests.cs +++ b/dotnet/tests/Microsoft.Agents.AI.Workflows.Declarative.UnitTests/Kit/RootExecutorTests.cs @@ -33,12 +33,13 @@ public async Task InitializeEnvironmentAsync_OnlyQueuesAllowedVariablesAsync() }; TestRootExecutor executor = new(options); Mock sourceContext = new(MockBehavior.Strict); + sourceContext.As().SetupGet(c => c.SessionId).Returns("test-session"); sourceContext.Setup(c => c.QueueStateUpdateAsync("ALLOWED", It.IsAny(), VariableScopeNames.Environment, It.IsAny())) .Returns(default(ValueTask)); sourceContext.Setup(c => c.QueueStateUpdateAsync("ALLOWED", SensitivityLevel.Sensitive, WorkflowFormulaState.GetSensitivityScopeName(VariableScopeNames.Environment), It.IsAny())) .Returns(default(ValueTask)); - DeclarativeWorkflowContext context = new(sourceContext.Object, executor.Session.State); + DeclarativeWorkflowContext context = await DeclarativeWorkflowContext.CreateAsync(sourceContext.Object, executor.Session.State); // Act await executor.InitializeAsync(context, "ALLOWED", "HIDDEN"); diff --git a/dotnet/tests/Microsoft.Agents.AI.Workflows.Declarative.UnitTests/ObjectModel/InvokeMcpToolExecutorTest.cs b/dotnet/tests/Microsoft.Agents.AI.Workflows.Declarative.UnitTests/ObjectModel/InvokeMcpToolExecutorTest.cs index bb34b6c9a26..be97b96aa82 100644 --- a/dotnet/tests/Microsoft.Agents.AI.Workflows.Declarative.UnitTests/ObjectModel/InvokeMcpToolExecutorTest.cs +++ b/dotnet/tests/Microsoft.Agents.AI.Workflows.Declarative.UnitTests/ObjectModel/InvokeMcpToolExecutorTest.cs @@ -1,5 +1,6 @@ // Copyright (c) Microsoft. All rights reserved. +using System; using System.Collections.Concurrent; using System.Collections.Generic; using System.Linq; @@ -127,6 +128,86 @@ public async Task InvokeMcpToolExecuteWithoutApprovalAsync() await this.ExecuteTestAsync(model); } + [Fact] + public async Task InvokeMcpToolSeparateWorkflowSessionsUseSeparateSessionScopesAsync() + { + // Arrange + this.State.InitializeSystem(); + RecordingScopedMcpToolHandler handler = new(); + InvokeMcpTool model = this.CreateModel( + displayName: nameof(InvokeMcpToolSeparateWorkflowSessionsUseSeparateSessionScopesAsync), + serverUrl: TestServerUrl, + toolName: TestToolName, + requireApproval: false); + MockAgentProvider agentProvider = new(); + InvokeMcpToolExecutor action = new(model, handler, agentProvider.Object, this.State); + Mock firstContext = CreateMockWorkflowContext(); + Mock secondContext = CreateMockWorkflowContext(); + firstContext.As().SetupGet(context => context.SessionId).Returns("workflow-a"); + secondContext.As().SetupGet(context => context.SessionId).Returns("workflow-b"); + + // Act + await action.HandleAsync(new ActionExecutorResult("first"), firstContext.Object, CancellationToken.None); + await action.HandleAsync(new ActionExecutorResult("second"), secondContext.Object, CancellationToken.None); + + // Assert + Assert.Equal(2, handler.WorkflowSessionIds.Count); + Assert.NotEqual(handler.WorkflowSessionIds[0], handler.WorkflowSessionIds[1]); + } + + [Fact] + public async Task InvokeMcpToolSameWorkflowSessionReusesSessionScopeAsync() + { + // Arrange + this.State.InitializeSystem(); + RecordingScopedMcpToolHandler handler = new(); + InvokeMcpTool model = this.CreateModel( + displayName: nameof(InvokeMcpToolSameWorkflowSessionReusesSessionScopeAsync), + serverUrl: TestServerUrl, + toolName: TestToolName, + requireApproval: false); + MockAgentProvider agentProvider = new(); + InvokeMcpToolExecutor action = new(model, handler, agentProvider.Object, this.State); + Mock context = CreateMockWorkflowContext(); + context.As().SetupGet(current => current.SessionId).Returns("workflow-a"); + + // Act + await action.HandleAsync(new ActionExecutorResult("first"), context.Object, CancellationToken.None); + await action.HandleAsync(new ActionExecutorResult("second"), context.Object, CancellationToken.None); + + // Assert + Assert.Equal(2, handler.WorkflowSessionIds.Count); + Assert.Equal(handler.WorkflowSessionIds[0], handler.WorkflowSessionIds[1]); + } + + [Fact] + public async Task InvokeMcpToolLegacyContextsPersistSeparateSessionScopesAsync() + { + // Arrange + this.State.InitializeSystem(); + RecordingScopedMcpToolHandler handler = new(); + InvokeMcpTool model = this.CreateModel( + displayName: nameof(InvokeMcpToolLegacyContextsPersistSeparateSessionScopesAsync), + serverUrl: TestServerUrl, + toolName: TestToolName, + requireApproval: false); + MockAgentProvider agentProvider = new(); + InvokeMcpToolExecutor action = new(model, handler, agentProvider.Object, this.State); + InvokeMcpToolExecutor reconstructedAction = new(model, handler, agentProvider.Object, this.State); + Mock firstContext = CreateMockWorkflowContextWithSessionState(); + Mock secondContext = CreateMockWorkflowContextWithSessionState(); + + // Act + await action.HandleAsync(new ActionExecutorResult("first"), firstContext.Object, CancellationToken.None); + await action.HandleAsync(new ActionExecutorResult("second"), secondContext.Object, CancellationToken.None); + await reconstructedAction.HandleAsync(new ActionExecutorResult("continued"), firstContext.Object, CancellationToken.None); + + // Assert + Assert.Equal(3, handler.WorkflowSessionIds.Count); + Assert.NotEqual(handler.WorkflowSessionIds[0], handler.WorkflowSessionIds[1]); + Assert.Equal(handler.WorkflowSessionIds[0], handler.WorkflowSessionIds[2]); + } + [Fact] public async Task InvokeMcpToolExecuteWithServerLabelAsync() { @@ -1550,6 +1631,28 @@ private static Mock CreateMockWorkflowContext(List CreateMockWorkflowContextWithSessionState() + { + Mock context = CreateMockWorkflowContext(); + Dictionary<(string? Scope, string Key), string> state = []; + context.Setup(current => current.ReadOrInitStateAsync( + It.IsAny(), + It.IsAny>(), + It.IsAny(), + It.IsAny())) + .Returns((string key, Func initialStateFactory, string? scopeName, CancellationToken _) => + { + if (!state.TryGetValue((scopeName, key), out string? value)) + { + value = initialStateFactory(); + state[(scopeName, key)] = value; + } + + return new ValueTask(value); + }); + return context; + } + /// /// Creates a mock workflow context that actually stores state values (for checkpoint/restore tests). /// Optionally accepts an externally-owned state store so callers can drive multi-step @@ -1778,6 +1881,35 @@ private InvokeMcpTool CreateModelWithVariableServerUrl(string displayName, strin #region Mock MCP Tool Provider + private sealed class RecordingScopedMcpToolHandler : IWorkflowScopedMcpToolHandler + { + public List WorkflowSessionIds { get; } = []; + + public Task InvokeToolAsync( + string serverUrl, + string? serverLabel, + string toolName, + IDictionary? arguments, + IDictionary? headers, + string? connectionName, + CancellationToken cancellationToken = default) => + throw new InvalidOperationException("The workflow-scoped overload must be used."); + + public Task InvokeToolInWorkflowSessionAsync( + string serverUrl, + string? serverLabel, + string toolName, + IDictionary? arguments, + IDictionary? headers, + string? connectionName, + string workflowSessionId, + CancellationToken cancellationToken = default) + { + this.WorkflowSessionIds.Add(workflowSessionId); + return Task.FromResult(new McpServerToolResultContent("mock-call-id") { Outputs = [] }); + } + } + /// /// Mock implementation of for unit testing purposes. /// diff --git a/dotnet/tests/Microsoft.Agents.AI.Workflows.Declarative.UnitTests/PowerFx/WorkflowFormulaStateTests.cs b/dotnet/tests/Microsoft.Agents.AI.Workflows.Declarative.UnitTests/PowerFx/WorkflowFormulaStateTests.cs index 0af699c501c..1ba4cb17031 100644 --- a/dotnet/tests/Microsoft.Agents.AI.Workflows.Declarative.UnitTests/PowerFx/WorkflowFormulaStateTests.cs +++ b/dotnet/tests/Microsoft.Agents.AI.Workflows.Declarative.UnitTests/PowerFx/WorkflowFormulaStateTests.cs @@ -1,8 +1,10 @@ // Copyright (c) Microsoft. All rights reserved. +using System; using System.Collections.Generic; using System.Threading; using System.Threading.Tasks; +using Microsoft.Agents.AI.Workflows.Declarative.Interpreter; using Microsoft.Agents.AI.Workflows.Declarative.PowerFx; using Microsoft.Agents.ObjectModel; using Microsoft.PowerFx.Types; @@ -86,6 +88,34 @@ public void SetOverwritesExistingValue() Assert.Equal(newValue, result); } + [Fact] + public async Task DeclarativeContextFallbackSessionId_IsScopedToPersistedRunStateAsync() + { + // Arrange + Dictionary<(string? ScopeName, string Key), string> firstRunState = []; + Dictionary<(string? ScopeName, string Key), string> secondRunState = []; + IWorkflowContext firstContext = CreateContext(firstRunState); + IWorkflowContext restoredContext = CreateContext(firstRunState); + IWorkflowContext secondContext = CreateContext(secondRunState); + + // Act + DeclarativeWorkflowContext first = + await DeclarativeWorkflowContext.CreateAsync(firstContext, this.State); + DeclarativeWorkflowContext continued = + await DeclarativeWorkflowContext.CreateAsync(firstContext, this.State); + DeclarativeWorkflowContext restored = + await DeclarativeWorkflowContext.CreateAsync(restoredContext, this.State); + DeclarativeWorkflowContext second = + await DeclarativeWorkflowContext.CreateAsync(secondContext, this.State); + + // Assert + Assert.Equal(first.SessionId, continued.SessionId); + Assert.Equal(first.SessionId, restored.SessionId); + Assert.NotEqual(first.SessionId, second.SessionId); + Assert.True(firstRunState.ContainsKey((VariableScopeNames.System, "__declarative_mcp_workflow_session_id"))); + Assert.False(firstRunState.ContainsKey((null, "__declarative_mcp_workflow_session_id"))); + } + [Fact] public async Task RestoreAsync_RestoresPersistedSensitivityAsync() { @@ -104,4 +134,27 @@ public async Task RestoreAsync_RestoresPersistedSensitivityAsync() // Assert Assert.Equal(SensitivityLevel.Sensitive, this.State.GetSensitivity("secret")); } + + private static IWorkflowContext CreateContext(Dictionary<(string? ScopeName, string Key), string> state) + { + Mock context = new(); + context + .Setup(current => current.ReadOrInitStateAsync( + It.IsAny(), + It.IsAny>(), + It.IsAny(), + It.IsAny())) + .Returns((string key, Func factory, string? scopeName, CancellationToken cancellationToken) => + { + var scopedKey = (scopeName, key); + if (!state.TryGetValue(scopedKey, out string? value)) + { + value = factory(); + state[scopedKey] = value; + } + + return new ValueTask(value); + }); + return context.Object; + } } diff --git a/dotnet/tests/Microsoft.Agents.AI.Workflows.UnitTests/RepresentationTests.cs b/dotnet/tests/Microsoft.Agents.AI.Workflows.UnitTests/RepresentationTests.cs index 1e92d130de9..bddef5e4f19 100644 --- a/dotnet/tests/Microsoft.Agents.AI.Workflows.UnitTests/RepresentationTests.cs +++ b/dotnet/tests/Microsoft.Agents.AI.Workflows.UnitTests/RepresentationTests.cs @@ -8,6 +8,7 @@ using System.Threading; using System.Threading.Tasks; using Microsoft.Agents.AI.Workflows.Checkpointing; +using Microsoft.Agents.AI.Workflows.InProc; using Microsoft.Agents.AI.Workflows.Sample; using Microsoft.Agents.AI.Workflows.Specialized; using Microsoft.Extensions.AI; @@ -91,6 +92,48 @@ public async Task Test_SpecializedExecutor_InfosAsync() await RunExecutorBindingInfoMatchTestAsync(new RequestInfoExecutor(TestRequestPort)); } + [Fact] + public void SubworkflowSessionId_IsStableAndHierarchicallyScoped() + { + // Arrange + const string ParentSessionId = "parent-session"; + + // Act + string firstChild = SubworkflowBinding.CreateSubworkflowSessionId(ParentSessionId, "first"); + string sameChild = SubworkflowBinding.CreateSubworkflowSessionId(ParentSessionId, "first"); + string siblingChild = SubworkflowBinding.CreateSubworkflowSessionId(ParentSessionId, "second"); + string nestedChild = SubworkflowBinding.CreateSubworkflowSessionId(firstChild, "nested"); + + // Assert + Assert.Equal(firstChild, sameChild); + Assert.NotEqual(ParentSessionId, firstChild); + Assert.NotEqual(firstChild, siblingChild); + Assert.NotEqual(firstChild, nestedChild); + Assert.NotEqual(siblingChild, nestedChild); + } + + [Fact] + public async Task SubworkflowRunner_PreservesLegacyCheckpointSessionIdAsync() + { + // Arrange + const string LegacySessionId = "legacy-session"; + string workflowSessionId = SubworkflowBinding.CreateSubworkflowSessionId(LegacySessionId, "child"); + TestExecutor executor = new(); + Workflow workflow = new WorkflowBuilder(executor).Build(); + + // Act + InProcessRunner runner = InProcessRunner.CreateSubworkflowRunner( + workflow, + checkpointManager: new InMemoryCheckpointManager(), + sessionId: LegacySessionId, + workflowSessionId: workflowSessionId); + + // Assert + Assert.Equal(LegacySessionId, runner.SessionId); + Assert.Equal(workflowSessionId, runner.WorkflowSessionId); + await runner.RequestEndRunAsync(); + } + private static string Source(int id) => $"Source/{id}"; private static string Sink(int id) => $"Sink/{id}"; diff --git a/python/packages/declarative/AGENTS.md b/python/packages/declarative/AGENTS.md index 9638b435337..0679795d516 100644 --- a/python/packages/declarative/AGENTS.md +++ b/python/packages/declarative/AGENTS.md @@ -15,6 +15,9 @@ YAML/JSON-based declarative agent and workflow definitions. ## MCP Handler Lifetimes `DefaultMCPToolHandler` caches/coalesces sessions only without a `client_provider`. +Cache identity includes a framework-owned workflow session ID in addition to +endpoint, label, connection, and headers, so separate fresh runs do not share a +stateful MCP protocol session while continuations and checkpoint restores do. With a provider, every invocation (including `tools/list`) gets a fresh tool/session, even if the provider returns `None` or a shared HTTP client. Invocation cleanup closes the session and any internally owned fallback client, never caller-owned HTTP clients. diff --git a/python/packages/declarative/agent_framework_declarative/_workflows/_executors_control_flow.py b/python/packages/declarative/agent_framework_declarative/_workflows/_executors_control_flow.py index 5db554574f4..a612f9aef08 100644 --- a/python/packages/declarative/agent_framework_declarative/_workflows/_executors_control_flow.py +++ b/python/packages/declarative/agent_framework_declarative/_workflows/_executors_control_flow.py @@ -410,6 +410,10 @@ async def handle_action( ) -> None: """Simply pass through to continue the workflow.""" await self._ensure_state_initialized(ctx, trigger) + if self._action_def.get("kind") == "Entry": + from ._mcp_handler import reset_workflow_session_id + + reset_workflow_session_id(ctx.state) await ctx.send_message(ActionComplete()) diff --git a/python/packages/declarative/agent_framework_declarative/_workflows/_executors_mcp.py b/python/packages/declarative/agent_framework_declarative/_workflows/_executors_mcp.py index 7d3647d00bd..fef61859393 100644 --- a/python/packages/declarative/agent_framework_declarative/_workflows/_executors_mcp.py +++ b/python/packages/declarative/agent_framework_declarative/_workflows/_executors_mcp.py @@ -48,7 +48,12 @@ DeclarativeWorkflowState, ) from ._executors_tools import ToolApprovalResponse -from ._mcp_handler import MCPToolHandler, MCPToolInvocation, MCPToolResult +from ._mcp_handler import ( + MCPToolHandler, + MCPToolInvocation, + MCPToolResult, + get_or_create_workflow_session_id, +) __all__ = [ "MCP_ACTION_EXECUTORS", @@ -236,6 +241,7 @@ async def handle_action( arguments=arguments, headers=headers, connection_name=connection_name, + workflow_session_id=get_or_create_workflow_session_id(ctx.state), ) if require_approval: await self._request_approval( @@ -342,6 +348,7 @@ async def handle_approval_response( arguments=original_request.arguments, headers=self._evaluate_headers(state, self._action_def.get("headers")), connection_name=getattr(original_request, "connection_name", None), + workflow_session_id=get_or_create_workflow_session_id(ctx.state), ) if invocation.headers or original_request.header_names: binding = getattr(original_request, "header_binding", None) diff --git a/python/packages/declarative/agent_framework_declarative/_workflows/_mcp_handler.py b/python/packages/declarative/agent_framework_declarative/_workflows/_mcp_handler.py index 64eb59750fb..a4f9b175bc6 100644 --- a/python/packages/declarative/agent_framework_declarative/_workflows/_mcp_handler.py +++ b/python/packages/declarative/agent_framework_declarative/_workflows/_mcp_handler.py @@ -30,8 +30,10 @@ import hashlib import json import logging +import uuid from collections import OrderedDict from collections.abc import Awaitable, Callable +from contextlib import suppress from contextvars import ContextVar, Token from dataclasses import dataclass, field from typing import TYPE_CHECKING, Any, ClassVar, Protocol, cast, runtime_checkable @@ -40,6 +42,7 @@ if TYPE_CHECKING: from agent_framework import Content + from agent_framework._workflows._state import State __all__ = [ "ClientProvider", @@ -52,6 +55,21 @@ logger = logging.getLogger(__name__) _DEFAULT_CACHE_MAX_SIZE = 32 +_WORKFLOW_SESSION_ID_KEY = "_declarative_mcp_workflow_session_id" + + +def get_or_create_workflow_session_id(state: State) -> str: + workflow_session_id = state.get(_WORKFLOW_SESSION_ID_KEY) + if workflow_session_id is None: + workflow_session_id = uuid.uuid4().hex + state.set(_WORKFLOW_SESSION_ID_KEY, workflow_session_id) + if not isinstance(workflow_session_id, str) or not workflow_session_id: + raise ValueError("Invalid MCP workflow session state.") + return workflow_session_id + + +def reset_workflow_session_id(state: State) -> None: + state.set(_WORKFLOW_SESSION_ID_KEY, uuid.uuid4().hex) @dataclass @@ -73,6 +91,9 @@ class MCPToolInvocation: - ``connection_name``: Optional Foundry connection name forwarded for handlers that resolve auth/credentials by connection. The default handler does not consume this field. + - ``workflow_session_id``: Framework-owned identifier for the current + workflow session. The default handler uses it to prevent separate + workflows from sharing one stateful MCP protocol session. """ server_url: str @@ -81,6 +102,7 @@ class MCPToolInvocation: arguments: dict[str, Any] = field(default_factory=dict) # type: ignore[reportUnknownVariableType] headers: dict[str, str] = field(default_factory=dict) # type: ignore[reportUnknownVariableType] connection_name: str | None = None + workflow_session_id: str | None = None def _empty_outputs() -> list[Any]: @@ -153,6 +175,15 @@ class _CacheEntry: tool: Any # MCPStreamableHTTPTool — typed Any to avoid import at module load owned_httpx_client: httpx.AsyncClient | None + active_users: int = 0 + evicted: bool = False + disposal_claimed: bool = False + closed: asyncio.Event = field(default_factory=asyncio.Event) + close_exception: BaseException | None = None + + +class _EntryCreationCancelled(Exception): + """Signal waiters to retry after the task creating their entry was cancelled.""" class DefaultMCPToolHandler: @@ -160,12 +191,15 @@ class DefaultMCPToolHandler: Without a ``client_provider``, caches one :class:`agent_framework.MCPStreamableHTTPTool` instance per - ``(server_url, server_label, connection_name, headers_hash)`` in a bounded - LRU. The cache prevents re-establishing an MCP session for every - invocation while ensuring different header sets (auth tokens) cannot - share a session — matches the .NET design intent while bounding - cardinality. ``server_label`` and ``connection_name`` also participate - in the key to distinguish logical connections. + ``(workflow_session_id, server_url, server_label, connection_name, + headers_hash)`` in a bounded LRU. The cache prevents re-establishing an MCP + session for every invocation while ensuring separate workflow sessions and + different header sets (auth tokens) cannot share a session — matches the + .NET design intent while bounding cardinality. ``server_label`` and + ``connection_name`` also participate in the key to distinguish logical + connections. Active invocations lease their cache entry, so eviction waits + for the final user before closing it. Concurrent entry creation is also + capped at ``cache_max_size``. Header *names* are lower-cased inside the hash payload only — the headers passed on the wire keep the caller's original casing — so two YAML actions that spell ``Authorization`` differently still share a @@ -199,8 +233,9 @@ class DefaultMCPToolHandler: client_provider: Optional per-invocation ``httpx.AsyncClient`` provider. cache_max_size: Maximum number of cached MCP clients in no-provider mode. When exceeded, the least-recently-used entry is evicted and its - owned client closed. Defaults to ``32``. Does not enable session - caching when a provider is configured. + owned client closed after any active invocation finishes. This also + limits simultaneous connection attempts. Defaults to ``32``. Does + not enable session caching when a provider is configured. """ LIST_TOOLS_TOOL_NAME: ClassVar[str] = "tools/list" @@ -228,14 +263,16 @@ def __init__( raise ValueError(f"cache_max_size must be positive, got {cache_max_size}") self._client_provider = client_provider self._cache_max_size = cache_max_size - self._cache: OrderedDict[tuple[str, str, str, str], _CacheEntry] = OrderedDict() + self._cache: OrderedDict[tuple[str, str, str, str, str], _CacheEntry] = OrderedDict() + self._retired: dict[int, _CacheEntry] = {} # Outer lock guards the cache + in-flight-future map only — never # held across network I/O. self._cache_lock = asyncio.Lock() + self._creation_semaphore = asyncio.Semaphore(cache_max_size) # Per-key in-flight futures: while one task is connecting, other # tasks awaiting the same key will await the same future and share # the resulting cache entry. - self._inflight: dict[tuple[str, str, str, str], asyncio.Future[_CacheEntry]] = {} + self._inflight: dict[tuple[str, str, str, str, str], asyncio.Future[_CacheEntry]] = {} # Completion signals only: provider-backed calls never share entries. self._active_invocations: set[asyncio.Future[None]] = set() # Keep ancestry so a completed nested call cannot hide an active parent @@ -246,6 +283,7 @@ def __init__( # Set by ``aclose`` to prevent post-close cache insertions and to # reject new ``invoke_tool`` calls. Once set, never cleared. self._closed = False + self._shutdown_task: asyncio.Task[None] | None = None async def invoke_tool(self, invocation: MCPToolInvocation) -> MCPToolResult: """Invoke ``invocation.tool_name`` on an MCP client for the server. @@ -303,7 +341,9 @@ async def invoke_tool(self, invocation: MCPToolInvocation) -> MCPToolResult: return await self._invoke_entry(entry, invocation) finally: - if completion is not None: + if self._client_provider is None and entry is not None: + await self._release_entry(entry) + elif completion is not None: try: if entry is not None: await self._close_invocation_entry(entry) @@ -428,7 +468,7 @@ async def aclose(self) -> None: are rejected after connecting and cleaned up by their invocation. Concurrent shutdown calls wait for those invocations as well. - Idempotent. In no-provider mode, a second call returns immediately. + Idempotent. Concurrent and subsequent calls await the same shutdown task. Drains any in-flight ``_create_entry`` tasks before returning so their resources are cleaned up; the in-flight tasks see ``self._closed`` in phase 3 of :meth:`_get_or_create_entry`, close their own entry, and resolve @@ -444,17 +484,45 @@ async def aclose(self) -> None: "DefaultMCPToolHandler.aclose() cannot be called from an active provider-backed invocation" ) async with self._cache_lock: - if self._closed and self._client_provider is None: - return - self._closed = True - entries = list(self._cache.values()) - self._cache.clear() - inflight_futures = list(self._inflight.values()) - active_invocations = list(self._active_invocations) - - for completion in active_invocations: - # Cancelling shutdown must not cancel an invocation's completion signal. - await asyncio.shield(completion) + if self._shutdown_task is None: + self._closed = True + entries = list(self._cache.values()) + entry_ids = {id(entry) for entry in entries} + entries.extend(entry for entry in self._retired.values() if id(entry) not in entry_ids) + self._cache.clear() + entries_to_close: list[_CacheEntry] = [] + for entry in entries: + entry.evicted = True + self._retired[id(entry)] = entry + if entry.active_users == 0 and not entry.disposal_claimed: + entry.disposal_claimed = True + entries_to_close.append(entry) + inflight_futures = list(self._inflight.values()) + active_invocations = list(self._active_invocations) + self._shutdown_task = asyncio.create_task( + self._drain_shutdown(entries, entries_to_close, inflight_futures, active_invocations) + ) + shutdown_task = self._shutdown_task + + await asyncio.shield(shutdown_task) + + async def _drain_shutdown( + self, + entries: list[_CacheEntry], + entries_to_close: list[_CacheEntry], + inflight_futures: list[asyncio.Future[_CacheEntry]], + active_invocations: list[asyncio.Future[None]], + ) -> None: + if active_invocations: + invocation_results = await asyncio.gather(*active_invocations, return_exceptions=True) + for result in invocation_results: + if isinstance(result, BaseException): + if isinstance(result, (KeyboardInterrupt, SystemExit, GeneratorExit)): + raise result + logger.debug( + "DefaultMCPToolHandler: active invocation raised during aclose", + exc_info=(type(result), result, result.__traceback__), + ) # Wait for in-flight creations to finish their self-cleanup. Each # in-flight task self-closes its entry under the closed-flag branch @@ -467,8 +535,18 @@ async def aclose(self) -> None: logger.debug("DefaultMCPToolHandler: in-flight future raised during aclose", exc_info=True) continue + close_results = await asyncio.gather( + *(self._close_claimed_entry(entry) for entry in entries_to_close), + return_exceptions=True, + ) + for entry in entries: + await entry.closed.wait() for entry in entries: - await self._close_entry(entry) + if entry.close_exception is not None: + raise entry.close_exception + for result in close_results: + if isinstance(result, BaseException): + raise result async def __aenter__(self) -> DefaultMCPToolHandler: return self @@ -483,6 +561,7 @@ async def __aexit__(self, exc_type: Any, exc: Any, tb: Any) -> None: async def _get_or_create_entry(self, invocation: MCPToolInvocation) -> _CacheEntry: """Look up (or create) the cached MCP client for this invocation.""" key = self._cache_key( + invocation.workflow_session_id, invocation.server_url, invocation.server_label, invocation.connection_name, @@ -498,6 +577,7 @@ async def _get_or_create_entry(self, invocation: MCPToolInvocation) -> _CacheEnt existing = self._cache.get(key) if existing is not None: self._cache.move_to_end(key) + existing.active_users += 1 return existing inflight = self._inflight.get(key) if inflight is None: @@ -506,20 +586,19 @@ async def _get_or_create_entry(self, invocation: MCPToolInvocation) -> _CacheEnt creating = True if not creating: - return await inflight + with suppress(_EntryCreationCancelled): + _ = await asyncio.shield(inflight) + return await self._get_or_create_entry(invocation) # Phase 2: we own creation. Build the entry outside the lock. try: - entry = await self._create_entry(invocation) + async with self._creation_semaphore: + async with self._cache_lock: + if self._closed: + raise RuntimeError("DefaultMCPToolHandler is closed") + entry = await self._create_entry(invocation) except BaseException as exc: - async with self._cache_lock: - self._inflight.pop(key, None) - if not inflight.done(): - inflight.set_exception(exc if isinstance(exc, BaseException) else RuntimeError(str(exc))) - # Mark the exception retrieved to suppress noisy "Future exception - # was never retrieved" warnings when there are no other awaiters - # (other awaiters still see the exception through their ``await``). - inflight.exception() + await self._complete_inflight_failure(key, inflight, exc) raise # Phase 3: insert with LRU eviction; resolve the in-flight future. @@ -530,43 +609,145 @@ async def _get_or_create_entry(self, invocation: MCPToolInvocation) -> _CacheEnt evicted: _CacheEntry | None = None duplicate: _CacheEntry | None = None handler_closed = False - async with self._cache_lock: - self._inflight.pop(key, None) - if self._closed: - handler_closed = True - else: - existing = self._cache.get(key) - if existing is not None: - # Another writer beat us; prefer the existing entry and - # discard ours after the lock is released. - self._cache.move_to_end(key) - duplicate = entry - entry = existing + try: + async with self._cache_lock: + self._inflight.pop(key, None) + if self._closed: + handler_closed = True else: - self._cache[key] = entry - self._cache.move_to_end(key) - if len(self._cache) > self._cache_max_size: - _evicted_key, evicted = self._cache.popitem(last=False) - if not inflight.done(): - inflight.set_result(entry) + existing = self._cache.get(key) + if existing is not None: + # Another writer beat us; prefer the existing entry and + # discard ours after the lock is released. + self._cache.move_to_end(key) + duplicate = entry + entry = existing + entry.active_users += 1 + else: + entry.active_users = 1 + self._cache[key] = entry + self._cache.move_to_end(key) + if len(self._cache) > self._cache_max_size: + _evicted_key, evicted = self._cache.popitem(last=False) + evicted.evicted = True + self._retired[id(evicted)] = evicted + if evicted.active_users == 0: + evicted.disposal_claimed = True + if not inflight.done(): + inflight.set_result(entry) + except BaseException as exc: + await self._abort_entry_creation(key, inflight, entry, exc) + raise if handler_closed: # Close our orphaned entry; resolve the future with a clear # error so the caller (and any other awaiters) surface a # consistent "handler is closed" failure rather than receiving # an entry we are about to close behind their back. - await self._close_entry(entry) err = RuntimeError("DefaultMCPToolHandler is closed") - if not inflight.done(): - inflight.set_exception(err) - inflight.exception() + try: + await self._close_invocation_entry(entry) + finally: + if not inflight.done(): + inflight.set_exception(err) + inflight.exception() raise err - if duplicate is not None: - await self._close_entry(duplicate) - if evicted is not None: - await self._close_entry(evicted) + try: + if duplicate is not None: + await self._close_entry(duplicate) + if evicted is not None and evicted.disposal_claimed: + await self._close_claimed_entry(evicted) + except BaseException: + await self._release_entry(entry) + raise return entry + async def _abort_entry_creation( + self, + key: tuple[str, str, str, str, str], + inflight: asyncio.Future[_CacheEntry], + entry: _CacheEntry, + exc: BaseException, + ) -> None: + async def cleanup() -> None: + try: + await self._close_entry(entry) + finally: + async with self._cache_lock: + self._inflight.pop(key, None) + if not inflight.done(): + inflight.set_exception(self._entry_creation_exception(exc)) + inflight.exception() + + cleanup_task = asyncio.create_task(cleanup()) + while not cleanup_task.done(): + try: + await asyncio.shield(cleanup_task) + except asyncio.CancelledError: + continue + cleanup_task.result() + + async def _complete_inflight_failure( + self, + key: tuple[str, str, str, str, str], + inflight: asyncio.Future[_CacheEntry], + exc: BaseException, + ) -> None: + async def cleanup() -> None: + async with self._cache_lock: + self._inflight.pop(key, None) + if not inflight.done(): + inflight.set_exception(self._entry_creation_exception(exc)) + inflight.exception() + + cleanup_task = asyncio.create_task(cleanup()) + while not cleanup_task.done(): + try: + await asyncio.shield(cleanup_task) + except asyncio.CancelledError: + continue + cleanup_task.result() + + @staticmethod + def _entry_creation_exception(exc: BaseException) -> BaseException: + if isinstance(exc, asyncio.CancelledError): + return _EntryCreationCancelled() + return exc + + async def _release_entry(self, entry: _CacheEntry) -> None: + release = asyncio.create_task(self._release_entry_core(entry)) + cancelled = False + while not release.done(): + try: + await asyncio.shield(release) + except asyncio.CancelledError: + cancelled = True + release.result() + if cancelled: + raise asyncio.CancelledError + + async def _release_entry_core(self, entry: _CacheEntry) -> None: + close_entry = False + async with self._cache_lock: + entry.active_users -= 1 + if entry.active_users == 0 and entry.evicted and not entry.disposal_claimed: + entry.disposal_claimed = True + close_entry = True + + if close_entry: + await self._close_claimed_entry(entry) + + async def _close_claimed_entry(self, entry: _CacheEntry) -> None: + try: + await self._close_invocation_entry(entry) + except BaseException as exc: + entry.close_exception = exc + raise + finally: + async with self._cache_lock: + self._retired.pop(id(entry), None) + entry.closed.set() + async def _create_entry(self, invocation: MCPToolInvocation) -> _CacheEntry: """Construct (and connect) a fresh MCP client for ``invocation``.""" from agent_framework import MCPStreamableHTTPTool @@ -648,16 +829,18 @@ async def _close_entry(self, entry: _CacheEntry) -> None: @staticmethod def _cache_key( + workflow_session_id: str | None, server_url: str, server_label: str | None, connection_name: str | None, headers: dict[str, str] | None, - ) -> tuple[str, str, str, str]: + ) -> tuple[str, str, str, str, str]: """Build an order-independent cache key for the invocation identity. Used only without a ``client_provider``. The key includes - ``server_label`` and ``connection_name`` to distinguish logical - connections. + ``workflow_session_id`` to isolate stateful MCP protocol sessions + between workflows, plus ``server_label`` and ``connection_name`` to + distinguish logical connections. Header *names* are lower-cased inside the hash payload only so that ``Authorization`` and ``authorization`` map to the same @@ -669,4 +852,10 @@ def _cache_key( normalized = sorted((k.lower(), v) for k, v in headers.items()) payload = json.dumps(normalized, ensure_ascii=False) headers_hash = hashlib.sha256(payload.encode("utf-8")).hexdigest() - return (server_url, server_label or "", connection_name or "", headers_hash) + return ( + workflow_session_id or "", + server_url, + server_label or "", + connection_name or "", + headers_hash, + ) diff --git a/python/packages/declarative/tests/test_default_mcp_tool_handler.py b/python/packages/declarative/tests/test_default_mcp_tool_handler.py index 0826302cf69..7661496ff27 100644 --- a/python/packages/declarative/tests/test_default_mcp_tool_handler.py +++ b/python/packages/declarative/tests/test_default_mcp_tool_handler.py @@ -23,11 +23,13 @@ import httpx import pytest from agent_framework import Content +from agent_framework._workflows._state import State from agent_framework.exceptions import ToolExecutionException from agent_framework_declarative._workflows._mcp_handler import ( DefaultMCPToolHandler, MCPToolInvocation, + get_or_create_workflow_session_id, ) pytestmark = pytest.mark.skipif( @@ -827,6 +829,22 @@ def test_invalid_cache_size_raises(self) -> None: DefaultMCPToolHandler(cache_max_size=-3) +class TestWorkflowSessionId: + def test_separate_workflow_states_get_separate_ids(self) -> None: + first = get_or_create_workflow_session_id(State()) + second = get_or_create_workflow_session_id(State()) + + assert first != second + + def test_same_workflow_state_reuses_id(self) -> None: + state = State() + + first = get_or_create_workflow_session_id(state) + second = get_or_create_workflow_session_id(state) + + assert first == second + + # ---------- Tool kwargs ---------------------------------------------------- @@ -879,12 +897,20 @@ class TestCache: async def test_same_url_and_headers_hit_cache(self) -> None: handler = DefaultMCPToolHandler() with _patch_tool(): - await handler.invoke_tool(_invocation(headers={"X": "1"})) - await handler.invoke_tool(_invocation(headers={"X": "1"})) + await handler.invoke_tool(_invocation(headers={"X": "1"}, workflow_session_id="workflow-a")) + await handler.invoke_tool(_invocation(headers={"X": "1"}, workflow_session_id="workflow-a")) # One tool created, connect called once. assert len(FakeTool.instances) == 1 assert FakeTool.instances[0].connect_count == 1 + @pytest.mark.asyncio + async def test_separate_workflow_sessions_use_separate_entries(self) -> None: + handler = DefaultMCPToolHandler() + with _patch_tool(): + await handler.invoke_tool(_invocation(workflow_session_id="workflow-a")) + await handler.invoke_tool(_invocation(workflow_session_id="workflow-b")) + assert len(FakeTool.instances) == 2 + @pytest.mark.asyncio async def test_different_headers_create_separate_entries(self) -> None: handler = DefaultMCPToolHandler() @@ -917,6 +943,232 @@ async def test_lru_eviction_closes_old_entry(self) -> None: assert FakeTool.instances[1].close_count == 0 assert FakeTool.instances[2].close_count == 0 + @pytest.mark.asyncio + async def test_lru_eviction_defers_close_until_active_invocation_finishes(self) -> None: + handler = DefaultMCPToolHandler(cache_max_size=1) + first_started = asyncio.Event() + release_first = asyncio.Event() + + async def gated_call(tool: FakeTool, tool_name: str, **arguments: Any) -> Any: + if tool.kwargs["url"] == "https://a/": + first_started.set() + await release_first.wait() + return [Content.from_text("ok")] + + with _patch_tool(), patch.object(FakeTool, "call_tool", gated_call): + first = asyncio.create_task(handler.invoke_tool(_invocation(server_url="https://a/"))) + await first_started.wait() + second = await handler.invoke_tool(_invocation(server_url="https://b/")) + + assert not second.is_error + assert FakeTool.instances[0].close_count == 0 + close_task = asyncio.create_task(handler.aclose()) + await asyncio.sleep(0) + assert not close_task.done() + + release_first.set() + first_result = await first + await asyncio.gather(close_task) + + assert not first_result.is_error + assert FakeTool.instances[0].close_count == 1 + assert FakeTool.instances[1].close_count == 1 + + @pytest.mark.asyncio + async def test_lru_eviction_cleanup_cancellation_releases_new_entry(self) -> None: + handler = DefaultMCPToolHandler(cache_max_size=1) + cancel_first_close = True + original_close = FakeTool.close + + async def close(tool: FakeTool) -> None: + if cancel_first_close and tool.kwargs["url"] == "https://a/": + tool.close_count += 1 + raise asyncio.CancelledError + await original_close(tool) + + with _patch_tool(), patch.object(FakeTool, "close", close): + await handler.invoke_tool(_invocation(server_url="https://a/")) + with pytest.raises(asyncio.CancelledError): + await handler.invoke_tool(_invocation(server_url="https://b/")) + + assert all(entry.active_users == 0 for entry in handler._cache.values()) + cancel_first_close = False + await asyncio.wait_for(handler.aclose(), timeout=1) + + @pytest.mark.asyncio + async def test_repeated_cancellation_while_releasing_entry_does_not_leak_lease(self) -> None: + handler = DefaultMCPToolHandler() + invocation_started = asyncio.Event() + release_invocation = asyncio.Event() + + async def gated_call(_tool: FakeTool, _tool_name: str, **_arguments: Any) -> Any: + invocation_started.set() + await release_invocation.wait() + return [Content.from_text("ok")] + + with _patch_tool(), patch.object(FakeTool, "call_tool", gated_call): + invocation = asyncio.create_task(handler.invoke_tool(_invocation(headers={"X": "1"}))) + await invocation_started.wait() + await handler._cache_lock.acquire() + invocation.cancel() + await asyncio.sleep(0) + invocation.cancel() + handler._cache_lock.release() + + with pytest.raises(asyncio.CancelledError): + await asyncio.gather(invocation) + await asyncio.wait_for(handler.aclose(), timeout=1) + + assert FakeTool.instances[0].close_count == 1 + assert all(entry.active_users == 0 for entry in handler._cache.values()) + + @pytest.mark.asyncio + async def test_entry_creation_is_bounded_and_cancelled_waiter_is_cleaned_up(self) -> None: + handler = DefaultMCPToolHandler(cache_max_size=2) + connecting = 0 + max_connecting = 0 + capacity_reached = asyncio.Event() + release = asyncio.Event() + original_connect = FakeTool.connect + + async def gated_connect(tool: FakeTool) -> None: + nonlocal connecting, max_connecting + connecting += 1 + max_connecting = max(max_connecting, connecting) + if connecting == 2: + capacity_reached.set() + try: + await release.wait() + await original_connect(tool) + finally: + connecting -= 1 + + with _patch_tool(), patch.object(FakeTool, "connect", gated_connect): + first = asyncio.create_task(handler.invoke_tool(_invocation(workflow_session_id="workflow-a"))) + second = asyncio.create_task(handler.invoke_tool(_invocation(workflow_session_id="workflow-b"))) + await capacity_reached.wait() + cancelled = asyncio.create_task(handler.invoke_tool(_invocation(workflow_session_id="workflow-c"))) + await asyncio.sleep(0) + + assert len(FakeTool.instances) == 2 + cancelled.cancel() + with pytest.raises(asyncio.CancelledError): + cancelled_result = await cancelled + assert cancelled_result is None + + release.set() + results = await asyncio.gather(first, second) + + assert all(not result.is_error for result in results) + assert max_connecting == 2 + assert not handler._inflight + + @pytest.mark.asyncio + async def test_cancelled_waiter_does_not_cancel_shared_inflight_creation(self) -> None: + handler = DefaultMCPToolHandler() + connect_started = asyncio.Event() + release_connect = asyncio.Event() + original_connect = FakeTool.connect + + async def gated_connect(tool: FakeTool) -> None: + connect_started.set() + await release_connect.wait() + await original_connect(tool) + + with _patch_tool(), patch.object(FakeTool, "connect", gated_connect): + creator = asyncio.create_task(handler.invoke_tool(_invocation(workflow_session_id="workflow-a"))) + await connect_started.wait() + waiter = asyncio.create_task(handler.invoke_tool(_invocation(workflow_session_id="workflow-a"))) + await asyncio.sleep(0) + + waiter.cancel() + with pytest.raises(asyncio.CancelledError): + waiter_result = await waiter + assert waiter_result is None + + release_connect.set() + creator_result = await creator + follow_up_result = await handler.invoke_tool(_invocation(workflow_session_id="workflow-a")) + + assert not creator_result.is_error + assert not follow_up_result.is_error + assert len(FakeTool.instances) == 1 + assert FakeTool.instances[0].connect_count == 1 + + @pytest.mark.asyncio + async def test_cancelled_creator_does_not_cancel_shared_inflight_waiter(self) -> None: + handler = DefaultMCPToolHandler() + first_connect_started = asyncio.Event() + release_first_connect = asyncio.Event() + original_connect = FakeTool.connect + + async def gated_connect(tool: FakeTool) -> None: + if len(FakeTool.instances) == 1: + first_connect_started.set() + await release_first_connect.wait() + await original_connect(tool) + + with _patch_tool(), patch.object(FakeTool, "connect", gated_connect): + creator = asyncio.create_task(handler.invoke_tool(_invocation(workflow_session_id="workflow-a"))) + await first_connect_started.wait() + waiter = asyncio.create_task(handler.invoke_tool(_invocation(workflow_session_id="workflow-a"))) + await asyncio.sleep(0) + + creator.cancel() + release_first_connect.set() + with pytest.raises(asyncio.CancelledError): + _ = await creator + + waiter_result = await waiter + follow_up_result = await handler.invoke_tool(_invocation(workflow_session_id="workflow-a")) + + assert not waiter_result.is_error + assert not follow_up_result.is_error + assert len(FakeTool.instances) == 2 + assert FakeTool.instances[0].close_count == 1 + assert FakeTool.instances[1].connect_count == 1 + + @pytest.mark.asyncio + async def test_repeated_creator_cancellation_completes_inflight_failure(self) -> None: + handler = DefaultMCPToolHandler() + + class GatedCreationSemaphore: + def __init__(self) -> None: + self.entered = asyncio.Event() + self.release = asyncio.Event() + + async def __aenter__(self) -> None: + self.entered.set() + await self.release.wait() + + async def __aexit__(self, *_args: Any) -> None: + pass + + creation_gate: Any = GatedCreationSemaphore() + handler._creation_semaphore = creation_gate + + with _patch_tool(): + creator = asyncio.create_task(handler.invoke_tool(_invocation(workflow_session_id="workflow-a"))) + await creation_gate.entered.wait() + assert handler._inflight + waiter = asyncio.create_task(handler.invoke_tool(_invocation(workflow_session_id="workflow-a"))) + await asyncio.sleep(0) + + await handler._cache_lock.acquire() + creator.cancel() + await asyncio.sleep(0) + creator.cancel() + handler._cache_lock.release() + creation_gate.release.set() + + with pytest.raises(asyncio.CancelledError): + await asyncio.gather(creator) + waiter_result = await asyncio.wait_for(waiter, timeout=1) + await asyncio.wait_for(handler.aclose(), timeout=1) + + assert not waiter_result.is_error + assert not handler._inflight + @pytest.mark.asyncio async def test_repeated_use_keeps_lru_alive(self) -> None: handler = DefaultMCPToolHandler(cache_max_size=2) @@ -950,10 +1202,10 @@ async def slow_connect(self: FakeTool) -> None: with _patch_tool(), patch.object(FakeTool, "connect", slow_connect): results = await asyncio.gather( - handler.invoke_tool(_invocation(headers={"X": "1"})), - handler.invoke_tool(_invocation(headers={"X": "1"})), - handler.invoke_tool(_invocation(headers={"X": "1"})), - handler.invoke_tool(_invocation(headers={"X": "1"})), + handler.invoke_tool(_invocation(headers={"X": "1"}, workflow_session_id="workflow-a")), + handler.invoke_tool(_invocation(headers={"X": "1"}, workflow_session_id="workflow-a")), + handler.invoke_tool(_invocation(headers={"X": "1"}, workflow_session_id="workflow-a")), + handler.invoke_tool(_invocation(headers={"X": "1"}, workflow_session_id="workflow-a")), ) assert all(not r.is_error for r in results) # Only one tool was created and connected, despite 4 concurrent calls. @@ -1061,6 +1313,93 @@ async def test_aclose_is_idempotent(self) -> None: await handler.aclose() assert FakeTool.instances[0].close_count == 1 + @pytest.mark.asyncio + async def test_cancelled_aclose_continues_draining_cached_entries(self) -> None: + handler = DefaultMCPToolHandler() + first_close_started = asyncio.Event() + release_first_close = asyncio.Event() + original_close = FakeTool.close + + async def gated_close(tool: FakeTool) -> None: + if tool.kwargs["url"] == "https://a/": + first_close_started.set() + await release_first_close.wait() + await original_close(tool) + + with _patch_tool(), patch.object(FakeTool, "close", gated_close): + await handler.invoke_tool(_invocation(server_url="https://a/", headers={"X": "1"})) + await handler.invoke_tool(_invocation(server_url="https://b/", headers={"X": "1"})) + + shutdown = asyncio.create_task(handler.aclose()) + await first_close_started.wait() + shutdown.cancel() + with pytest.raises(asyncio.CancelledError): + await asyncio.gather(shutdown) + + release_first_close.set() + await asyncio.wait_for(handler.aclose(), timeout=1) + + assert [tool.close_count for tool in FakeTool.instances] == [1, 1] + assert all(tool._httpx_client is not None and tool._httpx_client.is_closed for tool in FakeTool.instances) + + @pytest.mark.asyncio + async def test_cleanup_cancellation_continues_draining_cached_entries(self) -> None: + handler = DefaultMCPToolHandler() + original_close = FakeTool.close + + async def cancel_first_close(tool: FakeTool) -> None: + if tool.kwargs["url"] == "https://a/": + tool.close_count += 1 + raise asyncio.CancelledError + await original_close(tool) + + with _patch_tool(), patch.object(FakeTool, "close", cancel_first_close): + await handler.invoke_tool(_invocation(server_url="https://a/", headers={"X": "1"})) + await handler.invoke_tool(_invocation(server_url="https://b/", headers={"X": "1"})) + + with pytest.raises(asyncio.CancelledError): + await handler.aclose() + with pytest.raises(asyncio.CancelledError): + await handler.aclose() + + assert [tool.close_count for tool in FakeTool.instances] == [1, 1] + assert all(tool._httpx_client is not None and tool._httpx_client.is_closed for tool in FakeTool.instances) + + @pytest.mark.asyncio + async def test_active_entry_cleanup_cancellation_is_reported_by_aclose(self) -> None: + handler = DefaultMCPToolHandler() + invocation_started = asyncio.Event() + release_invocation = asyncio.Event() + + async def gated_call(_tool: FakeTool, _tool_name: str, **_arguments: Any) -> Any: + invocation_started.set() + await release_invocation.wait() + return [Content.from_text("ok")] + + async def cancelled_close(tool: FakeTool) -> None: + tool.close_count += 1 + raise asyncio.CancelledError + + with ( + _patch_tool(), + patch.object(FakeTool, "call_tool", gated_call), + patch.object(FakeTool, "close", cancelled_close), + ): + invocation = asyncio.create_task(handler.invoke_tool(_invocation(headers={"X": "1"}))) + await invocation_started.wait() + shutdown = asyncio.create_task(handler.aclose()) + await asyncio.sleep(0) + assert handler._closed + release_invocation.set() + + with pytest.raises(asyncio.CancelledError): + await asyncio.gather(invocation) + with pytest.raises(asyncio.CancelledError): + await asyncio.gather(shutdown) + + assert FakeTool.instances[0].close_count == 1 + assert not handler._retired + @pytest.mark.asyncio async def test_invoke_after_close_returns_error_result(self) -> None: """Post-close ``invoke_tool`` surfaces a tool error rather than crashing.""" @@ -1111,6 +1450,75 @@ async def gated_connect(self: FakeTool) -> None: assert result.is_error is True assert "closed" in (result.error_message or "").lower() + @pytest.mark.asyncio + async def test_cancelled_creator_does_not_block_aclose(self) -> None: + handler = DefaultMCPToolHandler() + connect_started = asyncio.Event() + release_connect = asyncio.Event() + close_started = asyncio.Event() + release_close = asyncio.Event() + original_connect = FakeTool.connect + original_close_entry = handler._close_entry + + async def gated_connect(self: FakeTool) -> None: + connect_started.set() + await release_connect.wait() + await original_connect(self) + + async def gated_close_entry(entry: Any) -> None: + close_started.set() + await release_close.wait() + await original_close_entry(entry) + + with ( + _patch_tool(), + patch.object(FakeTool, "connect", gated_connect), + patch.object(handler, "_close_entry", gated_close_entry), + ): + invoke_task = asyncio.create_task(handler.invoke_tool(_invocation(headers={"X": "1"}))) + await connect_started.wait() + close_task = asyncio.create_task(handler.aclose()) + await asyncio.sleep(0) + release_connect.set() + await close_started.wait() + invoke_task.cancel() + release_close.set() + + with pytest.raises(asyncio.CancelledError): + _ = await invoke_task + await asyncio.wait_for(close_task, timeout=1) + + assert FakeTool.instances[0].close_count == 1 + assert not handler._inflight + + @pytest.mark.asyncio + async def test_cancellation_while_waiting_for_phase_three_lock_cleans_up(self) -> None: + handler = DefaultMCPToolHandler() + entry_created = asyncio.Event() + release_creation = asyncio.Event() + original_create_entry = handler._create_entry + + async def gated_create_entry(invocation: MCPToolInvocation) -> Any: + entry = await original_create_entry(invocation) + entry_created.set() + await release_creation.wait() + return entry + + with _patch_tool(), patch.object(handler, "_create_entry", gated_create_entry): + invocation = asyncio.create_task(handler.invoke_tool(_invocation(headers={"X": "1"}))) + await entry_created.wait() + await handler._cache_lock.acquire() + release_creation.set() + await asyncio.sleep(0) + invocation.cancel() + handler._cache_lock.release() + + with pytest.raises(asyncio.CancelledError): + _ = await invocation + + assert FakeTool.instances[0].close_count == 1 + assert not handler._inflight + # ---------- Result normalisation ------------------------------------------ @@ -1235,38 +1643,43 @@ def boom(**_a: Any) -> Any: class TestCacheKey: def test_key_order_independent(self) -> None: - k1 = DefaultMCPToolHandler._cache_key("https://x/", None, None, {"A": "1", "B": "2"}) - k2 = DefaultMCPToolHandler._cache_key("https://x/", None, None, {"B": "2", "A": "1"}) + k1 = DefaultMCPToolHandler._cache_key("workflow", "https://x/", None, None, {"A": "1", "B": "2"}) + k2 = DefaultMCPToolHandler._cache_key("workflow", "https://x/", None, None, {"B": "2", "A": "1"}) assert k1 == k2 def test_key_distinguishes_values(self) -> None: - k1 = DefaultMCPToolHandler._cache_key("https://x/", None, None, {"A": "1"}) - k2 = DefaultMCPToolHandler._cache_key("https://x/", None, None, {"A": "2"}) + k1 = DefaultMCPToolHandler._cache_key("workflow", "https://x/", None, None, {"A": "1"}) + k2 = DefaultMCPToolHandler._cache_key("workflow", "https://x/", None, None, {"A": "2"}) assert k1 != k2 def test_empty_headers_use_fixed_hash(self) -> None: - k1 = DefaultMCPToolHandler._cache_key("https://x/", None, None, None) - k2 = DefaultMCPToolHandler._cache_key("https://x/", None, None, {}) + k1 = DefaultMCPToolHandler._cache_key("workflow", "https://x/", None, None, None) + k2 = DefaultMCPToolHandler._cache_key("workflow", "https://x/", None, None, {}) assert k1 == k2 + def test_key_distinguishes_workflow_session(self) -> None: + k1 = DefaultMCPToolHandler._cache_key("workflow-a", "https://x/", None, None, None) + k2 = DefaultMCPToolHandler._cache_key("workflow-b", "https://x/", None, None, None) + assert k1 != k2 + def test_key_distinguishes_connection_name(self) -> None: - k1 = DefaultMCPToolHandler._cache_key("https://x/", None, "conn-A", None) - k2 = DefaultMCPToolHandler._cache_key("https://x/", None, "conn-B", None) + k1 = DefaultMCPToolHandler._cache_key("workflow", "https://x/", None, "conn-A", None) + k2 = DefaultMCPToolHandler._cache_key("workflow", "https://x/", None, "conn-B", None) assert k1 != k2 def test_key_distinguishes_server_label(self) -> None: - k1 = DefaultMCPToolHandler._cache_key("https://x/", "Lbl-A", None, None) - k2 = DefaultMCPToolHandler._cache_key("https://x/", "Lbl-B", None, None) + k1 = DefaultMCPToolHandler._cache_key("workflow", "https://x/", "Lbl-A", None, None) + k2 = DefaultMCPToolHandler._cache_key("workflow", "https://x/", "Lbl-B", None, None) assert k1 != k2 def test_key_collapses_header_name_case(self) -> None: - k1 = DefaultMCPToolHandler._cache_key("https://x/", None, None, {"Authorization": "tk"}) - k2 = DefaultMCPToolHandler._cache_key("https://x/", None, None, {"authorization": "tk"}) + k1 = DefaultMCPToolHandler._cache_key("workflow", "https://x/", None, None, {"Authorization": "tk"}) + k2 = DefaultMCPToolHandler._cache_key("workflow", "https://x/", None, None, {"authorization": "tk"}) assert k1 == k2 def test_key_keeps_header_value_case(self) -> None: - k1 = DefaultMCPToolHandler._cache_key("https://x/", None, None, {"X": "Bearer-A"}) - k2 = DefaultMCPToolHandler._cache_key("https://x/", None, None, {"X": "bearer-a"}) + k1 = DefaultMCPToolHandler._cache_key("workflow", "https://x/", None, None, {"X": "Bearer-A"}) + k2 = DefaultMCPToolHandler._cache_key("workflow", "https://x/", None, None, {"X": "bearer-a"}) assert k1 != k2 diff --git a/python/packages/declarative/tests/test_invoke_mcp_tool_executor.py b/python/packages/declarative/tests/test_invoke_mcp_tool_executor.py index 28b3a1a0d67..3b450f4a340 100644 --- a/python/packages/declarative/tests/test_invoke_mcp_tool_executor.py +++ b/python/packages/declarative/tests/test_invoke_mcp_tool_executor.py @@ -164,6 +164,52 @@ async def test_basic_invocation_forwards_required_fields(self) -> None: assert inv.headers == {} assert inv.arguments == {} assert inv.connection_name is None + assert inv.workflow_session_id + + @pytest.mark.asyncio + async def test_separate_workflows_receive_separate_session_ids(self) -> None: + handler = StubMcpHandler(_ok()) + factory = WorkflowFactory(mcp_tool_handler=handler) + first = factory.create_workflow_from_definition(_yaml(_action())) + second = factory.create_workflow_from_definition(_yaml(_action())) + + await first.run({}) + await second.run({}) + + assert len(handler.invocations) == 2 + assert handler.invocations[0].workflow_session_id + assert handler.invocations[1].workflow_session_id + assert handler.invocations[0].workflow_session_id != handler.invocations[1].workflow_session_id + + @pytest.mark.asyncio + async def test_fresh_runs_on_same_workflow_receive_separate_session_ids(self) -> None: + handler = StubMcpHandler(_ok()) + factory = WorkflowFactory(mcp_tool_handler=handler) + workflow = factory.create_workflow_from_definition(_yaml(_action())) + + await workflow.run({}) + await workflow.run({}) + + assert len(handler.invocations) == 2 + assert handler.invocations[0].workflow_session_id + assert handler.invocations[0].workflow_session_id != handler.invocations[1].workflow_session_id + + @pytest.mark.asyncio + async def test_continuation_reuses_workflow_session_id(self) -> None: + from agent_framework_declarative._workflows import ToolApprovalResponse + from agent_framework_declarative._workflows._mcp_handler import get_or_create_workflow_session_id + + handler = StubMcpHandler(_ok()) + factory = WorkflowFactory(mcp_tool_handler=handler) + workflow = factory.create_workflow_from_definition(_yaml(_action(require_approval=True))) + + paused = await workflow.run({}) + [approval] = paused.get_request_info_events() + workflow_session_id = get_or_create_workflow_session_id(workflow._runner.state) # pyright: ignore[reportPrivateUsage] + await workflow.run(responses={approval.request_id: ToolApprovalResponse(approved=True)}) + + assert handler.last_invocation is not None + assert handler.last_invocation.workflow_session_id == workflow_session_id @pytest.mark.asyncio async def test_arguments_evaluated_and_preserves_none(self) -> None: