From cad66f4867a8e94348837ba79d96e9b62fcbe8c4 Mon Sep 17 00:00:00 2001 From: King Star Date: Mon, 20 Jul 2026 11:41:11 +0800 Subject: [PATCH 1/2] Cancel background task runners on server disposal --- .../Server/DestinationBoundMcpServer.cs | 8 +- .../Server/IMcpServerLifetimeFeature.cs | 22 ++ .../Server/McpServerImpl.cs | 67 +++++- .../TaskCancellationIntegrationTests.cs | 212 ++++++++++++++++++ 4 files changed, 305 insertions(+), 4 deletions(-) create mode 100644 src/ModelContextProtocol.Core/Server/IMcpServerLifetimeFeature.cs diff --git a/src/ModelContextProtocol.Core/Server/DestinationBoundMcpServer.cs b/src/ModelContextProtocol.Core/Server/DestinationBoundMcpServer.cs index 7aab34826..05dd78c53 100644 --- a/src/ModelContextProtocol.Core/Server/DestinationBoundMcpServer.cs +++ b/src/ModelContextProtocol.Core/Server/DestinationBoundMcpServer.cs @@ -5,7 +5,7 @@ namespace ModelContextProtocol.Server; #pragma warning disable MCPEXP002 -internal sealed class DestinationBoundMcpServer(McpServerImpl server, ITransport? transport, JsonRpcMessageContext? requestContext = null) : McpServer +internal sealed class DestinationBoundMcpServer(McpServerImpl server, ITransport? transport, JsonRpcMessageContext? requestContext = null) : McpServer, IMcpServerLifetimeFeature #pragma warning restore MCPEXP002 { private readonly bool _isJuly2026OrLaterRequest = server.IsJuly2026OrLaterProtocolRequest(requestContext); @@ -73,6 +73,12 @@ public override Implementation? ClientInfo public override bool IsMrtrSupported => server.IsMrtrSupported; + CancellationToken IMcpServerLifetimeFeature.BackgroundTaskCancellationToken => + ((IMcpServerLifetimeFeature)server).BackgroundTaskCancellationToken; + + void IMcpServerLifetimeFeature.RegisterBackgroundTask(Task backgroundTask) => + ((IMcpServerLifetimeFeature)server).RegisterBackgroundTask(backgroundTask); + public override ValueTask DisposeAsync() => server.DisposeAsync(); public override IAsyncDisposable RegisterNotificationHandler(string method, Func handler) => server.RegisterNotificationHandler(method, handler); diff --git a/src/ModelContextProtocol.Core/Server/IMcpServerLifetimeFeature.cs b/src/ModelContextProtocol.Core/Server/IMcpServerLifetimeFeature.cs new file mode 100644 index 000000000..a70f2aed4 --- /dev/null +++ b/src/ModelContextProtocol.Core/Server/IMcpServerLifetimeFeature.cs @@ -0,0 +1,22 @@ +using System.ComponentModel; + +namespace ModelContextProtocol.Server; + +/// +/// Provides server-lifetime services used by MCP extension infrastructure. +/// +[EditorBrowsable(EditorBrowsableState.Never)] +public interface IMcpServerLifetimeFeature +{ + /// Gets the token that should cancel background work owned by this server. + /// + /// The token is when background work intentionally outlives + /// the server instance, as it does for per-request servers in stateless HTTP mode. + /// + CancellationToken BackgroundTaskCancellationToken { get; } + + /// Registers background work that server disposal must await. + /// The background work to track. + /// This is a no-op when background work intentionally outlives the server instance. + void RegisterBackgroundTask(Task backgroundTask); +} diff --git a/src/ModelContextProtocol.Core/Server/McpServerImpl.cs b/src/ModelContextProtocol.Core/Server/McpServerImpl.cs index 2ce838713..c6e7a2d58 100644 --- a/src/ModelContextProtocol.Core/Server/McpServerImpl.cs +++ b/src/ModelContextProtocol.Core/Server/McpServerImpl.cs @@ -12,7 +12,7 @@ namespace ModelContextProtocol.Server; /// #pragma warning disable MCPEXP001, MCPEXP002 -internal sealed partial class McpServerImpl : McpServer +internal sealed partial class McpServerImpl : McpServer, IMcpServerLifetimeFeature { internal static Implementation DefaultImplementation { get; } = new() { @@ -31,6 +31,9 @@ internal sealed partial class McpServerImpl : McpServer private readonly string[] _initializeHandshakeProtocolVersions; private readonly string[] _perRequestMetadataProtocolVersions; private readonly SemaphoreSlim _disposeLock = new(1, 1); + private readonly CancellationTokenSource _serverLifetimeCts = new(); + private readonly object _backgroundTasksLock = new(); + private readonly ConcurrentDictionary _backgroundTasks = new(); private readonly ConcurrentDictionary _mrtrContinuations = new(); private readonly ConcurrentDictionary _mrtrContextsByRequestId = new(); private static readonly string[] s_perRequestMetadataKeys = @@ -55,6 +58,7 @@ internal sealed partial class McpServerImpl : McpServer private int _started; private bool _disposed; + private bool _backgroundTaskRegistrationClosed; /// Holds a boxed value for the server. /// @@ -597,6 +601,38 @@ public override Task SendMessageAsync(JsonRpcMessage message, CancellationToken public override IAsyncDisposable RegisterNotificationHandler(string method, Func handler) => _sessionHandler.RegisterNotificationHandler(method, handler); + CancellationToken IMcpServerLifetimeFeature.BackgroundTaskCancellationToken => + HasStatefulTransport() ? _serverLifetimeCts.Token : CancellationToken.None; + + void IMcpServerLifetimeFeature.RegisterBackgroundTask(Task backgroundTask) + { + Throw.IfNull(backgroundTask); + + // Stateless HTTP servers are request-scoped, while Tasks runners intentionally outlive + // the originating request and are governed by tasks/cancel and task-store retention. + if (!HasStatefulTransport()) + { + return; + } + + lock (_backgroundTasksLock) + { + if (_backgroundTaskRegistrationClosed) + { + throw new ObjectDisposedException(nameof(McpServer)); + } + + _backgroundTasks.TryAdd(backgroundTask, 0); + } + + _ = backgroundTask.ContinueWith( + static (task, state) => ((ConcurrentDictionary)state!).TryRemove(task, out _), + _backgroundTasks, + CancellationToken.None, + TaskContinuationOptions.ExecuteSynchronously, + TaskScheduler.Default); + } + /// public override async ValueTask DisposeAsync() { @@ -608,6 +644,7 @@ public override async ValueTask DisposeAsync() } _disposed = true; + _serverLifetimeCts.Cancel(); // Dispose the session handler - cancels message processing and waits for all // in-flight request handlers (including retries in AwaitMrtrHandlerAsync) to complete. @@ -616,6 +653,13 @@ public override async ValueTask DisposeAsync() _disposables.ForEach(d => d()); await _sessionHandler.DisposeAsync().ConfigureAwait(false); + Task[] backgroundTasks; + lock (_backgroundTasksLock) + { + _backgroundTaskRegistrationClosed = true; + backgroundTasks = [.. _backgroundTasks.Keys]; + } + // Cancel all orphaned MRTR handlers still suspended in continuations (waiting for // retries that will never arrive now that the session handler is disposed). int cancelledCount = _mrtrContinuations.Count; @@ -637,6 +681,11 @@ public override async ValueTask DisposeAsync() { await _allMrtrHandlersCompleted.Task.ConfigureAwait(false); } + + if (backgroundTasks.Length > 0) + { + await Task.WhenAll(backgroundTasks).ConfigureAwait(false); + } } private void ConfigureInitialize(McpServerOptions options) @@ -2379,6 +2428,9 @@ private void WrapHandlerWithMrtr(string method) // is thread-safe with itself, and not disposing avoids deadlock risks from // calling Cancel/Dispose inside locks or Interlocked guards. var handlerCts = CancellationTokenSource.CreateLinkedTokenSource(cancellationToken); + var serverLifetimeRegistration = _serverLifetimeCts.Token.Register( + static state => ((CancellationTokenSource)state!).Cancel(), + handlerCts); // Store the MrtrContext so CreateDestinationBoundServer can pick it up and set it // on the per-request DestinationBoundMcpServer. This is picked up synchronously @@ -2389,6 +2441,11 @@ private void WrapHandlerWithMrtr(string method) { handlerTask = originalHandler(request, handlerCts.Token); } + catch + { + serverLifetimeRegistration.Dispose(); + throw; + } finally { _mrtrContextsByRequestId.TryRemove(request.Id, out _); @@ -2401,7 +2458,7 @@ private void WrapHandlerWithMrtr(string method) // exceptions and decrements _mrtrInFlightCount when the handler completes, // mirroring how McpSessionHandler tracks in-flight handlers. Interlocked.Increment(ref _mrtrInFlightCount); - _ = ObserveHandlerCompletionAsync(handlerTask); + _ = ObserveHandlerCompletionAsync(handlerTask, serverLifetimeRegistration); return await AwaitMrtrHandlerAsync( handlerTask, continuation, mrtrContext.InitialExchangeTask, cancellationToken).ConfigureAwait(false); @@ -2460,7 +2517,9 @@ private void WrapHandlerWithMrtr(string method) /// double-reporting at Error) and decrements when the /// handler completes, following the same in-flight tracking pattern as . /// - private async Task ObserveHandlerCompletionAsync(Task handlerTask) + private async Task ObserveHandlerCompletionAsync( + Task handlerTask, + CancellationTokenRegistration serverLifetimeRegistration) { try { @@ -2480,6 +2539,8 @@ private async Task ObserveHandlerCompletionAsync(Task handlerTask) } finally { + serverLifetimeRegistration.Dispose(); + if (Interlocked.Decrement(ref _mrtrInFlightCount) == 0) { _allMrtrHandlersCompleted.TrySetResult(true); diff --git a/tests/ModelContextProtocol.Tests/Server/TaskCancellationIntegrationTests.cs b/tests/ModelContextProtocol.Tests/Server/TaskCancellationIntegrationTests.cs index e751bdcdc..e6041ceea 100644 --- a/tests/ModelContextProtocol.Tests/Server/TaskCancellationIntegrationTests.cs +++ b/tests/ModelContextProtocol.Tests/Server/TaskCancellationIntegrationTests.cs @@ -109,6 +109,218 @@ public async Task TaskTool_CancellationToken_GetTaskShowsWorkingBeforeCancel() } } +/// +/// Tests for task-store runner cleanup during server disposal. +/// +public class TaskRunnerLifecycleTests : ClientServerTestBase +{ + private readonly TaskCompletionSource _toolStarted = new(TaskCreationOptions.RunContinuationsAsynchronously); + private readonly TaskCompletionSource _toolCancellationFired = new(TaskCreationOptions.RunContinuationsAsynchronously); + private readonly TaskCompletionSource _releaseCancellationCleanup = new(TaskCreationOptions.RunContinuationsAsynchronously); + private readonly TaskCompletionSource _forceToolExit = new(TaskCreationOptions.RunContinuationsAsynchronously); + private readonly TaskCompletionSource _toolExited = new(TaskCreationOptions.RunContinuationsAsynchronously); + private readonly TaskCompletionSource _runnerRegistrationBlocked = new(TaskCreationOptions.RunContinuationsAsynchronously); + private readonly TaskCompletionSource _releaseRunnerRegistration = new(TaskCreationOptions.RunContinuationsAsynchronously); + private readonly BlockingCancellationTaskStore _taskStore = new(); + private bool _delayRunnerRegistration; + + public TaskRunnerLifecycleTests(ITestOutputHelper testOutputHelper) + : base(testOutputHelper) + { +#if !NET + Assert.SkipWhen(RuntimeInformation.IsOSPlatform(OSPlatform.Windows), "https://github.com/modelcontextprotocol/csharp-sdk/issues/587"); +#endif + } + + protected override void ConfigureServices(ServiceCollection services, IMcpServerBuilder mcpServerBuilder) + { +#pragma warning disable MCPEXP002 + services.Configure(options => + options.Filters.Request.CallToolWithAlternateFilters.Add(next => async (request, cancellationToken) => + { + if (_delayRunnerRegistration && request.Params?.Name == "lifecycle-tool") + { + _runnerRegistrationBlocked.TrySetResult(true); + await _releaseRunnerRegistration.Task; + } + + return await next(request, cancellationToken); + })); +#pragma warning restore MCPEXP002 + + mcpServerBuilder + .WithTasks(_taskStore) + .WithTools([McpServerTool.Create( + async (CancellationToken ct) => + { + _toolStarted.TrySetResult(true); + try + { + var cancellationTask = Task.Delay(Timeout.Infinite, ct); + var completedTask = await Task.WhenAny(cancellationTask, _forceToolExit.Task); + await completedTask; + return "forced test cleanup"; + } + catch (OperationCanceledException) when (ct.IsCancellationRequested) + { + _toolCancellationFired.TrySetResult(true); + await _releaseCancellationCleanup.Task; + throw; + } + finally + { + _toolExited.TrySetResult(true); + } + }, + new McpServerToolCreateOptions + { + Name = "lifecycle-tool", + Description = "A tool used to verify task runner lifecycle" + })]); + } + + [Fact] + public async Task DisposeAsync_CancelsAndWaitsForTaskStoreRunner() + { + await using var client = await CreateMcpClientForServer(); + var ct = TestContext.Current.CancellationToken; + + var augmented = await client.CallToolAsTaskAsync( + new CallToolRequestParams { Name = "lifecycle-tool" }, ct); + Assert.True(augmented.IsTask); + + await _toolStarted.Task.WaitAsync(TestConstants.DefaultTimeout, ct); + Task disposeTask = Server.DisposeAsync().AsTask(); + + try + { + Task firstCompleted = await Task.WhenAny(_toolCancellationFired.Task, disposeTask) + .WaitAsync(TestConstants.DefaultTimeout, ct); + + Assert.Same(_toolCancellationFired.Task, firstCompleted); + Assert.False(disposeTask.IsCompleted, "DisposeAsync should wait for the runner's cancellation cleanup."); + + _releaseCancellationCleanup.TrySetResult(true); + await disposeTask.WaitAsync(TestConstants.DefaultTimeout, ct); + } + finally + { + _releaseCancellationCleanup.TrySetResult(true); + _forceToolExit.TrySetResult(true); + await _toolExited.Task.WaitAsync(TestConstants.DefaultTimeout, ct); + } + } + + [Fact] + public async Task DisposeAsync_CancelsAndWaitsForRunnerRegisteredDuringDisposal() + { + await using var client = await CreateMcpClientForServer(); + var ct = TestContext.Current.CancellationToken; + _delayRunnerRegistration = true; + _taskStore.PauseCancellationRecording(); + + var callTask = client.CallToolAsTaskAsync( + new CallToolRequestParams { Name = "lifecycle-tool" }, ct).AsTask(); + _ = callTask.ContinueWith( + static task => _ = task.Exception, + CancellationToken.None, + TaskContinuationOptions.OnlyOnFaulted | TaskContinuationOptions.ExecuteSynchronously, + TaskScheduler.Default); + await _runnerRegistrationBlocked.Task.WaitAsync(TestConstants.DefaultTimeout, ct); + + var serverLifetime = Assert.IsAssignableFrom(Server); + var serverCancellationFired = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + using var registration = serverLifetime.BackgroundTaskCancellationToken.Register( + static state => ((TaskCompletionSource)state!).TrySetResult(true), serverCancellationFired); + + Task disposeTask = Server.DisposeAsync().AsTask(); + + try + { + await serverCancellationFired.Task.WaitAsync(TestConstants.DefaultTimeout, ct); + _releaseRunnerRegistration.TrySetResult(true); + + await _taskStore.CancellationRecordingStarted.WaitAsync(TestConstants.DefaultTimeout, ct); + Assert.False(disposeTask.IsCompleted, "DisposeAsync should wait for a runner registered during disposal."); + + _taskStore.ReleaseCancellationRecording(); + await disposeTask.WaitAsync(TestConstants.DefaultTimeout, ct); + } + finally + { + _releaseRunnerRegistration.TrySetResult(true); + _releaseCancellationCleanup.TrySetResult(true); + _taskStore.ReleaseCancellationRecording(); + _forceToolExit.TrySetResult(true); + + if (_toolStarted.Task.IsCompleted) + { + await _toolExited.Task.WaitAsync(TestConstants.DefaultTimeout, ct); + } + } + } + + private sealed class BlockingCancellationTaskStore : InMemoryMcpTaskStore, IMcpTaskStore + { + private readonly TaskCompletionSource _cancellationRecordingStarted = new(TaskCreationOptions.RunContinuationsAsynchronously); + private readonly TaskCompletionSource _releaseCancellationRecording = new(TaskCreationOptions.RunContinuationsAsynchronously); + private bool _pauseCancellationRecording; + + public Task CancellationRecordingStarted => _cancellationRecordingStarted.Task; + + public void PauseCancellationRecording() => _pauseCancellationRecording = true; + + public void ReleaseCancellationRecording() => _releaseCancellationRecording.TrySetResult(true); + + async Task IMcpTaskStore.SetCancelledAsync(string taskId, CancellationToken cancellationToken) + { + if (_pauseCancellationRecording) + { + _cancellationRecordingStarted.TrySetResult(true); + await _releaseCancellationRecording.Task; + } + + return await base.SetCancelledAsync(taskId, cancellationToken); + } + } +} + +public class McpServerLifetimeFeatureTests(ITestOutputHelper testOutputHelper) : LoggedTest(testOutputHelper) +{ + [Fact] + public async Task DisposeAsync_DoesNotCancelOrWaitForStatelessBackgroundTask() + { + await using var transport = new StreamableHttpServerTransport { Stateless = true }; + await using var statelessServer = McpServer.Create( + transport, + new McpServerOptions + { + ServerInfo = new Implementation { Name = "test-server", Version = "1.0" }, + }, + LoggerFactory); + var serverLifetime = Assert.IsAssignableFrom(statelessServer); + var releaseBackgroundTask = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + Task backgroundTask = releaseBackgroundTask.Task; + + serverLifetime.RegisterBackgroundTask(backgroundTask); + + try + { + await statelessServer.DisposeAsync().AsTask() + .WaitAsync(TestConstants.DefaultTimeout, TestContext.Current.CancellationToken); + + Assert.False(serverLifetime.BackgroundTaskCancellationToken.CanBeCanceled); + Assert.False(backgroundTask.IsCompleted, + "A stateless per-request server should not own background work that outlives the request."); + } + finally + { + releaseBackgroundTask.TrySetResult(true); + await backgroundTask; + } + } +} + /// /// Tests for task cancellation with multiple concurrent tasks. /// From db7f5c3042103defe06c3707fe5e6898cee75952 Mon Sep 17 00:00:00 2001 From: King Star Date: Mon, 27 Jul 2026 17:53:32 +0800 Subject: [PATCH 2/2] refactor(server): generalize lifetime registrations Signed-off-by: King Star --- .../Server/DestinationBoundMcpServer.cs | 8 +- .../Server/IMcpServerLifetimeFeature.cs | 20 ++-- .../Server/McpServerImpl.cs | 73 +++++++++---- .../Server/McpTasksBuilderExtensions.cs | 95 ++++++++++++++--- .../TaskCancellationIntegrationTests.cs | 100 +++++++++++++++--- 5 files changed, 234 insertions(+), 62 deletions(-) diff --git a/src/ModelContextProtocol.Core/Server/DestinationBoundMcpServer.cs b/src/ModelContextProtocol.Core/Server/DestinationBoundMcpServer.cs index 05dd78c53..e806b78ce 100644 --- a/src/ModelContextProtocol.Core/Server/DestinationBoundMcpServer.cs +++ b/src/ModelContextProtocol.Core/Server/DestinationBoundMcpServer.cs @@ -73,11 +73,11 @@ public override Implementation? ClientInfo public override bool IsMrtrSupported => server.IsMrtrSupported; - CancellationToken IMcpServerLifetimeFeature.BackgroundTaskCancellationToken => - ((IMcpServerLifetimeFeature)server).BackgroundTaskCancellationToken; + CancellationToken IMcpServerLifetimeFeature.ServerCancellationToken => + ((IMcpServerLifetimeFeature)server).ServerCancellationToken; - void IMcpServerLifetimeFeature.RegisterBackgroundTask(Task backgroundTask) => - ((IMcpServerLifetimeFeature)server).RegisterBackgroundTask(backgroundTask); + IDisposable IMcpServerLifetimeFeature.RegisterForDisposeAsync(IAsyncDisposable disposable) => + ((IMcpServerLifetimeFeature)server).RegisterForDisposeAsync(disposable); public override ValueTask DisposeAsync() => server.DisposeAsync(); diff --git a/src/ModelContextProtocol.Core/Server/IMcpServerLifetimeFeature.cs b/src/ModelContextProtocol.Core/Server/IMcpServerLifetimeFeature.cs index a70f2aed4..c4eb157d9 100644 --- a/src/ModelContextProtocol.Core/Server/IMcpServerLifetimeFeature.cs +++ b/src/ModelContextProtocol.Core/Server/IMcpServerLifetimeFeature.cs @@ -8,15 +8,19 @@ namespace ModelContextProtocol.Server; [EditorBrowsable(EditorBrowsableState.Never)] public interface IMcpServerLifetimeFeature { - /// Gets the token that should cancel background work owned by this server. + /// Gets the token that is cancelled when this server starts disposing. /// - /// The token is when background work intentionally outlives - /// the server instance, as it does for per-request servers in stateless HTTP mode. + /// The token is when work intentionally outlives the server, + /// as it does for per-request servers in stateless HTTP mode. /// - CancellationToken BackgroundTaskCancellationToken { get; } + CancellationToken ServerCancellationToken { get; } - /// Registers background work that server disposal must await. - /// The background work to track. - /// This is a no-op when background work intentionally outlives the server instance. - void RegisterBackgroundTask(Task backgroundTask); + /// Registers an asynchronously disposable resource that server disposal must await. + /// The resource to dispose when this server is disposed. + /// A handle that unregisters the resource without disposing it. + /// + /// Dispose the returned handle when the resource completes independently so the server does not + /// retain it until shutdown. Registration is a no-op when the server does not own the resource. + /// + IDisposable RegisterForDisposeAsync(IAsyncDisposable disposable); } diff --git a/src/ModelContextProtocol.Core/Server/McpServerImpl.cs b/src/ModelContextProtocol.Core/Server/McpServerImpl.cs index c6e7a2d58..9ed460954 100644 --- a/src/ModelContextProtocol.Core/Server/McpServerImpl.cs +++ b/src/ModelContextProtocol.Core/Server/McpServerImpl.cs @@ -32,8 +32,8 @@ internal sealed partial class McpServerImpl : McpServer, IMcpServerLifetimeFeatu private readonly string[] _perRequestMetadataProtocolVersions; private readonly SemaphoreSlim _disposeLock = new(1, 1); private readonly CancellationTokenSource _serverLifetimeCts = new(); - private readonly object _backgroundTasksLock = new(); - private readonly ConcurrentDictionary _backgroundTasks = new(); + private readonly object _serverLifetimeRegistrationsLock = new(); + private readonly HashSet _serverLifetimeRegistrations = []; private readonly ConcurrentDictionary _mrtrContinuations = new(); private readonly ConcurrentDictionary _mrtrContextsByRequestId = new(); private static readonly string[] s_perRequestMetadataKeys = @@ -58,7 +58,7 @@ internal sealed partial class McpServerImpl : McpServer, IMcpServerLifetimeFeatu private int _started; private bool _disposed; - private bool _backgroundTaskRegistrationClosed; + private bool _serverLifetimeRegistrationClosed; /// Holds a boxed value for the server. /// @@ -601,36 +601,40 @@ public override Task SendMessageAsync(JsonRpcMessage message, CancellationToken public override IAsyncDisposable RegisterNotificationHandler(string method, Func handler) => _sessionHandler.RegisterNotificationHandler(method, handler); - CancellationToken IMcpServerLifetimeFeature.BackgroundTaskCancellationToken => + CancellationToken IMcpServerLifetimeFeature.ServerCancellationToken => HasStatefulTransport() ? _serverLifetimeCts.Token : CancellationToken.None; - void IMcpServerLifetimeFeature.RegisterBackgroundTask(Task backgroundTask) + IDisposable IMcpServerLifetimeFeature.RegisterForDisposeAsync(IAsyncDisposable disposable) { - Throw.IfNull(backgroundTask); + Throw.IfNull(disposable); // Stateless HTTP servers are request-scoped, while Tasks runners intentionally outlive // the originating request and are governed by tasks/cancel and task-store retention. if (!HasStatefulTransport()) { - return; + return NoopRegistration.Instance; } - lock (_backgroundTasksLock) + var registration = new ServerLifetimeRegistration(this, disposable); + lock (_serverLifetimeRegistrationsLock) { - if (_backgroundTaskRegistrationClosed) + if (_serverLifetimeRegistrationClosed) { throw new ObjectDisposedException(nameof(McpServer)); } - _backgroundTasks.TryAdd(backgroundTask, 0); + _serverLifetimeRegistrations.Add(registration); } - _ = backgroundTask.ContinueWith( - static (task, state) => ((ConcurrentDictionary)state!).TryRemove(task, out _), - _backgroundTasks, - CancellationToken.None, - TaskContinuationOptions.ExecuteSynchronously, - TaskScheduler.Default); + return registration; + } + + private void UnregisterServerLifetime(ServerLifetimeRegistration registration) + { + lock (_serverLifetimeRegistrationsLock) + { + _serverLifetimeRegistrations.Remove(registration); + } } /// @@ -653,11 +657,11 @@ public override async ValueTask DisposeAsync() _disposables.ForEach(d => d()); await _sessionHandler.DisposeAsync().ConfigureAwait(false); - Task[] backgroundTasks; - lock (_backgroundTasksLock) + ServerLifetimeRegistration[] serverLifetimeRegistrations; + lock (_serverLifetimeRegistrationsLock) { - _backgroundTaskRegistrationClosed = true; - backgroundTasks = [.. _backgroundTasks.Keys]; + _serverLifetimeRegistrationClosed = true; + serverLifetimeRegistrations = [.. _serverLifetimeRegistrations]; } // Cancel all orphaned MRTR handlers still suspended in continuations (waiting for @@ -682,9 +686,34 @@ public override async ValueTask DisposeAsync() await _allMrtrHandlersCompleted.Task.ConfigureAwait(false); } - if (backgroundTasks.Length > 0) + if (serverLifetimeRegistrations.Length > 0) + { + await Task.WhenAll( + serverLifetimeRegistrations.Select(static registration => registration.DisposeResourceAsync().AsTask()) + ).ConfigureAwait(false); + } + } + + private sealed class ServerLifetimeRegistration( + McpServerImpl server, + IAsyncDisposable resource) : IDisposable + { + private McpServerImpl? _server = server; + + public ValueTask DisposeResourceAsync() => resource.DisposeAsync(); + + public void Dispose() + { + Interlocked.Exchange(ref _server, null)?.UnregisterServerLifetime(this); + } + } + + private sealed class NoopRegistration : IDisposable + { + public static NoopRegistration Instance { get; } = new(); + + public void Dispose() { - await Task.WhenAll(backgroundTasks).ConfigureAwait(false); } } diff --git a/src/ModelContextProtocol.Extensions.Tasks/Server/McpTasksBuilderExtensions.cs b/src/ModelContextProtocol.Extensions.Tasks/Server/McpTasksBuilderExtensions.cs index e61466a69..c5379c787 100644 --- a/src/ModelContextProtocol.Extensions.Tasks/Server/McpTasksBuilderExtensions.cs +++ b/src/ModelContextProtocol.Extensions.Tasks/Server/McpTasksBuilderExtensions.cs @@ -82,7 +82,7 @@ private sealed class McpTasksConfigureOptions( private readonly IServiceScopeFactory _serviceScopeFactory = serviceScopeFactory; private readonly ILogger _logger = (loggerFactory ?? NullLoggerFactory.Instance).CreateLogger(); private readonly McpTasksOptions _taskOptions = taskOptions; - private readonly ConcurrentDictionary _cancellationSources = new(StringComparer.Ordinal); + private readonly ConcurrentDictionary _cancellationStates = new(StringComparer.Ordinal); public void Configure(McpServerOptions options) { @@ -183,15 +183,27 @@ private async ValueTask> RunAsTaskAsync( var taskId = taskInfo.TaskId; executionRequest.Server = request.Server.WithMcpTaskOutgoingRequestInterceptor(taskId, _store); - var cts = new CancellationTokenSource(); - _cancellationSources[taskId] = cts; + var serverLifetime = request.Server as IMcpServerLifetimeFeature; + var cancellationState = new TaskCancellationState( + serverLifetime?.ServerCancellationToken ?? CancellationToken.None); + _cancellationStates[taskId] = cancellationState; - // Capture the token before dispatching. Cancellation can remove and dispose the source - // before the background delegate starts. - var taskCancellationToken = cts.Token; - _ = Task.Run( + var taskCancellationToken = cancellationState.Token; + var backgroundTask = Task.Run( () => ExecuteTaskAsync(next, executionRequest, taskId, taskCancellationToken, executionScope), CancellationToken.None); + cancellationState.SetBackgroundTask(backgroundTask); + try + { + cancellationState.SetServerLifetimeRegistration( + serverLifetime?.RegisterForDisposeAsync(cancellationState)); + } + catch + { + cancellationState.Cancel(); + await backgroundTask.ConfigureAwait(false); + throw; + } return ResultOrAlternate.FromAlternate( ToCreateTaskResult(taskInfo), @@ -235,9 +247,9 @@ private async Task ExecuteTaskAsync( } finally { - if (_cancellationSources.TryRemove(taskId, out var registeredCts)) + if (_cancellationStates.TryRemove(taskId, out var registeredState)) { - registeredCts.Dispose(); + registeredState.UnregisterServerLifetime(); } } } @@ -354,15 +366,74 @@ private async Task ExecuteToolPipelineAsync( await _store.SetCancelledAsync(requestParams.TaskId, cancellationToken).ConfigureAwait(false); - if (_cancellationSources.TryRemove(requestParams.TaskId, out var cts)) + if (_cancellationStates.TryGetValue(requestParams.TaskId, out var cancellationState)) { - cts.Cancel(); - cts.Dispose(); + cancellationState.Cancel(); } return JsonSerializer.SerializeToNode(new CancelTaskResult(), McpTasksJsonContext.Default.CancelTaskResult); } + private sealed class TaskCancellationState : IAsyncDisposable + { + private readonly CancellationTokenSource _source = new(); + private readonly CancellationTokenRegistration _serverLifetimeRegistration; + private Task? _backgroundTask; + private IDisposable? _serverLifetimeUnregistration; + private int _completed; + + public TaskCancellationState(CancellationToken serverLifetimeToken) + { + _serverLifetimeRegistration = serverLifetimeToken.Register( + static state => ((TaskCancellationState)state!).Cancel(), + this); + } + + public CancellationToken Token => _source.Token; + + public void Cancel() => _source.Cancel(); + + public void SetBackgroundTask(Task backgroundTask) => _backgroundTask = backgroundTask; + + public void SetServerLifetimeRegistration(IDisposable? registration) + { + if (registration is null) + { + return; + } + + if (Volatile.Read(ref _completed) != 0) + { + registration.Dispose(); + return; + } + + Interlocked.CompareExchange(ref _serverLifetimeUnregistration, registration, null); + if (Volatile.Read(ref _completed) != 0) + { + Interlocked.Exchange(ref _serverLifetimeUnregistration, null)?.Dispose(); + } + } + + public void UnregisterServerLifetime() + { + // Cancellation can arrive concurrently from tasks/cancel and server disposal. + // Once the dictionary entry and server registration are gone, the CTS is collectible. + Interlocked.Exchange(ref _completed, 1); + _serverLifetimeRegistration.Dispose(); + Interlocked.Exchange(ref _serverLifetimeUnregistration, null)?.Dispose(); + } + + public async ValueTask DisposeAsync() + { + Cancel(); + if (_backgroundTask is { } backgroundTask) + { + await backgroundTask.ConfigureAwait(false); + } + } + } + private static void GateToJuly2026OrLaterProtocol(JsonRpcRequest request, string method) { if (!IsJuly2026OrLaterProtocolRequest(request)) diff --git a/tests/ModelContextProtocol.Tests/Server/TaskCancellationIntegrationTests.cs b/tests/ModelContextProtocol.Tests/Server/TaskCancellationIntegrationTests.cs index e6041ceea..efdd6a590 100644 --- a/tests/ModelContextProtocol.Tests/Server/TaskCancellationIntegrationTests.cs +++ b/tests/ModelContextProtocol.Tests/Server/TaskCancellationIntegrationTests.cs @@ -136,7 +136,7 @@ protected override void ConfigureServices(ServiceCollection services, IMcpServer { #pragma warning disable MCPEXP002 services.Configure(options => - options.Filters.Request.CallToolWithAlternateFilters.Add(next => async (request, cancellationToken) => + options.Filters.Request.CallToolWithAlternateFilters.Add(async (request, next, cancellationToken) => { if (_delayRunnerRegistration && request.Params?.Name == "lifecycle-tool") { @@ -179,6 +179,45 @@ protected override void ConfigureServices(ServiceCollection services, IMcpServer })]); } + [Fact] + public async Task DisposeAsync_DisposesAndWaitsForRegisteredLifetimeResource() + { + await using var client = await CreateMcpClientForServer(); + var ct = TestContext.Current.CancellationToken; + var serverLifetime = Assert.IsAssignableFrom(Server); + var resource = new BlockingAsyncDisposable(); + using var registration = serverLifetime.RegisterForDisposeAsync(resource); + + Task disposeTask = Server.DisposeAsync().AsTask(); + + try + { + await resource.DisposeStarted.WaitAsync(TestConstants.DefaultTimeout, ct); + Assert.False(disposeTask.IsCompleted, "DisposeAsync should await the registered resource."); + + resource.Release(); + await disposeTask.WaitAsync(TestConstants.DefaultTimeout, ct); + } + finally + { + resource.Release(); + } + } + + [Fact] + public async Task LifetimeRegistration_DisposeUnregistersResource() + { + await using var client = await CreateMcpClientForServer(); + var serverLifetime = Assert.IsAssignableFrom(Server); + var resource = new RecordingAsyncDisposable(); + using var registration = serverLifetime.RegisterForDisposeAsync(resource); + + registration.Dispose(); + await Server.DisposeAsync(); + + Assert.False(resource.IsDisposed); + } + [Fact] public async Task DisposeAsync_CancelsAndWaitsForTaskStoreRunner() { @@ -230,7 +269,7 @@ public async Task DisposeAsync_CancelsAndWaitsForRunnerRegisteredDuringDisposal( var serverLifetime = Assert.IsAssignableFrom(Server); var serverCancellationFired = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); - using var registration = serverLifetime.BackgroundTaskCancellationToken.Register( + using var registration = serverLifetime.ServerCancellationToken.Register( static state => ((TaskCompletionSource)state!).TrySetResult(true), serverCancellationFired); Task disposeTask = Server.DisposeAsync().AsTask(); @@ -283,6 +322,33 @@ async Task IMcpTaskStore.SetCancelledAsync(string taskId, CancellationToke return await base.SetCancelledAsync(taskId, cancellationToken); } } + + private sealed class BlockingAsyncDisposable : IAsyncDisposable + { + private readonly TaskCompletionSource _disposeStarted = new(TaskCreationOptions.RunContinuationsAsynchronously); + private readonly TaskCompletionSource _release = new(TaskCreationOptions.RunContinuationsAsynchronously); + + public Task DisposeStarted => _disposeStarted.Task; + + public void Release() => _release.TrySetResult(true); + + public async ValueTask DisposeAsync() + { + _disposeStarted.TrySetResult(true); + await _release.Task; + } + } + + private sealed class RecordingAsyncDisposable : IAsyncDisposable + { + public bool IsDisposed { get; private set; } + + public ValueTask DisposeAsync() + { + IsDisposed = true; + return default; + } + } } public class McpServerLifetimeFeatureTests(ITestOutputHelper testOutputHelper) : LoggedTest(testOutputHelper) @@ -299,24 +365,26 @@ public async Task DisposeAsync_DoesNotCancelOrWaitForStatelessBackgroundTask() }, LoggerFactory); var serverLifetime = Assert.IsAssignableFrom(statelessServer); - var releaseBackgroundTask = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); - Task backgroundTask = releaseBackgroundTask.Task; + var backgroundResource = new RecordingAsyncDisposable(); - serverLifetime.RegisterBackgroundTask(backgroundTask); + using var registration = serverLifetime.RegisterForDisposeAsync(backgroundResource); - try - { - await statelessServer.DisposeAsync().AsTask() - .WaitAsync(TestConstants.DefaultTimeout, TestContext.Current.CancellationToken); + await statelessServer.DisposeAsync().AsTask() + .WaitAsync(TestConstants.DefaultTimeout, TestContext.Current.CancellationToken); - Assert.False(serverLifetime.BackgroundTaskCancellationToken.CanBeCanceled); - Assert.False(backgroundTask.IsCompleted, - "A stateless per-request server should not own background work that outlives the request."); - } - finally + Assert.False(serverLifetime.ServerCancellationToken.CanBeCanceled); + Assert.False(backgroundResource.IsDisposed, + "A stateless per-request server should not own background work that outlives the request."); + } + + private sealed class RecordingAsyncDisposable : IAsyncDisposable + { + public bool IsDisposed { get; private set; } + + public ValueTask DisposeAsync() { - releaseBackgroundTask.TrySetResult(true); - await backgroundTask; + IsDisposed = true; + return default; } } }