277 lines
10 KiB
Python
277 lines
10 KiB
Python
"""AISStream WebSocket collector for realtime vessel AIS observations."""
|
|
|
|
from datetime import UTC, datetime
|
|
import asyncio
|
|
import json
|
|
import os
|
|
from typing import Any
|
|
|
|
from sqlalchemy import select
|
|
from sqlalchemy.ext.asyncio import AsyncSession
|
|
|
|
from app.core.data_sources import get_data_sources_config
|
|
from app.models.datasource_config import DataSourceConfig
|
|
from app.services.collectors.base import BaseCollector
|
|
from app.services.vessel_ais_aggregation import (
|
|
AISSTREAM_DELIVERY_MODE,
|
|
AISSTREAM_TRANSPORT,
|
|
record_vessel_ais_observation,
|
|
update_ais_source_health,
|
|
)
|
|
from app.services.vessel_types import normalize_vessel_type_name
|
|
|
|
DEFAULT_AISSTREAM_URL = "wss://stream.aisstream.io/v0/stream"
|
|
DEFAULT_BOUNDING_BOXES = [[[-90, -180], [90, 180]]]
|
|
DEFAULT_MESSAGE_TYPES = ["PositionReport", "ShipStaticData"]
|
|
|
|
|
|
class AISStreamCollector(BaseCollector):
|
|
"""Collect AISStream WebSocket messages into the raw AIS observation layer."""
|
|
|
|
name = "aisstream_vessels"
|
|
priority = "P1"
|
|
module = "L4"
|
|
frequency_hours = 1
|
|
data_type = "vessel_ais"
|
|
fail_on_empty = False
|
|
|
|
async def _load_datasource_config(self) -> DataSourceConfig | None:
|
|
if self._db_session is None:
|
|
return None
|
|
result = await self._db_session.execute(
|
|
select(DataSourceConfig)
|
|
.where(DataSourceConfig.name == self.name)
|
|
.where(DataSourceConfig.is_active.is_(True))
|
|
)
|
|
return result.scalar_one_or_none()
|
|
|
|
async def _get_effective_config(self) -> dict[str, Any]:
|
|
datasource_config = await self._load_datasource_config()
|
|
config = dict(datasource_config.config or {}) if datasource_config else {}
|
|
auth_config = dict(datasource_config.auth_config or {}) if datasource_config else {}
|
|
endpoint = (
|
|
(datasource_config.endpoint if datasource_config else None)
|
|
or self._resolved_url
|
|
or get_data_sources_config().get_yaml_url(self.name)
|
|
or DEFAULT_AISSTREAM_URL
|
|
)
|
|
api_key = (
|
|
auth_config.get("api_key")
|
|
or config.get("api_key")
|
|
or os.getenv("AISSTREAM_API_KEY")
|
|
)
|
|
return {
|
|
"endpoint": endpoint,
|
|
"api_key": api_key,
|
|
"bounding_boxes": config.get("bounding_boxes") or DEFAULT_BOUNDING_BOXES,
|
|
"message_types": config.get("message_types") or DEFAULT_MESSAGE_TYPES,
|
|
"max_messages": int(config.get("max_messages") or 500),
|
|
"receive_timeout_seconds": float(config.get("receive_timeout_seconds") or 30),
|
|
}
|
|
|
|
async def fetch(self) -> list[dict[str, Any]]:
|
|
config = await self._get_effective_config()
|
|
if not config["api_key"]:
|
|
raise RuntimeError("AISStream API key is not configured")
|
|
|
|
try:
|
|
import websockets
|
|
except ImportError as exc:
|
|
raise RuntimeError("Python package 'websockets' is required for AISStream") from exc
|
|
|
|
subscription = {
|
|
"APIKey": config["api_key"],
|
|
"BoundingBoxes": config["bounding_boxes"],
|
|
"FilterMessageTypes": config["message_types"],
|
|
}
|
|
|
|
messages: list[dict[str, Any]] = []
|
|
try:
|
|
async with websockets.connect(config["endpoint"]) as websocket:
|
|
await websocket.send(json.dumps(subscription))
|
|
while len(messages) < config["max_messages"]:
|
|
try:
|
|
raw_message = await asyncio.wait_for(
|
|
websocket.recv(),
|
|
timeout=config["receive_timeout_seconds"],
|
|
)
|
|
except TimeoutError:
|
|
break
|
|
payload = json.loads(raw_message)
|
|
if isinstance(payload, dict):
|
|
messages.append(payload)
|
|
except Exception as exc:
|
|
if self._db_session is not None:
|
|
await update_ais_source_health(
|
|
self._db_session,
|
|
source=self.name,
|
|
connection_state="disconnected",
|
|
last_error=f"{exc.__class__.__name__}: {exc}",
|
|
)
|
|
await self._db_session.commit()
|
|
raise
|
|
|
|
return messages
|
|
|
|
def transform(self, raw_data: list[dict[str, Any]]) -> list[dict[str, Any]]:
|
|
records = []
|
|
for item in raw_data:
|
|
record = self._normalize_message(item)
|
|
if record:
|
|
records.append(record)
|
|
return records
|
|
|
|
async def _save_data(
|
|
self,
|
|
db: AsyncSession,
|
|
data: list[dict[str, Any]],
|
|
task_id: int | None = None,
|
|
snapshot_id: int | None = None,
|
|
) -> int:
|
|
now = datetime.now(UTC)
|
|
records_added = 0
|
|
latest_observed_at = now
|
|
for index, item in enumerate(data):
|
|
observed_at = item.get("received_at") or now
|
|
observation = await record_vessel_ais_observation(
|
|
db,
|
|
source=self.name,
|
|
normalized_payload=item,
|
|
raw_payload=item.get("_raw_payload") or item,
|
|
delivery_mode=AISSTREAM_DELIVERY_MODE,
|
|
transport=AISSTREAM_TRANSPORT,
|
|
message_type=item.get("_message_type") or "PositionReport",
|
|
source_message_id=item.get("_source_message_id"),
|
|
observed_at=observed_at,
|
|
collected_at=now,
|
|
)
|
|
if observation is not None:
|
|
records_added += 1
|
|
if isinstance(observed_at, datetime) and observed_at > latest_observed_at:
|
|
latest_observed_at = observed_at
|
|
if (index + 1) % 1000 == 0:
|
|
await self.update_progress(index + 1, commit=True)
|
|
|
|
await update_ais_source_health(
|
|
db,
|
|
source=self.name,
|
|
connection_state="connected",
|
|
observed_count=len(data),
|
|
last_seen_at=latest_observed_at,
|
|
last_success_at=now if data else None,
|
|
lag_seconds=max((now - latest_observed_at).total_seconds(), 0),
|
|
)
|
|
await db.commit()
|
|
await self.update_progress(records_added, force=True)
|
|
return records_added
|
|
|
|
def _normalize_message(self, item: dict[str, Any]) -> dict[str, Any] | None:
|
|
message_type = str(item.get("MessageType") or item.get("message_type") or "")
|
|
metadata = item.get("MetaData") if isinstance(item.get("MetaData"), dict) else {}
|
|
message = item.get("Message") if isinstance(item.get("Message"), dict) else {}
|
|
body = message.get(message_type) if isinstance(message.get(message_type), dict) else message
|
|
if not isinstance(body, dict):
|
|
body = {}
|
|
|
|
mmsi = _as_int(_pick(metadata, "MMSI", "mmsi") or _pick(body, "MMSI", "mmsi"))
|
|
if mmsi is None:
|
|
return None
|
|
|
|
received_at = _parse_datetime(
|
|
_pick(metadata, "time_utc", "Time_UTC", "timestamp")
|
|
or _pick(body, "Timestamp", "timestamp", "time")
|
|
)
|
|
ship_name = _clean_text(
|
|
_pick(body, "Name", "ShipName", "name")
|
|
or _pick(metadata, "ShipName", "ship_name", "name")
|
|
)
|
|
record: dict[str, Any] = {
|
|
"mmsi": mmsi,
|
|
"received_at": received_at,
|
|
"_message_type": message_type or None,
|
|
"_source_message_id": item.get("MessageID") or item.get("message_id"),
|
|
"_raw_payload": item,
|
|
}
|
|
|
|
lat = _as_float(_pick(body, "Latitude", "lat", "latitude"))
|
|
lon = _as_float(_pick(body, "Longitude", "lon", "lng", "longitude"))
|
|
if lat is not None and lon is not None:
|
|
if not (-90 <= lat <= 90 and -180 <= lon <= 180):
|
|
return None
|
|
record.update(
|
|
{
|
|
"lat": lat,
|
|
"lon": lon,
|
|
"sog": _as_float(_pick(body, "Sog", "SOG", "speedOverGround")),
|
|
"cog": _as_float(_pick(body, "Cog", "COG", "courseOverGround")),
|
|
"heading": _as_int(_pick(body, "TrueHeading", "Heading", "heading")),
|
|
"nav_status": _as_int(_pick(body, "NavigationalStatus", "nav_status")),
|
|
}
|
|
)
|
|
|
|
vessel_type = _as_int(_pick(body, "Type", "ShipType", "vessel_type"))
|
|
record.update(
|
|
{
|
|
"name": ship_name,
|
|
"callsign": _pick(body, "CallSign", "callsign"),
|
|
"imo": _as_int(_pick(body, "ImoNumber", "IMO", "imo")),
|
|
"vessel_type": vessel_type,
|
|
"vessel_type_name": _pick(body, "TypeName", "ShipTypeName", "vessel_type_name")
|
|
or normalize_vessel_type_name(vessel_type),
|
|
"length": _as_float(_pick(body, "DimensionToBow", "Length", "length")),
|
|
"width": _as_float(_pick(body, "DimensionToPort", "Width", "width")),
|
|
}
|
|
)
|
|
return record
|
|
|
|
|
|
def _pick(item: dict[str, Any], *keys: str) -> Any:
|
|
for key in keys:
|
|
if key in item and item[key] not in (None, ""):
|
|
return item[key]
|
|
return None
|
|
|
|
|
|
def _clean_text(value: Any) -> str | None:
|
|
if value in (None, ""):
|
|
return None
|
|
text = str(value).strip()
|
|
return text or None
|
|
|
|
|
|
def _as_float(value: Any) -> float | None:
|
|
try:
|
|
if value in (None, ""):
|
|
return None
|
|
return float(value)
|
|
except (TypeError, ValueError):
|
|
return None
|
|
|
|
|
|
def _as_int(value: Any) -> int | None:
|
|
try:
|
|
if value in (None, ""):
|
|
return None
|
|
return int(float(value))
|
|
except (TypeError, ValueError):
|
|
return None
|
|
|
|
|
|
def _parse_datetime(value: Any) -> datetime | None:
|
|
if isinstance(value, datetime):
|
|
return value if value.tzinfo else value.replace(tzinfo=UTC)
|
|
if not value:
|
|
return None
|
|
if isinstance(value, (int, float)):
|
|
timestamp = float(value)
|
|
if timestamp > 10_000_000_000:
|
|
timestamp /= 1000
|
|
return datetime.fromtimestamp(timestamp, UTC)
|
|
if isinstance(value, str):
|
|
try:
|
|
parsed = datetime.fromisoformat(value.replace("Z", "+00:00"))
|
|
return parsed if parsed.tzinfo else parsed.replace(tzinfo=UTC)
|
|
except ValueError:
|
|
return None
|
|
return None
|