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
62 changes: 49 additions & 13 deletions src/ModelContextProtocol.Core/McpSessionHandler.cs
Original file line number Diff line number Diff line change
Expand Up @@ -532,16 +532,41 @@ private void HandleMessageWithId(JsonRpcMessage message, JsonRpcMessageWithId me
throw new McpProtocolException($"Method '{request.Method}' is not available.", McpErrorCode.MethodNotFound);
}

JsonNode? result = await handler(request, cancellationToken).ConfigureAwait(false);

await SendMessageAsync(new JsonRpcResponse
RequestHandlerResult handlerResult = await handler(request, cancellationToken).ConfigureAwait(false);
var response = new JsonRpcResponse
{
Id = request.Id,
Result = result,
Result = handlerResult.Json,
Context = request.Context,
}, cancellationToken).ConfigureAwait(false);
};

Func<JsonRpcMessage, JsonRpcMessage>? prepareForEmission = handlerResult.IsProtocolResult ?
message => PrepareResponseForEmission(request, handlerResult, message) :
null;
await SendMessageAsync(response, prepareForEmission, cancellationToken).ConfigureAwait(false);

return result;
return handlerResult.Json;
}

private JsonRpcMessage PrepareResponseForEmission(
JsonRpcRequest request,
RequestHandlerResult handlerResult,
JsonRpcMessage message)
{
if (message is not JsonRpcResponse response || response.Id != request.Id)
{
return message;
}

return new JsonRpcResponse
{
JsonRpc = response.JsonRpc,
Id = response.Id,
Result = _requestHandlers.PrepareForEmission(
request,
handlerResult with { Json = response.Result?.DeepClone() }),
Context = response.Context,
};
}

/// <summary>
Expand Down Expand Up @@ -804,7 +829,13 @@ public async Task<JsonRpcResponse> SendRequestAsync(JsonRpcRequest request, Canc
}
}

public async Task SendMessageAsync(JsonRpcMessage message, CancellationToken cancellationToken = default)
public Task SendMessageAsync(JsonRpcMessage message, CancellationToken cancellationToken = default) =>
SendMessageAsync(message, prepareForEmission: null, cancellationToken);

private async Task SendMessageAsync(
JsonRpcMessage message,
Func<JsonRpcMessage, JsonRpcMessage>? prepareForEmission,
CancellationToken cancellationToken)
{
Throw.IfNull(message);

Expand Down Expand Up @@ -843,7 +874,7 @@ public async Task SendMessageAsync(JsonRpcMessage message, CancellationToken can
AddTags(ref tags, activity, message, method, target);
}

await SendToRelatedTransportAsync(message, cancellationToken).ConfigureAwait(false);
await SendToRelatedTransportAsync(message, cancellationToken, prepareForEmission).ConfigureAwait(false);

// If the sent notification was a cancellation notification, cancel the pending request's await, as either the
// server won't be sending a response, or per the specification, the response should be ignored. There are inherent
Expand All @@ -869,14 +900,19 @@ public async Task SendMessageAsync(JsonRpcMessage message, CancellationToken can
// The JsonRpcMessage should be sent over the RelatedTransport if set. This is used to support the
// Streamable HTTP transport where the specification states that the server SHOULD include JSON-RPC responses in
// the HTTP response body for the POST request containing the corresponding JSON-RPC request.
private Task SendToRelatedTransportAsync(JsonRpcMessage message, CancellationToken cancellationToken)
private Task SendToRelatedTransportAsync(
JsonRpcMessage message,
CancellationToken cancellationToken,
Func<JsonRpcMessage, JsonRpcMessage>? prepareForEmission = null)
=> _outgoingMessageFilter((msg, ct) =>
{
if (msg is JsonRpcRequest request)
JsonRpcMessage messageToSend = prepareForEmission?.Invoke(msg) ?? msg;

if (messageToSend is JsonRpcRequest request)
{
if (_logger.IsEnabled(LogLevel.Trace))
{
LogSendingRequestSensitive(EndpointName, request.Method, JsonSerializer.Serialize(msg, McpJsonUtilities.JsonContext.Default.JsonRpcMessage));
LogSendingRequestSensitive(EndpointName, request.Method, JsonSerializer.Serialize(messageToSend, McpJsonUtilities.JsonContext.Default.JsonRpcMessage));
}
else
{
Expand All @@ -887,15 +923,15 @@ private Task SendToRelatedTransportAsync(JsonRpcMessage message, CancellationTok
{
if (_logger.IsEnabled(LogLevel.Trace))
{
LogSendingMessageSensitive(EndpointName, JsonSerializer.Serialize(msg, McpJsonUtilities.JsonContext.Default.JsonRpcMessage));
LogSendingMessageSensitive(EndpointName, JsonSerializer.Serialize(messageToSend, McpJsonUtilities.JsonContext.Default.JsonRpcMessage));
}
else
{
LogSendingMessage(EndpointName);
}
}

return (msg.Context?.RelatedTransport ?? _transport).SendMessageAsync(msg, ct);
return (messageToSend.Context?.RelatedTransport ?? _transport).SendMessageAsync(messageToSend, ct);
})(message, cancellationToken);

private static CancelledNotificationParams? GetCancelledNotificationParams(JsonNode? notificationParams)
Expand Down
25 changes: 20 additions & 5 deletions src/ModelContextProtocol.Core/RequestHandlers.cs
Original file line number Diff line number Diff line change
Expand Up @@ -5,8 +5,18 @@

namespace ModelContextProtocol;

internal sealed class RequestHandlers : Dictionary<string, Func<JsonRpcRequest, CancellationToken, Task<JsonNode?>>>
internal readonly record struct RequestHandlerResult(JsonNode? Json, bool IsProtocolResult = false, bool IsCacheable = false);

internal sealed class RequestHandlers : Dictionary<string, Func<JsonRpcRequest, CancellationToken, Task<RequestHandlerResult>>>
{
private readonly Func<JsonRpcRequest, RequestHandlerResult, JsonNode?>? _prepareResponseForEmission;

public RequestHandlers(Func<JsonRpcRequest, RequestHandlerResult, JsonNode?>? prepareResponseForEmission = null) =>
_prepareResponseForEmission = prepareResponseForEmission;

public JsonNode? PrepareForEmission(JsonRpcRequest request, RequestHandlerResult result) =>
_prepareResponseForEmission?.Invoke(request, result) ?? result.Json;

/// <summary>
/// Registers a handler for incoming requests of a specific method in the MCP protocol.
/// </summary>
Expand Down Expand Up @@ -41,8 +51,9 @@ public void Set<TParams, TResult>(
this[method] = async (request, cancellationToken) =>
{
TParams typedRequest = JsonSerializer.Deserialize(request.Params, requestTypeInfo)!;
object? result = await handler(typedRequest, request, cancellationToken).ConfigureAwait(false);
return JsonSerializer.SerializeToNode(result, responseTypeInfo);
TResult result = await handler(typedRequest, request, cancellationToken).ConfigureAwait(false);
JsonNode? resultNode = JsonSerializer.SerializeToNode(result, responseTypeInfo);
return new(resultNode, result is Result, result is ICacheableResult);
};
}

Expand Down Expand Up @@ -70,10 +81,14 @@ public void SetWithAlternate<TParams, TResult>(

if (augmented.IsAlternate)
{
return JsonSerializer.SerializeToNode(augmented.Alternate!, augmented.AlternateTypeInfo!);
var result = augmented.Alternate!;
JsonNode? resultNode = JsonSerializer.SerializeToNode(result, augmented.AlternateTypeInfo!);
return new(resultNode, IsProtocolResult: true, IsCacheable: result is ICacheableResult);
}

return JsonSerializer.SerializeToNode(augmented.Result!, responseTypeInfo);
var immediateResult = augmented.Result!;
JsonNode? immediateResultNode = JsonSerializer.SerializeToNode(immediateResult, responseTypeInfo);
return new(immediateResultNode, IsProtocolResult: true, IsCacheable: immediateResult is ICacheableResult);
};
}
#pragma warning restore MCPEXP002
Expand Down
Loading