106 lines
6.3 KiB
C#
106 lines
6.3 KiB
C#
#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)
|
|
};
|
|
}
|
|
}
|
|
}
|