507 lines
20 KiB
Python
507 lines
20 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.core.time import to_iso8601_utc
|
|
from app.core.websocket.broadcaster import broadcaster
|
|
from app.models.datasource_config import DataSourceConfig
|
|
from app.models.task import CollectionTask
|
|
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),
|
|
"streaming_enabled": config.get("streaming_enabled", True) is not False,
|
|
"streaming_commit_interval": int(config.get("streaming_commit_interval") or 1),
|
|
"streaming_max_messages": int(config.get("streaming_max_messages") or 0),
|
|
"reconnect_delay_seconds": float(config.get("reconnect_delay_seconds") or 5),
|
|
"receive_timeout_seconds": float(config.get("receive_timeout_seconds") or 30),
|
|
}
|
|
|
|
def _build_subscription(self, config: dict[str, Any]) -> dict[str, Any]:
|
|
return {
|
|
"APIKey": config["api_key"],
|
|
"BoundingBoxes": config["bounding_boxes"],
|
|
"FilterMessageTypes": config["message_types"],
|
|
}
|
|
|
|
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 = self._build_subscription(config)
|
|
|
|
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
|
|
|
|
async def run(self, db: AsyncSession) -> dict[str, Any]:
|
|
"""Run AISStream as a long-lived streaming collector by default."""
|
|
config = await self._get_effective_config()
|
|
if not config.get("streaming_enabled", True):
|
|
return await super().run(db)
|
|
if not config["api_key"]:
|
|
return {"status": "failed", "error": "AISStream API key is not configured"}
|
|
|
|
from app.services.collectors.registry import collector_registry
|
|
|
|
if not collector_registry.is_active(self.name):
|
|
return {"status": "skipped", "reason": "Collector is disabled"}
|
|
|
|
try:
|
|
import websockets
|
|
except ImportError:
|
|
return {"status": "failed", "error": "Python package 'websockets' is required for AISStream"}
|
|
|
|
start_time = datetime.now(UTC)
|
|
task = CollectionTask(
|
|
datasource_id=getattr(self, "_datasource_id", 1),
|
|
status="running",
|
|
phase="connecting",
|
|
phase_message="正在连接 AISStream 实时流",
|
|
phase_unit="messages",
|
|
started_at=start_time,
|
|
)
|
|
db.add(task)
|
|
await db.commit()
|
|
self._current_task = task
|
|
self._db_session = db
|
|
self._last_broadcast_progress = None
|
|
await self.resolve_url(db)
|
|
await self._publish_task_update(force=True)
|
|
|
|
records_added = 0
|
|
messages_seen = 0
|
|
unique_mmsi: set[str] = set()
|
|
reconnect_delay = config["reconnect_delay_seconds"]
|
|
|
|
try:
|
|
while True:
|
|
config = await self._get_effective_config()
|
|
subscription = self._build_subscription(config)
|
|
try:
|
|
await update_ais_source_health(
|
|
db,
|
|
source=self.name,
|
|
connection_state="connecting",
|
|
)
|
|
await self.set_phase("connecting", message="正在连接 AISStream 实时流")
|
|
await db.commit()
|
|
|
|
async with websockets.connect(config["endpoint"]) as websocket:
|
|
await websocket.send(json.dumps(subscription))
|
|
await update_ais_source_health(
|
|
db,
|
|
source=self.name,
|
|
connection_state="connected",
|
|
last_success_at=datetime.now(UTC),
|
|
)
|
|
await self.set_phase(
|
|
"streaming",
|
|
message="正在接收 AISStream 实时消息",
|
|
reset_progress=False,
|
|
)
|
|
await db.commit()
|
|
|
|
while True:
|
|
try:
|
|
raw_message = await asyncio.wait_for(
|
|
websocket.recv(),
|
|
timeout=config["receive_timeout_seconds"],
|
|
)
|
|
except TimeoutError:
|
|
await update_ais_source_health(
|
|
db,
|
|
source=self.name,
|
|
connection_state="connected",
|
|
last_success_at=datetime.now(UTC),
|
|
)
|
|
await db.commit()
|
|
continue
|
|
|
|
payload = json.loads(raw_message)
|
|
if not isinstance(payload, dict):
|
|
continue
|
|
messages_seen += 1
|
|
record = self._normalize_message(payload)
|
|
if not record:
|
|
continue
|
|
unique_mmsi.add(str(record["mmsi"]))
|
|
created = await self._save_stream_record(db, record)
|
|
if created:
|
|
records_added += 1
|
|
|
|
task.records_processed = messages_seen
|
|
task.total_records = None
|
|
task.progress = None
|
|
task.phase = "streaming"
|
|
task.phase_message = "正在接收 AISStream 实时消息"
|
|
task.phase_current = messages_seen
|
|
task.phase_total = None
|
|
task.phase_unit = "messages"
|
|
await self._publish_task_update(force=True)
|
|
|
|
if config["streaming_max_messages"] and messages_seen >= config["streaming_max_messages"]:
|
|
task.status = "success"
|
|
task.phase = "stopped"
|
|
task.phase_message = "AISStream 测试流已停止"
|
|
task.completed_at = datetime.now(UTC)
|
|
await db.commit()
|
|
await self._publish_task_update(force=True)
|
|
return {
|
|
"status": "success",
|
|
"task_id": task.id,
|
|
"records_processed": records_added,
|
|
"messages_seen": messages_seen,
|
|
"unique_mmsi": len(unique_mmsi),
|
|
"execution_time_seconds": (datetime.now(UTC) - start_time).total_seconds(),
|
|
}
|
|
except asyncio.CancelledError:
|
|
raise
|
|
except Exception as exc:
|
|
await update_ais_source_health(
|
|
db,
|
|
source=self.name,
|
|
connection_state="reconnecting",
|
|
last_error=f"{exc.__class__.__name__}: {exc}",
|
|
)
|
|
task.phase = "reconnecting"
|
|
task.phase_message = "AISStream 连接中断,正在重连"
|
|
task.error_message = f"{exc.__class__.__name__}: {exc}"
|
|
await db.commit()
|
|
await self._publish_task_update(force=True)
|
|
await asyncio.sleep(reconnect_delay)
|
|
except asyncio.CancelledError:
|
|
task.status = "cancelled"
|
|
task.phase = "stopped"
|
|
task.phase_message = "AISStream 实时流已停止"
|
|
task.completed_at = datetime.now(UTC)
|
|
await update_ais_source_health(
|
|
db,
|
|
source=self.name,
|
|
connection_state="disconnected",
|
|
last_error=None,
|
|
)
|
|
await db.commit()
|
|
await self._publish_task_update(force=True)
|
|
raise
|
|
|
|
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),
|
|
)
|
|
if snapshot_id is not None:
|
|
from app.models.data_snapshot import DataSnapshot
|
|
|
|
snapshot = await db.get(DataSnapshot, snapshot_id)
|
|
if snapshot:
|
|
snapshot.record_count = records_added
|
|
snapshot.status = "success"
|
|
snapshot.completed_at = now
|
|
snapshot.summary = {
|
|
"created": records_added,
|
|
"updated": 0,
|
|
"unchanged": 0,
|
|
"deleted": 0,
|
|
"storage": "ais_raw_observations",
|
|
}
|
|
await db.commit()
|
|
await self.update_progress(records_added, force=True)
|
|
return records_added
|
|
|
|
async def _save_stream_record(self, db: AsyncSession, item: dict[str, Any]) -> bool:
|
|
now = datetime.now(UTC)
|
|
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,
|
|
)
|
|
await update_ais_source_health(
|
|
db,
|
|
source=self.name,
|
|
connection_state="connected",
|
|
observed_count=1,
|
|
last_seen_at=observed_at if isinstance(observed_at, datetime) else now,
|
|
last_success_at=now,
|
|
lag_seconds=max((now - observed_at).total_seconds(), 0) if isinstance(observed_at, datetime) else None,
|
|
)
|
|
await db.commit()
|
|
await self._broadcast_vessel_delta(item, created=observation is not None)
|
|
return observation is not None
|
|
|
|
async def _broadcast_vessel_delta(self, item: dict[str, Any], *, created: bool) -> None:
|
|
await broadcaster.broadcast_custom(
|
|
"vessels",
|
|
{
|
|
"action": "upsert",
|
|
"source": self.name,
|
|
"created": created,
|
|
"vessels": [
|
|
{
|
|
"mmsi": item.get("mmsi"),
|
|
"mmsi_display": str(item.get("mmsi")) if item.get("mmsi") is not None else None,
|
|
"name": item.get("name"),
|
|
"lat": item.get("lat"),
|
|
"lon": item.get("lon"),
|
|
"sog": item.get("sog"),
|
|
"cog": item.get("cog"),
|
|
"heading": item.get("heading"),
|
|
"nav_status": item.get("nav_status"),
|
|
"vessel_type": item.get("vessel_type"),
|
|
"vessel_type_name": item.get("vessel_type_name"),
|
|
"received_at": to_iso8601_utc(item.get("received_at")),
|
|
}
|
|
],
|
|
},
|
|
)
|
|
|
|
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
|