"""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 as exc: 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), ) 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