From 9ba50aa092ad965923318046a0517bd0d3fd4bc1 Mon Sep 17 00:00:00 2001 From: King Star Date: Mon, 20 Jul 2026 11:41:11 +0800 Subject: [PATCH] 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. ///