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
12 changes: 8 additions & 4 deletions docs/concepts/tasks/tasks.md
Original file line number Diff line number Diff line change
Expand Up @@ -210,6 +210,14 @@ Supported input request methods:
| --- | --- |
| `elicitation/create` | <xref:ModelContextProtocol.Client.McpClientHandlers.ElicitationHandler> |
| `sampling/createMessage` | <xref:ModelContextProtocol.Client.McpClientHandlers.SamplingHandler> |
| `roots/list` | <xref:ModelContextProtocol.Client.McpClientHandlers.RootsHandler> |

Tools that use MRTR directly by throwing <xref:ModelContextProtocol.Protocol.InputRequiredException>
also compose with task execution. When a task-enabled call throws with input requests, the SDK
publishes them through the task store and reruns the tool after `tasks/update`, with
<xref:ModelContextProtocol.Protocol.RequestParams.InputResponses> and
<xref:ModelContextProtocol.Protocol.RequestParams.RequestState> populated for the retry. Calls that
do not opt in to Tasks keep the normal MRTR behavior.

Per SEP-2663:

Expand Down Expand Up @@ -358,10 +366,6 @@ compatibility bridge for the previous experimental API.
synchronously and then transition its remaining work to a background task. Use a custom
<xref:ModelContextProtocol.Server.McpServerHandlers.CallToolWithAlternateHandler?displayProperty=nameWithType>
if you need that pattern.
- **`roots/list` as an input request**: the server SDK routes `RequestRootsAsync` through the
task channel when called from inside a task scope, but the client SDK does not currently
dispatch a handler for that method. Avoid calling `server.RequestRootsAsync` from within a
task scope until client-side support is added.
- **`ServerCapabilities.Extensions` round-trip**: the dictionary is typed as
`IDictionary<string, object>` so its values cannot be deserialized by the source generator.
The negotiated extension surfaces correctly at the wire level, but round-tripping arbitrary
Expand Down
179 changes: 179 additions & 0 deletions src/Common/InputRequiredRequestRunner.cs
Original file line number Diff line number Diff line change
@@ -0,0 +1,179 @@
using System.Text.Json;
using System.Text.Json.Nodes;
using System.Text.Json.Serialization.Metadata;

namespace ModelContextProtocol.Protocol;

/// <summary>
/// Runs request handlers that can require one or more rounds of additional input.
/// </summary>
internal static class InputRequiredRequestRunner
{
private const int MaxRetries = 10;

private static readonly JsonTypeInfo<IDictionary<string, InputResponse>> s_inputResponsesTypeInfo =
(JsonTypeInfo<IDictionary<string, InputResponse>>)McpJsonUtilities.DefaultOptions.GetTypeInfo(
typeof(IDictionary<string, InputResponse>));

private static readonly JsonTypeInfo<InputRequiredResult> s_inputRequiredResultTypeInfo =
(JsonTypeInfo<InputRequiredResult>)McpJsonUtilities.DefaultOptions.GetTypeInfo(typeof(InputRequiredResult));

/// <summary>
/// Invokes a handler until it returns a final result, normalizing both thrown and returned
/// <see cref="InputRequiredResult"/> values into the same retry flow.
/// </summary>
internal static async Task<TResult> RunAsync<TRequest, TResult>(
TRequest request,
Func<TRequest, CancellationToken, Task<TResult>> invoke,
Func<TResult, InputRequiredResult?> getReturnedInputRequiredResult,
Func<InputRequiredResult, TResult>? createDirectResult,
Func<TRequest, InputRequiredResult, Exception?, CancellationToken, Task<TRequest>> prepareRetry,
Func<string, Exception?, Exception> createFailure,
CancellationToken cancellationToken)
{
for (int retry = 0; ; retry++)
{
InputRequiredResult inputRequiredResult;
Exception? inputRequiredException = null;

try
{
TResult result = await invoke(request, cancellationToken).ConfigureAwait(false);
if (getReturnedInputRequiredResult(result) is not { } returnedInputRequiredResult)
{
return result;
}

inputRequiredResult = returnedInputRequiredResult;
}
catch (InputRequiredException ex)
{
inputRequiredResult = ex.Result;
inputRequiredException = ex;
}

if (createDirectResult is not null)
{
return createDirectResult(inputRequiredResult);
}

if (inputRequiredResult.InputRequests is not { Count: > 0 } &&
inputRequiredResult.RequestState is null)
{
throw createFailure(
"A tool returned an input-required result without input requests or request state.",
inputRequiredException);
}

if (retry >= MaxRetries)
{
throw createFailure(
$"MRTR-native tool exceeded {MaxRetries} retry rounds without completing.",
inputRequiredException);
}

request = await prepareRetry(
request,
inputRequiredResult,
inputRequiredException,
cancellationToken).ConfigureAwait(false);
}
}

/// <summary>
/// Resolves a batch concurrently, cancelling sibling requests when any resolver fails.
/// </summary>
internal static async Task<IDictionary<string, InputResponse>> ResolveInputRequestsAsync(
IDictionary<string, InputRequest> inputRequests,
Func<InputRequest, CancellationToken, Task<InputResponse>> resolveInputRequest,
CancellationToken cancellationToken)
{
using var linkedCts = CancellationTokenSource.CreateLinkedTokenSource(cancellationToken);
var keyedTasks = new (string Key, Task<InputResponse> ResponseTask)[inputRequests.Count];

int index = 0;
foreach (var pair in inputRequests)
{
keyedTasks[index++] = (pair.Key, ResolveAndCancelSiblingsAsync(pair.Value));
}

await Task.WhenAll(Array.ConvertAll(keyedTasks, static item => item.ResponseTask)).ConfigureAwait(false);

var responses = new Dictionary<string, InputResponse>(keyedTasks.Length);
foreach (var (key, responseTask) in keyedTasks)
{
responses[key] = responseTask.Result;
}

return responses;

async Task<InputResponse> ResolveAndCancelSiblingsAsync(InputRequest inputRequest)
{
try
{
return await resolveInputRequest(inputRequest, linkedCts.Token).ConfigureAwait(false);
}
catch
{
try
{
linkedCts.Cancel();
}
catch
{
// Preserve the resolver failure. Awaiting Task.WhenAll observes every sibling outcome.
}

throw;
}
}
}

/// <summary>
/// Clones request parameters and applies the response and state for the next round, removing
/// values left over from the previous round when the current result omits them.
/// </summary>
internal static JsonObject CreateRetryParams(
JsonNode? requestParams,
IDictionary<string, InputResponse>? inputResponses,
string? requestState)
{
var paramsObject = requestParams?.DeepClone() as JsonObject ?? new JsonObject();

if (inputResponses is not null)
{
paramsObject["inputResponses"] = JsonSerializer.SerializeToNode(inputResponses, s_inputResponsesTypeInfo);
}
else
{
paramsObject.Remove("inputResponses");
}

if (requestState is not null)
{
paramsObject["requestState"] = requestState;
}
else
{
paramsObject.Remove("requestState");
}

return paramsObject;
}

/// <summary>
/// Detects a serialized <see cref="InputRequiredResult"/> returned through an alternate-result path.
/// </summary>
internal static InputRequiredResult? GetReturnedInputRequiredResult(JsonNode? result)
{
if (result is JsonObject resultObject &&
resultObject.TryGetPropertyValue("resultType", out var resultTypeNode) &&
resultTypeNode?.GetValueKind() == JsonValueKind.String &&
resultTypeNode.GetValue<string>() == "input_required")
{
return JsonSerializer.Deserialize(result, s_inputRequiredResultTypeInfo);
}

return null;
}
}
136 changes: 37 additions & 99 deletions src/ModelContextProtocol.Core/Client/McpClientImpl.cs
Original file line number Diff line number Diff line change
Expand Up @@ -187,42 +187,10 @@ public override async ValueTask<IDictionary<string, InputResponse>> ResolveInput
IDictionary<string, InputRequest> inputRequests,
CancellationToken cancellationToken)
{
// Resolve all input requests concurrently. If any fails, cancel the rest so user-facing
// handlers (sampling/elicitation prompts) don't keep running for a request whose caller
// has already given up, and ensure exceptions from late-completing tasks are observed.
using var linkedCts = CancellationTokenSource.CreateLinkedTokenSource(cancellationToken);

var keyed = new (string Key, Task<InputResponse> Task)[inputRequests.Count];
int i = 0;
foreach (var kvp in inputRequests)
{
keyed[i++] = (kvp.Key, ResolveInputRequestAsync(kvp.Value, linkedCts.Token));
}

try
{
await Task.WhenAll(Array.ConvertAll(keyed, k => k.Task)).ConfigureAwait(false);
}
catch
{
linkedCts.Cancel();
try
{
await Task.WhenAll(Array.ConvertAll(keyed, k => k.Task)).ConfigureAwait(false);
}
catch
{
// Observed; the original exception is the one we want to surface.
}
throw;
}

var responses = new Dictionary<string, InputResponse>(keyed.Length);
foreach (var (key, task) in keyed)
{
responses[key] = task.Result;
}
return responses;
return await InputRequiredRequestRunner.ResolveInputRequestsAsync(
inputRequests,
ResolveInputRequestAsync,
cancellationToken).ConfigureAwait(false);
}

private async Task<InputResponse> ResolveInputRequestAsync(InputRequest inputRequest, CancellationToken cancellationToken)
Expand Down Expand Up @@ -693,74 +661,44 @@ request.Params is System.Text.Json.Nodes.JsonObject paramsObjForHeaders &&
}
}

const int maxRetries = 10;

InjectRequestMetaIfNeeded(request);

for (int attempt = 0; attempt <= maxRetries; attempt++)
{
JsonRpcResponse response = await _sessionHandler.SendRequestAsync(request, cancellationToken).ConfigureAwait(false);

// Check if the result is an InputRequiredResult by looking at result_type.
if (response.Result is JsonObject resultObj &&
resultObj.TryGetPropertyValue("resultType", out var resultTypeNode) &&
resultTypeNode?.GetValue<string>() is "input_required")
return await InputRequiredRequestRunner.RunAsync(
request,
(currentRequest, token) => _sessionHandler.SendRequestAsync(currentRequest, token),
static response => InputRequiredRequestRunner.GetReturnedInputRequiredResult(response.Result),
(Func<InputRequiredResult, JsonRpcResponse>?)null,
PrepareRetryAsync,
static (message, innerException) => new McpException(message, innerException),
cancellationToken).ConfigureAwait(false);

async Task<JsonRpcRequest> PrepareRetryAsync(
JsonRpcRequest currentRequest,
InputRequiredResult inputRequiredResult,
Exception? _,
CancellationToken retryCancellationToken)
{
WarnIfInputRequiredResultOnNonMrtrSession(currentRequest.Method);

IDictionary<string, InputResponse>? inputResponses = null;
if (inputRequiredResult.InputRequests is { Count: > 0 } inputRequests)
{
WarnIfInputRequiredResultOnNonMrtrSession(request.Method);

var inputRequiredResult = JsonSerializer.Deserialize(response.Result, McpJsonUtilities.JsonContext.Default.InputRequiredResult)
?? throw new JsonException("Failed to deserialize InputRequiredResult.");

if (inputRequiredResult.InputRequests is { Count: > 0 } inputRequests)
{
IDictionary<string, InputResponse> inputResponses =
await ResolveInputRequestsAsync(inputRequests, cancellationToken).ConfigureAwait(false);

// Clone the original request params and add inputResponses + requestState for the retry.
var paramsObj = request.Params?.DeepClone() as JsonObject ?? new JsonObject();

paramsObj["inputResponses"] = JsonSerializer.SerializeToNode(
inputResponses, McpJsonUtilities.JsonContext.Default.IDictionaryStringInputResponse);

if (inputRequiredResult.RequestState is { } requestState)
{
paramsObj["requestState"] = requestState;
}
else
{
// Strip any stale requestState carried over from the previous round's clone so
// the server doesn't see a continuation token the current round is not using.
paramsObj.Remove("requestState");
}

request = new JsonRpcRequest { Method = request.Method, Params = paramsObj, Context = request.Context };
InjectRequestMetaIfNeeded(request);
}
else if (inputRequiredResult.RequestState is not null)
{
// No input requests but has requestState (e.g., load shedding) - just retry with state.
var paramsObj = request.Params?.DeepClone() as JsonObject ?? new JsonObject();
paramsObj["requestState"] = inputRequiredResult.RequestState;
paramsObj.Remove("inputResponses");

request = new JsonRpcRequest { Method = request.Method, Params = paramsObj, Context = request.Context };
InjectRequestMetaIfNeeded(request);
}
else
{
// An input_required result carrying neither inputRequests nor requestState is
// malformed: there is nothing to resolve and nothing to continue, so retrying the
// unchanged request would just loop until maxRetries. Fail fast instead.
throw new McpException("Server returned an InputRequiredResult without inputRequests or requestState.");
}

continue; // retry with the updated request
inputResponses = await ResolveInputRequestsAsync(
inputRequests,
retryCancellationToken).ConfigureAwait(false);
}

return response;
var retryRequest = new JsonRpcRequest
{
Method = currentRequest.Method,
Params = InputRequiredRequestRunner.CreateRetryParams(
currentRequest.Params,
inputResponses,
inputRequiredResult.RequestState),
Context = currentRequest.Context,
};
InjectRequestMetaIfNeeded(retryRequest);
return retryRequest;
}

throw new McpException($"Server returned InputRequiredResult more than {maxRetries} times.");
}

/// <summary>
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -34,6 +34,7 @@
<Compile Include="..\Common\EncodingUtilities.cs" Link="EncodingUtilities.cs" />
<Compile Include="..\Common\McpHttpHeaders.cs" Link="McpHttpHeaders.cs" />
<Compile Include="..\Common\McpProtocolVersions.cs" Link="McpProtocolVersions.cs" />
<Compile Include="..\Common\InputRequiredRequestRunner.cs" Link="InputRequiredRequestRunner.cs" />
<Compile Include="..\Common\HttpResponseMessageExtensions.cs" Link="HttpResponseMessageExtensions.cs" />
<Compile Include="..\Common\ServerSentEvents\**\*.cs" Link="ServerSentEvents\%(RecursiveDir)%(FileName)%(Extension)" />
</ItemGroup>
Expand Down
Loading