282 lines
9.0 KiB
C#
282 lines
9.0 KiB
C#
using System.Net;
|
|
using System.Net.Sockets;
|
|
using System.Text;
|
|
|
|
namespace FinalFactory.Rendezvous.Contracts;
|
|
|
|
public static class RendezvousUdpCodec
|
|
{
|
|
public const byte MagicFirst = 0x52;
|
|
public const byte MagicSecond = 0x56;
|
|
public const byte FlagsNone = 0;
|
|
|
|
private const int FixedPrefixLength = 23;
|
|
private const int FixedSuffixLength = 3;
|
|
|
|
public static byte[] Encode(PresenceDatagram datagram)
|
|
{
|
|
if (datagram is null)
|
|
{
|
|
throw new ArgumentNullException(nameof(datagram));
|
|
}
|
|
|
|
if (ContractValidation.ValidateContractVersion(datagram.ContractVersion)
|
|
!= RendezvousErrorCode.None)
|
|
{
|
|
throw new ArgumentException("The UDP contract version is unsupported.", nameof(datagram));
|
|
}
|
|
|
|
if (datagram.MessageType is not UdpPresenceMessageType.HostPresence
|
|
and not UdpPresenceMessageType.ClientPresence)
|
|
{
|
|
throw new ArgumentException("The UDP presence message type is unknown.", nameof(datagram));
|
|
}
|
|
|
|
if (datagram.MediationHandle.Value == Guid.Empty)
|
|
{
|
|
throw new ArgumentException("The mediation handle cannot be empty.", nameof(datagram));
|
|
}
|
|
|
|
if (!TryGetAddressBytes(datagram.LocalAddress, datagram.AddressFamily, out byte[] addressBytes))
|
|
{
|
|
throw new ArgumentException("The local address does not match its address family.", nameof(datagram));
|
|
}
|
|
|
|
if (datagram.LocalPort is < 1 or > ushort.MaxValue)
|
|
{
|
|
throw new ArgumentOutOfRangeException(nameof(datagram), "The local port must be between 1 and 65535.");
|
|
}
|
|
|
|
if (!ContractValidation.IsCapabilityValid(datagram.Capability))
|
|
{
|
|
throw new ArgumentException("The UDP capability is invalid.", nameof(datagram));
|
|
}
|
|
|
|
byte[] capabilityBytes = Encoding.ASCII.GetBytes(datagram.Capability);
|
|
int encodedLength = FixedPrefixLength + addressBytes.Length + FixedSuffixLength
|
|
+ capabilityBytes.Length;
|
|
if (encodedLength > ContractLimits.UdpDatagramMaxBytes)
|
|
{
|
|
throw new ArgumentException("The encoded UDP datagram exceeds its size limit.", nameof(datagram));
|
|
}
|
|
|
|
byte[] encoded = new byte[encodedLength];
|
|
int offset = 0;
|
|
encoded[offset++] = MagicFirst;
|
|
encoded[offset++] = MagicSecond;
|
|
encoded[offset++] = checked((byte)datagram.ContractVersion);
|
|
encoded[offset++] = (byte)datagram.MessageType;
|
|
encoded[offset++] = FlagsNone;
|
|
WriteGuid(datagram.MediationHandle.Value, encoded, offset);
|
|
offset += 16;
|
|
encoded[offset++] = (byte)datagram.AddressFamily;
|
|
encoded[offset++] = checked((byte)addressBytes.Length);
|
|
addressBytes.CopyTo(encoded, offset);
|
|
offset += addressBytes.Length;
|
|
encoded[offset++] = checked((byte)(datagram.LocalPort >> 8));
|
|
encoded[offset++] = checked((byte)(datagram.LocalPort & 0xff));
|
|
encoded[offset++] = checked((byte)capabilityBytes.Length);
|
|
capabilityBytes.CopyTo(encoded, offset);
|
|
return encoded;
|
|
}
|
|
|
|
public static bool TryDecode(
|
|
ReadOnlySpan<byte> encoded,
|
|
out PresenceDatagram? datagram,
|
|
out UdpDecodeError error)
|
|
{
|
|
datagram = null;
|
|
error = UdpDecodeError.None;
|
|
|
|
if (encoded.Length > ContractLimits.UdpDatagramMaxBytes)
|
|
{
|
|
error = UdpDecodeError.DatagramTooLarge;
|
|
return false;
|
|
}
|
|
|
|
if (encoded.Length < FixedPrefixLength)
|
|
{
|
|
error = UdpDecodeError.Truncated;
|
|
return false;
|
|
}
|
|
|
|
int offset = 0;
|
|
if (encoded[offset++] != MagicFirst || encoded[offset++] != MagicSecond)
|
|
{
|
|
error = UdpDecodeError.InvalidMagic;
|
|
return false;
|
|
}
|
|
|
|
int version = encoded[offset++];
|
|
if (ContractValidation.ValidateContractVersion(version) != RendezvousErrorCode.None)
|
|
{
|
|
error = UdpDecodeError.UnsupportedVersion;
|
|
return false;
|
|
}
|
|
|
|
UdpPresenceMessageType messageType = (UdpPresenceMessageType)encoded[offset++];
|
|
if (messageType is not UdpPresenceMessageType.HostPresence
|
|
and not UdpPresenceMessageType.ClientPresence)
|
|
{
|
|
error = UdpDecodeError.UnknownMessageType;
|
|
return false;
|
|
}
|
|
|
|
if (encoded[offset++] != FlagsNone)
|
|
{
|
|
error = UdpDecodeError.InvalidFlags;
|
|
return false;
|
|
}
|
|
|
|
if (!TryReadGuid(encoded.Slice(offset, 16), out Guid handle)
|
|
|| handle == Guid.Empty)
|
|
{
|
|
error = UdpDecodeError.InvalidHandle;
|
|
return false;
|
|
}
|
|
|
|
offset += 16;
|
|
AddressFamilyKind addressFamily = (AddressFamilyKind)encoded[offset++];
|
|
int expectedAddressLength = addressFamily switch
|
|
{
|
|
AddressFamilyKind.Ipv4 => 4,
|
|
AddressFamilyKind.Ipv6 => 16,
|
|
_ => 0,
|
|
};
|
|
if (expectedAddressLength == 0)
|
|
{
|
|
error = UdpDecodeError.InvalidAddressFamily;
|
|
return false;
|
|
}
|
|
|
|
int addressLength = encoded[offset++];
|
|
if (addressLength != expectedAddressLength)
|
|
{
|
|
error = UdpDecodeError.InvalidAddress;
|
|
return false;
|
|
}
|
|
|
|
if (encoded.Length < offset + addressLength + FixedSuffixLength)
|
|
{
|
|
error = UdpDecodeError.Truncated;
|
|
return false;
|
|
}
|
|
|
|
string address;
|
|
try
|
|
{
|
|
address = new IPAddress(encoded.Slice(offset, addressLength).ToArray()).ToString();
|
|
}
|
|
catch (ArgumentException)
|
|
{
|
|
error = UdpDecodeError.InvalidAddress;
|
|
return false;
|
|
}
|
|
|
|
offset += addressLength;
|
|
int port = (encoded[offset++] << 8) | encoded[offset++];
|
|
if (port == 0)
|
|
{
|
|
error = UdpDecodeError.InvalidPort;
|
|
return false;
|
|
}
|
|
|
|
int capabilityLength = encoded[offset++];
|
|
if (capabilityLength == 0 || capabilityLength > ContractLimits.UdpCapabilityMaxCharacters)
|
|
{
|
|
error = UdpDecodeError.InvalidCapability;
|
|
return false;
|
|
}
|
|
|
|
if (encoded.Length < offset + capabilityLength)
|
|
{
|
|
error = UdpDecodeError.Truncated;
|
|
return false;
|
|
}
|
|
|
|
if (encoded.Length > offset + capabilityLength)
|
|
{
|
|
error = UdpDecodeError.TrailingData;
|
|
return false;
|
|
}
|
|
|
|
string capability = Encoding.ASCII.GetString(encoded.Slice(offset, capabilityLength).ToArray());
|
|
if (!ContractValidation.IsCapabilityValid(capability))
|
|
{
|
|
error = UdpDecodeError.InvalidCapability;
|
|
return false;
|
|
}
|
|
|
|
datagram = new PresenceDatagram
|
|
{
|
|
ContractVersion = version,
|
|
MessageType = messageType,
|
|
MediationHandle = new MediationHandle(handle),
|
|
AddressFamily = addressFamily,
|
|
LocalAddress = address,
|
|
LocalPort = port,
|
|
Capability = capability,
|
|
};
|
|
return true;
|
|
}
|
|
|
|
private static bool TryGetAddressBytes(
|
|
string value,
|
|
AddressFamilyKind addressFamily,
|
|
out byte[] addressBytes)
|
|
{
|
|
addressBytes = [];
|
|
if (!IPAddress.TryParse(value, out IPAddress? address))
|
|
{
|
|
return false;
|
|
}
|
|
|
|
bool familyMatches = addressFamily switch
|
|
{
|
|
AddressFamilyKind.Ipv4 => address.AddressFamily == AddressFamily.InterNetwork,
|
|
AddressFamilyKind.Ipv6 => address.AddressFamily == AddressFamily.InterNetworkV6,
|
|
_ => false,
|
|
};
|
|
if (!familyMatches)
|
|
{
|
|
return false;
|
|
}
|
|
|
|
addressBytes = address.GetAddressBytes();
|
|
return true;
|
|
}
|
|
|
|
private static void WriteGuid(Guid value, byte[] destination, int offset)
|
|
{
|
|
string hexadecimal = value.ToString("N");
|
|
for (int index = 0; index < 16; index++)
|
|
{
|
|
int high = ParseHexadecimal(hexadecimal[index * 2]);
|
|
int low = ParseHexadecimal(hexadecimal[(index * 2) + 1]);
|
|
destination[offset + index] = checked((byte)((high << 4) | low));
|
|
}
|
|
}
|
|
|
|
private static bool TryReadGuid(ReadOnlySpan<byte> encoded, out Guid value)
|
|
{
|
|
char[] hexadecimal = new char[32];
|
|
for (int index = 0; index < encoded.Length; index++)
|
|
{
|
|
hexadecimal[index * 2] = FormatHexadecimal(encoded[index] >> 4);
|
|
hexadecimal[(index * 2) + 1] = FormatHexadecimal(encoded[index] & 0x0f);
|
|
}
|
|
|
|
return Guid.TryParseExact(new string(hexadecimal), "N", out value);
|
|
}
|
|
|
|
private static int ParseHexadecimal(char value) => value switch
|
|
{
|
|
>= '0' and <= '9' => value - '0',
|
|
>= 'a' and <= 'f' => value - 'a' + 10,
|
|
_ => throw new FormatException("A GUID contained a non-hexadecimal character."),
|
|
};
|
|
|
|
private static char FormatHexadecimal(int value) =>
|
|
(char)(value < 10 ? '0' + value : 'a' + value - 10);
|
|
}
|