Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -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);
Expand Down Expand Up @@ -73,6 +73,12 @@ public override Implementation? ClientInfo

public override bool IsMrtrSupported => server.IsMrtrSupported;

CancellationToken IMcpServerLifetimeFeature.ServerCancellationToken =>
((IMcpServerLifetimeFeature)server).ServerCancellationToken;

IDisposable IMcpServerLifetimeFeature.RegisterForDisposeAsync(IAsyncDisposable disposable) =>
((IMcpServerLifetimeFeature)server).RegisterForDisposeAsync(disposable);

public override ValueTask DisposeAsync() => server.DisposeAsync();

public override IAsyncDisposable RegisterNotificationHandler(string method, Func<JsonRpcNotification, CancellationToken, ValueTask> handler) => server.RegisterNotificationHandler(method, handler);
Expand Down
26 changes: 26 additions & 0 deletions src/ModelContextProtocol.Core/Server/IMcpServerLifetimeFeature.cs
Original file line number Diff line number Diff line change
@@ -0,0 +1,26 @@
using System.ComponentModel;

namespace ModelContextProtocol.Server;

/// <summary>
/// Provides server-lifetime services used by MCP extension infrastructure.
/// </summary>
[EditorBrowsable(EditorBrowsableState.Never)]
public interface IMcpServerLifetimeFeature
{
/// <summary>Gets the token that is cancelled when this server starts disposing.</summary>
/// <remarks>
/// The token is <see cref="CancellationToken.None"/> when work intentionally outlives the server,
/// as it does for per-request servers in stateless HTTP mode.
/// </remarks>
CancellationToken ServerCancellationToken { get; }

/// <summary>Registers an asynchronously disposable resource that server disposal must await.</summary>
/// <param name="disposable">The resource to dispose when this server is disposed.</param>
/// <returns>A handle that unregisters the resource without disposing it.</returns>
/// <remarks>
/// 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.
/// </remarks>
IDisposable RegisterForDisposeAsync(IAsyncDisposable disposable);
}
96 changes: 93 additions & 3 deletions src/ModelContextProtocol.Core/Server/McpServerImpl.cs
Original file line number Diff line number Diff line change
Expand Up @@ -12,7 +12,7 @@ namespace ModelContextProtocol.Server;

/// <inheritdoc />
#pragma warning disable MCPEXP001, MCPEXP002
internal sealed partial class McpServerImpl : McpServer
internal sealed partial class McpServerImpl : McpServer, IMcpServerLifetimeFeature
{
internal static Implementation DefaultImplementation { get; } = new()
{
Expand All @@ -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 _serverLifetimeRegistrationsLock = new();
private readonly HashSet<ServerLifetimeRegistration> _serverLifetimeRegistrations = [];
private readonly ConcurrentDictionary<string, MrtrContinuation> _mrtrContinuations = new();
private readonly ConcurrentDictionary<RequestId, MrtrContext> _mrtrContextsByRequestId = new();
private static readonly string[] s_perRequestMetadataKeys =
Expand All @@ -55,6 +58,7 @@ internal sealed partial class McpServerImpl : McpServer
private int _started;

private bool _disposed;
private bool _serverLifetimeRegistrationClosed;

/// <summary>Holds a boxed <see cref="LoggingLevel"/> value for the server.</summary>
/// <remarks>
Expand Down Expand Up @@ -504,6 +508,42 @@ public override Task SendMessageAsync(JsonRpcMessage message, CancellationToken
public override IAsyncDisposable RegisterNotificationHandler(string method, Func<JsonRpcNotification, CancellationToken, ValueTask> handler)
=> _sessionHandler.RegisterNotificationHandler(method, handler);

CancellationToken IMcpServerLifetimeFeature.ServerCancellationToken =>
HasStatefulTransport() ? _serverLifetimeCts.Token : CancellationToken.None;

IDisposable IMcpServerLifetimeFeature.RegisterForDisposeAsync(IAsyncDisposable disposable)
{
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 NoopRegistration.Instance;
}

var registration = new ServerLifetimeRegistration(this, disposable);
lock (_serverLifetimeRegistrationsLock)
{
if (_serverLifetimeRegistrationClosed)
{
throw new ObjectDisposedException(nameof(McpServer));
}

_serverLifetimeRegistrations.Add(registration);
}

return registration;
}

private void UnregisterServerLifetime(ServerLifetimeRegistration registration)
{
lock (_serverLifetimeRegistrationsLock)
{
_serverLifetimeRegistrations.Remove(registration);
}
}

/// <inheritdoc/>
public override async ValueTask DisposeAsync()
{
Expand All @@ -515,6 +555,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.
Expand All @@ -523,6 +564,13 @@ public override async ValueTask DisposeAsync()
_disposables.ForEach(d => d());
await _sessionHandler.DisposeAsync().ConfigureAwait(false);

ServerLifetimeRegistration[] serverLifetimeRegistrations;
lock (_serverLifetimeRegistrationsLock)
{
_serverLifetimeRegistrationClosed = true;
serverLifetimeRegistrations = [.. _serverLifetimeRegistrations];
}

// 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;
Expand All @@ -544,6 +592,36 @@ public override async ValueTask DisposeAsync()
{
await _allMrtrHandlersCompleted.Task.ConfigureAwait(false);
}

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()
{
}
}

private void ConfigureInitialize(McpServerOptions options)
Expand Down Expand Up @@ -2204,6 +2282,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
Expand All @@ -2214,6 +2295,11 @@ private void WrapHandlerWithMrtr(string method)
{
handlerTask = originalHandler(request, handlerCts.Token);
}
catch
{
serverLifetimeRegistration.Dispose();
throw;
}
finally
{
_mrtrContextsByRequestId.TryRemove(request.Id, out _);
Expand All @@ -2226,7 +2312,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);
Expand Down Expand Up @@ -2285,7 +2371,9 @@ private void WrapHandlerWithMrtr(string method)
/// double-reporting at Error) and decrements <see cref="_mrtrInFlightCount"/> when the
/// handler completes, following the same in-flight tracking pattern as <see cref="McpSessionHandler"/>.
/// </summary>
private async Task ObserveHandlerCompletionAsync(Task<JsonNode?> handlerTask)
private async Task ObserveHandlerCompletionAsync(
Task<JsonNode?> handlerTask,
CancellationTokenRegistration serverLifetimeRegistration)
{
try
{
Expand All @@ -2305,6 +2393,8 @@ private async Task ObserveHandlerCompletionAsync(Task<JsonNode?> handlerTask)
}
finally
{
serverLifetimeRegistration.Dispose();

if (Interlocked.Decrement(ref _mrtrInFlightCount) == 0)
{
_allMrtrHandlersCompleted.TrySetResult(true);
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -57,7 +57,7 @@ private sealed class McpTasksConfigureOptions(
private readonly IMcpTaskStore _store = store;
private readonly IServiceScopeFactory _serviceScopeFactory = serviceScopeFactory;
private readonly ILogger _logger = (loggerFactory ?? NullLoggerFactory.Instance).CreateLogger<McpTasksConfigureOptions>();
private readonly ConcurrentDictionary<string, CancellationTokenSource> _cancellationSources = new(StringComparer.Ordinal);
private readonly ConcurrentDictionary<string, TaskCancellationState> _cancellationStates = new(StringComparer.Ordinal);

public void Configure(McpServerOptions options)
{
Expand Down Expand Up @@ -146,15 +146,27 @@ private async ValueTask<ResultOrAlternate<CallToolResult>> 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<CallToolResult>.FromAlternate(
ToCreateTaskResult(taskInfo),
Expand Down Expand Up @@ -198,9 +210,9 @@ private async Task ExecuteTaskAsync(
}
finally
{
if (_cancellationSources.TryRemove(taskId, out var registeredCts))
if (_cancellationStates.TryRemove(taskId, out var registeredState))
{
registeredCts.Dispose();
registeredState.UnregisterServerLifetime();
}
}
}
Expand Down Expand Up @@ -317,15 +329,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))
Expand Down
Loading