Files
planet/backend/tests/test_custom_datasource_runtime_live.py
2026-05-07 18:06:06 +08:00

150 lines
5.5 KiB
Python

"""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)