"""End-to-end integration test for the custom WebSocket datasource runner. Boots an in-process WebSocket server that mimics the bun mock AIS server (`scripts/mock-ais-ws-server.ts`) and runs the real `run_mapped_websocket_config` against it. Catches regressions where the runner stops connecting, fails to extract the configured message path, or quietly drops mapped records before broadcasting. """ from __future__ import annotations import asyncio import json from contextlib import asynccontextmanager from datetime import UTC, datetime from types import SimpleNamespace from unittest.mock import AsyncMock import pytest import websockets from app.models.datasource_config import DataSourceConfig from app.services import custom_datasource_runtime from app.services.custom_datasource_runtime import run_mapped_websocket_config def _make_payload(seq: int) -> str: return json.dumps( { "type": "vessel", "sequence": seq, "data": { "mmsi": str(999_000_000 + seq), "name": f"MOCK VESSEL {seq:03d}", "lat": 36.20 + seq * 0.001, "lon": 14.20 + seq * 0.001, "sog": 12.0, "cog": 90.0, "heading": 90, "vessel_type": 70, "vessel_type_name": "Cargo", "received_at": datetime.now(UTC).isoformat(), }, } ) @asynccontextmanager async def _mock_ais_server(emit_count: int): received_subscribe: list[str] = [] async def handler(ws): try: try: msg = await asyncio.wait_for(ws.recv(), timeout=0.5) received_subscribe.append(msg) except (asyncio.TimeoutError, websockets.ConnectionClosed): pass for seq in range(1, emit_count + 1): await ws.send(_make_payload(seq)) await asyncio.sleep(0.01) # keep the socket open briefly so the runner observes the messages await asyncio.sleep(0.05) except websockets.ConnectionClosed: return async with websockets.serve(handler, "127.0.0.1", 0) as server: port = next(iter(server.sockets)).getsockname()[1] yield port, received_subscribe @pytest.mark.asyncio async def test_websocket_runner_streams_from_live_mock(monkeypatch): mapping = SimpleNamespace( id=11, version=3, target_schema="vessel_ais", mapping_json={ "source": {"items_path": "$"}, "fields": { "mmsi": {"path": "$.mmsi", "type": "integer"}, "lat": {"path": "$.lat", "type": "float"}, "lon": {"path": "$.lon", "type": "float"}, "name": {"path": "$.name", "type": "string"}, "vessel_type": {"path": "$.vessel_type", "type": "integer", "default": None}, "vessel_type_name": {"path": "$.vessel_type_name", "type": "string", "default": None}, "sog": {"path": "$.sog", "type": "float", "default": None}, "cog": {"path": "$.cog", "type": "float", "default": None}, "heading": {"path": "$.heading", "type": "integer", "default": None}, "received_at": {"path": "$.received_at", "type": "datetime"}, }, }, ) class FakeResult: def scalar_one_or_none(self): return mapping class FakeDB: async def execute(self, _stmt): return FakeResult() persist = AsyncMock(return_value=1) monkeypatch.setattr(custom_datasource_runtime, "persist_mapped_records", persist) async with _mock_ais_server(emit_count=3) as (port, received_subscribe): result = await run_mapped_websocket_config( FakeDB(), DataSourceConfig( id=99, name="mock_ais_ws", source_type="websocket", endpoint=f"ws://127.0.0.1:{port}", auth_type="none", headers={}, config={ "ws_message_path": "$.data", "ws_subscribe_message": { "type": "subscribe", "anchor": {"lat": 36.2, "lon": 14.2}, "spread_km": 50, "rate_hz": 1, }, "debug_max_messages": 2, "delivery_mode": "realtime_stream", "ws_reconnect": False, }, ), use_config_debug_max_messages=True, ) assert result["status"] == "success" assert result["messages_seen"] == 2 assert result["written_count"] == 2 assert result["mapped_count"] == 2 assert result["target_schema"] == "vessel_ais" # subscribe message must reach the server unchanged assert received_subscribe, "runner did not forward ws_subscribe_message" parsed = json.loads(received_subscribe[0]) assert parsed["type"] == "subscribe" assert parsed["anchor"] == {"lat": 36.2, "lon": 14.2} assert parsed["rate_hz"] == 1 # mapped records carry the real MMSIs from the mock stream persisted_records = [] for call in persist.await_args_list: persisted_records.extend(call.kwargs["records"]) assert {record["mmsi"] for record in persisted_records} == {999_000_001, 999_000_002} assert all(record["vessel_type"] == 70 for record in persisted_records) assert all(record["vessel_type_name"] == "Cargo" for record in persisted_records)