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);
}
}