327 lines
10 KiB
Python
327 lines
10 KiB
Python
from __future__ import annotations
|
|
|
|
import ctypes
|
|
import ipaddress
|
|
import os
|
|
import sys
|
|
from dataclasses import dataclass
|
|
from enum import IntEnum
|
|
from pathlib import Path
|
|
from typing import Any
|
|
|
|
from packet_codec import EncodedPacket, MAX_MESSAGE_BYTES, encode_packet
|
|
from transport_policy import Delivery
|
|
|
|
|
|
class GnsTransportError(RuntimeError):
|
|
pass
|
|
|
|
|
|
class EventType(IntEnum):
|
|
CONNECTED = 1
|
|
DISCONNECTED = 2
|
|
MESSAGE = 3
|
|
OVERSIZE_MESSAGE = 4
|
|
|
|
|
|
class SendResult(IntEnum):
|
|
ERROR = -1
|
|
SENT = 0
|
|
DROPPED = 1
|
|
BACKPRESSURE = 2
|
|
NOT_CONNECTED = 3
|
|
TOO_LARGE = 4
|
|
|
|
|
|
_NATIVE_DELIVERY = {
|
|
Delivery.UNRELIABLE_SEQUENCED: 0,
|
|
Delivery.RELIABLE_ORDERED: 1,
|
|
}
|
|
|
|
|
|
class _CEvent(ctypes.Structure):
|
|
_fields_ = [
|
|
("type", ctypes.c_uint32),
|
|
("connection_id", ctypes.c_uint32),
|
|
("reason", ctypes.c_int32),
|
|
("payload_size", ctypes.c_uint32),
|
|
("debug", ctypes.c_char * 128),
|
|
]
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class GnsEvent:
|
|
type: EventType
|
|
connection_id: int
|
|
payload: bytes = b""
|
|
reason: int = 0
|
|
debug: str = ""
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class RemoteEndpoint:
|
|
host: str
|
|
port: int
|
|
|
|
|
|
def _default_library_names() -> tuple[str, ...]:
|
|
if sys.platform == "win32":
|
|
return ("commonwealth_online_gns_bridge.dll",)
|
|
if sys.platform == "darwin":
|
|
return ("libcommonwealth_online_gns_bridge.dylib",)
|
|
return ("libcommonwealth_online_gns_bridge.so",)
|
|
|
|
|
|
def _candidate_library_paths() -> list[Path]:
|
|
candidates: list[Path] = []
|
|
configured = os.environ.get("COMMONWEALTH_ONLINE_GNS_BRIDGE", "").strip()
|
|
if configured:
|
|
candidates.append(Path(configured).expanduser())
|
|
root = Path(__file__).resolve().parent
|
|
for name in _default_library_names():
|
|
candidates.extend(
|
|
(
|
|
root / name,
|
|
root / "native_transport" / name,
|
|
root.parent / name,
|
|
)
|
|
)
|
|
return candidates
|
|
|
|
|
|
def find_native_bridge() -> Path:
|
|
for candidate in _candidate_library_paths():
|
|
if candidate.is_file():
|
|
return candidate
|
|
searched = ", ".join(str(path) for path in _candidate_library_paths())
|
|
raise GnsTransportError(
|
|
"Commonwealth Online GNS native bridge was not found. "
|
|
f"Set COMMONWEALTH_ONLINE_GNS_BRIDGE or install it beside the server. Searched: {searched}"
|
|
)
|
|
|
|
|
|
class _NativeApi:
|
|
def __init__(self, library: Any) -> None:
|
|
self.library = library
|
|
self.create = library.co_gns_server_create
|
|
self.destroy = library.co_gns_server_destroy
|
|
self.local_port = library.co_gns_server_local_port
|
|
self.connection_count = library.co_gns_server_connection_count
|
|
self.poll = library.co_gns_server_poll
|
|
self.send = library.co_gns_server_send
|
|
self.disconnect = library.co_gns_server_disconnect
|
|
self.remote_ipv4 = library.co_gns_server_remote_ipv4
|
|
|
|
self.create.argtypes = [
|
|
ctypes.c_char_p,
|
|
ctypes.c_uint16,
|
|
ctypes.POINTER(ctypes.c_void_p),
|
|
ctypes.c_char_p,
|
|
ctypes.c_size_t,
|
|
]
|
|
self.create.restype = ctypes.c_int
|
|
self.destroy.argtypes = [ctypes.c_void_p]
|
|
self.destroy.restype = None
|
|
self.local_port.argtypes = [ctypes.c_void_p]
|
|
self.local_port.restype = ctypes.c_uint16
|
|
self.connection_count.argtypes = [ctypes.c_void_p]
|
|
self.connection_count.restype = ctypes.c_uint32
|
|
self.poll.argtypes = [
|
|
ctypes.c_void_p,
|
|
ctypes.POINTER(_CEvent),
|
|
ctypes.c_void_p,
|
|
ctypes.c_uint32,
|
|
]
|
|
self.poll.restype = ctypes.c_int
|
|
self.send.argtypes = [
|
|
ctypes.c_void_p,
|
|
ctypes.c_uint32,
|
|
ctypes.c_void_p,
|
|
ctypes.c_uint32,
|
|
ctypes.c_uint32,
|
|
]
|
|
self.send.restype = ctypes.c_int
|
|
self.disconnect.argtypes = [
|
|
ctypes.c_void_p,
|
|
ctypes.c_uint32,
|
|
ctypes.c_int32,
|
|
ctypes.c_char_p,
|
|
]
|
|
self.disconnect.restype = ctypes.c_int
|
|
self.remote_ipv4.argtypes = [
|
|
ctypes.c_void_p,
|
|
ctypes.c_uint32,
|
|
ctypes.POINTER(ctypes.c_uint32),
|
|
ctypes.POINTER(ctypes.c_uint16),
|
|
]
|
|
self.remote_ipv4.restype = ctypes.c_int
|
|
|
|
|
|
class GnsServerTransport:
|
|
def __init__(
|
|
self,
|
|
bind_host: str,
|
|
port: int,
|
|
*,
|
|
library_path: str | os.PathLike[str] | None = None,
|
|
native_library: Any | None = None,
|
|
) -> None:
|
|
if not isinstance(bind_host, str) or not bind_host:
|
|
raise ValueError("bind_host must be a non-empty string")
|
|
if not isinstance(port, int) or isinstance(port, bool) or not 0 <= port <= 65535:
|
|
raise ValueError("port must be between 0 and 65535")
|
|
|
|
if native_library is None:
|
|
path = Path(library_path).expanduser() if library_path is not None else find_native_bridge()
|
|
try:
|
|
native_library = ctypes.CDLL(str(path))
|
|
except OSError as error:
|
|
raise GnsTransportError(f"Could not load GNS native bridge at {path}: {error}") from error
|
|
|
|
self._api = _NativeApi(native_library)
|
|
self._handle = ctypes.c_void_p()
|
|
error_buffer = ctypes.create_string_buffer(512)
|
|
created = self._api.create(
|
|
bind_host.encode("utf-8"),
|
|
port,
|
|
ctypes.byref(self._handle),
|
|
error_buffer,
|
|
len(error_buffer),
|
|
)
|
|
if created != 1 or not self._handle.value:
|
|
message = error_buffer.value.decode("utf-8", errors="replace").strip()
|
|
raise GnsTransportError(message or "GNS native bridge failed to start")
|
|
self._closed = False
|
|
self._payload_buffer = ctypes.create_string_buffer(MAX_MESSAGE_BYTES)
|
|
|
|
@property
|
|
def local_port(self) -> int:
|
|
self._require_open()
|
|
return int(self._api.local_port(self._handle))
|
|
|
|
@property
|
|
def connection_count(self) -> int:
|
|
self._require_open()
|
|
return int(self._api.connection_count(self._handle))
|
|
|
|
@property
|
|
def is_closed(self) -> bool:
|
|
return self._closed
|
|
|
|
def _require_open(self) -> None:
|
|
if self._closed or not self._handle.value:
|
|
raise GnsTransportError("GNS transport is closed")
|
|
|
|
def remote_endpoint(self, connection_id: int) -> RemoteEndpoint | None:
|
|
self._require_open()
|
|
if not isinstance(connection_id, int) or isinstance(connection_id, bool) or connection_id <= 0:
|
|
raise ValueError("connection_id must be a positive integer")
|
|
ipv4 = ctypes.c_uint32()
|
|
port = ctypes.c_uint16()
|
|
found = int(
|
|
self._api.remote_ipv4(
|
|
self._handle,
|
|
connection_id,
|
|
ctypes.byref(ipv4),
|
|
ctypes.byref(port),
|
|
)
|
|
)
|
|
if found != 1:
|
|
return None
|
|
return RemoteEndpoint(str(ipaddress.IPv4Address(ipv4.value)), int(port.value))
|
|
|
|
def poll(self) -> GnsEvent | None:
|
|
self._require_open()
|
|
event = _CEvent()
|
|
result = int(
|
|
self._api.poll(
|
|
self._handle,
|
|
ctypes.byref(event),
|
|
self._payload_buffer,
|
|
MAX_MESSAGE_BYTES,
|
|
)
|
|
)
|
|
if result == 0:
|
|
return None
|
|
if result < 0:
|
|
raise GnsTransportError(f"GNS native poll failed with result {result}")
|
|
try:
|
|
event_type = EventType(event.type)
|
|
except ValueError as error:
|
|
raise GnsTransportError(f"GNS native bridge returned unknown event type {event.type}") from error
|
|
if event.payload_size > MAX_MESSAGE_BYTES and event_type is not EventType.OVERSIZE_MESSAGE:
|
|
raise GnsTransportError("GNS native bridge returned an oversized message payload")
|
|
payload = b""
|
|
if event_type is EventType.MESSAGE and event.payload_size:
|
|
payload = self._payload_buffer.raw[: event.payload_size]
|
|
debug = bytes(event.debug).split(b"\0", 1)[0].decode("utf-8", errors="replace")
|
|
return GnsEvent(
|
|
type=event_type,
|
|
connection_id=int(event.connection_id),
|
|
payload=payload,
|
|
reason=int(event.reason),
|
|
debug=debug,
|
|
)
|
|
|
|
def send_encoded(self, connection_id: int, encoded: EncodedPacket) -> SendResult:
|
|
self._require_open()
|
|
if not isinstance(connection_id, int) or isinstance(connection_id, bool) or connection_id <= 0:
|
|
raise ValueError("connection_id must be a positive integer")
|
|
payload = encoded.payload
|
|
if len(payload) > MAX_MESSAGE_BYTES:
|
|
return SendResult.TOO_LARGE
|
|
buffer = ctypes.create_string_buffer(payload, len(payload))
|
|
raw_result = int(
|
|
self._api.send(
|
|
self._handle,
|
|
connection_id,
|
|
buffer,
|
|
len(payload),
|
|
_NATIVE_DELIVERY[encoded.delivery],
|
|
)
|
|
)
|
|
try:
|
|
return SendResult(raw_result)
|
|
except ValueError as error:
|
|
raise GnsTransportError(f"GNS native send returned unknown result {raw_result}") from error
|
|
|
|
def send_packet(self, connection_id: int, packet: dict[str, Any]) -> SendResult:
|
|
return self.send_encoded(connection_id, encode_packet(packet))
|
|
|
|
def disconnect(self, connection_id: int, *, reason: int = 0, debug: str = "") -> bool:
|
|
self._require_open()
|
|
if not isinstance(connection_id, int) or isinstance(connection_id, bool) or connection_id <= 0:
|
|
raise ValueError("connection_id must be a positive integer")
|
|
if not isinstance(reason, int) or isinstance(reason, bool):
|
|
raise ValueError("reason must be an integer")
|
|
clean_debug = str(debug or "")[:127]
|
|
return bool(
|
|
self._api.disconnect(
|
|
self._handle,
|
|
connection_id,
|
|
reason,
|
|
clean_debug.encode("utf-8", errors="replace"),
|
|
)
|
|
)
|
|
|
|
def close(self) -> None:
|
|
if self._closed:
|
|
return
|
|
self._closed = True
|
|
if self._handle.value:
|
|
self._api.destroy(self._handle)
|
|
self._handle = ctypes.c_void_p()
|
|
|
|
def __enter__(self) -> "GnsServerTransport":
|
|
self._require_open()
|
|
return self
|
|
|
|
def __exit__(self, exc_type, exc, traceback) -> None:
|
|
self.close()
|
|
|
|
def __del__(self) -> None:
|
|
try:
|
|
self.close()
|
|
except Exception:
|
|
pass
|