Compare commits
3
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
54f8a06862 | ||
|
|
db7fb8e09a | ||
|
|
25953bea06 |
@@ -1,16 +1,17 @@
|
|||||||
#!/usr/bin/env python3
|
#!/usr/bin/env python3
|
||||||
import ast
|
import ast
|
||||||
import inspect
|
import inspect
|
||||||
import sys
|
|
||||||
import os
|
import os
|
||||||
|
import sys
|
||||||
|
|
||||||
# Add lib and tests to path
|
# Add lib and tests to path
|
||||||
sys.path.insert(0, os.path.join(os.path.dirname(__file__), 'lib'))
|
sys.path.insert(0, os.path.join(os.path.dirname(__file__), 'lib'))
|
||||||
sys.path.insert(0, os.path.dirname(__file__))
|
sys.path.insert(0, os.path.dirname(__file__))
|
||||||
|
|
||||||
from openttd import OpenTTDClient, OpenTTDAdminClient
|
from openttd import OpenTTDAdminClient, OpenTTDClient
|
||||||
from openttd.protocol import OpenTTDProtocol, OpenTTDAdminProtocol
|
from openttd.protocol import OpenTTDAdminProtocol, OpenTTDProtocol
|
||||||
import tests.test_e2e as test_e2e
|
|
||||||
|
from tests import test_e2e
|
||||||
|
|
||||||
# 1. Gather public functions dynamically at runtime using reflection
|
# 1. Gather public functions dynamically at runtime using reflection
|
||||||
classes = [OpenTTDClient, OpenTTDAdminClient, OpenTTDProtocol, OpenTTDAdminProtocol]
|
classes = [OpenTTDClient, OpenTTDAdminClient, OpenTTDProtocol, OpenTTDAdminProtocol]
|
||||||
|
|||||||
@@ -1,4 +1,4 @@
|
|||||||
|
from .client import OpenTTDAdminClient, OpenTTDClient
|
||||||
from .decorators import exclude_call_check
|
from .decorators import exclude_call_check
|
||||||
from .client import OpenTTDClient, OpenTTDAdminClient
|
|
||||||
|
|
||||||
__all__ = ['OpenTTDClient', 'OpenTTDAdminClient', 'exclude_call_check']
|
__all__ = ['OpenTTDAdminClient', 'OpenTTDClient', 'exclude_call_check']
|
||||||
|
|||||||
+46
-17
@@ -1,18 +1,42 @@
|
|||||||
import asyncio
|
import asyncio
|
||||||
import logging
|
|
||||||
import uuid
|
|
||||||
import monocypher
|
|
||||||
import os
|
|
||||||
import hashlib
|
import hashlib
|
||||||
from openttd_protocol.wire.write import write_init, write_string, write_uint8, write_uint16, write_uint32, write_presend, SEND_TCP_MTU
|
import logging
|
||||||
|
import os
|
||||||
|
import uuid
|
||||||
|
from typing import ClassVar
|
||||||
|
|
||||||
|
import monocypher
|
||||||
|
from openttd_protocol.wire.exceptions import SocketClosed
|
||||||
from openttd_protocol.wire.read import read_uint8, read_uint16
|
from openttd_protocol.wire.read import read_uint8, read_uint16
|
||||||
from .protocol import (
|
from openttd_protocol.wire.write import (
|
||||||
PacketGameType, OpenTTDProtocol, PacketAdminType, OpenTTDAdminProtocol, NetworkAuthenticationMethod,
|
SEND_TCP_MTU,
|
||||||
GameCommand, ModifyTimetableFlags, ModifyTimetableCtrlFlag,
|
write_init,
|
||||||
OrderType, OrderStopLocation, INVALID_VEH_ORDER_ID,
|
write_presend,
|
||||||
write_varuint, read_varuint, write_varuint_signed, read_varuint_signed
|
write_string,
|
||||||
|
write_uint8,
|
||||||
|
write_uint16,
|
||||||
|
write_uint32,
|
||||||
)
|
)
|
||||||
|
|
||||||
from .decorators import exclude_call_check
|
from .decorators import exclude_call_check
|
||||||
|
from .protocol import (
|
||||||
|
INVALID_VEH_ORDER_ID,
|
||||||
|
GameCommand,
|
||||||
|
ModifyTimetableCtrlFlag,
|
||||||
|
ModifyTimetableFlags,
|
||||||
|
NetworkAuthenticationMethod,
|
||||||
|
OpenTTDAdminProtocol,
|
||||||
|
OpenTTDProtocol,
|
||||||
|
OrderStopLocation,
|
||||||
|
OrderType,
|
||||||
|
PacketAdminType,
|
||||||
|
PacketGameType,
|
||||||
|
read_varuint,
|
||||||
|
read_varuint_signed,
|
||||||
|
write_varuint,
|
||||||
|
write_varuint_signed,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class OpenTTDClient:
|
class OpenTTDClient:
|
||||||
"""High-level OpenTTD client for easy integration."""
|
"""High-level OpenTTD client for easy integration."""
|
||||||
@@ -261,8 +285,9 @@ class OpenTTDClient:
|
|||||||
try:
|
try:
|
||||||
d = write_init(PacketGameType.ClientQuit)
|
d = write_init(PacketGameType.ClientQuit)
|
||||||
await self._protocol.send_packet(write_presend(d, SEND_TCP_MTU))
|
await self._protocol.send_packet(write_presend(d, SEND_TCP_MTU))
|
||||||
except Exception:
|
except (OSError, SocketClosed) as e:
|
||||||
pass
|
# Best-effort courtesy packet: the socket may already be gone.
|
||||||
|
self.log.debug(f"Could not send quit packet: {e}")
|
||||||
self._transport.close()
|
self._transport.close()
|
||||||
self.shutdown_event.set()
|
self.shutdown_event.set()
|
||||||
|
|
||||||
@@ -386,7 +411,7 @@ class OpenTTDClient:
|
|||||||
async def receive_ServerMapData(self, source, **kwargs): pass
|
async def receive_ServerMapData(self, source, **kwargs): pass
|
||||||
async def receive_ServerConfigurationUpdate(self, source, **kwargs): pass
|
async def receive_ServerConfigurationUpdate(self, source, **kwargs): pass
|
||||||
async def receive_ServerExternalChat(self, source, **kwargs): pass
|
async def receive_ServerExternalChat(self, source, **kwargs): pass
|
||||||
_TIMETABLE_FIELD_BY_FLAG = {
|
_TIMETABLE_FIELD_BY_FLAG: ClassVar[dict[ModifyTimetableFlags, str]] = {
|
||||||
ModifyTimetableFlags.WaitTime: "wait_time",
|
ModifyTimetableFlags.WaitTime: "wait_time",
|
||||||
ModifyTimetableFlags.TravelTime: "travel_time",
|
ModifyTimetableFlags.TravelTime: "travel_time",
|
||||||
ModifyTimetableFlags.TravelSpeed: "travel_speed",
|
ModifyTimetableFlags.TravelSpeed: "travel_speed",
|
||||||
@@ -395,7 +420,10 @@ class OpenTTDClient:
|
|||||||
ModifyTimetableFlags.SetLeaveType: "leave_type",
|
ModifyTimetableFlags.SetLeaveType: "leave_type",
|
||||||
ModifyTimetableFlags.AssignSchedule: "assigned_schedule",
|
ModifyTimetableFlags.AssignSchedule: "assigned_schedule",
|
||||||
}
|
}
|
||||||
_TIMETABLE_BOOL_FLAGS = {ModifyTimetableFlags.SetWaitFixed, ModifyTimetableFlags.SetTravelFixed}
|
_TIMETABLE_BOOL_FLAGS: ClassVar[set[ModifyTimetableFlags]] = {
|
||||||
|
ModifyTimetableFlags.SetWaitFixed,
|
||||||
|
ModifyTimetableFlags.SetTravelFixed,
|
||||||
|
}
|
||||||
|
|
||||||
async def receive_ServerCommand(self, source, cmd, payload, **kwargs):
|
async def receive_ServerCommand(self, source, cmd, payload, **kwargs):
|
||||||
if cmd == GameCommand.ChangeTimetable:
|
if cmd == GameCommand.ChangeTimetable:
|
||||||
@@ -512,8 +540,9 @@ class OpenTTDAdminClient:
|
|||||||
try:
|
try:
|
||||||
d = write_init(PacketAdminType.AdminQuit)
|
d = write_init(PacketAdminType.AdminQuit)
|
||||||
await self._protocol.send_packet(write_presend(d, SEND_TCP_MTU))
|
await self._protocol.send_packet(write_presend(d, SEND_TCP_MTU))
|
||||||
except Exception:
|
except (OSError, SocketClosed) as e:
|
||||||
pass
|
# Best-effort courtesy packet: the socket may already be gone.
|
||||||
|
self.log.debug(f"Could not send admin quit packet: {e}")
|
||||||
self._transport.close()
|
self._transport.close()
|
||||||
self.shutdown_event.set()
|
self.shutdown_event.set()
|
||||||
|
|
||||||
@@ -597,7 +626,7 @@ class OpenTTDAdminClient:
|
|||||||
Raises asyncio.TimeoutError if no reply arrives within `timeout`, ValueError on an error
|
Raises asyncio.TimeoutError if no reply arrives within `timeout`, ValueError on an error
|
||||||
reply, and ConnectionError if the admin connection drops while waiting.
|
reply, and ConnectionError if the admin connection drops while waiting.
|
||||||
"""
|
"""
|
||||||
from .protocol import AdminUpdateType, AdminUpdateFrequency
|
from .protocol import AdminUpdateFrequency, AdminUpdateType
|
||||||
if not self._gs_subscribed:
|
if not self._gs_subscribed:
|
||||||
await self.update_frequency(AdminUpdateType.Gamescript, AdminUpdateFrequency.Automatic)
|
await self.update_frequency(AdminUpdateType.Gamescript, AdminUpdateFrequency.Automatic)
|
||||||
self._gs_subscribed = True
|
self._gs_subscribed = True
|
||||||
|
|||||||
+10
-8
@@ -1,9 +1,11 @@
|
|||||||
import struct
|
import struct
|
||||||
import monocypher
|
|
||||||
from enum import IntEnum
|
from enum import IntEnum
|
||||||
from openttd_protocol.wire.tcp import TCPProtocol
|
|
||||||
from openttd_protocol.wire.read import read_uint8, read_string, read_uint16, read_uint32
|
import monocypher
|
||||||
from openttd_protocol.wire.exceptions import SocketClosed
|
from openttd_protocol.wire.exceptions import SocketClosed
|
||||||
|
from openttd_protocol.wire.read import read_string, read_uint8, read_uint16, read_uint32
|
||||||
|
from openttd_protocol.wire.tcp import TCPProtocol
|
||||||
|
|
||||||
|
|
||||||
def write_varuint(buffer, value):
|
def write_varuint(buffer, value):
|
||||||
"""Encode a non-negative integer using OpenTTD's UTF-8-like varuint scheme."""
|
"""Encode a non-negative integer using OpenTTD's UTF-8-like varuint scheme."""
|
||||||
@@ -232,21 +234,21 @@ class OpenTTDProtocol(TCPProtocol):
|
|||||||
if self.handler.encryption_enabled:
|
if self.handler.encryption_enabled:
|
||||||
if not self.handler._recv_aead:
|
if not self.handler._recv_aead:
|
||||||
self.handler._recv_aead = monocypher.IncrementalAuthenticatedEncryption(self.handler._session_key_recv, self.handler._encryption_nonce)
|
self.handler._recv_aead = monocypher.IncrementalAuthenticatedEncryption(self.handler._session_key_recv, self.handler._encryption_nonce)
|
||||||
length, rest = read_uint16(data)
|
_, rest = read_uint16(data)
|
||||||
payload = self.handler._recv_aead.unlock(bytes(rest[:16]), bytes(rest[16:]))
|
payload = self.handler._recv_aead.unlock(bytes(rest[:16]), bytes(rest[16:]))
|
||||||
if payload is None:
|
if payload is None:
|
||||||
raise SocketClosed("Decryption failed")
|
raise SocketClosed("Decryption failed")
|
||||||
data = memoryview(struct.pack("<H", len(payload) + 2) + payload)
|
data = memoryview(struct.pack("<H", len(payload) + 2) + payload)
|
||||||
|
|
||||||
return super().receive_packet(source, data)
|
return super().receive_packet(source, data)
|
||||||
except Exception:
|
except Exception: # noqa: BLE001 - untrusted wire data: any decode failure must degrade to a no-op packet rather than kill the connection
|
||||||
return PacketGameType.ServerUnused, {}
|
return PacketGameType.ServerUnused, {}
|
||||||
|
|
||||||
async def send_packet(self, data):
|
async def send_packet(self, data):
|
||||||
if self.handler.encryption_enabled:
|
if self.handler.encryption_enabled:
|
||||||
if not self.handler._send_aead:
|
if not self.handler._send_aead:
|
||||||
self.handler._send_aead = monocypher.IncrementalAuthenticatedEncryption(self.handler._session_key_send, self.handler._encryption_nonce)
|
self.handler._send_aead = monocypher.IncrementalAuthenticatedEncryption(self.handler._session_key_send, self.handler._encryption_nonce)
|
||||||
length, payload = read_uint16(memoryview(data))
|
_, payload = read_uint16(memoryview(data))
|
||||||
mac, ciphertext = self.handler._send_aead.lock(payload.tobytes())
|
mac, ciphertext = self.handler._send_aead.lock(payload.tobytes())
|
||||||
data = struct.pack("<H", 18 + len(ciphertext)) + mac + ciphertext
|
data = struct.pack("<H", 18 + len(ciphertext)) + mac + ciphertext
|
||||||
|
|
||||||
@@ -287,7 +289,7 @@ class OpenTTDProtocol(TCPProtocol):
|
|||||||
@staticmethod
|
@staticmethod
|
||||||
def receive_ServerFrame(source, data):
|
def receive_ServerFrame(source, data):
|
||||||
f, data = read_uint32(data)
|
f, data = read_uint32(data)
|
||||||
max_f, data = read_uint32(data)
|
_, data = read_uint32(data)
|
||||||
token = 0
|
token = 0
|
||||||
if len(data) > 0:
|
if len(data) > 0:
|
||||||
if len(data) >= 13:
|
if len(data) >= 13:
|
||||||
@@ -575,7 +577,7 @@ class OpenTTDAdminProtocol(OpenTTDProtocol):
|
|||||||
import json
|
import json
|
||||||
try:
|
try:
|
||||||
return {"data": json.loads(json_str)}
|
return {"data": json.loads(json_str)}
|
||||||
except Exception:
|
except json.JSONDecodeError:
|
||||||
return {"raw_data": json_str}
|
return {"raw_data": json_str}
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
|
|||||||
@@ -1,7 +1,7 @@
|
|||||||
import asyncio
|
import asyncio
|
||||||
import logging
|
import logging
|
||||||
import sys
|
|
||||||
import os
|
import os
|
||||||
|
import sys
|
||||||
|
|
||||||
# Add the lib directory to sys.path so we can import the openttd package
|
# Add the lib directory to sys.path so we can import the openttd package
|
||||||
sys.path.append(os.path.join(os.path.dirname(__file__), 'lib'))
|
sys.path.append(os.path.join(os.path.dirname(__file__), 'lib'))
|
||||||
@@ -171,7 +171,7 @@ async def run_client():
|
|||||||
print("--- Finished 10s stay, exiting gracefully ---")
|
print("--- Finished 10s stay, exiting gracefully ---")
|
||||||
await client.quit()
|
await client.quit()
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e: # noqa: BLE001 - top-level demo handler: report any failure instead of dumping a traceback
|
||||||
print(f"!!! Error: {e}")
|
print(f"!!! Error: {e}")
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
|
|||||||
+4
-4
@@ -1,13 +1,13 @@
|
|||||||
import asyncio
|
import asyncio
|
||||||
import logging
|
import logging
|
||||||
import sys
|
|
||||||
import os
|
import os
|
||||||
|
import sys
|
||||||
|
|
||||||
# Add the lib directory to sys.path so we can import the openttd package
|
# Add the lib directory to sys.path so we can import the openttd package
|
||||||
sys.path.append(os.path.join(os.path.dirname(__file__), 'lib'))
|
sys.path.append(os.path.join(os.path.dirname(__file__), 'lib'))
|
||||||
|
|
||||||
from openttd import OpenTTDAdminClient
|
from openttd import OpenTTDAdminClient
|
||||||
from openttd.protocol import AdminUpdateType, AdminUpdateFrequency
|
from openttd.protocol import AdminUpdateFrequency, AdminUpdateType
|
||||||
|
|
||||||
# Configuration
|
# Configuration
|
||||||
SERVER_HOST = "127.0.0.1"
|
SERVER_HOST = "127.0.0.1"
|
||||||
@@ -82,14 +82,14 @@ async def run_admin():
|
|||||||
print(f" waiting by next hop: {flow['waiting_by_via']}")
|
print(f" waiting by next hop: {flow['waiting_by_via']}")
|
||||||
print(f" planned by source: {flow['planned_by_from']}")
|
print(f" planned by source: {flow['planned_by_from']}")
|
||||||
print(f" planned by next hop: {flow['planned_by_via']}")
|
print(f" planned by next hop: {flow['planned_by_via']}")
|
||||||
except Exception as e:
|
except Exception as e: # noqa: BLE001 - demo script: one failed station query should not abort the walk
|
||||||
print(f"!!! station query failed: {e}")
|
print(f"!!! station query failed: {e}")
|
||||||
|
|
||||||
await asyncio.sleep(5)
|
await asyncio.sleep(5)
|
||||||
print("--- Quitting ---")
|
print("--- Quitting ---")
|
||||||
await admin.quit()
|
await admin.quit()
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e: # noqa: BLE001 - top-level demo handler: report any failure instead of dumping a traceback
|
||||||
print(f"!!! Error: {e}")
|
print(f"!!! Error: {e}")
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
|
|||||||
+3
-1
@@ -1,6 +1,8 @@
|
|||||||
import pytest
|
|
||||||
import os
|
import os
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture(scope="session")
|
@pytest.fixture(scope="session")
|
||||||
def server_config():
|
def server_config():
|
||||||
"""Provides server connection parameters from environment variables with defaults."""
|
"""Provides server connection parameters from environment variables with defaults."""
|
||||||
|
|||||||
+9
-4
@@ -1,10 +1,15 @@
|
|||||||
import pytest
|
|
||||||
import asyncio
|
import asyncio
|
||||||
import json
|
import json
|
||||||
import os
|
import os
|
||||||
|
|
||||||
import monocypher
|
import monocypher
|
||||||
|
import pytest
|
||||||
from openttd import OpenTTDAdminClient
|
from openttd import OpenTTDAdminClient
|
||||||
from openttd.protocol import PacketAdminType, AdminUpdateType, AdminUpdateFrequency
|
from openttd.protocol import AdminUpdateFrequency, AdminUpdateType, PacketAdminType
|
||||||
|
|
||||||
|
|
||||||
|
class FakeNetworkError(OSError):
|
||||||
|
"""Stand-in for a socket-level failure, so it matches the client's narrowed handlers."""
|
||||||
|
|
||||||
def decode_gamescript_payload(packet):
|
def decode_gamescript_payload(packet):
|
||||||
"""Decode the JSON payload of an AdminGamescript packet (2-byte length + 1-byte type + string)."""
|
"""Decode the JSON payload of an AdminGamescript packet (2-byte length + 1-byte type + string)."""
|
||||||
@@ -84,7 +89,7 @@ async def test_admin_client_connect_and_actions(monkeypatch):
|
|||||||
|
|
||||||
# 3. Connect Exception
|
# 3. Connect Exception
|
||||||
async def mock_fail(*args, **kwargs):
|
async def mock_fail(*args, **kwargs):
|
||||||
raise Exception("Connection Failed")
|
raise FakeNetworkError("Connection Failed")
|
||||||
monkeypatch.setattr(asyncio.get_running_loop(), "create_connection", mock_fail)
|
monkeypatch.setattr(asyncio.get_running_loop(), "create_connection", mock_fail)
|
||||||
with pytest.raises(Exception, match="Connection Failed"):
|
with pytest.raises(Exception, match="Connection Failed"):
|
||||||
await client.connect()
|
await client.connect()
|
||||||
@@ -196,7 +201,7 @@ async def test_admin_client_connect_and_actions(monkeypatch):
|
|||||||
# 8. Quit Exception
|
# 8. Quit Exception
|
||||||
class BadProtocol:
|
class BadProtocol:
|
||||||
async def send_packet(self, data):
|
async def send_packet(self, data):
|
||||||
raise Exception("Fail")
|
raise FakeNetworkError("Fail")
|
||||||
client._transport = MockTransport()
|
client._transport = MockTransport()
|
||||||
client._protocol = BadProtocol()
|
client._protocol = BadProtocol()
|
||||||
await client.quit()
|
await client.quit()
|
||||||
|
|||||||
@@ -1,8 +1,10 @@
|
|||||||
import pytest
|
|
||||||
import asyncio
|
import asyncio
|
||||||
import struct
|
import struct
|
||||||
from openttd.protocol import OpenTTDProtocol, PacketGameType
|
|
||||||
|
import pytest
|
||||||
from openttd.client import OpenTTDClient
|
from openttd.client import OpenTTDClient
|
||||||
|
from openttd.protocol import OpenTTDProtocol, PacketGameType
|
||||||
|
|
||||||
|
|
||||||
class MockTransport:
|
class MockTransport:
|
||||||
def __init__(self): self._closing = False
|
def __init__(self): self._closing = False
|
||||||
@@ -10,16 +12,19 @@ class MockTransport:
|
|||||||
def close(self): self._closing = True
|
def close(self): self._closing = True
|
||||||
def write(self, data): return len(data)
|
def write(self, data): return len(data)
|
||||||
|
|
||||||
|
class FakeNetworkError(OSError):
|
||||||
|
"""Stand-in for a socket-level failure, so it matches the client's narrowed handlers."""
|
||||||
|
|
||||||
class MockProtocol:
|
class MockProtocol:
|
||||||
async def send_packet(self, data):
|
async def send_packet(self, data):
|
||||||
raise Exception("Send failed")
|
raise FakeNetworkError("Send failed")
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_client_connect_exception(monkeypatch):
|
async def test_client_connect_exception(monkeypatch):
|
||||||
# Coverage for client.py:51-53
|
# Coverage for client.py:51-53
|
||||||
client = OpenTTDClient(host="127.0.0.1")
|
client = OpenTTDClient(host="127.0.0.1")
|
||||||
async def mock_fail(*args, **kwargs):
|
async def mock_fail(*args, **kwargs):
|
||||||
raise Exception("Async Failure")
|
raise FakeNetworkError("Async Failure")
|
||||||
monkeypatch.setattr(asyncio.get_running_loop(), "create_connection", mock_fail)
|
monkeypatch.setattr(asyncio.get_running_loop(), "create_connection", mock_fail)
|
||||||
with pytest.raises(Exception, match="Async Failure"):
|
with pytest.raises(Exception, match="Async Failure"):
|
||||||
await client.connect()
|
await client.connect()
|
||||||
|
|||||||
+12
-10
@@ -1,20 +1,22 @@
|
|||||||
import asyncio
|
import asyncio
|
||||||
import pytest
|
|
||||||
import pytest_asyncio
|
|
||||||
import sys
|
|
||||||
import os
|
import os
|
||||||
import random
|
import random
|
||||||
|
import sys
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
import pytest_asyncio
|
||||||
|
|
||||||
# Add lib to path
|
# Add lib to path
|
||||||
sys.path.insert(0, os.path.join(os.path.dirname(__file__), '..', 'lib'))
|
sys.path.insert(0, os.path.join(os.path.dirname(__file__), '..', 'lib'))
|
||||||
|
|
||||||
from openttd import OpenTTDClient, OpenTTDAdminClient
|
from openttd import OpenTTDAdminClient, OpenTTDClient
|
||||||
from openttd.protocol import (
|
from openttd.protocol import (
|
||||||
OpenTTDProtocol,
|
|
||||||
OpenTTDAdminProtocol,
|
|
||||||
AdminUpdateType,
|
|
||||||
AdminUpdateFrequency,
|
AdminUpdateFrequency,
|
||||||
|
AdminUpdateType,
|
||||||
|
ModifyTimetableFlags,
|
||||||
|
OpenTTDAdminProtocol,
|
||||||
|
OpenTTDProtocol,
|
||||||
PacketGameType,
|
PacketGameType,
|
||||||
ModifyTimetableFlags
|
|
||||||
)
|
)
|
||||||
|
|
||||||
# These identify a vehicle/order that already exists in the local dev server's persisted
|
# These identify a vehicle/order that already exists in the local dev server's persisted
|
||||||
@@ -738,11 +740,11 @@ async def test_e2e_protocol_public_functions_multiple_inputs(server_config):
|
|||||||
|
|
||||||
# 2. Test receive_packet() with multiple inputs
|
# 2. Test receive_packet() with multiple inputs
|
||||||
# Input 1: Valid packet structure (length >= 3)
|
# Input 1: Valid packet structure (length >= 3)
|
||||||
res_type1, res_data1 = proto_game.receive_packet(None, memoryview(b"\x03\x00\x05")) # type 5 is ServerUnused
|
res_type1, _ = proto_game.receive_packet(None, memoryview(b"\x03\x00\x05")) # type 5 is ServerUnused
|
||||||
assert res_type1 == PacketGameType.ServerUnused
|
assert res_type1 == PacketGameType.ServerUnused
|
||||||
|
|
||||||
# Input 2: Invalid/short packet structure (length < 3)
|
# Input 2: Invalid/short packet structure (length < 3)
|
||||||
res_type2, res_data2 = proto_game.receive_packet(None, memoryview(b"\x01"))
|
res_type2, _ = proto_game.receive_packet(None, memoryview(b"\x01"))
|
||||||
assert res_type2 == PacketGameType.ServerUnused # Falls back to ServerUnused on error
|
assert res_type2 == PacketGameType.ServerUnused # Falls back to ServerUnused on error
|
||||||
|
|
||||||
# 3. Test send_packet() with multiple inputs (using mock transports)
|
# 3. Test send_packet() with multiple inputs (using mock transports)
|
||||||
|
|||||||
+8
-3
@@ -1,8 +1,13 @@
|
|||||||
import pytest
|
|
||||||
import asyncio
|
import asyncio
|
||||||
import hashlib
|
import hashlib
|
||||||
from openttd.protocol import PacketGameType
|
|
||||||
|
import pytest
|
||||||
from openttd.client import OpenTTDClient
|
from openttd.client import OpenTTDClient
|
||||||
|
from openttd.protocol import PacketGameType
|
||||||
|
|
||||||
|
|
||||||
|
class FakeNetworkError(OSError):
|
||||||
|
"""Stand-in for a socket-level failure, so it matches the client's narrowed handlers."""
|
||||||
|
|
||||||
def test_packet_game_type_values():
|
def test_packet_game_type_values():
|
||||||
assert PacketGameType.ServerFull == 0
|
assert PacketGameType.ServerFull == 0
|
||||||
@@ -61,7 +66,7 @@ async def test_client_connect_success(monkeypatch):
|
|||||||
async def test_client_connect_failure(monkeypatch):
|
async def test_client_connect_failure(monkeypatch):
|
||||||
client = OpenTTDClient(host="127.0.0.1")
|
client = OpenTTDClient(host="127.0.0.1")
|
||||||
async def mock_fail(*args, **kwargs):
|
async def mock_fail(*args, **kwargs):
|
||||||
raise Exception("Async Failure")
|
raise FakeNetworkError("Async Failure")
|
||||||
monkeypatch.setattr(asyncio.get_running_loop(), "create_connection", mock_fail)
|
monkeypatch.setattr(asyncio.get_running_loop(), "create_connection", mock_fail)
|
||||||
with pytest.raises(Exception, match="Async Failure"):
|
with pytest.raises(Exception, match="Async Failure"):
|
||||||
await client.connect()
|
await client.connect()
|
||||||
|
|||||||
@@ -1,9 +1,12 @@
|
|||||||
import pytest
|
import contextlib
|
||||||
import struct
|
import struct
|
||||||
|
|
||||||
import monocypher
|
import monocypher
|
||||||
|
import pytest
|
||||||
from openttd.protocol import OpenTTDProtocol, PacketGameType
|
from openttd.protocol import OpenTTDProtocol, PacketGameType
|
||||||
from openttd_protocol.wire.exceptions import SocketClosed
|
from openttd_protocol.wire.exceptions import SocketClosed
|
||||||
|
|
||||||
|
|
||||||
class MockTransport:
|
class MockTransport:
|
||||||
def __init__(self): self._closing = False
|
def __init__(self): self._closing = False
|
||||||
def is_closing(self): return self._closing
|
def is_closing(self): return self._closing
|
||||||
@@ -59,10 +62,8 @@ def test_protocol_static_parsers():
|
|||||||
assert OpenTTDProtocol.receive_ServerNeedCompanyPassword(None, memoryview(struct.pack("<I", 1234) + b"sid\x00")) == {"seed": 1234, "server_id": "sid"}
|
assert OpenTTDProtocol.receive_ServerNeedCompanyPassword(None, memoryview(struct.pack("<I", 1234) + b"sid\x00")) == {"seed": 1234, "server_id": "sid"}
|
||||||
|
|
||||||
# Coverage for receive_ServerGameInfo
|
# Coverage for receive_ServerGameInfo
|
||||||
try:
|
with contextlib.suppress(Exception):
|
||||||
OpenTTDProtocol.receive_ServerGameInfo(None, memoryview(b"\x00" * 200))
|
OpenTTDProtocol.receive_ServerGameInfo(None, memoryview(b"\x00" * 200))
|
||||||
except Exception:
|
|
||||||
pass
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_protocol_exception_handling():
|
async def test_protocol_exception_handling():
|
||||||
@@ -71,7 +72,7 @@ async def test_protocol_exception_handling():
|
|||||||
proto.transport = MockTransport()
|
proto.transport = MockTransport()
|
||||||
|
|
||||||
# Passing data that causes struct.unpack to fail (too short for uint16)
|
# Passing data that causes struct.unpack to fail (too short for uint16)
|
||||||
ptype, kwargs = proto.receive_packet(None, memoryview(b"\x01"))
|
ptype, _ = proto.receive_packet(None, memoryview(b"\x01"))
|
||||||
assert ptype == PacketGameType.ServerUnused
|
assert ptype == PacketGameType.ServerUnused
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
@@ -112,7 +113,7 @@ async def test_protocol_decryption_failure():
|
|||||||
|
|
||||||
# Needs to be at least 18 bytes for read_uint16 + mac
|
# Needs to be at least 18 bytes for read_uint16 + mac
|
||||||
wire_data = memoryview(b"\x14\x00" + b"X" * 16 + b"junk")
|
wire_data = memoryview(b"\x14\x00" + b"X" * 16 + b"junk")
|
||||||
ptype, kwargs = proto.receive_packet(None, wire_data)
|
ptype, _ = proto.receive_packet(None, wire_data)
|
||||||
assert ptype == PacketGameType.ServerUnused
|
assert ptype == PacketGameType.ServerUnused
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
@@ -127,9 +128,10 @@ async def test_protocol_is_closing_failure():
|
|||||||
await proto.send_packet(b"\x02\x00")
|
await proto.send_packet(b"\x02\x00")
|
||||||
|
|
||||||
def test_admin_protocol_static_receives():
|
def test_admin_protocol_static_receives():
|
||||||
from openttd.protocol import OpenTTDAdminProtocol
|
|
||||||
import struct
|
import struct
|
||||||
|
|
||||||
|
from openttd.protocol import OpenTTDAdminProtocol
|
||||||
|
|
||||||
# 1. receive_ServerProtocol
|
# 1. receive_ServerProtocol
|
||||||
data = memoryview(struct.pack("<B B H H B", 3, 1, 10, 100, 0))
|
data = memoryview(struct.pack("<B B H H B", 3, 1, 10, 100, 0))
|
||||||
res = OpenTTDAdminProtocol.receive_ServerProtocol(None, data)
|
res = OpenTTDAdminProtocol.receive_ServerProtocol(None, data)
|
||||||
|
|||||||
+16
-7
@@ -1,9 +1,18 @@
|
|||||||
import pytest
|
import pytest
|
||||||
from openttd import OpenTTDClient
|
from openttd import OpenTTDClient
|
||||||
from openttd.protocol import (
|
from openttd.protocol import (
|
||||||
OpenTTDProtocol, GameCommand, ModifyTimetableFlags, ModifyTimetableCtrlFlag,
|
INVALID_VEH_ORDER_ID,
|
||||||
OrderType, OrderNonStopFlags, OrderStopLocation, INVALID_VEH_ORDER_ID,
|
GameCommand,
|
||||||
write_varuint, read_varuint, write_varuint_signed, read_varuint_signed
|
ModifyTimetableCtrlFlag,
|
||||||
|
ModifyTimetableFlags,
|
||||||
|
OpenTTDProtocol,
|
||||||
|
OrderNonStopFlags,
|
||||||
|
OrderStopLocation,
|
||||||
|
OrderType,
|
||||||
|
read_varuint,
|
||||||
|
read_varuint_signed,
|
||||||
|
write_varuint,
|
||||||
|
write_varuint_signed,
|
||||||
)
|
)
|
||||||
from openttd_protocol.wire.read import read_uint8, read_uint16
|
from openttd_protocol.wire.read import read_uint8, read_uint16
|
||||||
|
|
||||||
@@ -140,10 +149,10 @@ async def test_client_change_timetable_clear_field_sets_ctrl_flag():
|
|||||||
client = new_client()
|
client = new_client()
|
||||||
await client.change_timetable(7, 3, ModifyTimetableFlags.TravelTime, 0, clear_field=True)
|
await client.change_timetable(7, 3, ModifyTimetableFlags.TravelTime, 0, clear_field=True)
|
||||||
parsed = decode_sent_command(client._protocol.sent[0])
|
parsed = decode_sent_command(client._protocol.sent[0])
|
||||||
vehicle_id, rest = read_varuint(parsed["payload"])
|
_, rest = read_varuint(parsed["payload"])
|
||||||
order_position, rest = read_uint16(rest)
|
_, rest = read_uint16(rest)
|
||||||
flag, rest = read_uint8(rest)
|
flag, rest = read_uint8(rest)
|
||||||
value, rest = read_varuint(rest)
|
_, rest = read_varuint(rest)
|
||||||
ctrl_flags, _ = read_uint8(rest)
|
ctrl_flags, _ = read_uint8(rest)
|
||||||
assert flag == ModifyTimetableFlags.TravelTime
|
assert flag == ModifyTimetableFlags.TravelTime
|
||||||
assert ctrl_flags == ModifyTimetableCtrlFlag.ClearField
|
assert ctrl_flags == ModifyTimetableCtrlFlag.ClearField
|
||||||
@@ -212,7 +221,7 @@ async def test_client_add_order_insert_position_and_nonstop():
|
|||||||
client = new_client()
|
client = new_client()
|
||||||
await client.add_order(7, 6, before_position=1, non_stop=OrderNonStopFlags.NoStopAtIntermediate)
|
await client.add_order(7, 6, before_position=1, non_stop=OrderNonStopFlags.NoStopAtIntermediate)
|
||||||
parsed = decode_sent_command(client._protocol.sent[0])
|
parsed = decode_sent_command(client._protocol.sent[0])
|
||||||
vehicle_id, sel_ord, order_type, order_flags, station = _decode_insert_order_payload(parsed["payload"])
|
_, sel_ord, order_type, _, _ = _decode_insert_order_payload(parsed["payload"])
|
||||||
assert sel_ord == 1 # insert before order position 1
|
assert sel_ord == 1 # insert before order position 1
|
||||||
assert order_type == (OrderType.GotoStation
|
assert order_type == (OrderType.GotoStation
|
||||||
| (OrderStopLocation.PlatformFarEnd << 4)
|
| (OrderStopLocation.PlatformFarEnd << 4)
|
||||||
|
|||||||
Reference in New Issue
Block a user