fix: Fixes networking issues (#1958)

### Summary
- Puts back `EventSink.SocketConnect`.
- Reverts networking change to push the networking to a separate thread.
- Reverts changes to the firewall by removing the firewall queue.
- Fixes listeners not shutting down with the server.
- Fixes race condition causing connections to get stuck even after they are disposed.

> [!NOTE]
> **Developer Note**
> Networking has been reverted back to using the main thread instead of a background thread. This alleviated complexity and the requirement for concurrent queues all over the place.
This commit is contained in:
Kamron Batman 2024-09-19 16:54:49 -07:00 • committed by GitHub
parent 97c53e656e
commit e0fcde885c
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
10 changed files with 211 additions and 340 deletions

View file

@ -0,0 +1,41 @@
/*************************************************************************
* ModernUO *
* Copyright 2019-2023 - ModernUO Development Team *
* Email: hi@modernuo.com *
* File: SocketConnectionEvent.cs *
* *
* This program is free software: you can redistribute it and/or modify *
* it under the terms of the GNU General Public License as published by *
* the Free Software Foundation, either version 3 of the License, or *
* (at your option) any later version. *
* *
* You should have received a copy of the GNU General Public License *
* along with this program. If not, see <http://www.gnu.org/licenses/>. *
*************************************************************************/
using System;
using System.Net.Sockets;
using System.Runtime.CompilerServices;
namespace Server;
public class SocketConnectEventArgs
{
public SocketConnectEventArgs(Socket c)
{
Connection = c;
AllowConnection = true;
}
public Socket Connection { get; }
public bool AllowConnection { get; set; }
}
public static partial class EventSink
{
public static event Action<SocketConnectEventArgs> SocketConnect;
[MethodImpl(MethodImplOptions.AggressiveInlining)]
public static void InvokeSocketConnect(SocketConnectEventArgs e) => SocketConnect?.Invoke(e);
}

View file

@ -336,6 +336,8 @@ public static class Core
World.WaitForWriteCompletion();
World.ExitSerializationThreads();
PingServer.Shutdown();
TcpServer.Shutdown();
if (!_crashed)
{

View file

@ -14,83 +14,21 @@
*************************************************************************/
using System;
using System.Collections.Concurrent;
using System.Collections.Generic;
using System.Net;
using System.Runtime.CompilerServices;
using System.Runtime.InteropServices;
using Server.Logging;
namespace Server.Network;
public static class Firewall
{
private static readonly ILogger logger = LogFactory.GetLogger(typeof(Firewall));
private static InternalValidationEntry _validationEntry;
private static readonly Dictionary<IPAddress, bool> _isBlockedCache = new();
private static readonly ConcurrentQueue<(IFirewallEntry FirewallyEntry, bool Remove)> _firewallQueue = new();
private static readonly SortedSet<IFirewallEntry> _firewallSet = new();
[MethodImpl(MethodImplOptions.AggressiveInlining)]
public static IFirewallEntry RequestAddSingleIPEntry(string entry)
{
try
{
var firewallEntry = new SingleIpFirewallEntry(entry);
_firewallQueue.Enqueue((firewallEntry, false));
return firewallEntry;
}
catch (Exception e)
{
logger.Warning(e, "Failed to add firewall entry: {Pattern}", entry);
return null;
}
}
[MethodImpl(MethodImplOptions.AggressiveInlining)]
public static IFirewallEntry RequestAddCIDREntry(string entry)
{
try
{
var firewallEntry = new CidrFirewallEntry(entry);
_firewallQueue.Enqueue((firewallEntry, false));
return firewallEntry;
}
catch (Exception e)
{
logger.Warning(e, "Failed to add firewall entry: {Pattern}", entry);
return null;
}
}
[MethodImpl(MethodImplOptions.AggressiveInlining)]
public static void RequestAddEntry(IFirewallEntry entry)
{
_firewallQueue.Enqueue((entry, false));
}
[MethodImpl(MethodImplOptions.AggressiveInlining)]
public static void RequestRemoveEntry(IFirewallEntry entry)
{
_firewallQueue.Enqueue((entry, true));
}
internal static void ProcessQueue()
{
while (_firewallQueue.TryDequeue(out var entry))
{
if (entry.Remove)
{
RemoveEntry(entry.FirewallyEntry);
}
else
{
AddEntry(entry.FirewallyEntry);
}
}
}
public static SortedSet<IFirewallEntry> FirewallSet => _firewallSet;
internal static bool IsBlocked(IPAddress address)
{
@ -128,22 +66,32 @@ public static class Firewall
}
[MethodImpl(MethodImplOptions.AggressiveInlining)]
private static void AddEntry(IFirewallEntry firewallEntry)
public static bool Add(IFirewallEntry firewallEntry)
{
_firewallSet.Add(firewallEntry);
_isBlockedCache.Clear();
if (_firewallSet.Add(firewallEntry))
{
_isBlockedCache.Clear();
return true;
}
return false;
}
[MethodImpl(MethodImplOptions.AggressiveInlining)]
private static void RemoveEntry(IFirewallEntry entry)
public static bool Remove(IFirewallEntry entry)
{
if (entry == null)
{
return;
return false;
}
_firewallSet.Remove(entry);
_isBlockedCache.Clear();
if (_firewallSet.Remove(entry))
{
_isBlockedCache.Clear();
return true;
}
return false;
}
private class InternalValidationEntry : BaseFirewallEntry

View file

@ -51,22 +51,32 @@ public static class IPLimiter
var now = Core.Now;
CheckThrottledAddresses(now);
IPAccessLog accessLog;
while (_throttledAddresses.Count > 0)
{
accessLog = _throttledAddresses.Min;
if (now <= accessLog.Expiration)
{
break;
}
_throttledAddresses.Remove(accessLog);
}
_accessCheck.IPAddress = ourAddress;
if (_connectionAttempts.TryGetValue(_accessCheck, out var accessLog))
if (_connectionAttempts.TryGetValue(_accessCheck, out accessLog))
{
_connectionAttempts.Remove(accessLog);
accessLog.Count++;
accessLog.Expiration = now + ConnectionAttemptsDuration;
if (now <= accessLog.Expiration && accessLog.Count >= MaxConnections)
{
BlockConnection(now, accessLog);
_throttledAddresses.Add(accessLog);
return false;
}
accessLog.Expiration = now + ConnectionAttemptsDuration;
}
else
{
@ -79,26 +89,6 @@ public static class IPLimiter
return true;
}
private static void BlockConnection(DateTime now, IPAccessLog accessLog)
{
accessLog.Expiration = now + ConnectionAttemptsDuration;
_throttledAddresses.Add(accessLog);
}
private static void CheckThrottledAddresses(DateTime now)
{
while (_throttledAddresses.Count > 0)
{
var accessLog = _throttledAddresses.Min;
if (now <= accessLog.Expiration)
{
break;
}
_throttledAddresses.Remove(accessLog);
}
}
private class IPAccessLog : IComparable<IPAccessLog>
{
public IPAddress IPAddress;

View file

@ -32,8 +32,6 @@ using System.Runtime.InteropServices;
namespace Server.Network;
public delegate void NetStateCreatedCallback(NetState ns);
public delegate void DecodePacket(Span<byte> buffer, ref int length);
public delegate int EncodePacket(ReadOnlySpan<byte> inputBuffer, Span<byte> outputBuffer);
@ -56,8 +54,6 @@ public partial class NetState : IComparable<NetState>, IValueLinkListNode<NetSta
private static readonly Queue<NetState> _throttled = new(256);
private static readonly Queue<NetState> _throttledPending = new(256);
public static NetStateCreatedCallback CreatedCallback { get; set; }
private static readonly SortedSet<NetState> _connecting = new(NetStateConnectingComparer.Instance);
private static readonly HashSet<NetState> _instances = new(2048);
public static IReadOnlySet<NetState> Instances => _instances;
@ -68,8 +64,8 @@ public partial class NetState : IComparable<NetState>, IValueLinkListNode<NetSta
private volatile DecodePacket _packetDecoder;
private volatile EncodePacket _packetEncoder;
private bool _flushQueued;
private readonly long[] _packetThrottles = new long[0x100];
private readonly long[] _packetCounts = new long[0x100];
private long[] _packetThrottles;
private long[] _packetCounts;
private string _disconnectReason = string.Empty;
internal ParserState _parserState = ParserState.AwaitingNextPacket;
@ -143,8 +139,12 @@ public partial class NetState : IComparable<NetState>, IValueLinkListNode<NetSta
_toString = "(error)";
}
_instances.Add(this);
_connecting.Add(this);
_handle = GCHandle.Alloc(this);
LogInfo($"Connected. [{_instances.Count} Online]");
try
{
_pollGroup.Add(connection, _handle);
@ -254,22 +254,30 @@ public partial class NetState : IComparable<NetState>, IValueLinkListNode<NetSta
{
if (packetID is >= 0 and < 0x100)
{
_packetThrottles ??= new long[0x100];
_packetThrottles[packetID] = Core.TickCount;
}
}
public long GetPacketTime(int packetID) => packetID is >= 0 and < 0x100 ? _packetThrottles[packetID] : 0;
public long GetPacketTime(int packetID) =>
packetID is >= 0 and < 0x100 && _packetThrottles != null ? _packetThrottles[packetID] : 0;
private void UpdatePacketCount(int packetID)
{
if (packetID is >= 0 and < 0x100)
{
_packetCounts ??= new long[0x100];
_packetCounts[packetID]++;
}
}
public int CheckPacketCounts()
{
if (_packetCounts == null)
{
return 0;
}
for (int i = 0; i < _packetCounts.Length; i++)
{
long count = _packetCounts[i];
@ -917,7 +925,7 @@ public partial class NetState : IComparable<NetState>, IValueLinkListNode<NetSta
var ns = _connecting.Min;
var socketTime = ns.ConnectedOn;
// If the socket has been connected for less than 2 seconds, we can stop checking
// If the socket has been connected for less than the limit, we can stop checking
if (now - socketTime < ConnectingSocketIdleLimit)
{
break;
@ -946,17 +954,6 @@ public partial class NetState : IComparable<NetState>, IValueLinkListNode<NetSta
{
DisconnectUnattachedSockets();
const int maxEntriesPerLoop = 32;
var count = 0;
while (++count <= maxEntriesPerLoop && TcpServer.ConnectedQueue.TryDequeue(out var ns))
{
CreatedCallback?.Invoke(ns);
_instances.Add(ns);
_connecting.Add(ns); // Add to the connecting set, and remove them when they authenticated.
ns.LogInfo($"Connected. [{Instances.Count} Online]");
}
while (_throttled.Count > 0)
{
var ns = _throttled.Dequeue();
@ -972,7 +969,7 @@ public partial class NetState : IComparable<NetState>, IValueLinkListNode<NetSta
_throttled.Enqueue(_throttledPending.Dequeue());
}
count = _pollGroup.Poll(_polledStates);
var count = _pollGroup.Poll(_polledStates);
if (count > 0)
{

View file

@ -16,7 +16,6 @@
using System.Collections.Generic;
using System.Net;
using System.Net.Sockets;
using System.Threading;
using Server.Logging;
namespace Server.Network;
@ -47,8 +46,8 @@ public static class PingServer
return;
}
HashSet<IPEndPoint> listeningAddresses = new HashSet<IPEndPoint>();
List<UdpClient> listeners = new List<UdpClient>();
HashSet<IPEndPoint> listeningAddresses = [];
List<UdpClient> listeners = [];
foreach (var serverIpep in ServerConfiguration.Listeners)
{
@ -70,7 +69,7 @@ public static class PingServer
}
listeners.Add(listener);
new Thread(BeginAcceptingUdpRequest).Start(listener);
BeginAcceptingUdpRequest(listener);
}
foreach (var ipep in listeningAddresses)
@ -140,4 +139,12 @@ public static class PingServer
}
}
}
public static void Shutdown()
{
foreach (var listener in Listeners)
{
listener.Close();
}
}
}

View file

@ -13,8 +13,6 @@
* along with this program. If not, see <http://www.gnu.org/licenses/>. *
*************************************************************************/
using System;
using System.Collections.Concurrent;
using System.Collections.Generic;
using System.IO;
using System.Linq;
@ -22,7 +20,6 @@ using System.Net;
using System.Net.NetworkInformation;
using System.Net.Sockets;
using System.Runtime.CompilerServices;
using System.Threading;
using Server.Logging;
using Server.Misc;
@ -32,30 +29,52 @@ public static class TcpServer
{
private static readonly ILogger logger = LogFactory.GetLogger(typeof(TcpServer));
private const long ListenerErrorMessageDelay = 10000; // 10 seconds
private static long _nextMaximumSocketsReachedMessage;
private static readonly SemaphoreSlim _queueSemaphore = new(0);
private static readonly ConcurrentQueue<Socket> _connectingQueue = [];
private static Thread _processConnectionsThread;
// Sanity. 256 * 1024 * 4096 = ~1.3GB of ram
public static int MaxConnections { get; set; }
// AccountLoginReject BadComm
private static readonly byte[] _socketRejected = [0x82, 0xFF];
public static IPEndPoint[] ListeningAddresses { get; private set; }
public static Socket[] Listeners { get; private set; }
public static ConcurrentQueue<NetState> ConnectedQueue { get; } = [];
public static void Configure()
{
MaxConnections = ServerConfiguration.GetOrUpdateSetting("tcpServer.maxConnections", 4096);
}
public static void Start()
{
_processConnectionsThread = new Thread(ProcessConnections);
_processConnectionsThread.Start();
HashSet<IPEndPoint> listeningAddresses = new HashSet<IPEndPoint>();
List<Socket> listeners = new List<Socket>();
foreach (var ipep in ServerConfiguration.Listeners)
{
var listener = CreateListener(ipep);
if (listener == null)
{
continue;
}
if (ipep.Address.Equals(IPAddress.Any) || ipep.Address.Equals(IPAddress.IPv6Any))
{
listeningAddresses.UnionWith(GetListeningAddresses(ipep));
}
else
{
listeningAddresses.Add(ipep);
}
listeners.Add(listener);
BeginAcceptingSockets(listener);
}
foreach (var ipep in listeningAddresses)
{
logger.Information("Listening: {Address}:{Port}", ipep.Address, ipep.Port);
}
ListeningAddresses = listeningAddresses.ToArray();
Listeners = listeners.ToArray();
}
public static void Shutdown()
{
foreach (var listener in Listeners)
{
listener.Close();
}
}
public static IEnumerable<IPEndPoint> GetListeningAddresses(IPEndPoint ipep) =>
@ -104,104 +123,48 @@ public static class TcpServer
return null;
}
private static async void BeginAcceptingSockets(object state)
private static async void BeginAcceptingSockets(Socket listener)
{
if (state is not Socket listener)
{
return;
}
var cancellationToken = Core.ClosingTokenSource.Token;
while (!cancellationToken.IsCancellationRequested)
while (!Core.Closing)
{
Socket socket = null;
try
{
var socket = await listener.AcceptAsync(cancellationToken);
_connectingQueue.Enqueue(socket);
_queueSemaphore.Release();
}
catch (OperationCanceledException)
{
return;
socket = await listener.AcceptAsync();
var remoteIP = ((IPEndPoint)socket.RemoteEndPoint)!.Address;
if (!IPLimiter.Verify(remoteIP))
{
TraceDisconnect("Past IP limit threshold", remoteIP);
logger.Debug("{Address} Past IP limit threshold", remoteIP);
}
else if (Firewall.IsBlocked(remoteIP))
{
TraceDisconnect("Firewalled", remoteIP);
logger.Debug("{Address} Firewalled", remoteIP);
}
else
{
var args = new SocketConnectEventArgs(socket);
EventSink.InvokeSocketConnect(args);
if (args.AllowConnection)
{
_ = new NetState(socket);
continue;
}
TraceDisconnect("Rejected by socket event handler", remoteIP);
// Reject the connection
socket.Send(_socketRejected, SocketFlags.None);
}
CloseSocket(socket);
}
catch
{
// ignored
}
}
listener.Close();
}
private static void ProcessConnections()
{
var cancellationToken = Core.ClosingTokenSource.Token;
HashSet<IPEndPoint> listeningAddresses = [];
List<Socket> listeners = [];
foreach (var ipep in ServerConfiguration.Listeners)
{
var listener = CreateListener(ipep);
if (listener == null)
{
continue;
}
bool added;
if (ipep.Address.Equals(IPAddress.Any) || ipep.Address.Equals(IPAddress.IPv6Any))
{
var beforeCount = listeningAddresses.Count;
listeningAddresses.UnionWith(GetListeningAddresses(ipep));
added = listeningAddresses.Count > beforeCount;
}
else
{
added = listeningAddresses.Add(ipep);
}
if (added)
{
listeners.Add(listener);
new Thread(BeginAcceptingSockets).Start(listener);
}
}
foreach (var ipep in listeningAddresses)
{
logger.Information("Listening: {Address}:{Port}", ipep.Address, ipep.Port);
}
ListeningAddresses = listeningAddresses.ToArray();
Listeners = listeners.ToArray();
while (true)
{
try
{
while (!cancellationToken.IsCancellationRequested)
{
_queueSemaphore.Wait(cancellationToken);
Firewall.ProcessQueue();
if (_connectingQueue.TryDequeue(out var socket))
{
ProcessConnection(socket);
}
}
return;
}
catch (OperationCanceledException)
{
return;
}
catch (Exception e)
{
logger.Error(e, "Error occurred in ProcessConnections");
CloseSocket(socket);
}
}
}
@ -213,71 +176,9 @@ public static class TcpServer
{
socket.Shutdown(SocketShutdown.Both);
}
catch
finally
{
// ignored
}
socket.Close();
}
private static void ProcessConnection(Socket socket)
{
try
{
var remoteIP = ((IPEndPoint)socket.RemoteEndPoint)!.Address;
if (NetState.Instances.Count >= MaxConnections)
{
var ticks = Core.TickCount;
if (ticks - _nextMaximumSocketsReachedMessage > 0)
{
if (socket.RemoteEndPoint is IPEndPoint ipep)
{
var ip = ipep.Address.ToString();
logger.Warning("{Address} Failed (Maximum connections reached)", ip);
}
_nextMaximumSocketsReachedMessage = ticks + ListenerErrorMessageDelay;
}
CloseSocket(socket);
return;
}
var firewalled = Firewall.IsBlocked(remoteIP);
if (!firewalled)
{
var socketConnectedArgs = new SocketConnectedEventArgs(socket);
EventSink.InvokeSocketConnected(socketConnectedArgs);
firewalled = !socketConnectedArgs.ConnectionAllowed;
}
if (firewalled)
{
TraceDisconnect("Firewalled", remoteIP);
logger.Debug("{Address} Firewalled", remoteIP);
CloseSocket(socket);
return;
}
if (!IPLimiter.Verify(remoteIP))
{
TraceDisconnect("Past IP limit threshold", remoteIP);
logger.Debug("{Address} Past IP limit threshold", remoteIP);
CloseSocket(socket);
return;
}
var ns = new NetState(socket);
ConnectedQueue.Enqueue(ns);
}
catch
{
// ignored
socket.Close();
}
}
@ -299,22 +200,4 @@ public static class TcpServer
// ignored
}
}
public static class EventSink
{
// IMPORTANT: This is executed asynchronously! Do not run any game thread code on these delegates!
public static event Action<SocketConnectedEventArgs> SocketConnected;
internal static void InvokeSocketConnected(SocketConnectedEventArgs context) =>
SocketConnected?.Invoke(context);
}
public class SocketConnectedEventArgs
{
public Socket Socket { get; }
public bool ConnectionAllowed { get; set; } = true;
internal SocketConnectedEventArgs(Socket socket) => Socket = socket;
}
}

View file

@ -169,7 +169,7 @@ namespace Server.Gumps
AddLabel(150, 150, LabelHue, banned.ToString());
AddLabel(20, 170, LabelHue, "Firewalled:");
AddLabel(150, 170, LabelHue, AdminFirewall.Set.Count.ToString());
AddLabel(150, 170, LabelHue, Firewall.FirewallSet.Count.ToString());
AddLabel(20, 190, LabelHue, "Clients:");
AddLabel(150, 190, LabelHue, NetState.Instances.Count.ToString());
@ -1161,7 +1161,7 @@ namespace Server.Gumps
{
AddFirewallHeader();
m_List ??= AdminFirewall.Set.ToList<object>();
m_List ??= Firewall.FirewallSet.ToList<object>();
AddLabelCropped(12, 120, 358, 20, LabelHue, "IP Address");
@ -3464,7 +3464,7 @@ namespace Server.Gumps
}
else
{
foreach (var check in AdminFirewall.Set)
foreach (var check in Firewall.FirewallSet)
{
var checkStr = check.ToString();

View file

@ -1,6 +1,5 @@
using System;
using System.Buffers;
using System.Collections.Generic;
using System.IO;
using System.Net;
using System.Runtime.CompilerServices;
@ -13,7 +12,6 @@ public static class AdminFirewall
{
private static readonly ILogger logger = LogFactory.GetLogger(typeof(AdminFirewall));
private static readonly HashSet<IFirewallEntry> _firewallSet = [];
private const string firewallConfigPath = "firewall.cfg";
public static void Configure()
@ -44,9 +42,6 @@ public static class AdminFirewall
}
}
// Note: This is not optimized, so do not use this in hot paths
public static IReadOnlySet<IFirewallEntry> Set => _firewallSet;
[MethodImpl(MethodImplOptions.AggressiveInlining)]
public static IFirewallEntry ToFirewallEntry(object entry)
{
@ -91,41 +86,49 @@ public static class AdminFirewall
}
}
public static void Remove(object obj, bool save = true)
public static bool Remove(object obj, bool save = true)
{
var entry = ToFirewallEntry(obj);
if (entry != null)
if (entry == null)
{
_firewallSet.Remove(entry);
Firewall.RequestRemoveEntry(entry); // Request that the TcpServer also remove the entry
if (save)
{
Save();
}
return false;
}
}
public static bool Add(object obj) => Add(ToFirewallEntry(obj));
public static bool Add(IFirewallEntry entry, bool save = true)
{
var added = _firewallSet.Add(entry);
Firewall.RequestAddEntry(entry); // Request that the TcpServer also add the entry
if (!Firewall.Remove(entry))
{
return false;
}
if (save)
{
Save();
}
return added;
return true;
}
public static void Add(object obj) => Add(ToFirewallEntry(obj));
public static bool Add(IFirewallEntry entry, bool save = true)
{
if (!Firewall.Add(entry))
{
return false;
}
if (save)
{
Save();
}
return true;
}
public static void Save()
{
using var op = new StreamWriter(firewallConfigPath);
foreach (var entry in Set)
foreach (var entry in Firewall.FirewallSet)
{
op.WriteLine(entry);
}

View file

@ -1,4 +1,4 @@
{
"$schema": "https://raw.githubusercontent.com/dotnet/Nerdbank.GitVersioning/master/src/NerdBank.GitVersioning/version.schema.json",
"version": "0.13.5"
"version": "0.13.6"
}