392 lines
16 KiB
Python
392 lines
16 KiB
Python
from __future__ import annotations
|
|
|
|
import threading
|
|
import time
|
|
from dataclasses import dataclass
|
|
from typing import TYPE_CHECKING
|
|
|
|
from client_session import ClientSession
|
|
from gns_snapshot_envelope import SnapshotEnvelopeError, decode_snapshot, encode_snapshot
|
|
from gns_transport import EventType, GnsEvent, GnsServerTransport, GnsTransportError, SendResult
|
|
from packet_codec import EncodedPacket, PacketCodecError, decode_packet
|
|
from snapshot_sequence import SequenceCounter, SequenceWindow
|
|
from transport_policy import delivery_for_packet_type, is_snapshot_packet
|
|
|
|
if TYPE_CHECKING:
|
|
from server_core import FalloutTogetherServer
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class GnsAdapterStats:
|
|
connections: int
|
|
thread_running: bool
|
|
|
|
|
|
class GnsConnectionAdapter:
|
|
def __init__(self, transport: GnsServerTransport, connection_id: int) -> None:
|
|
self.transport = transport
|
|
self.connection_id = connection_id
|
|
self._closed = False
|
|
self._lock = threading.Lock()
|
|
self._snapshot_counters = {
|
|
"transform": SequenceCounter(),
|
|
"npcState": SequenceCounter(),
|
|
}
|
|
|
|
def fileno(self) -> int:
|
|
with self._lock:
|
|
return -1 if self._closed else self.connection_id
|
|
|
|
def mark_remote_closed(self) -> None:
|
|
with self._lock:
|
|
self._closed = True
|
|
|
|
def _handle_send_result(self, packet_type: str, result: SendResult) -> None:
|
|
if result is SendResult.SENT:
|
|
return
|
|
if result is SendResult.DROPPED and is_snapshot_packet(packet_type):
|
|
return
|
|
if result is SendResult.DROPPED:
|
|
raise OSError("GNS refused to queue a reliable message")
|
|
if result is SendResult.BACKPRESSURE and is_snapshot_packet(packet_type):
|
|
return
|
|
if result is SendResult.BACKPRESSURE:
|
|
raise OSError("GNS reliable send queue is under backpressure")
|
|
if result is SendResult.NOT_CONNECTED:
|
|
self.mark_remote_closed()
|
|
raise OSError("GNS connection is no longer active")
|
|
if result is SendResult.TOO_LARGE:
|
|
raise ValueError("GNS outbound packet exceeds maximum message size")
|
|
raise OSError("GNS outbound send failed")
|
|
|
|
def send_encoded(self, encoded: EncodedPacket) -> None:
|
|
with self._lock:
|
|
if self._closed:
|
|
raise OSError("GNS connection is closed")
|
|
|
|
packet_type = encoded.packet_type
|
|
expected_delivery = delivery_for_packet_type(packet_type)
|
|
if encoded.delivery is not expected_delivery:
|
|
raise OSError("GNS packet delivery metadata does not match protocol policy")
|
|
|
|
wire_payload = encoded.payload
|
|
if is_snapshot_packet(packet_type):
|
|
with self._lock:
|
|
sequence = self._snapshot_counters[packet_type].advance()
|
|
try:
|
|
wire_payload = encode_snapshot(packet_type, encoded.payload, sequence)
|
|
except SnapshotEnvelopeError as error:
|
|
raise ValueError(str(error)) from error
|
|
|
|
wire_packet = EncodedPacket(packet_type, wire_payload, encoded.delivery)
|
|
result = self.transport.send_encoded(self.connection_id, wire_packet)
|
|
self._handle_send_result(packet_type, result)
|
|
|
|
def sendall(self, framed_payload: bytes) -> None:
|
|
"""Temporary compatibility for callers that still emit TCP line framing."""
|
|
raw = bytes(framed_payload)
|
|
if not raw.endswith(b"\n") or raw.count(b"\n") != 1:
|
|
raise OSError("GNS compatibility adapter received invalid TCP framing")
|
|
payload = raw[:-1]
|
|
try:
|
|
packet = decode_packet(payload)
|
|
except PacketCodecError as error:
|
|
raise OSError(f"invalid outbound GNS packet: {error}") from error
|
|
packet_type = packet["type"]
|
|
self.send_encoded(EncodedPacket(packet_type, payload, delivery_for_packet_type(packet_type)))
|
|
|
|
def close(self) -> None:
|
|
with self._lock:
|
|
if self._closed:
|
|
return
|
|
self._closed = True
|
|
try:
|
|
self.transport.disconnect(self.connection_id, debug="Commonwealth Online disconnect")
|
|
except Exception:
|
|
pass
|
|
|
|
|
|
class GnsGameplayAdapter:
|
|
def __init__(self, server: "FalloutTogetherServer", transport: GnsServerTransport) -> None:
|
|
self.server = server
|
|
self.transport = transport
|
|
self._clients: dict[int, ClientSession] = {}
|
|
self._connected_monotonic: dict[int, float] = {}
|
|
self._incoming_snapshot_windows: dict[tuple[int, str], SequenceWindow] = {}
|
|
self._lock = threading.RLock()
|
|
self._thread: threading.Thread | None = None
|
|
self._thread_running = False
|
|
|
|
def get_stats(self) -> GnsAdapterStats:
|
|
with self._lock:
|
|
return GnsAdapterStats(len(self._clients), self._thread_running)
|
|
|
|
def _send_pre_session_end(self, connection_id: int, code: str, reason: str) -> None:
|
|
packet = self.server._build_session_ended_packet(code, reason)
|
|
self.transport.send_packet(connection_id, packet)
|
|
|
|
def _welcome_packet(self, client: ClientSession) -> dict:
|
|
from server_core import PROTOCOL_VERSION
|
|
|
|
return {
|
|
"type": "welcome",
|
|
"playerId": client.player_id,
|
|
"serverTime": time.time(),
|
|
"serverName": self.server.server_name,
|
|
"serverDescription": self.server.server_description,
|
|
"protocolVersion": PROTOCOL_VERSION,
|
|
"capabilities": [
|
|
"interest-v1",
|
|
"hello-v2",
|
|
"bounded-framing",
|
|
"rate-limit-v1",
|
|
"movement-correction-v1",
|
|
"npc-authority-epoch-v1",
|
|
"player-state-v1",
|
|
"gns-message-transport-v1",
|
|
"gns-snapshot-sequence-v1",
|
|
],
|
|
}
|
|
|
|
def _handle_connected(self, event: GnsEvent) -> None:
|
|
from server_core import SESSION_ENDED_BANNED, SESSION_ENDED_RATE_LIMITED
|
|
|
|
endpoint = self.transport.remote_endpoint(event.connection_id)
|
|
if endpoint is None:
|
|
self.transport.disconnect(event.connection_id, debug="Remote endpoint unavailable")
|
|
return
|
|
|
|
ban_entry = self.server._ban_store.get_ban(endpoint.host)
|
|
if ban_entry is not None:
|
|
with self.server._lock:
|
|
self.server._stats["bannedConnectionsRejected"] += 1
|
|
self._send_pre_session_end(event.connection_id, SESSION_ENDED_BANNED, ban_entry.reason)
|
|
self.transport.disconnect(event.connection_id, debug="Banned")
|
|
return
|
|
|
|
if not self.server._allow_connect_attempt(endpoint.host):
|
|
with self.server._lock:
|
|
self.server._stats["pendingConnectionsRejected"] += 1
|
|
self._send_pre_session_end(
|
|
event.connection_id,
|
|
SESSION_ENDED_RATE_LIMITED,
|
|
"Too many connection attempts.",
|
|
)
|
|
self.transport.disconnect(event.connection_id, debug="Connection attempt rate limited")
|
|
return
|
|
|
|
connection = GnsConnectionAdapter(self.transport, event.connection_id)
|
|
client = self.server._assign_client(connection, (endpoint.host, endpoint.port))
|
|
if client is None:
|
|
self._send_pre_session_end(
|
|
event.connection_id,
|
|
SESSION_ENDED_RATE_LIMITED,
|
|
"Too many pending connections.",
|
|
)
|
|
connection.close()
|
|
return
|
|
|
|
connected_mono = time.monotonic()
|
|
self.server._init_rate_state(client, connected_mono)
|
|
with self._lock:
|
|
self._clients[event.connection_id] = client
|
|
self._connected_monotonic[event.connection_id] = connected_mono
|
|
try:
|
|
self.server._send_packet(client, self._welcome_packet(client))
|
|
except (OSError, ValueError):
|
|
self.server._disconnect_client(client)
|
|
self._purge_closed_clients()
|
|
|
|
def _account_transport_reject(self, client: ClientSession, reason: str, *, warning: bool = False) -> None:
|
|
if not self.server._allow_packet(client):
|
|
return
|
|
client.record_received(time.time())
|
|
with self.server._lock:
|
|
self.server._stats["packetsReceived"] += 1
|
|
self.server._reject_packet(client, reason, warning=warning)
|
|
|
|
def _handle_message(self, event: GnsEvent) -> None:
|
|
client = self._client_for_connection(event.connection_id)
|
|
if client is None:
|
|
self.transport.disconnect(event.connection_id, debug="Message before GNS admission")
|
|
return
|
|
|
|
payload = event.payload
|
|
try:
|
|
envelope = decode_snapshot(payload)
|
|
except SnapshotEnvelopeError as error:
|
|
self._account_transport_reject(client, f"Malformed GNS snapshot envelope: {error}")
|
|
return
|
|
|
|
if envelope is not None:
|
|
window_key = (event.connection_id, envelope.packet_type)
|
|
with self._lock:
|
|
window = self._incoming_snapshot_windows.setdefault(window_key, SequenceWindow())
|
|
accepted = window.accept(envelope.sequence)
|
|
if not accepted:
|
|
self._account_transport_reject(client, "Stale or duplicate GNS snapshot sequence", warning=False)
|
|
return
|
|
try:
|
|
packet = decode_packet(envelope.payload)
|
|
except PacketCodecError as error:
|
|
self._account_transport_reject(client, f"Invalid GNS snapshot payload: {error}")
|
|
return
|
|
if packet["type"] != envelope.packet_type:
|
|
self._account_transport_reject(client, "GNS snapshot envelope family does not match packet type")
|
|
return
|
|
payload = envelope.payload
|
|
else:
|
|
try:
|
|
packet = decode_packet(payload)
|
|
except PacketCodecError:
|
|
packet = None
|
|
if packet is not None and is_snapshot_packet(packet["type"]):
|
|
self._account_transport_reject(client, "GNS snapshot missing required sequence envelope")
|
|
return
|
|
|
|
try:
|
|
line = payload.decode("utf-8", errors="strict")
|
|
except UnicodeDecodeError:
|
|
self._account_transport_reject(client, "Packet is not valid UTF-8")
|
|
return
|
|
self.server._handle_line(client, line)
|
|
self._purge_closed_clients()
|
|
|
|
def _handle_oversize(self, event: GnsEvent) -> None:
|
|
from server_core import SESSION_ENDED_PACKET_TOO_LARGE
|
|
|
|
client = self._client_for_connection(event.connection_id)
|
|
if client is None:
|
|
self.transport.disconnect(event.connection_id, debug="Oversized pre-session packet")
|
|
return
|
|
self.server._end_client_session(
|
|
client,
|
|
code=SESSION_ENDED_PACKET_TOO_LARGE,
|
|
reason="Packet exceeded maximum message size.",
|
|
)
|
|
self._purge_closed_clients()
|
|
|
|
def _handle_disconnected(self, event: GnsEvent) -> None:
|
|
client = self._client_for_connection(event.connection_id)
|
|
if client is None:
|
|
return
|
|
if isinstance(client.connection, GnsConnectionAdapter):
|
|
client.connection.mark_remote_closed()
|
|
self.server._disconnect_client(client)
|
|
self._remove_mapping(event.connection_id)
|
|
|
|
def _client_for_connection(self, connection_id: int) -> ClientSession | None:
|
|
with self._lock:
|
|
return self._clients.get(connection_id)
|
|
|
|
def _remove_mapping(self, connection_id: int) -> None:
|
|
with self._lock:
|
|
self._clients.pop(connection_id, None)
|
|
self._connected_monotonic.pop(connection_id, None)
|
|
stale_windows = [key for key in self._incoming_snapshot_windows if key[0] == connection_id]
|
|
for key in stale_windows:
|
|
self._incoming_snapshot_windows.pop(key, None)
|
|
|
|
def _purge_closed_clients(self) -> None:
|
|
with self._lock:
|
|
stale = [
|
|
connection_id
|
|
for connection_id, client in self._clients.items()
|
|
if client.connection.fileno() < 0
|
|
]
|
|
for connection_id in stale:
|
|
self._remove_mapping(connection_id)
|
|
|
|
def _enforce_timeouts(self) -> None:
|
|
from server_core import CLIENT_HANDSHAKE_TIMEOUT_SECONDS, CLIENT_IDLE_TIMEOUT_SECONDS
|
|
|
|
now_mono = time.monotonic()
|
|
now_wall = time.time()
|
|
with self._lock:
|
|
snapshot = [
|
|
(connection_id, client, self._connected_monotonic.get(connection_id, now_mono))
|
|
for connection_id, client in self._clients.items()
|
|
]
|
|
for connection_id, client, connected_mono in snapshot:
|
|
if not client.gameplay_active and now_mono - connected_mono > CLIENT_HANDSHAKE_TIMEOUT_SECONDS:
|
|
self.transport.disconnect(connection_id, debug="Handshake timeout")
|
|
if isinstance(client.connection, GnsConnectionAdapter):
|
|
client.connection.mark_remote_closed()
|
|
self.server._disconnect_client(client)
|
|
self._remove_mapping(connection_id)
|
|
continue
|
|
if (
|
|
client.gameplay_active
|
|
and client.last_packet_at is not None
|
|
and now_wall - client.last_packet_at > CLIENT_IDLE_TIMEOUT_SECONDS
|
|
):
|
|
self.transport.disconnect(connection_id, debug="Idle timeout")
|
|
if isinstance(client.connection, GnsConnectionAdapter):
|
|
client.connection.mark_remote_closed()
|
|
self.server._disconnect_client(client)
|
|
self._remove_mapping(connection_id)
|
|
|
|
def pump_once(self, max_events: int = 128) -> int:
|
|
processed = 0
|
|
while processed < max_events:
|
|
event = self.transport.poll()
|
|
if event is None:
|
|
break
|
|
processed += 1
|
|
if event.type is EventType.CONNECTED:
|
|
self._handle_connected(event)
|
|
elif event.type is EventType.MESSAGE:
|
|
self._handle_message(event)
|
|
elif event.type is EventType.OVERSIZE_MESSAGE:
|
|
self._handle_oversize(event)
|
|
elif event.type is EventType.DISCONNECTED:
|
|
self._handle_disconnected(event)
|
|
self._enforce_timeouts()
|
|
self._purge_closed_clients()
|
|
return processed
|
|
|
|
def _run(self) -> None:
|
|
while True:
|
|
with self._lock:
|
|
if not self._thread_running:
|
|
break
|
|
try:
|
|
processed = self.pump_once()
|
|
except GnsTransportError as error:
|
|
self.server._log(f"GNS transport loop stopped: {error}", level="error")
|
|
with self._lock:
|
|
self._thread_running = False
|
|
break
|
|
if processed == 0:
|
|
time.sleep(0.002)
|
|
|
|
def start(self) -> None:
|
|
with self._lock:
|
|
if self._thread_running:
|
|
return
|
|
self._thread_running = True
|
|
self._thread = threading.Thread(target=self._run, daemon=True, name="CommonwealthOnlineGNS")
|
|
self._thread.start()
|
|
|
|
def stop(self) -> None:
|
|
with self._lock:
|
|
self._thread_running = False
|
|
thread = self._thread
|
|
self._thread = None
|
|
clients = list(self._clients.items())
|
|
if thread is not None and thread is not threading.current_thread():
|
|
thread.join(timeout=1.0)
|
|
for connection_id, client in clients:
|
|
if isinstance(client.connection, GnsConnectionAdapter):
|
|
client.connection.mark_remote_closed()
|
|
try:
|
|
self.transport.disconnect(connection_id, debug="GNS adapter stopped")
|
|
except Exception:
|
|
pass
|
|
self.server._disconnect_client(client)
|
|
with self._lock:
|
|
self._clients.clear()
|
|
self._connected_monotonic.clear()
|
|
self._incoming_snapshot_windows.clear()
|
|
self.transport.close()
|