From 9ba50aa092ad965923318046a0517bd0d3fd4bc1 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 +++++- .../Server/McpTasksBuilderExtensions.cs | 58 +++-- .../TaskCancellationIntegrationTests.cs | 212 ++++++++++++++++++ 5 files changed, 345 insertions(+), 22 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 8b38421c4..e6febfea0 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(); @@ -56,6 +59,7 @@ internal sealed partial class McpServerImpl : McpServer private int _started; private bool _disposed; + private bool _backgroundTaskRegistrationClosed; /// Holds a boxed value for the server. /// @@ -505,6 +509,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() { @@ -516,6 +552,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. @@ -524,6 +561,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; @@ -545,6 +589,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) @@ -2124,6 +2173,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 @@ -2134,6 +2186,11 @@ private void WrapHandlerWithMrtr(string method) { handlerTask = originalHandler(request, handlerCts.Token); } + catch + { + serverLifetimeRegistration.Dispose(); + throw; + } finally { _mrtrContextsByRequestId.TryRemove(request.Id, out _); @@ -2146,7 +2203,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); @@ -2205,7 +2262,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 { @@ -2225,6 +2284,8 @@ private async Task ObserveHandlerCompletionAsync(Task handlerTask) } finally { + serverLifetimeRegistration.Dispose(); + if (Interlocked.Decrement(ref _mrtrInFlightCount) == 0) { _allMrtrHandlersCompleted.TrySetResult(true); diff --git a/src/ModelContextProtocol.Extensions.Tasks/Server/McpTasksBuilderExtensions.cs b/src/ModelContextProtocol.Extensions.Tasks/Server/McpTasksBuilderExtensions.cs index 06073ce17..fb45cf34d 100644 --- a/src/ModelContextProtocol.Extensions.Tasks/Server/McpTasksBuilderExtensions.cs +++ b/src/ModelContextProtocol.Extensions.Tasks/Server/McpTasksBuilderExtensions.cs @@ -44,7 +44,7 @@ private sealed class McpTasksPostConfigureOptions(IMcpTaskStore store, ILoggerFa { private readonly IMcpTaskStore _store = store; private readonly ILogger _logger = (loggerFactory ?? NullLoggerFactory.Instance).CreateLogger(); - private readonly ConcurrentDictionary _cancellationSources = new(StringComparer.Ordinal); + private readonly ConcurrentDictionary _cancellationStates = new(StringComparer.Ordinal); public void PostConfigure(string? name, McpServerOptions options) { @@ -75,11 +75,13 @@ public void PostConfigure(string? name, McpServerOptions options) { var taskInfo = await _store.CreateTaskAsync(cancellationToken).ConfigureAwait(false); var taskId = taskInfo.TaskId; - var cts = new CancellationTokenSource(); - _cancellationSources[taskId] = cts; - var taskCancellationToken = cts.Token; + var serverLifetime = request.Server as IMcpServerLifetimeFeature; + var cancellationState = new TaskCancellationState( + serverLifetime?.BackgroundTaskCancellationToken ?? CancellationToken.None); + _cancellationStates[taskId] = cancellationState; + var taskCancellationToken = cancellationState.Token; - _ = Task.Run(async () => + var backgroundTask = Task.Run(async () => { try { @@ -142,13 +144,6 @@ public void PostConfigure(string? name, McpServerOptions options) var resultJson = JsonSerializer.SerializeToElement(errorResult, McpJsonUtilities.DefaultOptions.GetTypeInfo()); await _store.SetCompletedAsync(taskId, resultJson).ConfigureAwait(false); } - finally - { - if (_cancellationSources.TryRemove(taskId, out var registeredCts)) - { - registeredCts.Dispose(); - } - } } } catch (Exception outer) @@ -169,14 +164,18 @@ public void PostConfigure(string? name, McpServerOptions options) { _logger.LogError(storeEx, "Failed to record the failure of background task '{TaskId}'.", taskId); } - - if (_cancellationSources.TryRemove(taskId, out var leftoverCts)) + } + finally + { + if (_cancellationStates.TryRemove(taskId, out var registeredState)) { - leftoverCts.Dispose(); + registeredState.UnregisterServerLifetime(); } } }, CancellationToken.None); + serverLifetime?.RegisterBackgroundTask(backgroundTask); + return ResultOrAlternate.FromAlternate( ToCreateTaskResult(taskInfo), McpTasksJsonContext.Default.CreateTaskResult); @@ -230,15 +229,38 @@ public void PostConfigure(string? name, McpServerOptions options) 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 + { + private readonly CancellationTokenSource _source = new(); + private readonly CancellationTokenRegistration _serverLifetimeRegistration; + + public TaskCancellationState(CancellationToken serverLifetimeToken) + { + _serverLifetimeRegistration = serverLifetimeToken.Register( + static state => ((CancellationTokenSource)state!).Cancel(), + _source); + } + + public CancellationToken Token => _source.Token; + + public void Cancel() => _source.Cancel(); + + 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. + _serverLifetimeRegistration.Dispose(); + } + } + 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 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 bde2e3e4ff46d7e3fb302e6e8f56597db7987147 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 | 57 ++++++++++- .../TaskCancellationIntegrationTests.cs | 98 ++++++++++++++++--- 5 files changed, 202 insertions(+), 54 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 ce9ca6e75..79a262ada 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. /// @@ -508,36 +508,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); + } } /// @@ -560,11 +564,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 @@ -589,9 +593,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 7a7e9816c..6ba8536f7 100644 --- a/src/ModelContextProtocol.Extensions.Tasks/Server/McpTasksBuilderExtensions.cs +++ b/src/ModelContextProtocol.Extensions.Tasks/Server/McpTasksBuilderExtensions.cs @@ -148,14 +148,25 @@ private async ValueTask> RunAsTaskAsync( executionRequest.Server = request.Server.WithMcpTaskOutgoingRequestInterceptor(taskId, _store); var serverLifetime = request.Server as IMcpServerLifetimeFeature; var cancellationState = new TaskCancellationState( - serverLifetime?.BackgroundTaskCancellationToken ?? CancellationToken.None); + serverLifetime?.ServerCancellationToken ?? CancellationToken.None); _cancellationStates[taskId] = cancellationState; var taskCancellationToken = cancellationState.Token; var backgroundTask = Task.Run( () => ExecuteTaskAsync(next, executionRequest, taskId, taskCancellationToken, executionScope), CancellationToken.None); - serverLifetime?.RegisterBackgroundTask(backgroundTask); + cancellationState.SetBackgroundTask(backgroundTask); + try + { + cancellationState.SetServerLifetimeRegistration( + serverLifetime?.RegisterForDisposeAsync(cancellationState)); + } + catch + { + cancellationState.Cancel(); + await backgroundTask.ConfigureAwait(false); + throw; + } return ResultOrAlternate.FromAlternate( ToCreateTaskResult(taskInfo), @@ -326,27 +337,63 @@ private async Task ExecuteToolPipelineAsync( return JsonSerializer.SerializeToNode(new CancelTaskResult(), McpTasksJsonContext.Default.CancelTaskResult); } - private sealed class TaskCancellationState + 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 => ((CancellationTokenSource)state!).Cancel(), - _source); + 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); + } } } diff --git a/tests/ModelContextProtocol.Tests/Server/TaskCancellationIntegrationTests.cs b/tests/ModelContextProtocol.Tests/Server/TaskCancellationIntegrationTests.cs index 6cc54be88..efdd6a590 100644 --- a/tests/ModelContextProtocol.Tests/Server/TaskCancellationIntegrationTests.cs +++ b/tests/ModelContextProtocol.Tests/Server/TaskCancellationIntegrationTests.cs @@ -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; } } }