diff --git a/dotnet/src/Client.cs b/dotnet/src/Client.cs index dff4681a80..58f7eb7053 100644 --- a/dotnet/src/Client.cs +++ b/dotnet/src/Client.cs @@ -2834,16 +2834,12 @@ public async ValueTask DisposeAsync() private class RpcHandler(CopilotClient client) { - public void OnSessionEvent(string sessionId, JsonElement? @event) + public void OnSessionEvent(string sessionId, SessionEvent? @event) { var session = client.GetSession(sessionId); if (session != null && @event != null) { - var evt = SessionEvent.FromJson(@event.Value.GetRawText()); - if (evt != null) - { - session.DispatchEvent(evt); - } + session.DispatchEvent(@event); } } @@ -2863,9 +2859,7 @@ public void OnSessionLifecycle(string type, string sessionId, JsonElement? metad evt.SessionId = sessionId; if (metadata is not null) { - evt.Metadata = JsonSerializer.Deserialize( - metadata.Value.GetRawText(), - TypesJsonContext.Default.SessionLifecycleEventMetadata); + evt.Metadata = metadata.Value.Deserialize(TypesJsonContext.Default.SessionLifecycleEventMetadata); } client.DispatchLifecycleEvent(evt); diff --git a/dotnet/src/JsonRpc.cs b/dotnet/src/JsonRpc.cs index c5c444ea70..d133053d76 100644 --- a/dotnet/src/JsonRpc.cs +++ b/dotnet/src/JsonRpc.cs @@ -280,11 +280,22 @@ private async Task ReadLoopAsync(CancellationToken cancellationToken) // Parse the raw JSON. Body is at buffer[0..contentLength], carried bytes // for the next message are at buffer[contentLength..contentLength+carried]. - JsonElement? message = null; try { using var doc = JsonDocument.Parse(buffer.AsMemory(0, contentLength)); - message = doc.RootElement.Clone(); + var parsed = doc.RootElement; + + // Route while the document is alive. Incoming method arguments are + // materialized synchronously before dispatch can become asynchronous. + if (parsed.TryGetProperty("id", out var idProp) && !parsed.TryGetProperty("method", out _)) + { + // It's a response to one of our requests. + HandleResponse(parsed, idProp); + } + else if (parsed.TryGetProperty("method", out var methodProp) && methodProp.GetString() is string methodName) + { + _ = HandleIncomingMethodAsync(methodName, parsed, cancellationToken); + } } catch (JsonException ex) { @@ -311,21 +322,6 @@ private async Task ReadLoopAsync(CancellationToken cancellationToken) buffer = retainedBuffer; } - if (message is not { } parsed) - { - continue; - } - - // Route the message - if (parsed.TryGetProperty("id", out var idProp) && !parsed.TryGetProperty("method", out _)) - { - // It's a response to one of our requests - HandleResponse(parsed, idProp); - } - else if (parsed.TryGetProperty("method", out var methodProp) && methodProp.GetString() is string methodName) - { - _ = HandleIncomingMethodAsync(methodName, parsed, cancellationToken); - } } } catch (OperationCanceledException) when (cancellationToken.IsCancellationRequested) @@ -469,7 +465,7 @@ private async Task ReadLoopAsync(CancellationToken cancellationToken) private void HandleResponse(JsonElement message, JsonElement idProp) { - if (!idProp.TryGetInt64(out long id)) + if (idProp.ValueKind != JsonValueKind.Number || !idProp.TryGetInt64(out long id)) { return; } @@ -528,7 +524,8 @@ private async Task HandleIncomingMethodAsync(string methodName, JsonElement mess JsonElement? requestId = null; if (message.TryGetProperty("id", out var idProp)) { - requestId = idProp; + // Requests may outlive the parsed message while an asynchronous handler runs. + requestId = idProp.Clone(); } if (!_methods.TryGetValue(methodName, out var registration)) @@ -544,7 +541,10 @@ private async Task HandleIncomingMethodAsync(string methodName, JsonElement mess try { - var result = await InvokeHandlerAsync(registration, paramsProp, cancellationToken).ConfigureAwait(false); + // Materialize arguments before the first possible suspension so none of + // them borrow from the JsonDocument owned by the read loop. + var invokeArgs = DeserializeHandlerArguments(registration, paramsProp, cancellationToken); + var result = await InvokeHandlerAsync(registration, invokeArgs).ConfigureAwait(false); if (requestId.HasValue) { @@ -599,11 +599,13 @@ await SendResultResponseAsync( } } - private async ValueTask InvokeHandlerAsync(MethodRegistration registration, JsonElement paramsProp, CancellationToken cancellationToken) + private object?[] DeserializeHandlerArguments( + MethodRegistration registration, + JsonElement paramsProp, + CancellationToken cancellationToken) { var parameters = registration.Parameters; - // Build argument list var invokeArgs = new object?[parameters.Length]; if (registration.SingleObjectParam) @@ -681,7 +683,13 @@ await SendResultResponseAsync( $"Unsupported JSON-RPC params shape '{paramsProp.ValueKind}' for handler with positional parameters."); } - // Invoke + return invokeArgs; + } + + private static async ValueTask InvokeHandlerAsync( + MethodRegistration registration, + object?[] invokeArgs) + { var result = registration.Handler.DynamicInvoke(invokeArgs); // Handlers return one of: a synchronous value, Task (void async), or ValueTask. diff --git a/dotnet/test/Unit/ClientSessionLifetimeTests.cs b/dotnet/test/Unit/ClientSessionLifetimeTests.cs index e6764c0ef4..3ae2493a8d 100644 --- a/dotnet/test/Unit/ClientSessionLifetimeTests.cs +++ b/dotnet/test/Unit/ClientSessionLifetimeTests.cs @@ -1566,6 +1566,59 @@ public async Task Raw_SendAsync_MessageSource_Remains_Available(string? source) AssertMessageSource(Assert.Single(server.Requests, request => request.Method == "session.send").Params, source); } + [Fact] + public async Task SessionEvents_Recover_From_Malformed_Input_And_Isolate_Multiple_Handlers() + { + await using var server = await FakeCopilotServer.StartAsync(); + await using var client = new CopilotClient(new CopilotClientOptions { Connection = RuntimeConnection.ForUri(server.Url) }); + await using var session = await client.CreateSessionAsync(new SessionConfig()); + var firstHandlerEvents = new List(); + var secondHandlerEvents = new List(); + var received = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + + using var firstSubscription = session.On(@event => + { + firstHandlerEvents.Add(@event.Type); + throw new InvalidOperationException("Expected test handler failure."); + }); + using var secondSubscription = session.On(@event => + { + secondHandlerEvents.Add(@event); + if (secondHandlerEvents.Count == 2) + { + received.TrySetResult(); + } + }); + + await server.SendSessionEventPayloadAsync(session.SessionId, 42); + await server.SendSessionEventPayloadAsync(session.SessionId, new Dictionary + { + ["id"] = Guid.NewGuid().ToString(), + ["timestamp"] = DateTimeOffset.UtcNow.ToString("O"), + ["parentId"] = null, + ["type"] = "future.event", + ["data"] = new object?[] { null, false, 42, "text", new Dictionary { ["nested"] = true } } + }); + await server.SendSessionEventAsync(session.SessionId, "tool.execution_start", new() + { + ["toolCallId"] = "tool-1", + ["toolName"] = "view", + ["arguments"] = new Dictionary + { + ["path"] = "README.md", + ["nested"] = new object?[] { null, false, 42, "text", new Dictionary { ["value"] = true } } + } + }); + + await received.Task.WaitAsync(TimeSpan.FromSeconds(5)); + + Assert.Equal(["unknown", "tool.execution_start"], firstHandlerEvents); + Assert.IsType(secondHandlerEvents[0]); + var toolEvent = Assert.IsType(secondHandlerEvents[1]); + Assert.Equal("README.md", toolEvent.Data.Arguments?.GetProperty("path").GetString()); + Assert.True(toolEvent.Data.Arguments?.GetProperty("nested")[4].GetProperty("value").GetBoolean()); + } + public static IEnumerable MessageSourcesAndOutcomes { get @@ -2510,6 +2563,21 @@ public Task SendSessionEventAsync(string sessionId, string type, Dictionary + { + ["jsonrpc"] = "2.0", + ["method"] = "session.event", + ["params"] = new Dictionary + { + ["sessionId"] = sessionId, + ["event"] = @event + } + }, _cts.Token); + } + public async Task SendAndDrainSessionEventAsync( CopilotSession session, string type, diff --git a/dotnet/test/Unit/JsonRpcTests.cs b/dotnet/test/Unit/JsonRpcTests.cs index f4acfd3555..08b0cf352c 100644 --- a/dotnet/test/Unit/JsonRpcTests.cs +++ b/dotnet/test/Unit/JsonRpcTests.cs @@ -122,6 +122,106 @@ public async Task JsonRpc_Does_Not_Retain_Oversized_Receive_Buffer() Assert.InRange(await receiveStream.PostFrameReadBufferSize, 1, 1024 * 1024); } + [Theory] + [InlineData("null")] + [InlineData("\"server-id\"")] + public async Task JsonRpc_Invalid_Response_Id_Does_Not_End_Read_Loop(string invalidId) + { + var invalidFrame = CreateFrame( + $$"""{"jsonrpc":"2.0","id":{{invalidId}},"result":"ignored"}"""); + var validFrame = CreateResponseFrame(1, "carried"); + using var receiveStream = new MemoryStream(CombineFrames([invalidFrame, validFrame])); + using var rpc = new JsonRpcReflection(Stream.Null, receiveStream); + + var response = rpc.InvokeAsync("pending", args: null); + rpc.StartListening(); + + Assert.Equal("carried", await response.WaitAsync(TimeSpan.FromSeconds(5))); + await rpc.Completion.WaitAsync(TimeSpan.FromSeconds(5)); + } + + [Fact] + public async Task JsonRpc_JsonElement_Params_Remain_Valid_After_Message_Disposal() + { + var payloads = new[] + { + "null", + "false", + "42", + "\"text\"", + """[1,{"nested":[null,true]}]""", + """{"result":{"rows":[{"content":"preserved"}]}}""", + }; + using var receiveStream = new MemoryStream(CombineFrames( + payloads.Select(payload => CreateNotificationFrame("payload", $$"""{"payload":{{payload}}}""")))); + using var rpc = new JsonRpcReflection(Stream.Null, receiveStream); + var collector = new JsonElementCollector(payloads.Length); + rpc.SetLocalRpcMethod("payload", (Action)collector.Handle); + + rpc.StartListening(); + await collector.Completion.WaitAsync(TimeSpan.FromSeconds(5)); + await rpc.Completion.WaitAsync(TimeSpan.FromSeconds(5)); + + Assert.Equal(payloads, collector.Payloads.Select(payload => payload?.GetRawText() ?? "null")); + } + + [Fact] + public async Task JsonRpc_Malformed_Session_Event_Does_Not_Block_Unknown_Or_Known_Events() + { + var malformed = CreateNotificationFrame( + "session.event", + """{"sessionId":"session-1","event":42}"""); + var unknown = CreateNotificationFrame( + "session.event", + """ + { + "sessionId":"session-1", + "event":{ + "id":"11111111-1111-1111-1111-111111111111", + "timestamp":"2026-09-19T00:00:00Z", + "parentId":null, + "type":"future.event", + "data":{"nested":[null,true,{"value":42}]} + } + } + """); + var known = CreateNotificationFrame( + "session.event", + """ + { + "sessionId":"session-1", + "event":{ + "id":"22222222-2222-2222-2222-222222222222", + "timestamp":"2026-09-19T00:00:01Z", + "parentId":null, + "type":"tool.execution_start", + "data":{ + "toolCallId":"tool-1", + "toolName":"view", + "arguments":{"path":"README.md","nested":[1,{"value":true}]} + } + } + } + """); + + using var receiveStream = new MemoryStream(CombineFrames([malformed, unknown, known])); + using var rpc = new JsonRpcReflection(Stream.Null, receiveStream); + var collector = new SessionEventCollector(expectedCount: 2); + rpc.SetLocalRpcMethod("session.event", (Action)collector.Handle); + + rpc.StartListening(); + await collector.Completion.WaitAsync(TimeSpan.FromSeconds(5)); + await rpc.Completion.WaitAsync(TimeSpan.FromSeconds(5)); + + Assert.All(collector.SessionIds, sessionId => Assert.Equal("session-1", sessionId)); + var unknownEvent = Assert.IsType(collector.Events[0]); + Assert.Equal("unknown", unknownEvent.Type); + + var toolEvent = Assert.IsType(collector.Events[1]); + Assert.Equal("README.md", toolEvent.Data.Arguments?.GetProperty("path").GetString()); + Assert.True(toolEvent.Data.Arguments?.GetProperty("nested")[1].GetProperty("value").GetBoolean()); + } + private static byte[] CreateResponseFrame(long id, string result, int headerPaddingLength = 0) { using var bodyStream = new MemoryStream(); @@ -145,6 +245,29 @@ private static byte[] CreateResponseFrame(long id, string result, int headerPadd return frame; } + private static byte[] CreateNotificationFrame(string method, string paramsJson) + => CreateFrame($$"""{"jsonrpc":"2.0","method":"{{method}}","params":{{paramsJson}}}"""); + + private static byte[] CreateFrame(string json) + { + var body = Encoding.UTF8.GetBytes(json); + var header = Encoding.ASCII.GetBytes($"Content-Length: {body.Length}\r\n\r\n"); + var frame = new byte[header.Length + body.Length]; + header.CopyTo(frame, 0); + body.CopyTo(frame, header.Length); + return frame; + } + + private static byte[] CombineFrames(IEnumerable frames) + { + using var stream = new MemoryStream(); + foreach (var frame in frames) + { + stream.Write(frame); + } + return stream.ToArray(); + } + private static int GetRemoteErrorCode(Exception exception) { var property = exception.GetType().GetProperty("ErrorCode", BindingFlags.Instance | BindingFlags.Public); @@ -216,6 +339,7 @@ private sealed class JsonRpcReflection : IDisposable private static readonly JsonSerializerOptions SerializerOptions = new(JsonSerializerDefaults.Web) { + AllowOutOfOrderMetadataProperties = true, TypeInfoResolver = new DefaultJsonTypeInfoResolver(), }; @@ -238,6 +362,8 @@ public JsonRpcReflection(Stream sendStream, Stream receiveStream) public void StartListening() => JsonRpcType.GetMethod(nameof(StartListening))!.Invoke(_instance, null); + public Task Completion => (Task)JsonRpcType.GetProperty(nameof(Completion))!.GetValue(_instance)!; + public void SetLocalRpcMethod(string methodName, Delegate handler, bool singleObjectParam = false) => JsonRpcType.GetMethod("SetLocalRpcMethod")!.Invoke(_instance, [methodName, handler, singleObjectParam]); @@ -255,6 +381,47 @@ public async Task InvokeAsync(string methodName, object?[]? args, Cancella public void Dispose() => ((IDisposable)_instance).Dispose(); } + private sealed class JsonElementCollector(int expectedCount) + { + private readonly TaskCompletionSource _completion = + new(TaskCreationOptions.RunContinuationsAsynchronously); + + public List Payloads { get; } = []; + + public Task Completion => _completion.Task; + + public void Handle(JsonElement? payload) + { + Payloads.Add(payload); + if (Payloads.Count == expectedCount) + { + _completion.TrySetResult(); + } + } + } + + private sealed class SessionEventCollector(int expectedCount) + { + private readonly TaskCompletionSource _completion = + new(TaskCreationOptions.RunContinuationsAsynchronously); + + public List Events { get; } = []; + + public List SessionIds { get; } = []; + + public Task Completion => _completion.Task; + + public void Handle(string sessionId, SessionEvent? @event) + { + SessionIds.Add(sessionId); + Events.Add(@event); + if (Events.Count == expectedCount) + { + _completion.TrySetResult(); + } + } + } + private sealed class CoalescedFramesThenWaitStream : Stream { private readonly TaskCompletionSource _postFrameReadBufferSize =