using System.Buffers; using System.Buffers.Binary; using Cysharp.Threading.Tasks; using MessagePack; using NUnit.Framework; using ShrinkNetwork; using ShrinkNetwork.MessagePack; [MessagePackObject] public partial class CodecMessage : IShrinkNetworkMessage { [Key(0)] public int Value { get; set; } [Key(1)] public string Text { get; set; } = ""; } [GeneratedMessagePackResolver] public partial class TestMessageResolver { } public class NetworkCodecTests { [TestCase(false)] [TestCase(true)] public void EnvelopeRoundTripsPayloadWithoutEncodingItAgain(bool messagePack) { IShrinkNetworkBufferSerializer serializer; if (messagePack) { var mp = new ShrinkMessagePackNetworkSerializer(TestMessageResolver.Instance); mp.Register(); serializer = mp; } else serializer = new ShrinkJsonNetworkSerializer(); var counting = new CountingSerializer(serializer); var message = new CodecMessage { Value = 42, Text = "跨宿主消息" }; using var encoded = ShrinkPacketCodec.Encode(new ShrinkNetworkPacket { Opcode = 123, Route = "中文/route", SessionToken = "secret", RequestToken = new(9) }, counting, message); var packet = ShrinkPacketCodec.Decode(encoded.WrittenMemory); var decoded = (CodecMessage)serializer.Deserialize(packet.Payload, typeof(CodecMessage)); Assert.That(counting.Count, Is.EqualTo(1)); Assert.That(decoded.Text, Is.EqualTo(message.Text)); Assert.That(decoded.Value, Is.EqualTo(42)); Assert.That(packet.Route, Is.EqualTo("中文/route")); Assert.That(packet.SessionToken, Is.EqualTo("secret")); Assert.That(packet.RequestToken.Value, Is.EqualTo(9)); Assert.That(encoded.WrittenCount, Is.EqualTo(ShrinkPacketCodec.HeaderSize + System.Text.Encoding.UTF8.GetByteCount(packet.Route!) + 6 + packet.Payload.Length)); } [Test] public void RejectsOldVersionMalformedLengthsAndTrailingData() { Assert.Throws(() => ShrinkPacketCodec.Decode(System.Text.Encoding.UTF8.GetBytes("{\"ProtocolVersion\":1}"))); using var encoded = ShrinkPacketCodec.Encode(new ShrinkNetworkPacket { Payload = new byte[] { 1, 2, 3 } }); var bytes = encoded.WrittenMemory.ToArray(); BinaryPrimitives.WriteUInt16LittleEndian(bytes.AsSpan(4), 1); Assert.Throws(() => ShrinkPacketCodec.Decode(bytes)); BinaryPrimitives.WriteUInt16LittleEndian(bytes.AsSpan(4), 2); BinaryPrimitives.WriteInt32LittleEndian(bytes.AsSpan(29), int.MaxValue); Assert.Throws(() => ShrinkPacketCodec.Decode(bytes)); Assert.Throws(() => ShrinkPacketCodec.Decode(encoded.WrittenMemory[..^1])); Assert.Throws(() => ShrinkPacketCodec.Decode(encoded.WrittenMemory.ToArray().Concat(new byte[1]).ToArray())); } [Test] public void InvalidUtf8HeaderIsAProtocolRejection() { using var encoded = ShrinkPacketCodec.Encode(new ShrinkNetworkPacket { Route = "r", SessionToken = "t" }); foreach (var offset in new[] { ShrinkPacketCodec.HeaderSize, ShrinkPacketCodec.HeaderSize + 1 }) { var bytes = encoded.WrittenMemory.ToArray(); bytes[offset] = 0xff; Assert.Throws(() => ShrinkPacketCodec.Decode(bytes)); } } [Test] public void BorrowedDecodeAndWriterDisposeHaveExplicitLifetimes() { var writer = ShrinkPacketCodec.Encode(new ShrinkNetworkPacket { Payload = new byte[] { 7 } }); var packet = ShrinkPacketCodec.Decode(writer.WrittenMemory); Assert.That(packet.Payload.Span[0], Is.EqualTo(7)); writer.Dispose(); writer.Dispose(); Assert.Throws(() => writer.GetMemory()); } [TestCase(false)] [TestCase(true)] public async Task ServiceKeepsBufferAliveUntilTransportCompletes(bool fail) { var transport = new DeferredTransport(); var service = new ShrinkNetworkService(); service.RegisterMessage(1); service.BindTransport(transport); transport.Connect(); var sending = service.SendAsync(service.Sessions[1], new CodecMessage { Text = "alive" }).AsTask(); Assert.That(sending.IsCompleted, Is.False); Assert.That(((CodecMessage)new ShrinkJsonNetworkSerializer().Deserialize(ShrinkPacketCodec.Decode(transport.Memory).Payload, typeof(CodecMessage))).Text, Is.EqualTo("alive")); if (fail) transport.Completion.TrySetException(new IOException("disconnect")); else transport.Completion.TrySetResult(); if (fail) Assert.ThrowsAsync(async () => await sending); else await sending; Assert.That(transport.Memory.Span.ToArray(), Is.All.EqualTo(0), "pooled buffer is cleared and released after completion/failure"); } [Test] public async Task QueuedLoopbackReceiverOwnsPacketAfterSenderReturns() { using var queue = new ShrinkNetworkDispatchQueue(4); var client = new ShrinkNetworkService(); var server = new ShrinkNetworkService { DispatchScheduler = queue }; client.RegisterMessage(1); server.RegisterMessage(1); var received = ""; server.RegisterHandler((ctx, message) => { received = message.Text; return UniTask.CompletedTask; }); var left = new ShrinkLoopbackTransport(); var right = new ShrinkLoopbackTransport(); left.LinkPeer(right); client.BindTransport(left); server.BindTransport(right); left.OpenSession(1); await client.SendAsync(client.Sessions[1], new CodecMessage { Text = "retained" }); Assert.That(received, Is.Empty); await queue.PumpAsync(4); Assert.That(received, Is.EqualTo("retained")); } private sealed class CountingSerializer(IShrinkNetworkBufferSerializer inner) : IShrinkNetworkBufferSerializer { public int Count; public void Serialize(IBufferWriter writer, object value) { Count++; inner.Serialize(writer, value); } public object Deserialize(ReadOnlyMemory payload, Type type) => inner.Deserialize(payload, type); public byte[] Serialize(object value) => throw new AssertionException("Legacy allocation path used"); public object Deserialize(byte[] payload, Type type) => inner.Deserialize(payload, type); public T Deserialize(byte[] payload) => inner.Deserialize(payload); } private sealed class DeferredTransport : IShrinkNetworkMemoryTransport { public bool IsStarted { get; private set; } public event Action? OnEvent; public UniTaskCompletionSource Completion = new(); public ReadOnlyMemory Memory; public void Start() => IsStarted = true; public void Stop() => IsStarted = false; public void Connect() => OnEvent?.Invoke(ShrinkNetworkTransportEvent.Connected(1, "test")); public void Send(long id, byte[] data) => throw new AssertionException("Legacy transport used"); public UniTask SendAsync(long id, byte[] data) => throw new AssertionException("Legacy transport used"); public UniTask SendAsync(long id, ReadOnlyMemory data, CancellationToken ct = default) { Memory = data; return Completion.Task; } } }