diff --git a/src/modules/cmdpal/Microsoft.CmdPal.UI.ViewModels/Services/JsonRpc/JsonRpcConnection.cs b/src/modules/cmdpal/Microsoft.CmdPal.UI.ViewModels/Services/JsonRpc/JsonRpcConnection.cs index c433e1776b..6e44084cae 100644 --- a/src/modules/cmdpal/Microsoft.CmdPal.UI.ViewModels/Services/JsonRpc/JsonRpcConnection.cs +++ b/src/modules/cmdpal/Microsoft.CmdPal.UI.ViewModels/Services/JsonRpc/JsonRpcConnection.cs @@ -38,10 +38,24 @@ public sealed partial class JsonRpcConnection : IDisposable // The reader never blocks on this queue: when it is full the oldest notification is dropped. private const int NotificationQueueCapacity = 1024; + // Number of worker tasks that service inbound requests. This caps how many inbound request + // handlers can run at once so a flood of inbound requests cannot spawn unbounded work. + internal const int InboundRequestWorkerCount = 16; + + // Upper bound on the number of inbound requests buffered ahead of the workers. When the queue + // is full the read loop blocks on the enqueue, applying backpressure to the peer rather than + // buffering without limit. + private const int InboundRequestQueueCapacity = 256; + + // Upper bound on how many characters of an offending payload are written to the log. Malformed + // or oversized bodies are truncated so a single bad frame cannot flood the log up to the frame cap. + internal const int MaxLoggedBodyChars = 1024; + private readonly Stream _input; private readonly Stream _output; private readonly Stream? _errorStream; private readonly TimeSpan _requestTimeout; + private readonly TimeSpan _writeTimeout; private readonly ConcurrentDictionary> _pendingRequests = new(); private readonly ConcurrentDictionary> _notificationHandlers = new(StringComparer.Ordinal); @@ -55,6 +69,17 @@ public sealed partial class JsonRpcConnection : IDisposable FullMode = BoundedChannelFullMode.Wait, }); + // Bounded queue of inbound requests drained by a fixed pool of workers. Only the read loop writes + // to it (SingleWriter), the workers read (multiple readers), and a full queue blocks the writer so + // the peer is throttled instead of letting the host buffer inbound requests without limit. + private readonly Channel _inboundRequestQueue = Channel.CreateBounded( + new BoundedChannelOptions(InboundRequestQueueCapacity) + { + SingleReader = false, + SingleWriter = true, + FullMode = BoundedChannelFullMode.Wait, + }); + private readonly SemaphoreSlim _writeLock = new(1, 1); private readonly CancellationTokenSource _disposalCts = new(); @@ -64,6 +89,7 @@ public sealed partial class JsonRpcConnection : IDisposable private Task? _readLoopTask; private Task? _errorPumpTask; private Task? _notificationConsumerTask; + private Task[]? _inboundRequestWorkers; private volatile bool _disposed; private int _disconnectedRaised; @@ -74,12 +100,14 @@ public sealed partial class JsonRpcConnection : IDisposable /// The stream to write outgoing framed messages to (for example, a process's standard input). /// An optional stream carrying out-of-band diagnostics (for example, a process's standard error). It is logged but is never part of the protocol. /// The per-request timeout. Defaults to 10 seconds when null. - public JsonRpcConnection(Stream input, Stream output, Stream? errorStream = null, TimeSpan? requestTimeout = null) + /// The maximum time a single outbound frame may take to reach the peer before the write is abandoned and the connection is torn down. Defaults to 10 seconds when null. + public JsonRpcConnection(Stream input, Stream output, Stream? errorStream = null, TimeSpan? requestTimeout = null, TimeSpan? writeTimeout = null) { _input = input ?? throw new ArgumentNullException(nameof(input)); _output = output ?? throw new ArgumentNullException(nameof(output)); _errorStream = errorStream; _requestTimeout = requestTimeout ?? TimeSpan.FromSeconds(10); + _writeTimeout = writeTimeout ?? TimeSpan.FromSeconds(10); } /// @@ -107,6 +135,14 @@ public sealed partial class JsonRpcConnection : IDisposable _readLoopTask = Task.Run(ReadLoopAsync); _notificationConsumerTask = Task.Run(ConsumeNotificationsAsync); + var workers = new Task[InboundRequestWorkerCount]; + for (var i = 0; i < workers.Length; i++) + { + workers[i] = Task.Run(ProcessInboundRequestsAsync); + } + + _inboundRequestWorkers = workers; + if (_errorStream is not null) { _errorPumpTask = Task.Run(PumpErrorStreamAsync); @@ -276,8 +312,9 @@ public sealed partial class JsonRpcConnection : IDisposable // Disconnected exactly once. Close("The JSON-RPC connection was disposed."); - // (c) Complete the notification queue so its consumer drains and exits. + // (c) Complete the notification and inbound-request queues so their consumers drain and exit. _notificationQueue.Writer.TryComplete(); + _inboundRequestQueue.Writer.TryComplete(); // (d) Drain any writer that is mid-frame by acquiring the write lock once. Writes emit their // header, body, and flush under CancellationToken.None, so an in-flight write completes the @@ -295,6 +332,14 @@ public sealed partial class JsonRpcConnection : IDisposable WaitForBackgroundTask(_errorPumpTask); WaitForBackgroundTask(_notificationConsumerTask); + if (_inboundRequestWorkers is { } workers) + { + foreach (var worker in workers) + { + WaitForBackgroundTask(worker); + } + } + if (acquiredWriteLock) { try @@ -349,6 +394,7 @@ public sealed partial class JsonRpcConnection : IDisposable } _notificationQueue.Writer.TryComplete(); + _inboundRequestQueue.Writer.TryComplete(); FailAllPending(reason); RaiseDisconnected(); @@ -382,15 +428,30 @@ public sealed partial class JsonRpcConnection : IDisposable throw new JsonRpcException("The JSON-RPC connection is closed."); } - // Once frame emission begins the header, body, and flush must be written as one unit. - // Passing the caller's cancellation token here would let a cancellation between the header - // and body write a corrupt frame and leave the connection reusable, so a non-cancellable - // token is used. Cancellation is still honored while acquiring the lock above. + // Once frame emission begins the header, body, and flush must be written as one unit, so + // the caller's cancellation token is deliberately not honored here: cancelling between the + // header and body would leave a corrupt partial frame on a still-open connection. Instead + // the emission is bounded by a dedicated write timeout and by disposal. If the peer stops + // draining stdin the write cannot block forever: when the timeout or disposal fires the + // write is abandoned and the connection is torn down (never reused), so a partial frame can + // only ever appear on a connection that is already closing. + using var writeCts = CancellationTokenSource.CreateLinkedTokenSource(_disposalCts.Token); + writeCts.CancelAfter(_writeTimeout); + try { - await _output.WriteAsync(header, CancellationToken.None).ConfigureAwait(false); - await _output.WriteAsync(body, CancellationToken.None).ConfigureAwait(false); - await _output.FlushAsync(CancellationToken.None).ConfigureAwait(false); + await _output.WriteAsync(header, writeCts.Token).ConfigureAwait(false); + await _output.WriteAsync(body, writeCts.Token).ConfigureAwait(false); + await _output.FlushAsync(writeCts.Token).ConfigureAwait(false); + } + catch (OperationCanceledException) when (writeCts.IsCancellationRequested) + { + // The write did not complete within the write timeout, or disposal cancelled it. The + // peer is no longer draining stdin (or is gone), so the stream can no longer be trusted: + // enter the terminal closed state, which raises Disconnected so the owner tears the child + // process down, and fail this write instead of blocking the write lock indefinitely. + Close("The JSON-RPC connection failed because a write did not complete in time."); + throw new JsonRpcException("The JSON-RPC connection failed because a write did not complete in time."); } catch (Exception ex) when (ex is not JsonRpcException) { @@ -439,7 +500,7 @@ public sealed partial class JsonRpcConnection : IDisposable } var json = Encoding.UTF8.GetString(body); - DispatchMessage(json); + await DispatchMessageAsync(json).ConfigureAwait(false); } } catch (OperationCanceledException) when (_disposalCts.IsCancellationRequested) @@ -544,7 +605,7 @@ public sealed partial class JsonRpcConnection : IDisposable return buffer; } - private void DispatchMessage(string json) + private async Task DispatchMessageAsync(string json) { JsonDocument document; try @@ -553,7 +614,10 @@ public sealed partial class JsonRpcConnection : IDisposable } catch (JsonException ex) { - Logger.LogError($"Failed to parse an inbound JSON-RPC message: {json}", ex); + // Bound the logged payload: a malformed body can be as large as the frame cap, so only a + // short prefix is recorded (with a truncation marker) to keep a single bad frame from + // flooding the log. + Logger.LogError($"Failed to parse an inbound JSON-RPC message: {TruncateForLog(json)}", ex); RaiseError(ex); return; } @@ -570,7 +634,7 @@ public sealed partial class JsonRpcConnection : IDisposable } else if (hasMethod && hasId) { - DispatchInboundRequest(methodElement.GetString() ?? string.Empty, idElement, root); + await EnqueueInboundRequestAsync(methodElement.GetString() ?? string.Empty, idElement, root).ConfigureAwait(false); } else if (hasId) { @@ -583,6 +647,16 @@ public sealed partial class JsonRpcConnection : IDisposable } } + internal static string TruncateForLog(string value) + { + if (value.Length <= MaxLoggedBodyChars) + { + return value; + } + + return $"{value.Substring(0, MaxLoggedBodyChars)}... [truncated; {value.Length} total characters]"; + } + private void DispatchResponse(JsonElement idElement, string json) { if (idElement.ValueKind != JsonValueKind.Number || !idElement.TryGetInt32(out var id)) @@ -691,18 +765,60 @@ public sealed partial class JsonRpcConnection : IDisposable } } - private void DispatchInboundRequest(string method, JsonElement idElement, JsonElement root) + private async Task EnqueueInboundRequestAsync(string method, JsonElement idElement, JsonElement root) { - var id = idElement.Clone(); - var parameters = root.TryGetProperty("params", out var p) ? p.Clone() : default; + // Clone the id and params so the buffered envelope stays valid after the source document is + // disposed by the read loop. + var envelope = new InboundRequestEnvelope( + method, + idElement.Clone(), + root.TryGetProperty("params", out var p) ? p.Clone() : default); - if (!_requestHandlers.TryGetValue(method, out var handler)) + try { - _ = SendErrorResponseAsync(id, JsonRpcError.MethodNotFound, $"The method '{method}' is not supported."); - return; + // Hand the request to the bounded worker pool. When the queue is full this awaits, which + // throttles the read loop (backpressure) instead of spawning unbounded handler tasks. The + // workers, not the read loop, run the handler so a slow handler cannot stall framing. + await _inboundRequestQueue.Writer.WriteAsync(envelope, _disposalCts.Token).ConfigureAwait(false); + } + catch (OperationCanceledException) when (_disposalCts.IsCancellationRequested) + { + // The connection is closing; the request is dropped along with everything else in flight. + } + catch (ChannelClosedException) + { + // The queue was completed because the connection is closing; nothing further to do. + } + } + + private async Task ProcessInboundRequestsAsync() + { + try + { + while (await _inboundRequestQueue.Reader.WaitToReadAsync(_disposalCts.Token).ConfigureAwait(false)) + { + while (_inboundRequestQueue.Reader.TryRead(out var envelope)) + { + await DispatchInboundRequestAsync(envelope).ConfigureAwait(false); + } + } + } + catch (OperationCanceledException) when (_disposalCts.IsCancellationRequested) + { + } + catch (ChannelClosedException) + { + } + } + + private Task DispatchInboundRequestAsync(InboundRequestEnvelope envelope) + { + if (!_requestHandlers.TryGetValue(envelope.Method, out var handler)) + { + return SendErrorResponseAsync(envelope.Id, JsonRpcError.MethodNotFound, $"The method '{envelope.Method}' is not supported."); } - _ = HandleInboundRequestAsync(method, id, parameters, handler); + return HandleInboundRequestAsync(envelope.Method, envelope.Id, envelope.Parameters, handler); } private async Task HandleInboundRequestAsync(string method, JsonElement id, JsonElement parameters, Func> handler) @@ -806,4 +922,10 @@ public sealed partial class JsonRpcConnection : IDisposable /// A buffered inbound notification: the method name plus a detached clone of its parameters. /// private readonly record struct NotificationEnvelope(string Method, JsonElement Parameters); + + /// + /// A buffered inbound request: the method name plus detached clones of its id and parameters, so + /// the envelope remains valid after the source document is disposed by the read loop. + /// + private readonly record struct InboundRequestEnvelope(string Method, JsonElement Id, JsonElement Parameters); } diff --git a/src/modules/cmdpal/Tests/Microsoft.CmdPal.UI.ViewModels.UnitTests/JsonRpcConnectionTests.Hardening.cs b/src/modules/cmdpal/Tests/Microsoft.CmdPal.UI.ViewModels.UnitTests/JsonRpcConnectionTests.Hardening.cs index f1c4507e4b..45e437b2dc 100644 --- a/src/modules/cmdpal/Tests/Microsoft.CmdPal.UI.ViewModels.UnitTests/JsonRpcConnectionTests.Hardening.cs +++ b/src/modules/cmdpal/Tests/Microsoft.CmdPal.UI.ViewModels.UnitTests/JsonRpcConnectionTests.Hardening.cs @@ -229,12 +229,16 @@ public partial class JsonRpcConnectionTests await gated.AfterFirstWrite.WaitAsync(cts.Token); // Dispose concurrently while the writer is blocked mid-frame. + var stopwatch = Stopwatch.StartNew(); var disposeTask = Task.Run(host.Dispose); - // Let the body finish so the writer releases the lock and Dispose can drain. + // Let the body write proceed so the writer can observe the disposal and release the lock. gated.ReleaseBody(); await disposeTask.WaitAsync(cts.Token); + stopwatch.Stop(); + + Assert.IsTrue(stopwatch.Elapsed < TimeSpan.FromSeconds(10), "Dispose during an active write must not block on the write indefinitely."); Exception? sendException = null; try @@ -246,12 +250,10 @@ public partial class JsonRpcConnectionTests sendException = ex; } + // Disposal abandons the in-flight frame and tears the connection down, so the send either + // completes or fails with a transport error, but must never surface ObjectDisposedException. Assert.IsFalse(sendException is ObjectDisposedException, "Dispose during an active write must not surface ObjectDisposedException."); - - // The frame that was mid-flight is intact, not a corrupt partial frame. - var (_, body) = await ReadFramedAsync(outputPipe.Reader.AsStream(), cts.Token); - using var document = JsonDocument.Parse(body); - Assert.AreEqual("during-dispose", document.RootElement.GetProperty("method").GetString()); + Assert.IsTrue(sendException is null or JsonRpcException, $"An abandoned write must fail with a transport error, not '{sendException?.GetType().Name}'."); } finally { @@ -261,6 +263,231 @@ public partial class JsonRpcConnectionTests } } + [TestMethod] + public async Task WriteThatNeverDrains_TimesOut_AndClosesConnection() + { + using var cts = new CancellationTokenSource(TestTimeout); + var input = new Pipe(); + var stuck = new BlockingWriteStream(); + var host = new JsonRpcConnection( + input.Reader.AsStream(), + stuck, + errorStream: null, + requestTimeout: TimeSpan.FromSeconds(30), + writeTimeout: TimeSpan.FromMilliseconds(500)); + + var disconnected = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + host.Disconnected += (_, _) => disconnected.TrySetResult(); + host.StartListening(); + + try + { + var stopwatch = Stopwatch.StartNew(); + + // The child never drains stdin, so the frame write blocks. It must be abandoned at the + // write timeout instead of hanging forever. + await Assert.ThrowsExceptionAsync(async () => + await host.SendNotificationAsync("stuck", null, cts.Token)); + stopwatch.Stop(); + + Assert.IsTrue(stopwatch.Elapsed < TimeSpan.FromSeconds(10), "A write that never drains must be abandoned at the write timeout, not hang."); + + // The stalled write tears the connection down so the owner can terminate the child. + await disconnected.Task.WaitAsync(cts.Token); + } + finally + { + host.Dispose(); + input.Writer.Complete(); + } + } + + [TestMethod] + public async Task Dispose_WhileWriteNeverDrains_CompletesAndFailsTheWrite() + { + using var cts = new CancellationTokenSource(TestTimeout); + var input = new Pipe(); + var stuck = new BlockingWriteStream(); + + // A long write timeout guarantees it is disposal, not the timeout, that unblocks the write. + var host = new JsonRpcConnection( + input.Reader.AsStream(), + stuck, + errorStream: null, + requestTimeout: TimeSpan.FromSeconds(30), + writeTimeout: TimeSpan.FromSeconds(30)); + host.StartListening(); + + var sendTask = host.SendNotificationAsync("stuck", null, cts.Token); + + try + { + // Wait until the frame write has begun and is blocked on the never-draining stream. + await stuck.WriteStarted.WaitAsync(cts.Token); + + var stopwatch = Stopwatch.StartNew(); + var disposeTask = Task.Run(host.Dispose); + await disposeTask.WaitAsync(cts.Token); + stopwatch.Stop(); + + Assert.IsTrue(stopwatch.Elapsed < TimeSpan.FromSeconds(10), "Dispose must abandon a stuck write rather than block on it indefinitely."); + + // Disposal cancels the in-flight write, which fails rather than hanging forever. + await Assert.ThrowsExceptionAsync(async () => await sendTask.WaitAsync(cts.Token)); + } + finally + { + input.Writer.Complete(); + } + } + + [TestMethod] + public void TruncateForLog_BoundsOversizedPayloads() + { + var oversized = new string('x', 5_000_000); + var truncated = JsonRpcConnection.TruncateForLog(oversized); + + Assert.IsTrue(truncated.Length < oversized.Length, "An oversized payload must be shortened before logging."); + Assert.IsTrue(truncated.Length <= JsonRpcConnection.MaxLoggedBodyChars + 128, "The truncated payload must stay near the logging cap."); + StringAssert.Contains(truncated, "truncated"); + + const string small = "{\"jsonrpc\":\"2.0\"}"; + Assert.AreEqual(small, JsonRpcConnection.TruncateForLog(small), "A payload under the cap must be logged verbatim."); + } + + [TestMethod] + public async Task OversizedMalformedBody_IsTruncatedInLogPath_AndRaisesError() + { + using var cts = new CancellationTokenSource(TestTimeout); + var harness = CreateHarness(); + + var errorRaised = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + harness.Host.Error += (_, e) => errorRaised.TrySetResult(e.Exception); + + try + { + // A large, validly framed but unparseable body. It exercises the truncated log path and + // must be reported through the Error event without hanging. + var malformed = "{" + new string('a', 2_000_000); + await WriteFramedAsync(harness.ExtensionWrites, malformed, cts.Token); + + var exception = await errorRaised.Task.WaitAsync(cts.Token); + Assert.IsInstanceOfType(exception, typeof(JsonException)); + } + finally + { + harness.Host.Dispose(); + } + } + + [TestMethod] + public async Task InboundRequests_AreConcurrencyBounded_AndDisposeCompletes() + { + using var cts = new CancellationTokenSource(TestTimeout); + var harness = CreateHarness(TimeSpan.FromSeconds(30)); + + var release = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + var saturated = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + var running = 0; + var maxObserved = 0; + + harness.Host.RegisterRequestHandler("block", async (_, token) => + { + var current = Interlocked.Increment(ref running); + + int observed; + do + { + observed = Volatile.Read(ref maxObserved); + } + while (current > observed && Interlocked.CompareExchange(ref maxObserved, current, observed) != observed); + + if (current >= JsonRpcConnection.InboundRequestWorkerCount) + { + saturated.TrySetResult(); + } + + try + { + await release.Task.WaitAsync(token).ConfigureAwait(false); + } + finally + { + Interlocked.Decrement(ref running); + } + + return new JsonObject { ["ok"] = true }; + }); + + try + { + // Flood the host with far more inbound requests than there are workers. + var total = JsonRpcConnection.InboundRequestWorkerCount * 2; + for (var i = 0; i < total; i++) + { + await WriteFramedAsync(harness.ExtensionWrites, BuildRequest(i + 1, "block", null), cts.Token); + } + + // Wait until every worker is busy, then give any excess a chance to (wrongly) start. + await saturated.Task.WaitAsync(cts.Token); + await Task.Delay(250, cts.Token); + + Assert.IsTrue( + Volatile.Read(ref maxObserved) <= JsonRpcConnection.InboundRequestWorkerCount, + $"Inbound request concurrency {Volatile.Read(ref maxObserved)} exceeded the bound of {JsonRpcConnection.InboundRequestWorkerCount}."); + } + finally + { + // Release the handlers so the workers drain, then confirm disposal completes cleanly. + release.TrySetResult(); + harness.Host.Dispose(); + } + } + + private sealed class BlockingWriteStream : Stream + { + private readonly TaskCompletionSource _writeStarted = new(TaskCreationOptions.RunContinuationsAsynchronously); + + public Task WriteStarted => _writeStarted.Task; + + public override bool CanRead => false; + + public override bool CanSeek => false; + + public override bool CanWrite => true; + + public override long Length => throw new NotSupportedException(); + + public override long Position + { + get => throw new NotSupportedException(); + set => throw new NotSupportedException(); + } + + public override async ValueTask WriteAsync(ReadOnlyMemory buffer, CancellationToken cancellationToken = default) + { + _writeStarted.TrySetResult(); + + // Never accept the bytes, but honor cancellation the way a real cancellable stream does, + // so the write timeout and disposal can abandon the write. + await Task.Delay(Timeout.Infinite, cancellationToken).ConfigureAwait(false); + } + + public override Task FlushAsync(CancellationToken cancellationToken) => Task.Delay(Timeout.Infinite, cancellationToken); + + public override void Flush() + { + } + + public override void Write(byte[] buffer, int offset, int count) => throw new NotSupportedException(); + + public override int Read(byte[] buffer, int offset, int count) => throw new NotSupportedException(); + + public override long Seek(long offset, SeekOrigin origin) => throw new NotSupportedException(); + + public override void SetLength(long value) => throw new NotSupportedException(); + } + private sealed class ThrowingWriteStream : Stream { public override bool CanRead => false;