Remove ConnectToHost auto try NetID feature, Reimplemented as server send new NetID

This commit is contained in:
kmyuhkyuk
2026-03-10 23:01:51 +08:00
parent f758fa36e3
commit 3d9018ab88
5 changed files with 308 additions and 90 deletions
@@ -0,0 +1,59 @@
using System.Buffers.Binary;
using MegaCrit.Sts2.Core.Multiplayer.Transport.ENet;
using SlayTheSpire2.LAN.Multiplayer.Models;
namespace SlayTheSpire2.LAN.Multiplayer.Helpers
{
internal class LanHandshakeResponseHelper
{
public static ENetPacket FromLanHandshakeResponse(ENetLanHandshakeResponse response)
{
var array = new byte[]
{
1,
(byte)response.status,
0,
0,
0,
0,
0,
0,
0,
0,
0,
0,
0,
0,
0,
0,
0,
0
};
var span = array.AsSpan();
BinaryPrimitives.WriteUInt64BigEndian(span.Slice(2, 8), response.netId);
BinaryPrimitives.WriteUInt64BigEndian(span.Slice(10, 8), response.newNetId);
return new ENetPacket(array);
}
public static ENetLanHandshakeResponse AsLanHandshakeResponse(ENetPacket eNetPacket)
{
if (eNetPacket.PacketType != ENetPacketType.HandshakeResponse)
{
throw new InvalidOperationException(
$"Attempted to interpret ENet packet of type {eNetPacket.PacketType} as handshake response");
}
var span = eNetPacket.AllBytes.AsSpan();
var status = (ENetHandshakeStatus)span[1];
var netId = BinaryPrimitives.ReadUInt64BigEndian(span.Slice(2, 8));
var newNetId = BinaryPrimitives.ReadUInt64BigEndian(span.Slice(10, 8));
return new ENetLanHandshakeResponse
{
netId = netId,
newNetId = newNetId,
status = status
};
}
}
}
@@ -0,0 +1,15 @@
using MegaCrit.Sts2.Core.Multiplayer.Transport.ENet;
// ReSharper disable InconsistentNaming
namespace SlayTheSpire2.LAN.Multiplayer.Models
{
public struct ENetLanHandshakeResponse
{
public ENetHandshakeStatus status;
public ulong netId;
public ulong newNetId;
}
}
@@ -1,73 +0,0 @@
using HarmonyLib;
using MegaCrit.Sts2.Core.Entities.Multiplayer;
using MegaCrit.Sts2.Core.Helpers;
using MegaCrit.Sts2.Core.Logging;
using MegaCrit.Sts2.Core.Multiplayer;
using MegaCrit.Sts2.Core.Multiplayer.Connection;
using MegaCrit.Sts2.Core.Multiplayer.Transport.ENet;
using MegaCrit.Sts2.Core.Platform;
// ReSharper disable UnusedMember.Global
// ReSharper disable UnusedType.Global
namespace SlayTheSpire2.LAN.Multiplayer.Patchs
{
[HarmonyPatch(typeof(ENetClientConnectionInitializer), "Connect")]
internal class ENetClientConnectionInitializerPatch
{
private static bool Prefix(NetClientGameService gameService,
CancellationToken cancelToken, ulong ____netId, string ____ip, ushort ____port, ref Task __result)
{
__result = TaskHelper.RunSafely(Connect(gameService, cancelToken, ____netId, ____ip, ____port));
return false;
}
private static async Task<NetErrorInfo?> Connect(NetClientGameService gameService,
CancellationToken cancelToken, ulong netId, string ip, ushort port)
{
if (gameService.IsConnected)
{
throw new InvalidOperationException(
"NetClientGameService must not be connected when passed to ENetClientConnectionInitializer!");
}
var eNetClient = new ENetClient(gameService);
gameService.Initialize(eNetClient, PlatformType.None);
var count = 0;
const int tryCount = 10;
NetErrorInfo? netErrorInfo = null;
while (count < tryCount)
{
if (cancelToken.IsCancellationRequested)
return netErrorInfo;
netErrorInfo = await eNetClient.ConnectToHost(netId, ip, port, cancelToken);
if (!netErrorInfo.HasValue)
{
Log.Info($"Connect {ip}:{port} HostGame NetID:{netId}");
return null;
}
if (netErrorInfo.Value.GetReason() != NetError.Kicked)
return netErrorInfo;
var nextNetId = netId + 1000u;
if (count < tryCount - 1)
{
Log.Warn($"{ip}:{port} HostGame NetID:{netId} already occupied, Next will try NetID:{nextNetId}");
}
netId = nextNetId;
count++;
}
return netErrorInfo;
}
}
}
@@ -0,0 +1,157 @@
using Godot;
using HarmonyLib;
using MegaCrit.Sts2.Core.Entities.Multiplayer;
using MegaCrit.Sts2.Core.Helpers;
using MegaCrit.Sts2.Core.Multiplayer.Transport;
using MegaCrit.Sts2.Core.Multiplayer.Transport.ENet;
using SlayTheSpire2.LAN.Multiplayer.Helpers;
using Logger = MegaCrit.Sts2.Core.Logging.Logger;
// ReSharper disable UnusedMember.Global
// ReSharper disable UnusedType.Global
namespace SlayTheSpire2.LAN.Multiplayer.Patchs
{
[HarmonyPatch(typeof(ENetClient), "ConnectToHost")]
internal class ENetClientConnectToHostPatch
{
private static bool Prefix(ENetClient __instance, ulong netId, string ip, ushort port,
CancellationToken cancelToken, Logger ____logger, INetClientHandler ____handler, ref Task __result)
{
__result = TaskHelper.RunSafely(ConnectToHost(__instance, ____logger, ____handler, netId, ip, port,
cancelToken));
return false;
}
private static async Task<NetErrorInfo?> ConnectToHost(ENetClient eNetClient, Logger logger, INetClientHandler handler, ulong netId, string ip, ushort port, CancellationToken cancelToken)
{
while (true)
{
var connection = new ENetConnection();
Traverse.Create(eNetClient).Field("_connection").SetValue(connection);
connection.CreateHost();
var peer = connection.ConnectToHost(ip, port);
Traverse.Create(eNetClient).Field("_peer").SetValue(peer);
var timeoutTimer = 0;
while (!connection.TryService(out var output) || output is not { type: ENetConnection.EventType.Connect })
{
await Task.Delay(100, cancelToken);
if (cancelToken.IsCancellationRequested)
{
eNetClient.DisconnectFromHost(NetError.CancelledJoin);
logger.Warn("User cancelled join flow");
return null;
}
timeoutTimer += 100;
if (timeoutTimer > 10000)
{
peer.Reset();
logger.Error("Connection timed out!");
return new NetErrorInfo(NetError.Timeout, selfInitiated: false);
}
}
if (peer.GetState() != ENetPacketPeer.PeerState.Connected)
{
logger.Error($"Connection to {ip}:{port} failed!");
return new NetErrorInfo(NetError.UnknownNetworkError, selfInitiated: false);
}
var bufferedPackets = new List<ENetServiceData>();
var (result, newNetId) = await SendAndWaitForNetIdAck(eNetClient, logger, peer, connection, netId, bufferedPackets, cancelToken);
if (result.HasValue)
{
peer.PeerDisconnect();
if (result.Value.GetReason() == NetError.Kicked)
{
netId = newNetId;
continue;
}
return result;
}
Traverse.Create(eNetClient).Field("_netId").SetValue(newNetId);
Traverse.Create(eNetClient).Field("_isConnected").SetValue(true);
handler.OnConnectedToHost();
foreach (var item in bufferedPackets)
{
Traverse.Create(eNetClient).Method("HandleMessageReceived", item).GetValue();
}
return null;
}
}
private static async Task<(NetErrorInfo?, ulong)> SendAndWaitForNetIdAck(ENetClient eNetClient, Logger logger,
ENetPacketPeer peer, ENetConnection connection, ulong netId, List<ENetServiceData> bufferedPackets,
CancellationToken cancelToken)
{
logger.Info($"Sending handshake with net ID {netId}");
var eNetPacket = ENetPacket.FromHandshakeRequest(new ENetHandshakeRequest
{
netId = netId
});
peer.Send(0, eNetPacket.AllBytes, 1);
var receivedAck = false;
var timeoutTimer = 0;
while (!receivedAck)
{
await Task.Delay(100, cancelToken);
if (cancelToken.IsCancellationRequested)
{
logger.Warn("User cancelled join flow");
eNetClient.DisconnectFromHost(NetError.CancelledJoin);
return (null, netId);
}
if (connection.TryService(out var output) && output is { type: ENetConnection.EventType.Receive })
{
var packetData = output.Value.packetData;
var eNetPacket2 = new ENetPacket(packetData);
if (eNetPacket2.PacketType == ENetPacketType.ApplicationMessage)
{
bufferedPackets.Add(output.Value);
continue;
}
var eNetHandshakeResponse = LanHandshakeResponseHelper.AsLanHandshakeResponse(eNetPacket2);
if (eNetHandshakeResponse.netId != netId)
{
logger.Error(
$"Received net ID ({eNetHandshakeResponse.netId}) during handshake that did not match ours!");
return (new NetErrorInfo(NetError.InternalError, selfInitiated: false), netId);
}
if (eNetHandshakeResponse.status == ENetHandshakeStatus.IdCollision)
{
logger.Warn(
$"NetID:{netId} already occupied, Next try server send new NetID:{eNetHandshakeResponse.newNetId}");
return (new NetErrorInfo(NetError.Kicked, selfInitiated: false),
eNetHandshakeResponse.newNetId);
}
else if (eNetHandshakeResponse.status != ENetHandshakeStatus.Success)
{
logger.Error($"Received non-success code during handshake ({eNetHandshakeResponse.status})!");
return (new NetErrorInfo(NetError.Kicked, selfInitiated: false), netId);
}
receivedAck = true;
}
timeoutTimer += 100;
if (timeoutTimer > 10000)
{
logger.Error("Timed out waiting for handshake ack!");
eNetClient.DisconnectFromHost(NetError.Timeout);
return (new NetErrorInfo(NetError.Timeout, selfInitiated: false), netId);
}
}
return (null, netId);
}
}
}
@@ -1,39 +1,99 @@
using System.Reflection; using System.Collections;
using System.Reflection.Emit;
using Godot; using Godot;
using HarmonyLib; using HarmonyLib;
using MegaCrit.Sts2.Core.Helpers;
using MegaCrit.Sts2.Core.Multiplayer.Transport;
using MegaCrit.Sts2.Core.Multiplayer.Transport.ENet; using MegaCrit.Sts2.Core.Multiplayer.Transport.ENet;
using SlayTheSpire2.LAN.Multiplayer.Helpers;
using SlayTheSpire2.LAN.Multiplayer.Models;
using Logger = MegaCrit.Sts2.Core.Logging.Logger;
// ReSharper disable UnusedMember.Global // ReSharper disable UnusedMember.Global
// ReSharper disable UnusedType.Global // ReSharper disable UnusedType.Global
namespace SlayTheSpire2.LAN.Multiplayer.Patchs namespace SlayTheSpire2.LAN.Multiplayer.Patchs
{ {
[HarmonyPatch] [HarmonyPatch(typeof(ENetHost), "DoClientHandshake")]
internal class ENetHostDoClientHandshakePatch internal class ENetHostDoClientHandshakePatch
{ {
private static MethodInfo TargetMethod() private static bool Prefix(ENetHost __instance, ENetPacketPeer peer, Logger ____logger,
IList ____receivedHandshakes, IList ____connectedPeers, INetHostHandler ____handler, ref Task __result)
{ {
return AccessTools.AsyncMoveNext(typeof(ENetHost).GetMethod("DoClientHandshake", __result = TaskHelper.RunSafely(DoClientHandshake(__instance, ____logger, ____receivedHandshakes,
BindingFlags.Instance | BindingFlags.NonPublic)); ____connectedPeers, ____handler, peer));
return false;
} }
private static IEnumerable<CodeInstruction> Transpiler(IEnumerable<CodeInstruction> instructions) private static async Task DoClientHandshake(ENetHost eNetHost, Logger logger, IList receivedHandshakes,
IList connectedPeers, INetHostHandler handler, ENetPacketPeer peer)
{ {
//Fixed IdCollision HandshakeResponse will not send peer.SetTimeout(24, 20000, 20000);
var timeoutTimer = 0;
var peerDisconnectMethod = AccessTools.Method(typeof(ENetPacketPeer), "PeerDisconnect"); object? handshake = null;
var peerDisconnectLaterMethod = AccessTools.Method(typeof(ENetPacketPeer), "PeerDisconnectLater"); while (handshake == null)
foreach (var instruction in instructions)
{ {
if (instruction.opcode == OpCodes.Callvirt && instruction.operand as MethodInfo == peerDisconnectMethod) foreach (var receivedHandshake in receivedHandshakes)
{ {
yield return new CodeInstruction(OpCodes.Callvirt, peerDisconnectLaterMethod); if (Traverse.Create(receivedHandshake).Field("conn").Field("peer").GetValue() == peer)
continue; {
handshake = receivedHandshake;
break;
}
} }
yield return instruction; if (handshake == null)
{
await Task.Delay(100);
timeoutTimer += 100;
if (timeoutTimer >= 10000)
{
logger.Error("Timed out waiting for handshake!");
peer.Reset();
return;
}
}
}
var handshakeConn = Traverse.Create(handshake).Field("conn").GetValue();
var handshakeNetId = Traverse.Create(handshakeConn).Field("netId").GetValue<ulong>();
var handshakePeer = Traverse.Create(handshakeConn).Field("peer").GetValue<ENetPacketPeer>();
var connectedPeerIdHashSet = new HashSet<ulong>(eNetHost.ConnectedPeerIds);
if (connectedPeerIdHashSet.Contains(handshakeNetId))
{
var newNetId = handshakeNetId;
while (connectedPeerIdHashSet.Contains(newNetId))
{
newNetId += 1000;
}
logger.Info(
$"Second client attempted to connect with peer ID {handshakeNetId}, disconnecting them and return new NetId:{newNetId}");
var eNetPacket = LanHandshakeResponseHelper.FromLanHandshakeResponse(new ENetLanHandshakeResponse
{
netId = handshakeNetId,
newNetId = newNetId,
status = ENetHandshakeStatus.IdCollision
});
handshakePeer.Send(0, eNetPacket.AllBytes, 1);
handshakePeer.PeerDisconnectLater();
}
else
{
logger.Debug($"Acknowledging handshake for peer with ID {handshakeNetId}");
var eNetPacket2 = LanHandshakeResponseHelper.FromLanHandshakeResponse(new ENetLanHandshakeResponse
{
netId = handshakeNetId,
newNetId = handshakeNetId,
status = ENetHandshakeStatus.Success
});
handshakePeer.Send(0, eNetPacket2.AllBytes, 1);
connectedPeers.Add(handshakeConn);
handler.OnPeerConnected(handshakeNetId);
} }
} }
} }