fix: Fixes stalled connections and infinite throttle (#1796)

> [!Warning]
> **Developer Warning**
> The `PacketThrottle` callback return value is now reversed. `true` indicates the connection is _throttled_.

### Summary
- Fixes an issue where connections get stalled forever
- Fixes an issue where the throttler is not working properly
- Removes account attack limiter
- Rewrites IP limiter
- Removes IP restrictions (they weren't used, and not practical)
- Fixes issue where IP limiter was counting before firewall was blocking.

View without whitespace:
https://github.com/modernuo/ModernUO/pull/1796/files?diff=split&w=1
This commit is contained in:
Kamron Batman 2024-05-27 21:32:55 -07:00 committed by GitHub
parent ccad915464
commit a4522b9d43
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
13 changed files with 1538 additions and 1783 deletions

View file

@ -16,100 +16,103 @@
using System;
using System.Collections.Generic;
using System.Net;
using System.Runtime.InteropServices;
namespace Server.Misc;
public static class IPLimiter
{
private static readonly Dictionary<IPAddress, int> _connectionAttempts = new(128);
private static readonly HashSet<IPAddress> _throttledAddresses = new();
private static readonly SortedSet<IPAccessLog> _connectionAttempts = [];
private static readonly SortedSet<IPAccessLog> _throttledAddresses = [];
private static long _lastClearedThrottles;
private static long _lastClearedAttempts;
private static readonly IPAddress _localHost = IPAddress.Parse("127.0.0.1");
public static readonly IPAddress[] Exemptions =
{
IPAddress.Parse( "127.0.0.1" )
};
public static TimeSpan ClearConnectionAttemptsDuration { get; private set; }
public static TimeSpan ClearThrottledDuration { get; private set; }
public static TimeSpan ConnectionAttemptsDuration { get; private set; }
public static TimeSpan ConnectionThrottleDuration { get; private set; }
public static bool Enabled { get; private set; }
public static int MaxAddresses { get; private set; }
public static int MaxConnections { get; private set; }
public static void Configure()
{
Enabled = ServerConfiguration.GetOrUpdateSetting("ipLimiter.enable", true);
MaxAddresses = ServerConfiguration.GetOrUpdateSetting("ipLimiter.maxConnectionsPerIP", 10);
ClearConnectionAttemptsDuration = ServerConfiguration.GetOrUpdateSetting("ipLimiter.clearConnectionAttemptsDuration", TimeSpan.FromSeconds(10));
ClearThrottledDuration = ServerConfiguration.GetOrUpdateSetting("ipLimiter.clearThrottledDuration", TimeSpan.FromMinutes(2));
MaxConnections = ServerConfiguration.GetOrUpdateSetting("ipLimiter.maxConnectionsPerIP", 5);
ConnectionAttemptsDuration = ServerConfiguration.GetOrUpdateSetting("ipLimiter.clearConnectionAttemptsDuration", TimeSpan.FromSeconds(10));
ConnectionThrottleDuration = ServerConfiguration.GetOrUpdateSetting("ipLimiter.connectionThrottleDuration", TimeSpan.FromMinutes(5));
}
public static bool IsExempt(IPAddress ip)
{
for (int i = 0; i < Exemptions.Length; i++)
{
if (ip.Equals(Exemptions[i]))
{
return true;
}
}
return false;
}
private static readonly IPAccessLog _accessCheck = new(IPAddress.None, DateTime.MinValue);
public static bool Verify(IPAddress ourAddress)
{
if (!Enabled || IsExempt(ourAddress))
if (!Enabled || ourAddress.Equals(_localHost))
{
return true;
}
var now = Core.TickCount;
var now = Core.Now;
if (_throttledAddresses.Count > 0)
CheckThrottledAddresses(now);
_accessCheck.IPAddress = ourAddress;
if (_connectionAttempts.TryGetValue(_accessCheck, out var accessLog))
{
if (now - _lastClearedThrottles > ClearThrottledDuration.TotalMilliseconds)
{
_lastClearedThrottles = now;
ClearThrottledAddresses();
}
else if (_throttledAddresses.Contains(ourAddress))
_connectionAttempts.Remove(accessLog);
accessLog.Count++;
if (now <= accessLog.Expiration && accessLog.Count >= MaxConnections)
{
BlockConnection(now, accessLog);
return false;
}
}
if (_connectionAttempts.Count > 0 && now - _lastClearedAttempts > ClearConnectionAttemptsDuration.TotalMilliseconds)
accessLog.Expiration = now + ConnectionAttemptsDuration;
}
else
{
_lastClearedAttempts = now;
ClearConnectionAttempts();
accessLog = new IPAccessLog(ourAddress, now + ConnectionAttemptsDuration);
}
ref var count = ref CollectionsMarshal.GetValueRefOrAddDefault(_connectionAttempts, ourAddress, out _);
count++;
if (count > MaxAddresses)
{
_connectionAttempts.Remove(ourAddress);
_throttledAddresses.Add(ourAddress);
return false;
}
// Add it back so it is sorted properly
_connectionAttempts.Add(accessLog);
return true;
}
private static void ClearThrottledAddresses()
private static void BlockConnection(DateTime now, IPAccessLog accessLog)
{
_throttledAddresses.Clear();
accessLog.Expiration = now + ConnectionAttemptsDuration;
_throttledAddresses.Add(accessLog);
}
private static void ClearConnectionAttempts()
private static void CheckThrottledAddresses(DateTime now)
{
_connectionAttempts.Clear();
_connectionAttempts.TrimExcess(128);
while (_throttledAddresses.Count > 0)
{
var accessLog = _throttledAddresses.Min;
if (now <= accessLog.Expiration)
{
break;
}
_throttledAddresses.Remove(accessLog);
}
}
private class IPAccessLog : IComparable<IPAccessLog>
{
public IPAddress IPAddress;
public DateTime Expiration;
public int Count;
public IPAccessLog(IPAddress ipAddress, DateTime expiration)
{
IPAddress = ipAddress;
Expiration = expiration;
Count = 1;
}
public int CompareTo(IPAccessLog other) =>
IPAddress.Equals(other.IPAddress) ? 0 : Expiration.CompareTo(other.Expiration);
}
}

View file

@ -33,15 +33,13 @@ public static class MovementThrottle
IncomingPackets.RegisterThrottler(0x02, &Throttle);
}
public static bool Throttle(int packetId, NetState ns, out bool drop)
public static bool Throttle(int packetId, NetState ns)
{
drop = false;
var from = ns.Mobile;
if (from?.Deleted != false || from.AccessLevel > AccessLevel.Player)
{
return true;
return false;
}
long now = Core.TickCount;
@ -53,7 +51,7 @@ public static class MovementThrottle
{
ns._movementCredit = 0;
ns._nextMovementTime = now;
return true;
return false;
}
long cost = nextMove - now;
@ -61,11 +59,11 @@ public static class MovementThrottle
if (credit < cost)
{
// Not enough credit, therefore throttled
return false;
return true;
}
// On the next event loop, the player receives up to 400ms in grace latency
ns._movementCredit = Math.Min(_throttleThreshold, credit - cost);
return true;
return false;
}
}

View file

@ -823,9 +823,9 @@ public partial class NetState : IComparable<NetState>, IValueLinkListNode<NetSta
var throttler = handler.ThrottleCallback;
if (throttler != null)
{
if (!throttler(packetId, this, out bool drop))
if (throttler(packetId, this))
{
return drop ? ParserState.Throttled : ParserState.AwaitingNextPacket;
return ParserState.Throttled;
}
SetPacketTime(packetId);

View file

@ -35,7 +35,7 @@ public unsafe class PacketHandler
public delegate*<NetState, SpanReader, void> OnReceive { get; }
public delegate*<int, NetState, out bool, bool> ThrottleCallback { get; set; }
public delegate*<int, NetState, bool> ThrottleCallback { get; set; }
public bool Ingame { get; }
}

View file

@ -57,7 +57,7 @@ public static class IncomingPackets
}
}
public static unsafe void RegisterThrottler(int packetID, delegate*<int, NetState, out bool, bool> t)
public static unsafe void RegisterThrottler(int packetID, delegate*<int, NetState, bool> t)
{
var ph = GetHandler(packetID);

View file

@ -223,7 +223,6 @@ public static class TcpServer
private static void ProcessConnection(Socket socket)
{
var ipLimiter = IPLimiter.Enabled;
try
{
var remoteIP = ((IPEndPoint)socket.RemoteEndPoint)!.Address;
@ -247,15 +246,6 @@ public static class TcpServer
return;
}
if (ipLimiter && !IPLimiter.Verify(remoteIP))
{
TraceDisconnect("Past IP limit threshold", remoteIP);
logger.Debug("{Address} Past IP limit threshold", remoteIP);
CloseSocket(socket);
return;
}
var firewalled = Firewall.IsBlocked(remoteIP);
if (!firewalled)
{
@ -273,6 +263,15 @@ public static class TcpServer
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);
}