using System; using System.Collections.Concurrent; using System.Globalization; using System.IO; using System.Net; using System.Net.Sockets; using System.Text; using System.Threading; using System.Threading.Tasks; using StackExchange.Redis; namespace Jellyfin.Server.Tests.HighAvailability; /// /// A loopback TCP proxy in front of a Redis server. Cutting it drops every connection through it and /// refuses new ones, so a test can take Redis away from one instance mid-flow - and give it back - the /// way a restarted valkey does, and watch what a real StackExchange.Redis client makes of it. /// public sealed class RedisFaultProxy : IAsyncDisposable { private readonly ConcurrentDictionary _live = new(); private readonly CancellationTokenSource _cts = new(); private readonly TcpListener _listener; private readonly string _targetHost; private readonly int _targetPort; private readonly int _port; private volatile bool _cut; private volatile byte[]? _cutAfterMarker; private RedisFaultProxy(TcpListener listener, int port, string targetHost, int targetPort) { _listener = listener; _port = port; _targetHost = targetHost; _targetPort = targetPort; } /// /// Gets a connection string pointing at the proxy. The timeouts are short so a cut surfaces as a /// failure in seconds rather than in the library's minute-scale defaults. /// public string ConnectionString => string.Create( CultureInfo.InvariantCulture, $"127.0.0.1:{_port},abortConnect=false,connectTimeout=500,syncTimeout=2000,connectRetry=1"); /// /// Starts a proxy in front of the server named by . /// /// The connection string of the server to forward to. /// The running proxy. public static RedisFaultProxy Start(string target) { var endpoint = ConfigurationOptions.Parse(target).EndPoints[0]; var (host, port) = endpoint switch { DnsEndPoint dns => (dns.Host, dns.Port), IPEndPoint ip => (ip.Address.ToString(), ip.Port), _ => throw new NotSupportedException("Unsupported endpoint " + endpoint) }; var listener = new TcpListener(IPAddress.Loopback, 0); listener.Start(); var proxy = new RedisFaultProxy(listener, ((IPEndPoint)listener.LocalEndpoint).Port, host, port); _ = Task.Run(proxy.AcceptAsync); return proxy; } /// /// Takes Redis away from everything connected through the proxy. /// public void Cut() { _cut = true; DropLiveConnections(); } /// /// Arms a cut for the moment after a command containing has been forwarded /// and answered, so a test can take Redis away between two round trips of one operation rather than /// only before or after all of them. /// /// Text that identifies the command to cut after. public void CutAfterForwarding(string marker) => _cutAfterMarker = Encoding.UTF8.GetBytes(marker); /// /// Lets connections through again. Clients reconnect on their own schedule, so callers have to wait /// for the connection to come back rather than assume it already has. /// public void Restore() { _cutAfterMarker = null; _cut = false; } /// public async ValueTask DisposeAsync() { _cut = true; await _cts.CancelAsync().ConfigureAwait(false); _listener.Stop(); DropLiveConnections(); _cts.Dispose(); } private void DropLiveConnections() { foreach (var client in _live.Keys) { if (_live.TryRemove(client, out _)) { client.Dispose(); } } } private async Task AcceptAsync() { while (!_cts.IsCancellationRequested) { TcpClient client; try { client = await _listener.AcceptTcpClientAsync(_cts.Token).ConfigureAwait(false); } catch (Exception exception) when (exception is OperationCanceledException or SocketException or ObjectDisposedException) { return; } if (_cut) { client.Dispose(); continue; } _ = Task.Run(() => ForwardAsync(client)); } } private async Task ForwardAsync(TcpClient client) { TcpClient? upstream = null; try { upstream = new TcpClient(); await upstream.ConnectAsync(_targetHost, _targetPort, _cts.Token).ConfigureAwait(false); _live[client] = 0; _live[upstream] = 0; // Registered first, then rechecked: a cut concurrent with this connect would otherwise drop // the live connections before this pair joined them and leave it running through the outage. if (_cut) { return; } var clientStream = client.GetStream(); var upstreamStream = upstream.GetStream(); await Task.WhenAny( CopyFromClientAsync(clientStream, upstreamStream), CopyAsync(upstreamStream, clientStream)).ConfigureAwait(false); } catch (Exception exception) when (exception is IOException or SocketException or OperationCanceledException or ObjectDisposedException) { } finally { _live.TryRemove(client, out _); client.Dispose(); if (upstream is not null) { _live.TryRemove(upstream, out _); upstream.Dispose(); } } } private async Task CopyFromClientAsync(NetworkStream from, NetworkStream to) { var buffer = new byte[16 * 1024]; try { while (true) { var read = await from.ReadAsync(buffer, _cts.Token).ConfigureAwait(false); if (read == 0) { return; } await to.WriteAsync(buffer.AsMemory(0, read), _cts.Token).ConfigureAwait(false); var marker = _cutAfterMarker; if (marker is not null && buffer.AsSpan(0, read).IndexOf(marker) >= 0) { _cutAfterMarker = null; // Long enough for the server to have applied the command that was just forwarded. await Task.Delay(TimeSpan.FromMilliseconds(250), _cts.Token).ConfigureAwait(false); Cut(); return; } } } catch (Exception exception) when (exception is IOException or SocketException or OperationCanceledException or ObjectDisposedException) { } } private async Task CopyAsync(NetworkStream from, NetworkStream to) { try { await from.CopyToAsync(to, _cts.Token).ConfigureAwait(false); } catch (Exception exception) when (exception is IOException or SocketException or OperationCanceledException or ObjectDisposedException) { } } }