diff --git a/src/Grpc.Net.Client/Internal/GrpcCall.NonGeneric.cs b/src/Grpc.Net.Client/Internal/GrpcCall.NonGeneric.cs index 9876cf262..1b5af62cd 100644 --- a/src/Grpc.Net.Client/Internal/GrpcCall.NonGeneric.cs +++ b/src/Grpc.Net.Client/Internal/GrpcCall.NonGeneric.cs @@ -37,9 +37,34 @@ internal abstract class GrpcCall public bool ResponseFinished { get; protected set; } public HttpResponseMessage? HttpResponse { get; protected set; } - public GrpcCallSerializationContext SerializationContext + /// + /// Rents a for a single serialization and write operation. + /// At most one idle context should be cached per + /// + internal SerializationContextLease RentSerializationContext(CallOptions callOptions) { - get { return _serializationContext ??= new GrpcCallSerializationContext(this); } + var context = Interlocked.Exchange(ref _serializationContext, null) ?? new GrpcCallSerializationContext(this); + + try + { + context.CallOptions = callOptions; + context.Initialize(); + return new SerializationContextLease(this, context); + } + catch + { + // Don't cache a context after initialization fails + context.Reset(); + throw; + } + } + + /// + /// Returns a context to the cache + /// + internal void ReturnSerializationContext(GrpcCallSerializationContext context) + { + Interlocked.CompareExchange(ref _serializationContext, context, null); } public DefaultDeserializationContext DeserializationContext diff --git a/src/Grpc.Net.Client/Internal/Http/WinHttpUnaryContent.cs b/src/Grpc.Net.Client/Internal/Http/WinHttpUnaryContent.cs index 39e946a85..f07816f42 100644 --- a/src/Grpc.Net.Client/Internal/Http/WinHttpUnaryContent.cs +++ b/src/Grpc.Net.Client/Internal/Http/WinHttpUnaryContent.cs @@ -77,19 +77,19 @@ protected override bool TryComputeLength(out long length) private int GetPayloadLength() { - var serializationContext = _call.SerializationContext; - serializationContext.CallOptions = _call.Options; - serializationContext.Initialize(); + var lease = _call.RentSerializationContext(_call.Options); try { - _call.Method.RequestMarshaller.ContextualSerializer(_request, serializationContext); + _call.Method.RequestMarshaller.ContextualSerializer(_request, lease.Context); - return serializationContext.GetWrittenPayload().Length; + var length = lease.Context.GetWrittenPayload().Length; + lease.MarkReusable(); + return length; } finally { - serializationContext.Reset(); + lease.Dispose(); } } } diff --git a/src/Grpc.Net.Client/Internal/Retry/RetryCallBase.cs b/src/Grpc.Net.Client/Internal/Retry/RetryCallBase.cs index 15769bb69..891888eef 100644 --- a/src/Grpc.Net.Client/Internal/Retry/RetryCallBase.cs +++ b/src/Grpc.Net.Client/Internal/Retry/RetryCallBase.cs @@ -301,20 +301,21 @@ protected bool IsDeadlineExceeded() protected byte[] SerializePayload(GrpcCall call, CallOptions callOptions, TRequest request) { - var serializationContext = call.SerializationContext; - serializationContext.CallOptions = callOptions; - serializationContext.Initialize(); + var lease = call.RentSerializationContext(callOptions); try { - call.Method.RequestMarshaller.ContextualSerializer(request, serializationContext); + call.Method.RequestMarshaller.ContextualSerializer(request, lease.Context); // Need to take a copy because the serialization context will returned a rented buffer. - return serializationContext.GetWrittenPayload().ToArray(); + var payload = lease.Context.GetWrittenPayload().ToArray(); + + lease.MarkReusable(); + return payload; } finally { - serializationContext.Reset(); + lease.Dispose(); } } diff --git a/src/Grpc.Net.Client/Internal/SerializationContextLease.cs b/src/Grpc.Net.Client/Internal/SerializationContextLease.cs new file mode 100644 index 000000000..76ba5a9f6 --- /dev/null +++ b/src/Grpc.Net.Client/Internal/SerializationContextLease.cs @@ -0,0 +1,65 @@ +#region Copyright notice and license + +// Copyright 2019 The gRPC Authors +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#endregion + +namespace Grpc.Net.Client.Internal; + +/// +/// A temporary handle to a for a single usage. +/// +/// Callers should:
+/// - call only after successful operation. +/// Serialization context will be reused by next renter.
+/// - call once completed usage. +///
+internal struct SerializationContextLease : IDisposable +{ + private readonly GrpcCall _call; + private GrpcCallSerializationContext? _context; + private bool _reusable; + + internal SerializationContextLease(GrpcCall call, GrpcCallSerializationContext context) + { + _call = call; + _context = context; + _reusable = false; + } + + public readonly GrpcCallSerializationContext Context => _context!; + + /// + /// Marks the context as safe to hand back for reuse once disposed. + /// + public void MarkReusable() => _reusable = true; + + public void Dispose() + { + var context = _context; + if (context == null) + { + return; + } + + _context = null; + context.Reset(); // Always release the rented payload buffer. + + if (_reusable) + { + _call.ReturnSerializationContext(context); + } + } +} diff --git a/src/Grpc.Net.Client/Internal/StreamExtensions.cs b/src/Grpc.Net.Client/Internal/StreamExtensions.cs index 7048a6fba..0cf84e35c 100644 --- a/src/Grpc.Net.Client/Internal/StreamExtensions.cs +++ b/src/Grpc.Net.Client/Internal/StreamExtensions.cs @@ -295,21 +295,22 @@ public static async Task WriteMessageAsync( CallOptions callOptions) { // Sync relevant changes here with other WriteMessageAsync - var serializationContext = call.SerializationContext; - serializationContext.CallOptions = callOptions; - serializationContext.Initialize(); + var lease = call.RentSerializationContext(callOptions); try { GrpcCallLog.SendingMessage(call.Logger); // Serialize message first. Need to know size to prefix the length in the header - serializer(message, serializationContext); + serializer(message, lease.Context); // Sending the header+content in a single WriteAsync call has significant performance benefits // https://github.com/dotnet/runtime/issues/35184#issuecomment-626304981 - await stream.WriteAsync(serializationContext.GetWrittenPayload(), call.CancellationToken).ConfigureAwait(false); + await stream.WriteAsync(lease.Context.GetWrittenPayload(), call.CancellationToken).ConfigureAwait(false); GrpcCallLog.MessageSent(call.Logger); + + // The write fully completed - safe to hand this context back for reuse. + lease.MarkReusable(); } catch (Exception ex) { @@ -328,7 +329,7 @@ public static async Task WriteMessageAsync( } finally { - serializationContext.Reset(); + lease.Dispose(); } } diff --git a/test/Grpc.Net.Client.Tests/GrpcCallSerializationContextTests.cs b/test/Grpc.Net.Client.Tests/GrpcCallSerializationContextTests.cs index 1f803483f..5107edf89 100644 --- a/test/Grpc.Net.Client.Tests/GrpcCallSerializationContextTests.cs +++ b/test/Grpc.Net.Client.Tests/GrpcCallSerializationContextTests.cs @@ -21,6 +21,7 @@ using Grpc.Core; using Grpc.Net.Client.Internal; using Grpc.Net.Client.Tests.Infrastructure; +using Grpc.Tests.Shared; using NUnit.Framework; namespace Grpc.Net.Client.Tests; @@ -309,6 +310,118 @@ public void Reset_AfterGetBufferWriter_RemovesPayload() Assert.AreEqual("Serialization did not return a payload.", ex.Message); } + [Test] + public void RentSerializationContext_SequentialUse_ReusesSameInstance() + { + // Arrange + var call = CreateCall(); + + // Act - first, fully completed lease. + var lease1 = call.RentSerializationContext(new CallOptions()); + var context1 = lease1.Context; + lease1.MarkReusable(); + lease1.Dispose(); + + // A second, later (non-overlapping) lease. + var lease2 = call.RentSerializationContext(new CallOptions()); + var context2 = lease2.Context; + lease2.MarkReusable(); + lease2.Dispose(); + + // Assert - context is cached and reused + Assert.AreSame(context1, context2); + } + + [Test] + public void RentSerializationContext_OverlappingUse_DoesNotShareInstance() + { + // Arrange + var call = CreateCall(); + + // Act - first lease is rented but not yet returned, simulating a write still in flight + var lease1 = call.RentSerializationContext(new CallOptions()); + + // A second lease overlaps with the first. + var lease2 = call.RentSerializationContext(new CallOptions()); + + try + { + // Assert - the two overlapping operations must never share the same mutable context/buffer + Assert.AreNotSame(lease1.Context, lease2.Context); + } + finally + { + lease1.Dispose(); + lease2.Dispose(); + } + } + + [Test] + public void RentSerializationContext_NotMarkedReusable_IsNotCached() + { + // Arrange + var call = CreateCall(); + + // Act - a lease that is disposed without MarkReusable must not be handed back for reuse. + var lease1 = call.RentSerializationContext(new CallOptions()); + var context1 = lease1.Context; + lease1.Dispose(); + + var lease2 = call.RentSerializationContext(new CallOptions()); + + try + { + // Assert + Assert.AreNotSame(context1, lease2.Context); + } + finally + { + lease2.MarkReusable(); + lease2.Dispose(); + } + } + + [Test] + public async Task WriteMessageAsync_ConcurrentWrites_PayloadsAreIndependent() + { + // Arrange + var call = CreateCall(); + var stream1 = new BlockingWriteStream(); + var stream2 = new BlockingWriteStream(); + + // Act + var writeTask1 = stream1.WriteMessageAsync(call, (byte)1, SerializeByte, new CallOptions()); + var payload1 = await stream1.WriteStartedTask.DefaultTimeout(); + + var writeTask2 = stream2.WriteMessageAsync(call, (byte)2, SerializeByte, new CallOptions()); + var payload2 = await stream2.WriteStartedTask.DefaultTimeout(); + + try + { + // Assert + Assert.AreEqual(1, payload1.Span[GrpcProtocolConstants.HeaderSize]); + Assert.AreEqual(2, payload2.Span[GrpcProtocolConstants.HeaderSize]); + } + finally + { + stream1.Continue(); + stream2.Continue(); + } + + await Task.WhenAll(writeTask1, writeTask2).DefaultTimeout(); + + static void SerializeByte(byte value, SerializationContext serializationContext) + { + serializationContext.SetPayloadLength(1); + + var bufferWriter = serializationContext.GetBufferWriter(); + bufferWriter.GetSpan(1)[0] = value; + bufferWriter.Advance(1); + + serializationContext.Complete(); + } + } + private class TestGrpcCall : GrpcCall { public TestGrpcCall(CallOptions options, GrpcChannel channel) : base(options, channel) @@ -321,7 +434,56 @@ public TestGrpcCall(CallOptions options, GrpcChannel channel) : base(options, ch public override Task CallTask => Task.FromResult(Status.DefaultCancelled); } + private sealed class BlockingWriteStream : Stream + { + private readonly TaskCompletionSource> _writeStartedTcs = new TaskCompletionSource>(TaskCreationOptions.RunContinuationsAsynchronously); + private readonly TaskCompletionSource _continueTcs = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + + public Task> WriteStartedTask => _writeStartedTcs.Task; + + public void Continue() => _continueTcs.TrySetResult(null); + + 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 void Flush() + { + } + + 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(); + public override void Write(byte[] buffer, int offset, int count) => throw new NotSupportedException(); + +#if NET462 + public override async Task WriteAsync(byte[] buffer, int offset, int count, CancellationToken cancellationToken) + { + _writeStartedTcs.TrySetResult(buffer.AsMemory(offset, count)); + await _continueTcs.Task.ConfigureAwait(false); + } +#else + public override async ValueTask WriteAsync(ReadOnlyMemory buffer, CancellationToken cancellationToken = default) + { + _writeStartedTcs.TrySetResult(buffer); + await _continueTcs.Task.ConfigureAwait(false); + } +#endif + } + private GrpcCallSerializationContext CreateSerializationContext(string? requestGrpcEncoding = null, int? maxSendMessageSize = null) + { + var call = CreateCall(requestGrpcEncoding, maxSendMessageSize); + return new GrpcCallSerializationContext(call); + } + + private TestGrpcCall CreateCall(string? requestGrpcEncoding = null, int? maxSendMessageSize = null) { var channelOptions = new GrpcChannelOptions(); channelOptions.MaxSendMessageSize = maxSendMessageSize; @@ -330,7 +492,7 @@ private GrpcCallSerializationContext CreateSerializationContext(string? requestG var call = new TestGrpcCall(new CallOptions(), GrpcChannel.ForAddress("http://localhost", channelOptions)); call.RequestGrpcEncoding = requestGrpcEncoding ?? "identity"; - return new GrpcCallSerializationContext(call); + return call; } private static (bool Compressed, int Length) DecodeHeader(ReadOnlySpan buffer)