267 lines
9.2 KiB
Plaintext
267 lines
9.2 KiB
Plaintext
using System.Buffers.Binary;
|
|
using System.Collections.Concurrent;
|
|
using System.IO;
|
|
using System.Net;
|
|
using System.Net.Security;
|
|
using System.Net.Sockets;
|
|
using System.Security.Cryptography.X509Certificates;
|
|
using System.Threading;
|
|
using Cysharp.Threading.Tasks;
|
|
|
|
namespace ShrinkNetwork.ServerHost;
|
|
|
|
using ShrinkNetwork;
|
|
|
|
public sealed class TcpServerTransport : IShrinkNetworkAsyncTransport, IShrinkNetworkSessionControlTransport
|
|
{
|
|
private readonly ConcurrentDictionary<long, TcpClient> _clients = new();
|
|
private readonly ConcurrentDictionary<long, SemaphoreSlim> _sendLocks = new();
|
|
private readonly ConcurrentDictionary<long, Stream> _streams = new();
|
|
private readonly TcpListener _listener;
|
|
private readonly int _maxPacketSize;
|
|
private readonly ShrinkTcpTlsOptions? _tlsOptions;
|
|
private long _sessionIdGenerator;
|
|
private CancellationTokenSource? _cts;
|
|
private X509Certificate2? _serverCertificate;
|
|
|
|
public TcpServerTransport(IPAddress ipAddress, int port, int maxPacketSize = 64 * 1024,
|
|
ShrinkTcpTlsOptions? tlsOptions = null)
|
|
{
|
|
if (maxPacketSize <= 0)
|
|
throw new ArgumentOutOfRangeException(nameof(maxPacketSize));
|
|
|
|
_listener = new TcpListener(ipAddress, port);
|
|
_maxPacketSize = maxPacketSize;
|
|
_tlsOptions = tlsOptions?.Clone();
|
|
}
|
|
|
|
public bool IsStarted { get; private set; }
|
|
|
|
public event Action<ShrinkNetworkTransportEvent>? OnEvent;
|
|
|
|
public void Start()
|
|
{
|
|
if (IsStarted)
|
|
return;
|
|
|
|
IsStarted = true;
|
|
_cts = new CancellationTokenSource();
|
|
EnsureTlsCertificateLoaded();
|
|
_listener.Start();
|
|
_ = AcceptLoopAsync(_cts.Token);
|
|
}
|
|
|
|
public void Stop()
|
|
{
|
|
if (!IsStarted)
|
|
return;
|
|
|
|
IsStarted = false;
|
|
_cts?.Cancel();
|
|
_listener.Stop();
|
|
|
|
foreach (var pair in _clients)
|
|
{
|
|
pair.Value.Close();
|
|
}
|
|
|
|
_clients.Clear();
|
|
|
|
foreach (var pair in _streams)
|
|
pair.Value.Dispose();
|
|
|
|
_streams.Clear();
|
|
|
|
foreach (var pair in _sendLocks)
|
|
pair.Value.Dispose();
|
|
|
|
_sendLocks.Clear();
|
|
}
|
|
|
|
public void Send(long sessionId, byte[] packetData)
|
|
{
|
|
SendAsync(sessionId, packetData).Forget();
|
|
}
|
|
|
|
public bool DisconnectSession(long sessionId, string? reason = null)
|
|
{
|
|
if (!_clients.TryGetValue(sessionId, out var client))
|
|
return false;
|
|
|
|
try
|
|
{
|
|
if (_streams.TryRemove(sessionId, out var stream))
|
|
stream.Dispose();
|
|
client.Close();
|
|
}
|
|
catch
|
|
{
|
|
}
|
|
|
|
if (!string.IsNullOrWhiteSpace(reason))
|
|
ShrinkNetworkLogger.Warn($"[ShrinkNetwork][TCP-Server] Disconnect session {sessionId}: {reason}");
|
|
return true;
|
|
}
|
|
|
|
public async UniTask SendAsync(long sessionId, byte[] packetData)
|
|
{
|
|
if (!_clients.TryGetValue(sessionId, out var client))
|
|
throw new InvalidOperationException($"Session {sessionId} is not connected.");
|
|
if (!_streams.TryGetValue(sessionId, out var stream))
|
|
throw new InvalidOperationException($"Session {sessionId} stream is not initialized.");
|
|
|
|
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}.");
|
|
|
|
await SendInternalAsync(stream, sendLock, packetData ?? Array.Empty<byte>(), CancellationToken.None);
|
|
}
|
|
|
|
private async Task AcceptLoopAsync(CancellationToken cancellationToken)
|
|
{
|
|
while (!cancellationToken.IsCancellationRequested)
|
|
{
|
|
TcpClient client;
|
|
try
|
|
{
|
|
client = await _listener.AcceptTcpClientAsync(cancellationToken);
|
|
}
|
|
catch (OperationCanceledException)
|
|
{
|
|
break;
|
|
}
|
|
catch
|
|
{
|
|
if (!IsStarted)
|
|
break;
|
|
|
|
ShrinkNetworkLogger.Warn("[ShrinkNetwork][TCP-Server] Accept failed, retrying.");
|
|
continue;
|
|
}
|
|
|
|
var sessionId = Interlocked.Increment(ref _sessionIdGenerator);
|
|
client.NoDelay = true;
|
|
var stream = await CreateServerStreamAsync(client, cancellationToken);
|
|
_clients[sessionId] = client;
|
|
_streams[sessionId] = stream;
|
|
_sendLocks[sessionId] = new SemaphoreSlim(1, 1);
|
|
var remoteAddress = client.Client.RemoteEndPoint?.ToString() ?? "unknown";
|
|
OnEvent?.Invoke(ShrinkNetworkTransportEvent.Connected(sessionId, remoteAddress));
|
|
_ = ReceiveLoopAsync(sessionId, client, stream, cancellationToken);
|
|
}
|
|
}
|
|
|
|
private async Task<Stream> CreateServerStreamAsync(TcpClient client, CancellationToken cancellationToken)
|
|
{
|
|
Stream stream = client.GetStream();
|
|
var tlsOptions = _tlsOptions;
|
|
if (tlsOptions == null || !tlsOptions.Enabled)
|
|
return stream;
|
|
|
|
if (_serverCertificate == null)
|
|
throw new InvalidOperationException("TCP TLS is enabled, but the server certificate is not loaded.");
|
|
|
|
var sslStream = new SslStream(stream, false);
|
|
await sslStream.AuthenticateAsServerAsync(_serverCertificate, false, tlsOptions.EnabledProtocols,
|
|
tlsOptions.CheckCertificateRevocation);
|
|
return sslStream;
|
|
}
|
|
|
|
private void EnsureTlsCertificateLoaded()
|
|
{
|
|
var tlsOptions = _tlsOptions;
|
|
if (tlsOptions == null || !tlsOptions.Enabled)
|
|
{
|
|
_serverCertificate = null;
|
|
return;
|
|
}
|
|
|
|
if (string.IsNullOrWhiteSpace(tlsOptions.ServerCertificatePath))
|
|
throw new InvalidOperationException("TCP TLS is enabled, but ServerCertificatePath is empty.");
|
|
|
|
_serverCertificate = new X509Certificate2(tlsOptions.ServerCertificatePath, tlsOptions.ServerCertificatePassword);
|
|
}
|
|
|
|
private async Task ReceiveLoopAsync(long sessionId, TcpClient client, Stream stream, CancellationToken cancellationToken)
|
|
{
|
|
Exception? disconnectException = null;
|
|
try
|
|
{
|
|
var lengthBuffer = new byte[4];
|
|
|
|
while (!cancellationToken.IsCancellationRequested)
|
|
{
|
|
await ReadExactlyAsync(stream, lengthBuffer, cancellationToken);
|
|
var length = BinaryPrimitives.ReadInt32LittleEndian(lengthBuffer);
|
|
if (length <= 0)
|
|
throw new InvalidDataException($"Invalid packet length: {length}");
|
|
if (length > _maxPacketSize)
|
|
throw new InvalidDataException($"TCP packet length exceeded limit. Length={length}, Limit={_maxPacketSize}");
|
|
|
|
var payload = new byte[length];
|
|
await ReadExactlyAsync(stream, payload, cancellationToken);
|
|
OnEvent?.Invoke(ShrinkNetworkTransportEvent.Packet(sessionId, payload));
|
|
}
|
|
}
|
|
catch (OperationCanceledException)
|
|
{
|
|
}
|
|
catch (Exception ex)
|
|
{
|
|
disconnectException = ex;
|
|
}
|
|
finally
|
|
{
|
|
if (_clients.TryRemove(sessionId, out var removed))
|
|
{
|
|
if (_streams.TryRemove(sessionId, out var ownedStream))
|
|
ownedStream.Dispose();
|
|
if (_sendLocks.TryRemove(sessionId, out var sendLock))
|
|
sendLock.Dispose();
|
|
|
|
var remoteAddress = removed.Client.RemoteEndPoint?.ToString() ?? "unknown";
|
|
if (disconnectException != null)
|
|
{
|
|
ShrinkNetworkLogger.Warn(
|
|
$"[ShrinkNetwork][TCP-Server] Session {sessionId} {remoteAddress} disconnected: {disconnectException.GetType().Name}: {disconnectException.Message}");
|
|
}
|
|
|
|
removed.Close();
|
|
OnEvent?.Invoke(ShrinkNetworkTransportEvent.Disconnected(sessionId, remoteAddress));
|
|
}
|
|
}
|
|
}
|
|
|
|
private static async Task SendInternalAsync(Stream stream, SemaphoreSlim sendLock, 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);
|
|
}
|
|
finally
|
|
{
|
|
sendLock.Release();
|
|
}
|
|
}
|
|
|
|
private static async Task ReadExactlyAsync(Stream stream, byte[] buffer, CancellationToken cancellationToken)
|
|
{
|
|
var offset = 0;
|
|
while (offset < buffer.Length)
|
|
{
|
|
var read = await stream.ReadAsync(buffer.AsMemory(offset, buffer.Length - offset), cancellationToken);
|
|
if (read <= 0)
|
|
throw new IOException("Remote closed.");
|
|
offset += read;
|
|
}
|
|
}
|
|
}
|