using System.Collections.Concurrent; using System.Net; using System.Net.Sockets; using System.Net.WebSockets; using System.Threading.Channels; using LumaTunnel.Shared.Protocol; namespace LumaTunnel.Server.Core.Runtime; public sealed class TunnelConnection : IAsyncDisposable { private readonly CancellationTokenSource _lifetime = new(); private readonly Channel _incoming; private readonly Channel _outgoing; private readonly ConcurrentDictionary _streams = new(); private readonly ServerRuntime _runtime; private readonly Action _terminated; private Task? _incomingTask; private Task? _outgoingTask; private int _closed; public TunnelConnection(WebSocket webSocket, string deviceId, ServerRuntime runtime, Action terminated) { WebSocket = webSocket; DeviceId = deviceId; _runtime = runtime; _terminated = terminated; _incoming = Channel.CreateBounded(new BoundedChannelOptions(runtime.Settings.MaxConnectionsPerDevice * 2) { FullMode = BoundedChannelFullMode.Wait, SingleReader = true, SingleWriter = false }); _outgoing = Channel.CreateBounded(new BoundedChannelOptions(runtime.Settings.MaxConnectionsPerDevice * 4) { FullMode = BoundedChannelFullMode.Wait, SingleReader = true, SingleWriter = false }); } public string DeviceId { get; } public WebSocket WebSocket { get; } public int ActiveStreams => _streams.Count; public void Start() { _incomingTask = ProcessIncomingAsync(_lifetime.Token); _outgoingTask = SendLoopAsync(_lifetime.Token); } public bool TryEnqueue(ReadOnlySpan bytes) { try { return _incoming.Writer.TryWrite(TunnelProtocol.Decode(bytes)); } catch (TunnelProtocolException) { return false; } } public ValueTask SendAsync(TunnelFrame frame, CancellationToken cancellationToken = default) => _outgoing.Writer.WriteAsync(frame, cancellationToken); public async Task CloseAsync(WebSocketCloseStatus status, string description) { if (Interlocked.Exchange(ref _closed, 1) != 0) return; _lifetime.Cancel(); _incoming.Writer.TryComplete(); _outgoing.Writer.TryComplete(); foreach (var pair in _streams.ToArray()) { if (_streams.TryRemove(pair.Key, out var stream)) { _runtime.Traffic.ConnectionClosed(DeviceId); await stream.DisposeAsync().ConfigureAwait(false); } } try { if (WebSocket.State is WebSocketState.Open or WebSocketState.CloseReceived) await WebSocket.CloseAsync(status, description, CancellationToken.None).ConfigureAwait(false); } catch (WebSocketException) { WebSocket.Abort(); } finally { _terminated(this); } } public async ValueTask DisposeAsync() { await CloseAsync(WebSocketCloseStatus.NormalClosure, "Tunnel connection closed.").ConfigureAwait(false); _lifetime.Dispose(); } internal async Task RemoveStreamAsync(uint streamId) { if (_streams.TryRemove(streamId, out var stream)) { _runtime.Traffic.ConnectionClosed(DeviceId); await stream.DisposeAsync().ConfigureAwait(false); } } private async Task ProcessIncomingAsync(CancellationToken cancellationToken) { try { await foreach (var frame in _incoming.Reader.ReadAllAsync(cancellationToken).ConfigureAwait(false)) { switch (frame.Type) { case TunnelFrameType.Open: await OpenStreamAsync(frame, cancellationToken).ConfigureAwait(false); break; case TunnelFrameType.Data: if (_streams.TryGetValue(frame.StreamId, out var dataStream)) await dataStream.WriteFromClientAsync(frame.Payload, cancellationToken).ConfigureAwait(false); break; case TunnelFrameType.HalfClose: if (_streams.TryGetValue(frame.StreamId, out var halfCloseStream)) await halfCloseStream.ClientFinishedWritingAsync().ConfigureAwait(false); break; case TunnelFrameType.Close: case TunnelFrameType.Reset: await RemoveStreamAsync(frame.StreamId).ConfigureAwait(false); break; case TunnelFrameType.Ping: await SendAsync(new TunnelFrame(TunnelFrameType.Pong, 0, frame.Payload), cancellationToken).ConfigureAwait(false); break; } } } catch (OperationCanceledException) when (cancellationToken.IsCancellationRequested) { } catch (Exception) { await CloseAsync(WebSocketCloseStatus.InternalServerError, "Tunnel receive loop failed.").ConfigureAwait(false); } } private async Task OpenStreamAsync(TunnelFrame frame, CancellationToken cancellationToken) { if (frame.StreamId == 0 || (frame.StreamId & 1) == 0 || _streams.ContainsKey(frame.StreamId)) { await SendOpenErrorAsync(frame.StreamId, TunnelErrorCode.MalformedFrame, "Stream id must be a new, non-zero odd number.", cancellationToken).ConfigureAwait(false); return; } if (_streams.Count >= _runtime.Settings.MaxConnectionsPerDevice || _runtime.Traffic.ActiveConnections >= _runtime.Settings.MaxConnections) { await SendOpenErrorAsync(frame.StreamId, TunnelErrorCode.TooManyConnections, "Connection limit reached.", cancellationToken).ConfigureAwait(false); return; } try { var request = TunnelMessageSerializer.Deserialize(frame.Payload); var address = await _runtime.TargetPolicy.ResolveAndValidateAsync(request.Host, request.Port, cancellationToken).ConfigureAwait(false); using var timeout = CancellationTokenSource.CreateLinkedTokenSource(cancellationToken); timeout.CancelAfter(TimeSpan.FromSeconds(_runtime.Settings.ConnectTimeoutSeconds)); var client = new TcpClient(address.AddressFamily) { NoDelay = true }; try { await client.ConnectAsync(address, request.Port, timeout.Token).ConfigureAwait(false); } catch { client.Dispose(); throw; } var stream = new ServerTunnelStream(frame.StreamId, client, this, _runtime); if (!_streams.TryAdd(frame.StreamId, stream)) { await stream.DisposeAsync().ConfigureAwait(false); await SendOpenErrorAsync(frame.StreamId, TunnelErrorCode.MalformedFrame, "Stream id is already open.", cancellationToken).ConfigureAwait(false); return; } _runtime.Traffic.ConnectionOpened(DeviceId); stream.Start(); await SendAsync(TunnelFrame.Empty(TunnelFrameType.OpenOk, frame.StreamId), cancellationToken).ConfigureAwait(false); } catch (TunnelProtocolException exception) { await SendOpenErrorAsync(frame.StreamId, exception.Code, exception.Message, cancellationToken).ConfigureAwait(false); } catch (OperationCanceledException) when (!cancellationToken.IsCancellationRequested) { await SendOpenErrorAsync(frame.StreamId, TunnelErrorCode.ConnectTimeout, "Target connection timed out.", cancellationToken).ConfigureAwait(false); } catch (SocketException exception) { await SendOpenErrorAsync(frame.StreamId, TunnelErrorCode.ConnectRefused, $"Target connection failed: {exception.SocketErrorCode}.", cancellationToken).ConfigureAwait(false); } catch (Exception) { await SendOpenErrorAsync(frame.StreamId, TunnelErrorCode.InternalError, "Target connection failed.", cancellationToken).ConfigureAwait(false); } } private ValueTask SendOpenErrorAsync(uint streamId, TunnelErrorCode code, string message, CancellationToken cancellationToken) => SendAsync(new TunnelFrame(TunnelFrameType.OpenError, streamId, TunnelMessageSerializer.Serialize(new TunnelOpenError(code, message))), cancellationToken); private async Task SendLoopAsync(CancellationToken cancellationToken) { try { await foreach (var frame in _outgoing.Reader.ReadAllAsync(cancellationToken).ConfigureAwait(false)) { var bytes = TunnelProtocol.Encode(frame); await WebSocket.SendAsync(bytes, WebSocketMessageType.Binary, true, cancellationToken).ConfigureAwait(false); } } catch (OperationCanceledException) when (cancellationToken.IsCancellationRequested) { } catch (Exception) { await CloseAsync(WebSocketCloseStatus.InternalServerError, "Tunnel send loop failed.").ConfigureAwait(false); } } }