feat(network)!: add v2 framing, pooled buffers and bounded dispatch
Publish UPM package / publish (push) Successful in 2s
Publish NuGet packages / publish (push) Successful in 3m5s

This commit is contained in:
2026-09-29 10:14:47 +08:00
parent c7c45b26f8
commit 93c2b8520b
36 changed files with 752 additions and 398 deletions
+105 -177
View File
@@ -54,7 +54,7 @@ namespace ShrinkNetwork
Serializer = serializer ?? throw new ArgumentNullException(nameof(serializer));
MessageRegistry = messageRegistry ?? throw new ArgumentNullException(nameof(messageRegistry));
Router = router ?? throw new ArgumentNullException(nameof(router));
_dispatchScheduler = dispatchScheduler ?? ShrinkNetworkDispatchSchedulers.Inline;
_dispatchScheduler = dispatchScheduler ?? new ShrinkNetworkInlineDispatchScheduler();
}
public IShrinkNetworkSerializer Serializer { get; }
@@ -74,6 +74,8 @@ namespace ShrinkNetwork
set => _dispatchScheduler = value ?? throw new ArgumentNullException(nameof(value));
}
/// <summary>Reliable receive work could not be processed. Session-control transports also disconnect the peer.</summary>
public event Action<long>? OnDispatchRejected;
public event Action<ShrinkNetworkSession>? OnSessionConnected;
public event Action<ShrinkNetworkSession>? OnSessionDisconnected;
@@ -273,11 +275,11 @@ namespace ShrinkNetwork
if (messageType == null)
throw new ArgumentNullException(nameof(messageType));
return SendPacketAsync(session, messageType, kind, requestToken, route, Serializer.Serialize(message));
return SendPacketAsync(session, messageType, kind, requestToken, route, Array.Empty<byte>(), message);
}
private async UniTask SendPacketAsync(ShrinkNetworkSession session, Type messageType,
ShrinkNetworkPacketKind kind, ShrinkRequestToken requestToken, string? route, byte[] payload)
ShrinkNetworkPacketKind kind, ShrinkRequestToken requestToken, string? route, byte[] payload, object? message = null)
{
if (session == null)
throw new ArgumentNullException(nameof(session));
@@ -299,16 +301,22 @@ namespace ShrinkNetwork
Payload = payload
};
var packetData = Serializer.Serialize(packet);
using var encoded = ShrinkPacketCodec.Encode(packet, Serializer, message);
var packetData = encoded.WrittenMemory;
Interlocked.Increment(ref _packetsSent);
Interlocked.Add(ref _bytesSent, packetData.Length);
if (transport is IShrinkNetworkMemoryTransport memoryTransport)
{
await memoryTransport.SendAsync(session.SessionId, packetData);
return;
}
if (transport is IShrinkNetworkAsyncTransport asyncTransport)
{
await asyncTransport.SendAsync(session.SessionId, packetData);
await asyncTransport.SendAsync(session.SessionId, packetData.ToArray());
return;
}
transport.Send(session.SessionId, packetData);
transport.Send(session.SessionId, packetData.ToArray());
}
private void OnTransportEvent(ShrinkNetworkTransportEvent evt)
@@ -330,6 +338,12 @@ namespace ShrinkNetwork
return;
}
// Responses complete pending RPCs independently of the serial handler queue, including nested calls.
if (evt.Type == ShrinkNetworkTransportEventType.Packet && ShrinkPacketCodec.IsResponse(evt.PacketData))
{
HandlePacketAsync(evt).Forget();
return;
}
ScheduleTransportEventAsync(evt).Forget();
}
@@ -337,9 +351,23 @@ namespace ShrinkNetwork
{
try
{
var scheduled = await DispatchScheduler.ScheduleAsync(() => HandleTransportEventAsync(evt));
var scheduler = DispatchScheduler;
_sessions.TryGetValue(evt.SessionId, out var receivedSession);
UniTask DispatchCurrent()
{
return receivedSession != null && _sessions.TryGetValue(evt.SessionId, out var current) && ReferenceEquals(current, receivedSession)
? HandleTransportEventAsync(evt) : UniTask.CompletedTask;
}
var scheduled = scheduler is IShrinkNetworkPacketDispatchScheduler queue
? await queue.ScheduleAsync(evt.SessionId, evt.PacketData?.Length ?? 0, DispatchCurrent)
: await scheduler.ScheduleAsync(DispatchCurrent);
if (!scheduled)
{
Interlocked.Increment(ref _dispatchQueueRejectedCount);
OnDispatchRejected?.Invoke(evt.SessionId);
if (_transport is IShrinkNetworkSessionControlTransport control)
control.DisconnectSession(evt.SessionId, "SHRINK-NET-CONGESTION: reliable receive queue rejected work.");
}
}
catch (Exception ex)
{
@@ -395,12 +423,12 @@ namespace ShrinkNetwork
{
Interlocked.Increment(ref _packetsReceived);
Interlocked.Add(ref _bytesReceived, evt.PacketData.Length);
var packet = Serializer.Deserialize<ShrinkNetworkPacket>(evt.PacketData);
var packet = ShrinkPacketCodec.Decode(evt.PacketData);
if (!_sessions.TryGetValue(evt.SessionId, out var session))
{
session = _sessions.GetOrAdd(evt.SessionId,
id => new ShrinkNetworkSession(id, evt.RemoteAddress, this));
// Queued packets from a disconnected session cannot resurrect its state.
return;
}
if (!ValidatePacketCompatibility(packet, evt.SessionId))
@@ -410,7 +438,7 @@ namespace ShrinkNetwork
if (packet.Kind == ShrinkNetworkPacketKind.Response)
{
HandleResponse(packet);
HandleResponse(session, packet);
return;
}
@@ -425,7 +453,7 @@ namespace ShrinkNetwork
}
var resolvedMeta = meta!;
var message = Serializer.Deserialize(packet.Payload, resolvedMeta.MessageType);
var message = DeserializePayload(packet.Payload, resolvedMeta.MessageType);
if (message == null)
{
ShrinkNetworkLogger.Warn($"[ShrinkNetwork] Failed to deserialize message for opcode {packet.Opcode}.");
@@ -442,6 +470,12 @@ namespace ShrinkNetwork
ShrinkNetworkLogger.Warn($"[ShrinkNetwork] No handler found for {resolvedMeta.MessageType.FullName}");
}
}
catch (ShrinkProtocolException ex)
{
Interlocked.Increment(ref _protocolViolations);
if (DisconnectOnProtocolViolation && _transport is IShrinkNetworkSessionControlTransport control)
control.DisconnectSession(evt.SessionId, ex.Message);
}
catch (Exception ex)
{
Interlocked.Increment(ref _serializationErrorCount);
@@ -450,9 +484,10 @@ namespace ShrinkNetwork
}
}
private void HandleResponse(ShrinkNetworkPacket packet)
private void HandleResponse(ShrinkNetworkSession session, ShrinkNetworkPacket packet)
{
if (!_pendingRequests.TryRemove(packet.RequestToken, out var pending))
if (!_pendingRequests.TryGetValue(packet.RequestToken, out var expected) || expected.SessionId != session.SessionId ||
!_pendingRequests.TryRemove(packet.RequestToken, out var pending))
{
ShrinkNetworkLogger.Warn($"[ShrinkNetwork] Pending request not found. RequestToken={packet.RequestToken}");
return;
@@ -460,7 +495,7 @@ namespace ShrinkNetwork
try
{
var response = Serializer.Deserialize(packet.Payload, pending.ResponseType);
var response = DeserializePayload(packet.Payload, pending.ResponseType);
pending.CompletionSource.TrySetResult(response);
}
catch (Exception ex)
@@ -469,6 +504,10 @@ namespace ShrinkNetwork
}
}
private object DeserializePayload(ReadOnlyMemory<byte> payload, Type type) =>
Serializer is IShrinkNetworkBufferSerializer buffered
? buffered.Deserialize(payload, type) : Serializer.Deserialize(payload.ToArray(), type);
private async UniTask<TResponse> WaitForPendingResponse<TResponse>(ShrinkRequestToken requestToken, PendingRequest pending,
ShrinkRpcCallOptions? options)
where TResponse : class, IShrinkNetworkResponse
@@ -627,20 +666,42 @@ namespace ShrinkNetwork
public static class ShrinkNetworkDispatchSchedulers
{
public static IShrinkNetworkDispatchScheduler Inline { get; } =
new ShrinkNetworkInlineDispatchScheduler();
public static IShrinkNetworkDispatchScheduler Inline => new ShrinkNetworkInlineDispatchScheduler();
}
public sealed class ShrinkNetworkInlineDispatchScheduler : IShrinkNetworkDispatchScheduler
public interface IShrinkNetworkPacketDispatchScheduler : IShrinkNetworkDispatchScheduler
{
public async UniTask<bool> ScheduleAsync(Func<UniTask> callback)
{
if (callback == null)
throw new ArgumentNullException(nameof(callback));
UniTask<bool> ScheduleAsync(long sessionId, int bytes, Func<UniTask> callback);
}
await callback();
return true;
/// <summary>Automatically drains a bounded serial queue. Callbacks start on the initiating transport/continuation thread.</summary>
public sealed class ShrinkNetworkInlineDispatchScheduler : IShrinkNetworkPacketDispatchScheduler, IDisposable
{
private readonly object _gate = new();
private readonly ShrinkNetworkWorkQueue _queue = new(4096, 64 * 1024 * 1024);
private bool _draining;
public ShrinkNetworkQueueDiagnostics CaptureDiagnostics() => _queue.CaptureDiagnostics();
public UniTask<bool> ScheduleAsync(Func<UniTask> callback) => ScheduleAsync(0, 0, callback);
public async UniTask<bool> ScheduleAsync(long sessionId, int bytes, Func<UniTask> callback)
{
var completion = _queue.EnqueueAsync(sessionId, "receive", bytes, callback);
var start = false;
lock (_gate) { if (!_draining) { _draining = true; start = true; } }
if (start) DrainAsync().Forget();
return await completion == ShrinkNetworkQueueResult.Completed;
}
private async UniTask DrainAsync()
{
while (true)
{
await _queue.PumpAsync(int.MaxValue);
lock (_gate)
{
if (_queue.PendingCount == 0) { _draining = false; return; }
}
}
}
public void Dispose() => _queue.Dispose();
}
public enum ShrinkNetworkDispatchOverflowPolicy
@@ -650,163 +711,30 @@ namespace ShrinkNetwork
DropOldest = 2
}
/// <summary>
/// A caller-pumped, bounded dispatch queue. Unity can pump it from Update
/// while a dedicated server can keep the default inline scheduler.
/// </summary>
public sealed class ShrinkNetworkDispatchQueue : IShrinkNetworkDispatchScheduler, IDisposable
/// <summary>Caller-pumped serial receive queue; rejects overflow without silently dropping reliable packets.</summary>
public sealed class ShrinkNetworkDispatchQueue : IShrinkNetworkPacketDispatchScheduler, IDisposable
{
private sealed class WorkItem
{
public Func<UniTask> Callback = null!;
public UniTaskCompletionSource<bool> Completion = null!;
}
private readonly ConcurrentQueue<WorkItem> _queue = new();
private readonly object _lifecycleLock = new();
private readonly int _capacity;
private readonly ShrinkNetworkDispatchOverflowPolicy _overflowPolicy;
private int _queuedCount;
private int _pumping;
private int _disposed;
private long _rejectedCount;
private long _droppedCount;
private readonly ShrinkNetworkWorkQueue _queue;
public ShrinkNetworkDispatchQueue(int capacity,
ShrinkNetworkDispatchOverflowPolicy overflowPolicy = ShrinkNetworkDispatchOverflowPolicy.Reject)
ShrinkNetworkDispatchOverflowPolicy overflowPolicy = ShrinkNetworkDispatchOverflowPolicy.Reject,
long byteCapacity = 64 * 1024 * 1024, int perSessionCapacity = int.MaxValue)
{
if (capacity <= 0)
throw new ArgumentOutOfRangeException(nameof(capacity));
_capacity = capacity;
_overflowPolicy = overflowPolicy;
}
public int Capacity => _capacity;
public int PendingCount => Volatile.Read(ref _queuedCount);
public long RejectedCount => Volatile.Read(ref _rejectedCount);
public long DroppedCount => Volatile.Read(ref _droppedCount);
public UniTask<bool> ScheduleAsync(Func<UniTask> callback)
{
if (callback == null)
throw new ArgumentNullException(nameof(callback));
if (Volatile.Read(ref _disposed) != 0)
return UniTask.FromException<bool>(new ObjectDisposedException(nameof(ShrinkNetworkDispatchQueue)));
var item = new WorkItem
{
Callback = callback,
Completion = new UniTaskCompletionSource<bool>()
};
while (true)
{
if (Volatile.Read(ref _queuedCount) >= _capacity)
{
switch (_overflowPolicy)
{
case ShrinkNetworkDispatchOverflowPolicy.Reject:
Interlocked.Increment(ref _rejectedCount);
return UniTask.FromResult(false);
case ShrinkNetworkDispatchOverflowPolicy.DropNewest:
Interlocked.Increment(ref _droppedCount);
return UniTask.FromResult(false);
case ShrinkNetworkDispatchOverflowPolicy.DropOldest:
if (_queue.TryDequeue(out var dropped))
{
Interlocked.Decrement(ref _queuedCount);
Interlocked.Increment(ref _droppedCount);
dropped.Completion.TrySetResult(false);
continue;
}
Thread.Yield();
continue;
default:
throw new ArgumentOutOfRangeException();
}
}
var currentCount = Volatile.Read(ref _queuedCount);
if (currentCount >= _capacity ||
Interlocked.CompareExchange(ref _queuedCount, currentCount + 1, currentCount) != currentCount)
{
continue;
}
lock (_lifecycleLock)
{
if (Volatile.Read(ref _disposed) != 0)
{
Interlocked.Decrement(ref _queuedCount);
item.Completion.TrySetResult(false);
return item.Completion.Task;
}
_queue.Enqueue(item);
return item.Completion.Task;
}
}
}
public UniTask<int> PumpAsync(int maxItems)
{
if (maxItems <= 0)
throw new ArgumentOutOfRangeException(nameof(maxItems));
if (Interlocked.Exchange(ref _pumping, 1) == 1)
return UniTask.FromResult(0);
return PumpCoreAsync(maxItems);
}
public void Dispose()
{
lock (_lifecycleLock)
{
if (Interlocked.Exchange(ref _disposed, 1) != 0)
return;
while (_queue.TryDequeue(out var item))
{
Interlocked.Decrement(ref _queuedCount);
item.Completion.TrySetResult(false);
}
}
}
private async UniTask<int> PumpCoreAsync(int maxItems)
{
var processed = 0;
try
{
while (processed < maxItems && _queue.TryDequeue(out var item))
{
Interlocked.Decrement(ref _queuedCount);
await ExecuteItemAsync(item);
processed++;
}
return processed;
}
finally
{
Volatile.Write(ref _pumping, 0);
}
}
private static async UniTask ExecuteItemAsync(WorkItem item)
{
try
{
await item.Callback();
item.Completion.TrySetResult(true);
}
catch (Exception ex)
{
item.Completion.TrySetException(ex);
}
if (overflowPolicy != ShrinkNetworkDispatchOverflowPolicy.Reject)
throw new ArgumentException("Protocol v2 reliable dispatch requires Reject. Use an explicit state key on ShrinkNetworkWorkQueue for replaceable state.", nameof(overflowPolicy));
Capacity = capacity;
_queue = new ShrinkNetworkWorkQueue(capacity, byteCapacity, perSessionCapacity);
}
public int Capacity { get; }
public int PendingCount => _queue.PendingCount;
public long RejectedCount => CaptureDiagnostics().Rejected;
public long DroppedCount => 0;
public ShrinkNetworkQueueDiagnostics CaptureDiagnostics() => _queue.CaptureDiagnostics();
public UniTask<bool> ScheduleAsync(Func<UniTask> callback) => ScheduleAsync(0, 0, callback);
public async UniTask<bool> ScheduleAsync(long sessionId, int bytes, Func<UniTask> callback) =>
await _queue.EnqueueAsync(sessionId, "receive", bytes, callback) == ShrinkNetworkQueueResult.Completed;
public UniTask<int> PumpAsync(int maxItems) => _queue.PumpAsync(maxItems);
public UniTask<int> PumpAsync(int maxItems, long maxBytes, TimeSpan timeBudget) => _queue.PumpAsync(maxItems, maxBytes, timeBudget);
public void Dispose() => _queue.Dispose();
}
public sealed class ShrinkNetworkServiceDiagnosticsSnapshot
+159
View File
@@ -0,0 +1,159 @@
#nullable enable
using System;
using System.Collections.Generic;
using System.Diagnostics;
using System.Threading;
using Cysharp.Threading.Tasks;
namespace ShrinkNetwork
{
public enum ShrinkNetworkQueueResult { Completed, Rejected, Replaced, Canceled }
public sealed class ShrinkNetworkQueueDiagnostics
{
public int PendingCount { get; internal set; }
public long PendingBytes { get; internal set; }
public long Rejected { get; internal set; }
public long Replaced { get; internal set; }
public long Completed { get; internal set; }
public double OldestWaitMilliseconds { get; internal set; }
public double LastWaitMilliseconds { get; internal set; }
}
/// <summary>Caller-pumped serial execution, fair across session/channel partitions. Only queued state is replaceable.</summary>
public sealed class ShrinkNetworkWorkQueue : IDisposable
{
private sealed class Work
{
public Func<UniTask> Callback = null!;
public UniTaskCompletionSource<ShrinkNetworkQueueResult> Completion = new();
public string? StateKey;
public int Bytes;
public long Enqueued = Stopwatch.GetTimestamp();
public CancellationToken Cancellation;
}
private sealed class Partition
{
public readonly LinkedList<Work> Items = new();
public readonly Dictionary<string, LinkedListNode<Work>> States = new(StringComparer.Ordinal);
}
private readonly object _gate = new();
private readonly Dictionary<(long, string), Partition> _partitions = new();
private readonly Queue<(long, string)> _ready = new();
private readonly int _capacity;
private readonly long _byteCapacity;
private readonly int _partitionCapacity;
private int _count, _pumping;
private long _bytes, _rejected, _replaced, _completed;
private double _lastWait;
private bool _disposed;
public int PendingCount { get { lock (_gate) return _count; } }
public ShrinkNetworkWorkQueue(int capacity, long byteCapacity, int perPartitionCapacity = int.MaxValue)
{
if (capacity <= 0 || byteCapacity <= 0 || perPartitionCapacity <= 0) throw new ArgumentOutOfRangeException(nameof(capacity));
_capacity = capacity; _byteCapacity = byteCapacity; _partitionCapacity = perPartitionCapacity;
}
/// <summary>A nonempty stateKey explicitly permits replacing an unsent state in this session/channel. Never use for RPC or snapshot fragments.</summary>
public UniTask<ShrinkNetworkQueueResult> EnqueueAsync(long sessionId, string channel, int byteCount,
Func<UniTask> callback, string? stateKey = null, CancellationToken cancellationToken = default)
{
if (callback == null) throw new ArgumentNullException(nameof(callback));
if (channel == null) throw new ArgumentNullException(nameof(channel));
if (byteCount < 0) throw new ArgumentOutOfRangeException(nameof(byteCount));
Work? replaced = null;
Work item;
lock (_gate)
{
if (_disposed || cancellationToken.IsCancellationRequested) return UniTask.FromResult(ShrinkNetworkQueueResult.Canceled);
var key = (sessionId, channel);
_partitions.TryGetValue(key, out var partition);
LinkedListNode<Work>? old = null;
if (!string.IsNullOrEmpty(stateKey)) partition?.States.TryGetValue(stateKey!, out old);
var nextBytes = _bytes - (old?.Value.Bytes ?? 0) + byteCount;
if (nextBytes > _byteCapacity || (old == null && (_count >= _capacity || (partition?.Items.Count ?? 0) >= _partitionCapacity)))
{ _rejected++; return UniTask.FromResult(ShrinkNetworkQueueResult.Rejected); }
item = new Work { Callback = callback, Bytes = byteCount, StateKey = string.IsNullOrEmpty(stateKey) ? null : stateKey, Cancellation = cancellationToken };
if (partition == null) { partition = new Partition(); _partitions.Add(key, partition); _ready.Enqueue(key); }
if (old != null)
{
replaced = old.Value;
// Move a replacement to the tail: later state must not jump ahead of intervening reliable operations.
partition.Items.Remove(old); _replaced++;
}
else _count++;
var node = partition.Items.AddLast(item);
if (item.StateKey != null) partition.States[item.StateKey] = node;
_bytes = nextBytes;
}
replaced?.Completion.TrySetResult(ShrinkNetworkQueueResult.Replaced);
return item.Completion.Task;
}
public async UniTask<int> PumpAsync(int maxItems, long maxBytes = long.MaxValue, TimeSpan? timeBudget = null)
{
if (maxItems <= 0 || maxBytes <= 0 || (timeBudget.HasValue && timeBudget.Value <= TimeSpan.Zero)) throw new ArgumentOutOfRangeException(nameof(maxItems));
if (Interlocked.Exchange(ref _pumping, 1) != 0) return 0;
var started = Stopwatch.GetTimestamp();
var processed = 0;
long bytes = 0;
try
{
while (processed < maxItems && (!timeBudget.HasValue || Elapsed(started) < timeBudget.Value.TotalMilliseconds))
{
Work item;
lock (_gate)
{
if (_disposed || _ready.Count == 0) break;
var key = _ready.Peek();
var partition = _partitions[key];
item = partition.Items.First!.Value;
// Allow one oversized item so a byte budget cannot permanently starve a valid packet.
if (processed > 0 && item.Bytes > maxBytes - bytes) break;
_ready.Dequeue(); partition.Items.RemoveFirst();
if (item.StateKey != null) partition.States.Remove(item.StateKey);
if (partition.Items.Count == 0) _partitions.Remove(key); else _ready.Enqueue(key);
_count--; _bytes -= item.Bytes; _lastWait = Elapsed(item.Enqueued);
}
try
{
if (item.Cancellation.IsCancellationRequested) item.Completion.TrySetResult(ShrinkNetworkQueueResult.Canceled);
else { await item.Callback(); item.Completion.TrySetResult(ShrinkNetworkQueueResult.Completed); lock (_gate) _completed++; }
}
catch (OperationCanceledException) { item.Completion.TrySetResult(ShrinkNetworkQueueResult.Canceled); }
catch (Exception ex) { item.Completion.TrySetException(ex); }
processed++; bytes += item.Bytes;
}
return processed;
}
finally { Volatile.Write(ref _pumping, 0); }
}
public ShrinkNetworkQueueDiagnostics CaptureDiagnostics()
{
lock (_gate)
{
long oldest = Stopwatch.GetTimestamp();
foreach (var partition in _partitions.Values)
if (partition.Items.First != null) oldest = Math.Min(oldest, partition.Items.First.Value.Enqueued);
return new ShrinkNetworkQueueDiagnostics { PendingCount = _count, PendingBytes = _bytes, Rejected = _rejected, Replaced = _replaced,
Completed = _completed, LastWaitMilliseconds = _lastWait, OldestWaitMilliseconds = _count == 0 ? 0 : Elapsed(oldest) };
}
}
private static double Elapsed(long start) => (Stopwatch.GetTimestamp() - start) * 1000d / Stopwatch.Frequency;
public void Dispose()
{
List<Work> canceled = new();
lock (_gate)
{
if (_disposed) return;
_disposed = true;
foreach (var partition in _partitions.Values) canceled.AddRange(partition.Items);
_partitions.Clear(); _ready.Clear(); _count = 0; _bytes = 0;
}
foreach (var item in canceled) item.Completion.TrySetResult(ShrinkNetworkQueueResult.Canceled);
}
}
}
@@ -1,5 +1,5 @@
fileFormatVersion: 2
guid: 88c6239e9fc049f43b2999d7311eec1d
guid: 3b35a44b897396549b5a1dec58ae05c0
MonoImporter:
externalObjects: {}
serializedVersion: 2
+2 -2
View File
@@ -6,7 +6,7 @@ namespace ShrinkNetwork
{
public static class ShrinkNetworkProtocol
{
public const int CurrentProtocolVersion = 1;
public const int CurrentProtocolVersion = 2;
public const int CurrentSchemaVersion = 1;
}
@@ -28,6 +28,6 @@ namespace ShrinkNetwork
public long SessionTokenExpiresAtUnixTimeSeconds;
public string? Route;
public ShrinkNetworkPacketKind Kind;
public byte[] Payload = Array.Empty<byte>();
public ReadOnlyMemory<byte> Payload = ReadOnlyMemory<byte>.Empty;
}
}
@@ -9,6 +9,8 @@ namespace ShrinkNetwork
{
private readonly Dictionary<int, ShrinkNetworkMessageMeta> _opcodeToMeta = new();
private readonly Dictionary<Type, ShrinkNetworkMessageMeta> _typeToMeta = new();
private readonly HashSet<string> _routes = new(StringComparer.Ordinal);
public IReadOnlyCollection<ShrinkNetworkMessageMeta> Registrations => _typeToMeta.Values;
public void Register<TMessage>(int opcode, string? route = null) where TMessage : IShrinkNetworkMessage
=> Register(typeof(TMessage), opcode, route);
@@ -24,6 +26,8 @@ namespace ShrinkNetwork
if (_typeToMeta.ContainsKey(messageType))
throw new InvalidOperationException($"Message type {messageType.FullName} is already registered.");
route = string.IsNullOrWhiteSpace(route) ? null : route.Trim();
if (route != null && !_routes.Add(route)) throw new InvalidOperationException($"SHRINK002 Route '{route}' is already registered.");
var meta = new ShrinkNetworkMessageMeta(opcode, messageType, route);
_opcodeToMeta.Add(opcode, meta);
_typeToMeta.Add(messageType, meta);
+9
View File
@@ -2,6 +2,7 @@
using System;
using System.Collections.Generic;
using System.Reflection;
using Cysharp.Threading.Tasks;
namespace ShrinkNetwork
@@ -21,6 +22,8 @@ namespace ShrinkNetwork
public Func<ShrinkNetworkContext, object, UniTask<object?>> Handler = null!;
}
private readonly Dictionary<Type, MethodInfo> _bindingMethods = new();
public IReadOnlyDictionary<Type, MethodInfo> CaptureBindings() => new Dictionary<Type, MethodInfo>(_bindingMethods);
private readonly Dictionary<Type, MessageHandlerRegistration> _messageHandlers = new();
private readonly Dictionary<Type, RequestHandlerRegistration> _requestHandlers = new();
@@ -28,7 +31,9 @@ namespace ShrinkNetwork
ShrinkNetworkPermissionRequirement requirement = default)
where TMessage : IShrinkNetworkMessage
{
if (handler == null) throw new ArgumentNullException(nameof(handler));
RegisterHandler(typeof(TMessage), (context, message) => handler(context, (TMessage)message), requirement);
_bindingMethods[typeof(TMessage)] = handler.Method;
}
public void RegisterHandler(Type messageType, Func<ShrinkNetworkContext, object, UniTask> handler,
@@ -41,6 +46,7 @@ namespace ShrinkNetwork
if (_messageHandlers.ContainsKey(messageType) || _requestHandlers.ContainsKey(messageType))
throw new InvalidOperationException($"Handler already exists for {messageType.FullName}.");
_bindingMethods[messageType] = handler.Method;
_messageHandlers.Add(messageType, new MessageHandlerRegistration
{
Requirement = requirement,
@@ -53,8 +59,10 @@ namespace ShrinkNetwork
where TRequest : IShrinkNetworkRequest
where TResponse : class, IShrinkNetworkResponse
{
if (handler == null) throw new ArgumentNullException(nameof(handler));
RegisterRequestHandler(typeof(TRequest), typeof(TResponse),
async (context, message) => await handler(context, (TRequest)message), requirement);
_bindingMethods[typeof(TRequest)] = handler.Method;
}
public void RegisterRequestHandler(Type requestType, Type responseType,
@@ -70,6 +78,7 @@ namespace ShrinkNetwork
if (_messageHandlers.ContainsKey(requestType) || _requestHandlers.ContainsKey(requestType))
throw new InvalidOperationException($"Handler already exists for {requestType.FullName}.");
_bindingMethods[requestType] = handler.Method;
_requestHandlers.Add(requestType, new RequestHandlerRegistration
{
ResponseType = responseType,
@@ -1,4 +1,5 @@
using System;
using System.Buffers;
namespace ShrinkNetwork
{
@@ -8,4 +9,10 @@ namespace ShrinkNetwork
object Deserialize(byte[] payload, Type type);
T Deserialize<T>(byte[] payload);
}
public interface IShrinkNetworkBufferSerializer : IShrinkNetworkSerializer
{
void Serialize(IBufferWriter<byte> writer, object value);
object Deserialize(ReadOnlyMemory<byte> payload, Type type);
}
}
@@ -0,0 +1,40 @@
#nullable enable
using System;
using System.Buffers;
namespace ShrinkNetwork
{
/// <summary>Owns rented memory until Dispose. A sender must await completion before disposing.</summary>
public sealed class ShrinkBufferWriter : IBufferWriter<byte>, IDisposable
{
private byte[]? _buffer;
public ShrinkBufferWriter(int initialCapacity = 256) => _buffer = ArrayPool<byte>.Shared.Rent(Math.Max(1, initialCapacity));
public int WrittenCount { get; private set; }
public ReadOnlyMemory<byte> WrittenMemory => Buffer.AsMemory(0, WrittenCount);
internal Span<byte> WrittenSpan => Buffer.AsSpan(0, WrittenCount);
private byte[] Buffer => _buffer ?? throw new ObjectDisposedException(nameof(ShrinkBufferWriter));
public void Advance(int count)
{
if (count < 0 || count > Buffer.Length - WrittenCount) throw new ArgumentOutOfRangeException(nameof(count));
WrittenCount += count;
}
public Memory<byte> GetMemory(int sizeHint = 0) { Ensure(sizeHint); return Buffer.AsMemory(WrittenCount); }
public Span<byte> GetSpan(int sizeHint = 0) { Ensure(sizeHint); return Buffer.AsSpan(WrittenCount); }
private void Ensure(int sizeHint)
{
if (sizeHint < 0) throw new ArgumentOutOfRangeException(nameof(sizeHint));
sizeHint = Math.Max(1, sizeHint);
if (sizeHint <= Buffer.Length - WrittenCount) return;
var next = ArrayPool<byte>.Shared.Rent(checked(Math.Max(Buffer.Length * 2, WrittenCount + sizeHint)));
Buffer.AsSpan(0, WrittenCount).CopyTo(next);
ArrayPool<byte>.Shared.Return(Buffer, clearArray: true);
_buffer = next;
}
public void Dispose()
{
var buffer = _buffer;
_buffer = null;
if (buffer != null) ArrayPool<byte>.Shared.Return(buffer, clearArray: true);
}
}
}
@@ -0,0 +1,11 @@
fileFormatVersion: 2
guid: 9ae3dee10fcf29f4da1b7ace3063bf92
MonoImporter:
externalObjects: {}
serializedVersion: 2
defaultReferences: []
executionOrder: 0
icon: {instanceID: 0}
userData:
assetBundleName:
assetBundleVariant:
@@ -2,11 +2,13 @@
using System;
using System.Text;
using System.Buffers;
using System.IO;
using Newtonsoft.Json;
namespace ShrinkNetwork
{
public sealed class ShrinkJsonNetworkSerializer : IShrinkNetworkSerializer
public sealed class ShrinkJsonNetworkSerializer : IShrinkNetworkBufferSerializer
{
private static readonly JsonSerializerSettings Settings = new()
{
@@ -25,5 +27,33 @@ namespace ShrinkNetwork
public T Deserialize<T>(byte[] payload)
=> JsonConvert.DeserializeObject<T>(Encoding.UTF8.GetString(payload), Settings)
?? throw new JsonSerializationException($"Failed to deserialize payload into {typeof(T).FullName}.");
public void Serialize(IBufferWriter<byte> writer, object value)
{
using var stream = new BufferStream(writer);
using var text = new StreamWriter(stream, new UTF8Encoding(false), 1024, leaveOpen: true);
using var json = new JsonTextWriter(text);
JsonSerializer.Create(Settings).Serialize(json, value);
}
public object Deserialize(ReadOnlyMemory<byte> payload, Type type) =>
JsonConvert.DeserializeObject(Encoding.UTF8.GetString(payload.Span), type, Settings)
?? throw new JsonSerializationException($"Failed to deserialize payload into {type.FullName}.");
private sealed class BufferStream : Stream
{
private readonly IBufferWriter<byte> _writer;
public BufferStream(IBufferWriter<byte> writer) => _writer = writer;
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 void Write(byte[] buffer, int offset, int count) { buffer.AsSpan(offset, count).CopyTo(_writer.GetSpan(count)); _writer.Advance(count); }
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();
}
}
}
@@ -1,150 +0,0 @@
#nullable enable
using System;
using System.Linq;
using System.Reflection;
namespace ShrinkNetwork
{
public sealed class ShrinkMessagePackNetworkSerializer : IShrinkNetworkSerializer
{
private readonly MethodInfo _serializeMethod;
private readonly MethodInfo _deserializeMethod;
private readonly object? _serializerOptions;
public ShrinkMessagePackNetworkSerializer()
{
var serializerType = Type.GetType("MessagePack.MessagePackSerializer, MessagePack");
if (serializerType == null)
{
throw new InvalidOperationException(
"MessagePack assembly was not found. Please install MessagePack-CSharp before using ShrinkMessagePackNetworkSerializer.");
}
_serializerOptions = ResolveSerializerOptions(serializerType.Assembly);
var serializeMethod = serializerType
.GetMethods(BindingFlags.Public | BindingFlags.Static)
.FirstOrDefault(m =>
{
if (m.Name != "Serialize")
return false;
var parameters = m.GetParameters();
return parameters.Length >= 2 &&
parameters[0].ParameterType == typeof(Type) &&
parameters[1].ParameterType == typeof(object);
});
var deserializeMethod = serializerType
.GetMethods(BindingFlags.Public | BindingFlags.Static)
.FirstOrDefault(m =>
{
if (m.Name != "Deserialize")
return false;
var parameters = m.GetParameters();
return parameters.Length >= 2 &&
parameters[0].ParameterType == typeof(Type) &&
(parameters[1].ParameterType == typeof(byte[]) ||
parameters[1].ParameterType == typeof(ReadOnlyMemory<byte>));
});
if (serializeMethod == null || deserializeMethod == null)
throw new MissingMethodException("MessagePack serialize/deserialize API not found.");
_serializeMethod = serializeMethod;
_deserializeMethod = deserializeMethod;
}
public byte[] Serialize(object value)
{
if (value == null)
return Array.Empty<byte>();
var parameters = BuildParameters(_serializeMethod, value.GetType(), value, _serializerOptions);
return (byte[])_serializeMethod.Invoke(null, parameters)!;
}
public object Deserialize(byte[] payload, Type type)
{
payload ??= Array.Empty<byte>();
var parameters = BuildParameters(_deserializeMethod, type, payload, _serializerOptions);
return _deserializeMethod.Invoke(null, parameters)
?? throw new InvalidOperationException($"MessagePack returned null for type {type.FullName}.");
}
public T Deserialize<T>(byte[] payload)
{
return (T)Deserialize(payload, typeof(T));
}
private static object?[] BuildParameters(MethodInfo method, Type type, object valueOrBytes, object? serializerOptions)
{
var parameters = method.GetParameters();
var args = new object?[parameters.Length];
if (parameters.Length > 0)
args[0] = type;
if (parameters.Length > 1)
args[1] = ConvertPrimaryArgument(parameters[1].ParameterType, valueOrBytes);
for (var i = 2; i < parameters.Length; i++)
{
args[i] = ResolveAdditionalArgument(parameters[i], serializerOptions);
}
return args;
}
private static object? ResolveAdditionalArgument(ParameterInfo parameter, object? serializerOptions)
{
if (serializerOptions != null && parameter.ParameterType.IsInstanceOfType(serializerOptions))
return serializerOptions;
return parameter.HasDefaultValue
? parameter.DefaultValue
: GetDefault(parameter.ParameterType);
}
private static object ConvertPrimaryArgument(Type parameterType, object value)
{
if (parameterType == typeof(ReadOnlyMemory<byte>) && value is byte[] bytes)
return new ReadOnlyMemory<byte>(bytes);
return value;
}
private static object? ResolveSerializerOptions(Assembly serializerAssembly)
{
var contractlessResolverType = serializerAssembly.GetType("MessagePack.Resolvers.ContractlessStandardResolver");
if (contractlessResolverType != null)
{
var optionsField = contractlessResolverType.GetField("Options",
BindingFlags.Public | BindingFlags.NonPublic | BindingFlags.Static);
var options = optionsField?.GetValue(null);
if (options != null)
return options;
var instanceField = contractlessResolverType.GetField("Instance",
BindingFlags.Public | BindingFlags.NonPublic | BindingFlags.Static);
var instance = instanceField?.GetValue(null);
if (instance != null)
{
var optionsType = serializerAssembly.GetType("MessagePack.MessagePackSerializerOptions");
var standardProperty = optionsType?.GetProperty("Standard", BindingFlags.Public | BindingFlags.Static);
var standardOptions = standardProperty?.GetValue(null);
var withResolverMethod = optionsType?.GetMethod("WithResolver", BindingFlags.Public | BindingFlags.Instance);
var resolvedOptions = withResolverMethod?.Invoke(standardOptions, new[] { instance });
if (resolvedOptions != null)
return resolvedOptions;
}
}
return null;
}
private static object? GetDefault(Type type)
{
return type.IsValueType ? Activator.CreateInstance(type) : null;
}
}
}
+105
View File
@@ -0,0 +1,105 @@
#nullable enable
using System;
using System.Buffers.Binary;
using System.IO;
using System.Text;
namespace ShrinkNetwork
{
public sealed class ShrinkProtocolException : IOException
{
public ShrinkProtocolException(string message) : base(message) { }
public ShrinkProtocolException(string message, Exception inner) : base(message, inner) { }
}
/// <summary>V2 little-endian framing, independent of the payload serializer. Decode borrows input memory.</summary>
public static class ShrinkPacketCodec
{
private const uint Magic = 0x324B4853; // SHK2
public const int HeaderSize = 33;
public const int MaximumPacketBytes = 16 * 1024 * 1024;
private static readonly UTF8Encoding Utf8 = new(false, true);
public static ShrinkBufferWriter Encode(ShrinkNetworkPacket packet, IShrinkNetworkSerializer? serializer = null, object? message = null)
{
var route = packet.Route ?? string.Empty;
var token = packet.SessionToken ?? string.Empty;
var routeBytes = Utf8.GetByteCount(route);
var tokenBytes = Utf8.GetByteCount(token);
if (routeBytes > ushort.MaxValue || tokenBytes > ushort.MaxValue) throw new ShrinkProtocolException("Route or session token exceeds framing limit.");
if (packet.ProtocolVersion != ShrinkNetworkProtocol.CurrentProtocolVersion || packet.SchemaVersion is < 0 or > ushort.MaxValue || packet.Kind < 0 || packet.Kind > ShrinkNetworkPacketKind.Response)
throw new ShrinkProtocolException("Invalid packet header.");
var writer = new ShrinkBufferWriter(HeaderSize + routeBytes + tokenBytes);
try
{
var header = writer.GetSpan(HeaderSize + routeBytes + tokenBytes);
BinaryPrimitives.WriteUInt32LittleEndian(header, Magic);
BinaryPrimitives.WriteUInt16LittleEndian(header.Slice(4), (ushort)packet.ProtocolVersion);
BinaryPrimitives.WriteUInt16LittleEndian(header.Slice(6), (ushort)packet.SchemaVersion);
BinaryPrimitives.WriteInt32LittleEndian(header.Slice(8), packet.Opcode);
BinaryPrimitives.WriteInt32LittleEndian(header.Slice(12), packet.RequestToken.Value);
header[16] = (byte)packet.Kind;
BinaryPrimitives.WriteInt64LittleEndian(header.Slice(17), packet.SessionTokenExpiresAtUnixTimeSeconds);
BinaryPrimitives.WriteUInt16LittleEndian(header.Slice(25), (ushort)routeBytes);
BinaryPrimitives.WriteUInt16LittleEndian(header.Slice(27), (ushort)tokenBytes);
Utf8.GetBytes(route.AsSpan(), header.Slice(HeaderSize, routeBytes));
Utf8.GetBytes(token.AsSpan(), header.Slice(HeaderSize + routeBytes, tokenBytes));
writer.Advance(HeaderSize + routeBytes + tokenBytes);
var payloadStart = writer.WrittenCount;
if (message != null)
{
if (serializer is IShrinkNetworkBufferSerializer buffered) buffered.Serialize(writer, message);
else
{
var bytes = (serializer ?? throw new ArgumentNullException(nameof(serializer))).Serialize(message);
bytes.CopyTo(writer.GetSpan(bytes.Length)); writer.Advance(bytes.Length);
}
}
else
{
packet.Payload.Span.CopyTo(writer.GetSpan(packet.Payload.Length));
writer.Advance(packet.Payload.Length);
}
if (writer.WrittenCount > MaximumPacketBytes) throw new ShrinkProtocolException("Packet exceeds framing limit.");
BinaryPrimitives.WriteInt32LittleEndian(writer.WrittenSpan.Slice(29), writer.WrittenCount - payloadStart);
return writer;
}
catch { writer.Dispose(); throw; }
}
internal static bool IsResponse(ReadOnlyMemory<byte> memory) => memory.Length >= HeaderSize &&
BinaryPrimitives.ReadUInt32LittleEndian(memory.Span) == Magic && memory.Span[16] == (byte)ShrinkNetworkPacketKind.Response;
public static ShrinkNetworkPacket Decode(ReadOnlyMemory<byte> memory)
{
var span = memory.Span;
if (span.Length < HeaderSize || span.Length > MaximumPacketBytes || BinaryPrimitives.ReadUInt32LittleEndian(span) != Magic)
throw new ShrinkProtocolException("Expected ShrinkNetwork protocol v2 binary envelope. Legacy JSON/MessagePack envelopes are unsupported.");
var protocol = BinaryPrimitives.ReadUInt16LittleEndian(span.Slice(4));
if (protocol != ShrinkNetworkProtocol.CurrentProtocolVersion) throw new ShrinkProtocolException("Unsupported protocol version " + protocol);
var routeLength = BinaryPrimitives.ReadUInt16LittleEndian(span.Slice(25));
var tokenLength = BinaryPrimitives.ReadUInt16LittleEndian(span.Slice(27));
var payloadLength = BinaryPrimitives.ReadInt32LittleEndian(span.Slice(29));
var start = HeaderSize + routeLength + tokenLength;
if (payloadLength < 0 || start > span.Length || payloadLength != span.Length - start || span[16] > (byte)ShrinkNetworkPacketKind.Response)
throw new ShrinkProtocolException("Invalid packet lengths or kind.");
string route, token;
try
{
route = Utf8.GetString(span.Slice(HeaderSize, routeLength));
token = Utf8.GetString(span.Slice(HeaderSize + routeLength, tokenLength));
}
catch (DecoderFallbackException exception)
{
throw new ShrinkProtocolException("Invalid UTF-8 in packet header.", exception);
}
return new ShrinkNetworkPacket {
ProtocolVersion = protocol, SchemaVersion = BinaryPrimitives.ReadUInt16LittleEndian(span.Slice(6)),
Opcode = BinaryPrimitives.ReadInt32LittleEndian(span.Slice(8)), RequestToken = new ShrinkRequestToken(BinaryPrimitives.ReadInt32LittleEndian(span.Slice(12))),
Kind = (ShrinkNetworkPacketKind)span[16], SessionTokenExpiresAtUnixTimeSeconds = BinaryPrimitives.ReadInt64LittleEndian(span.Slice(17)),
Route = route, SessionToken = token,
Payload = memory.Slice(start, payloadLength)
};
}
}
}
@@ -0,0 +1,11 @@
fileFormatVersion: 2
guid: 8a01225f12864f142b0d0c6395ef381b
MonoImporter:
externalObjects: {}
serializedVersion: 2
defaultReferences: []
executionOrder: 0
icon: {instanceID: 0}
userData:
assetBundleName:
assetBundleVariant:
@@ -0,0 +1,49 @@
#nullable enable
using System;
using System.Buffers;
using System.Collections.Generic;
namespace ShrinkNetwork
{
public interface IShrinkMessageCodec<T>
{
void Write(IBufferWriter<byte> writer, T value);
T Read(ReadOnlyMemory<byte> payload);
}
/// <summary>Explicit codecs work with AOT and external modules without runtime generic reflection.</summary>
public class ShrinkRegisteredNetworkSerializer : IShrinkNetworkBufferSerializer
{
private interface ICodec
{
void Write(IBufferWriter<byte> writer, object value);
object Read(ReadOnlyMemory<byte> payload);
}
private sealed class Codec<T> : ICodec
{
private readonly IShrinkMessageCodec<T> _codec;
public Codec(IShrinkMessageCodec<T> codec) => _codec = codec;
public void Write(IBufferWriter<byte> writer, object value) => _codec.Write(writer, (T)value);
public object Read(ReadOnlyMemory<byte> payload) => _codec.Read(payload)!;
}
private readonly Dictionary<Type, ICodec> _codecs = new();
public void Register<T>(IShrinkMessageCodec<T> codec)
{
if (codec == null) throw new ArgumentNullException(nameof(codec));
_codecs.Add(typeof(T), new Codec<T>(codec));
}
public bool Unregister<T>() => _codecs.Remove(typeof(T));
private ICodec Resolve(Type type) => _codecs.TryGetValue(type, out var codec) ? codec :
throw new InvalidOperationException($"SHRINK-NET-CODEC: No codec registered for {type.FullName}. Register a generated formatter before binding the transport.");
public void Serialize(IBufferWriter<byte> writer, object value) => Resolve(value.GetType()).Write(writer, value);
public object Deserialize(ReadOnlyMemory<byte> payload, Type type) => Resolve(type).Read(payload);
public byte[] Serialize(object value)
{
using var writer = new ShrinkBufferWriter();
Serialize(writer, value);
return writer.WrittenMemory.ToArray();
}
public object Deserialize(byte[] payload, Type type) => Deserialize((ReadOnlyMemory<byte>)payload, type);
public T Deserialize<T>(byte[] payload) => (T)Deserialize(payload, typeof(T));
}
}
@@ -0,0 +1,11 @@
fileFormatVersion: 2
guid: b0225bad560778f43827ddefb85fd61c
MonoImporter:
externalObjects: {}
serializedVersion: 2
defaultReferences: []
executionOrder: 0
icon: {instanceID: 0}
userData:
assetBundleName:
assetBundleVariant:
@@ -1,4 +1,6 @@
using Cysharp.Threading.Tasks;
using System;
using System.Threading;
namespace ShrinkNetwork
{
@@ -6,4 +8,10 @@ namespace ShrinkNetwork
{
UniTask SendAsync(long sessionId, byte[] packetData);
}
/// <summary>Borrowed memory remains owned by caller until SendAsync completes, faults or cancels.</summary>
public interface IShrinkNetworkMemoryTransport : IShrinkNetworkAsyncTransport
{
UniTask SendAsync(long sessionId, ReadOnlyMemory<byte> packetData, CancellationToken cancellationToken = default);
}
}
@@ -163,9 +163,9 @@ namespace ShrinkNetwork
{
while (!cancellationToken.IsCancellationRequested)
{
var sessions = _sessions.Values.ToArray();
foreach (var session in sessions)
foreach (var pair in _sessions)
{
var session = pair.Value;
var shouldDisconnect = false;
lock (session.SyncRoot)
{
@@ -1,11 +1,12 @@
#nullable enable
using System;
using System.Collections.Generic;
using System.Threading;
using Cysharp.Threading.Tasks;
namespace ShrinkNetwork
{
public sealed class ShrinkLoopbackTransport : IShrinkNetworkAsyncTransport, IShrinkNetworkSessionControlTransport
public sealed class ShrinkLoopbackTransport : IShrinkNetworkMemoryTransport, IShrinkNetworkSessionControlTransport
{
private readonly HashSet<long> _openedSessions = new();
private ShrinkLoopbackTransport? _peer;
@@ -73,5 +74,12 @@ namespace ShrinkNetwork
_peer.OnEvent?.Invoke(ShrinkNetworkTransportEvent.Packet(sessionId, packetData));
return UniTask.CompletedTask;
}
public UniTask SendAsync(long sessionId, ReadOnlyMemory<byte> packetData, CancellationToken cancellationToken = default)
{
cancellationToken.ThrowIfCancellationRequested();
// Receivers may enqueue the event beyond this call; transfer a dedicated copy.
return SendAsync(sessionId, packetData.ToArray());
}
}
}
@@ -1,6 +1,7 @@
#nullable enable
using System;
using System.Buffers.Binary;
using System.Buffers;
using System.IO;
using System.Net.Sockets;
using System.Net.Security;
@@ -10,7 +11,7 @@ using Cysharp.Threading.Tasks;
namespace ShrinkNetwork
{
public sealed class ShrinkTcpClientTransport : IShrinkNetworkAsyncTransport
public sealed class ShrinkTcpClientTransport : IShrinkNetworkMemoryTransport
{
private readonly string _host;
private readonly int _port;
@@ -86,7 +87,9 @@ namespace ShrinkNetwork
SendAsync(sessionId, packetData).Forget();
}
public UniTask SendAsync(long sessionId, byte[] packetData)
public UniTask SendAsync(long sessionId, byte[] packetData) => SendAsync(sessionId, (ReadOnlyMemory<byte>)(packetData ?? Array.Empty<byte>()));
public UniTask SendAsync(long sessionId, ReadOnlyMemory<byte> packetData, CancellationToken cancellationToken = default)
{
if (!IsStarted)
throw new InvalidOperationException("Transport is not started.");
@@ -94,11 +97,11 @@ namespace ShrinkNetwork
throw new InvalidOperationException($"Unsupported session id {sessionId}. This transport only supports {_sessionId}.");
if (_client == null || !_client.Connected || _stream == null)
throw new InvalidOperationException("TCP client is not connected.");
if ((packetData?.Length ?? 0) > _maxPacketSize)
throw new InvalidOperationException($"TCP packet is too large. Size={(packetData?.Length ?? 0)}, Limit={_maxPacketSize}.");
if (packetData.Length > _maxPacketSize)
throw new InvalidOperationException($"TCP packet is too large. Size={packetData.Length}, Limit={_maxPacketSize}.");
return SendInternalAsync(_stream, _sendLock, packetData ?? Array.Empty<byte>(),
_cts?.Token ?? CancellationToken.None);
return SendInternalAsync(_stream, _sendLock, packetData,
cancellationToken.CanBeCanceled ? cancellationToken : _cts?.Token ?? CancellationToken.None);
}
private async UniTaskVoid ConnectAsync(CancellationToken cancellationToken)
@@ -209,18 +212,20 @@ namespace ShrinkNetwork
ex.SocketErrorCode == SocketError.Shutdown;
}
private static async UniTask SendInternalAsync(Stream stream, SemaphoreSlim sendLock, byte[] packetData,
private static async UniTask SendInternalAsync(Stream stream, SemaphoreSlim sendLock, ReadOnlyMemory<byte> packetData,
CancellationToken cancellationToken)
{
await sendLock.WaitAsync(cancellationToken);
try
{
var header = new byte[4];
BinaryPrimitives.WriteInt32LittleEndian(header, packetData.Length);
await stream.WriteAsync(header, cancellationToken);
if (packetData.Length > 0)
await stream.WriteAsync(packetData, cancellationToken);
await stream.FlushAsync(cancellationToken);
var frame = ArrayPool<byte>.Shared.Rent(checked(4 + packetData.Length));
try
{
BinaryPrimitives.WriteInt32LittleEndian(frame, packetData.Length);
packetData.CopyTo(frame.AsMemory(4));
await stream.WriteAsync(frame.AsMemory(0, 4 + packetData.Length), cancellationToken);
}
finally { ArrayPool<byte>.Shared.Return(frame, clearArray: true); }
}
finally
{
@@ -1,6 +1,7 @@
#nullable enable
using System;
using System.Buffers.Binary;
using System.Buffers;
using System.Collections.Concurrent;
using System.IO;
using System.Net;
@@ -13,7 +14,7 @@ using Cysharp.Threading.Tasks;
namespace ShrinkNetwork
{
public sealed class ShrinkTcpServerTransport : IShrinkNetworkAsyncTransport, IShrinkNetworkSessionControlTransport
public sealed class ShrinkTcpServerTransport : IShrinkNetworkMemoryTransport, IShrinkNetworkSessionControlTransport
{
private readonly ConcurrentDictionary<long, TcpClient> _clients = new();
private readonly ConcurrentDictionary<long, SemaphoreSlim> _sendLocks = new();
@@ -92,11 +93,8 @@ namespace ShrinkNetwork
_streams.Clear();
foreach (var pair in _sendLocks)
{
pair.Value.Dispose();
}
// In-flight writers still release these managed semaphores after the stream closes.
// No WaitHandle is allocated; let their final users release and then collect them.
_sendLocks.Clear();
}
@@ -125,7 +123,9 @@ namespace ShrinkNetwork
return true;
}
public async UniTask SendAsync(long sessionId, byte[] packetData)
public UniTask SendAsync(long sessionId, byte[] packetData) => SendAsync(sessionId, (ReadOnlyMemory<byte>)(packetData ?? Array.Empty<byte>()));
public async UniTask SendAsync(long sessionId, ReadOnlyMemory<byte> packetData, CancellationToken cancellationToken = default)
{
if (!_clients.TryGetValue(sessionId, out var client))
throw new InvalidOperationException($"Session {sessionId} is not connected.");
@@ -134,10 +134,10 @@ namespace ShrinkNetwork
if (!_sendLocks.TryGetValue(sessionId, out var sendLock))
throw new InvalidOperationException($"Session {sessionId} send lock is not initialized.");
if ((packetData?.Length ?? 0) > _maxPacketSize)
throw new InvalidOperationException($"TCP packet is too large. Size={(packetData?.Length ?? 0)}, Limit={_maxPacketSize}.");
if (packetData.Length > _maxPacketSize)
throw new InvalidOperationException($"TCP packet is too large. Size={packetData.Length}, Limit={_maxPacketSize}.");
await SendInternalAsync(stream, sendLock, packetData ?? Array.Empty<byte>(), CancellationToken.None);
await SendInternalAsync(stream, sendLock, packetData, cancellationToken);
}
private async Task AcceptLoopAsync(CancellationToken cancellationToken)
@@ -244,8 +244,7 @@ namespace ShrinkNetwork
{
if (_streams.TryRemove(sessionId, out var ownedStream))
ownedStream.Dispose();
if (_sendLocks.TryRemove(sessionId, out var sendLock))
sendLock.Dispose();
_sendLocks.TryRemove(sessionId, out _);
var remoteAddress = removed.Client.RemoteEndPoint?.ToString() ?? "unknown";
try
@@ -261,18 +260,20 @@ namespace ShrinkNetwork
}
}
private static async Task SendInternalAsync(Stream stream, SemaphoreSlim sendLock, byte[] packetData,
private static async Task SendInternalAsync(Stream stream, SemaphoreSlim sendLock, ReadOnlyMemory<byte> packetData,
CancellationToken cancellationToken)
{
await sendLock.WaitAsync(cancellationToken);
try
{
var header = new byte[4];
BinaryPrimitives.WriteInt32LittleEndian(header, packetData.Length);
await stream.WriteAsync(header, cancellationToken);
if (packetData.Length > 0)
await stream.WriteAsync(packetData, cancellationToken);
await stream.FlushAsync(cancellationToken);
var frame = ArrayPool<byte>.Shared.Rent(checked(4 + packetData.Length));
try
{
BinaryPrimitives.WriteInt32LittleEndian(frame, packetData.Length);
packetData.CopyTo(frame.AsMemory(4));
await stream.WriteAsync(frame.AsMemory(0, 4 + packetData.Length), cancellationToken);
}
finally { ArrayPool<byte>.Shared.Return(frame, clearArray: true); }
}
finally
{