release: bump version to 0.48.0
This commit is contained in:
@@ -10,7 +10,10 @@ 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,
|
||||
@@ -66,9 +69,20 @@ class AISStreamCollector(BaseCollector):
|
||||
"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"]:
|
||||
@@ -79,11 +93,7 @@ class AISStreamCollector(BaseCollector):
|
||||
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"],
|
||||
}
|
||||
subscription = self._build_subscription(config)
|
||||
|
||||
messages: list[dict[str, Any]] = []
|
||||
try:
|
||||
@@ -113,6 +123,157 @@ class AISStreamCollector(BaseCollector):
|
||||
|
||||
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:
|
||||
@@ -165,6 +326,60 @@ class AISStreamCollector(BaseCollector):
|
||||
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 {}
|
||||
|
||||
@@ -1,13 +1,13 @@
|
||||
"""BarentsWatch AIS collector for vessel tracking."""
|
||||
|
||||
from datetime import UTC, datetime, timedelta
|
||||
from datetime import UTC, datetime
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
from sqlalchemy import delete
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.models.vessel import VesselPosition, VesselStatic
|
||||
from app.core.time import to_iso8601_utc
|
||||
from app.core.websocket.broadcaster import broadcaster
|
||||
from app.services.barentswatch import (
|
||||
BARENTSWATCH_LATEST_URL,
|
||||
fetch_barentswatch_access_token,
|
||||
@@ -101,40 +101,6 @@ class VesselAISCollector(BaseCollector):
|
||||
observed_at=observed_at,
|
||||
collected_at=now,
|
||||
)
|
||||
|
||||
static = await db.get(VesselStatic, item["mmsi"])
|
||||
if static is None:
|
||||
static = VesselStatic(mmsi=item["mmsi"])
|
||||
db.add(static)
|
||||
|
||||
for field in (
|
||||
"name",
|
||||
"callsign",
|
||||
"vessel_type",
|
||||
"vessel_type_name",
|
||||
"flag",
|
||||
"length",
|
||||
"width",
|
||||
"draught",
|
||||
"imo",
|
||||
):
|
||||
value = item.get(field)
|
||||
if value not in (None, ""):
|
||||
setattr(static, field, value)
|
||||
static.updated_at = now
|
||||
|
||||
db.add(
|
||||
VesselPosition(
|
||||
mmsi=item["mmsi"],
|
||||
lat=item["lat"],
|
||||
lon=item["lon"],
|
||||
sog=item.get("sog"),
|
||||
cog=item.get("cog"),
|
||||
heading=item.get("heading"),
|
||||
nav_status=item.get("nav_status"),
|
||||
received_at=observed_at,
|
||||
)
|
||||
)
|
||||
records_added += 1
|
||||
|
||||
if (index + 1) % 1000 == 0:
|
||||
@@ -153,13 +119,46 @@ class VesselAISCollector(BaseCollector):
|
||||
last_success_at=now if data else None,
|
||||
lag_seconds=max((now - latest_observed_at).total_seconds(), 0),
|
||||
)
|
||||
await db.execute(
|
||||
delete(VesselPosition).where(VesselPosition.received_at < now - timedelta(hours=24))
|
||||
)
|
||||
await db.commit()
|
||||
await self._broadcast_vessel_snapshot(data)
|
||||
await self.update_progress(records_added, force=True)
|
||||
return records_added
|
||||
|
||||
async def _broadcast_vessel_snapshot(self, data: list[dict[str, Any]]) -> None:
|
||||
"""Push REST collector updates through the same realtime vessel channel."""
|
||||
if not data:
|
||||
return
|
||||
|
||||
batch_size = 500
|
||||
for offset in range(0, len(data), batch_size):
|
||||
batch = data[offset : offset + batch_size]
|
||||
await broadcaster.broadcast_custom(
|
||||
"vessels",
|
||||
{
|
||||
"action": "upsert",
|
||||
"source": self.name,
|
||||
"created": True,
|
||||
"vessels": [
|
||||
{
|
||||
"mmsi": item.get("mmsi"),
|
||||
"mmsi_display": str(item.get("mmsi")) if item.get("mmsi") is not None else None,
|
||||
"name": item.get("name"),
|
||||
"callsign": item.get("callsign"),
|
||||
"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")),
|
||||
}
|
||||
for item in batch
|
||||
],
|
||||
},
|
||||
)
|
||||
|
||||
def _normalize_record(self, item: dict[str, Any]) -> dict[str, Any] | None:
|
||||
mmsi = _as_int(_pick(item, "mmsi", "MMSI", "Mmsi"))
|
||||
lat = _as_float(_pick(item, "lat", "latitude", "Latitude"))
|
||||
|
||||
391
backend/app/services/custom_datasource_runtime.py
Normal file
391
backend/app/services/custom_datasource_runtime.py
Normal file
@@ -0,0 +1,391 @@
|
||||
"""Runtime helpers for mapped custom data sources."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import base64
|
||||
import json
|
||||
from datetime import UTC, datetime
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.core.target_schema_registry import TARGET_SCHEMAS
|
||||
from app.db.session import async_session_factory
|
||||
from app.models.datasource_config import DataSourceConfig
|
||||
from app.models.datasource_mapping import DataSourceMappingTemplate
|
||||
from app.services.datasource_mapping import (
|
||||
MappingError,
|
||||
execute_mapping,
|
||||
extract_path,
|
||||
persist_mapped_records,
|
||||
)
|
||||
|
||||
DEFAULT_MAPPING_TEMPLATES: dict[str, dict[str, Any]] = {
|
||||
"vessel_ais": {
|
||||
"source": {"items_path": "$"},
|
||||
"fields": {
|
||||
"mmsi": {"path": "$.mmsi", "type": "integer"},
|
||||
"name": {"path": "$.name", "type": "string", "default": None},
|
||||
"lat": {"path": "$.lat", "type": "float"},
|
||||
"lon": {"path": "$.lon", "type": "float"},
|
||||
"sog": {"path": "$.sog", "type": "float", "default": None},
|
||||
"cog": {"path": "$.cog", "type": "float", "default": None},
|
||||
"heading": {"path": "$.heading", "type": "integer", "default": None},
|
||||
"nav_status": {"path": "$.nav_status", "type": "integer", "default": None},
|
||||
"callsign": {"path": "$.callsign", "type": "string", "default": None},
|
||||
"vessel_type": {"path": "$.vessel_type", "type": "string", "default": None},
|
||||
"vessel_type_name": {"path": "$.vessel_type_name", "type": "string", "default": None},
|
||||
"received_at": {"path": "$.received_at", "type": "datetime", "default": None},
|
||||
},
|
||||
"meta": {"generated_by": "default_template", "requires_review": False},
|
||||
},
|
||||
}
|
||||
|
||||
RUNNING_CUSTOM_STREAM_TASKS: dict[int, asyncio.Task[Any]] = {}
|
||||
|
||||
|
||||
class CustomDatasourceRuntimeError(RuntimeError):
|
||||
"""Raised when a custom datasource cannot run."""
|
||||
|
||||
|
||||
def build_request_headers(auth_type: str, auth_config: dict, headers: dict) -> dict[str, str]:
|
||||
request_headers = {str(key): str(value) for key, value in (headers or {}).items()}
|
||||
auth_type = str(auth_type or "none").lower()
|
||||
auth_config = auth_config or {}
|
||||
|
||||
if auth_type == "bearer" and auth_config.get("token"):
|
||||
request_headers["Authorization"] = f"Bearer {auth_config['token']}"
|
||||
elif auth_type == "api_key" and auth_config.get("api_key"):
|
||||
location = str(auth_config.get("in") or auth_config.get("location") or "header").lower()
|
||||
if location != "query":
|
||||
key_name = auth_config.get("key_name", "X-API-Key")
|
||||
request_headers[str(key_name)] = str(auth_config["api_key"])
|
||||
elif auth_type == "basic":
|
||||
username = auth_config.get("username", "")
|
||||
password = auth_config.get("password", "")
|
||||
credentials = f"{username}:{password}"
|
||||
encoded = base64.b64encode(credentials.encode()).decode()
|
||||
request_headers["Authorization"] = f"Basic {encoded}"
|
||||
return request_headers
|
||||
|
||||
|
||||
def build_query_params(auth_type: str, auth_config: dict, config: dict) -> dict[str, Any]:
|
||||
params: dict[str, Any] = {}
|
||||
candidate = (config or {}).get("params") or (config or {}).get("query_params")
|
||||
if isinstance(candidate, dict):
|
||||
params.update(candidate)
|
||||
|
||||
auth_type = str(auth_type or "none").lower()
|
||||
auth_config = auth_config or {}
|
||||
if auth_type == "api_key" and auth_config.get("api_key"):
|
||||
location = str(auth_config.get("in") or auth_config.get("location") or "header").lower()
|
||||
if location == "query":
|
||||
key_name = auth_config.get("key_name") or auth_config.get("param_name") or "api_key"
|
||||
params[str(key_name)] = auth_config["api_key"]
|
||||
return params
|
||||
|
||||
|
||||
async def load_active_mapping(
|
||||
db: AsyncSession,
|
||||
datasource_config_id: int,
|
||||
) -> DataSourceMappingTemplate:
|
||||
result = await db.execute(
|
||||
select(DataSourceMappingTemplate)
|
||||
.where(DataSourceMappingTemplate.datasource_config_id == datasource_config_id)
|
||||
.where(DataSourceMappingTemplate.is_active.is_(True))
|
||||
.order_by(DataSourceMappingTemplate.version.desc())
|
||||
.limit(1)
|
||||
)
|
||||
mapping = result.scalar_one_or_none()
|
||||
if mapping is not None:
|
||||
return mapping
|
||||
|
||||
datasource = await db.get(DataSourceConfig, datasource_config_id)
|
||||
if datasource is None:
|
||||
raise CustomDatasourceRuntimeError("Configuration not found")
|
||||
target_schema = (datasource.config or {}).get("target_schema")
|
||||
template_body = DEFAULT_MAPPING_TEMPLATES.get(str(target_schema or "")) if target_schema else None
|
||||
if not template_body or target_schema not in TARGET_SCHEMAS:
|
||||
raise CustomDatasourceRuntimeError(
|
||||
"No active mapping template found and no default template available for this target schema"
|
||||
)
|
||||
|
||||
mapping = DataSourceMappingTemplate(
|
||||
datasource_config_id=datasource_config_id,
|
||||
target_schema=str(target_schema),
|
||||
mapping_json=template_body,
|
||||
sample_payload_hash=None,
|
||||
validation_status="valid",
|
||||
version=1,
|
||||
is_active=True,
|
||||
)
|
||||
db.add(mapping)
|
||||
await db.commit()
|
||||
await db.refresh(mapping)
|
||||
return mapping
|
||||
|
||||
|
||||
async def fetch_rest_payload(config: DataSourceConfig, limit_bytes: int) -> Any:
|
||||
request_config = config.config or {}
|
||||
method = str(request_config.get("method") or request_config.get("request_method") or "GET").upper()
|
||||
if method not in {"GET", "POST"}:
|
||||
raise CustomDatasourceRuntimeError("Only GET and POST sample requests are supported.")
|
||||
|
||||
headers = build_request_headers(config.auth_type, config.auth_config or {}, config.headers or {})
|
||||
params = build_query_params(config.auth_type, config.auth_config or {}, request_config)
|
||||
timeout = float(request_config.get("timeout", 30))
|
||||
json_body = request_config.get("json_body")
|
||||
if json_body is None and str(request_config.get("body_type") or "").lower() in {"json", ""}:
|
||||
candidate = request_config.get("body")
|
||||
if isinstance(candidate, (dict, list)):
|
||||
json_body = candidate
|
||||
|
||||
async with httpx.AsyncClient(timeout=timeout, follow_redirects=True) as client:
|
||||
response = await client.request(
|
||||
method,
|
||||
config.endpoint,
|
||||
headers=headers,
|
||||
params=params or None,
|
||||
json=json_body,
|
||||
)
|
||||
response.raise_for_status()
|
||||
content = response.content[:limit_bytes]
|
||||
if "application/json" in response.headers.get("content-type", ""):
|
||||
return json.loads(content.decode(response.encoding or "utf-8"))
|
||||
return {"text": content.decode(response.encoding or "utf-8", errors="replace")}
|
||||
|
||||
|
||||
async def run_mapped_rest_config(
|
||||
db: AsyncSession,
|
||||
datasource: DataSourceConfig,
|
||||
) -> dict[str, Any]:
|
||||
mapping = await load_active_mapping(db, datasource.id)
|
||||
sample = await fetch_rest_payload(datasource, 5_000_000)
|
||||
mapped = execute_mapping(sample, mapping.mapping_json, mapping.target_schema)
|
||||
if mapped["failed_count"] > 0:
|
||||
return {
|
||||
"status": "failed",
|
||||
"datasource_config_id": datasource.id,
|
||||
"mapping_id": mapping.id,
|
||||
"mapping_version": mapping.version,
|
||||
"target_schema": mapping.target_schema,
|
||||
"mapped_count": mapped["mapped_count"],
|
||||
"failed_count": mapped["failed_count"],
|
||||
"errors": mapped["errors"][:20],
|
||||
}
|
||||
|
||||
request_config = datasource.config or {}
|
||||
written_count = await persist_mapped_records(
|
||||
db,
|
||||
datasource_name=datasource.name,
|
||||
datasource_config_id=datasource.id,
|
||||
target_schema=mapping.target_schema,
|
||||
records=mapped["records"],
|
||||
mapping_version=mapping.version,
|
||||
delivery_mode=request_config.get("delivery_mode") or "polling",
|
||||
transport="http",
|
||||
)
|
||||
return {
|
||||
"status": "success",
|
||||
"datasource_config_id": datasource.id,
|
||||
"mapping_id": mapping.id,
|
||||
"mapping_version": mapping.version,
|
||||
"target_schema": mapping.target_schema,
|
||||
"fetched_count": mapped["total_items"],
|
||||
"mapped_count": mapped["mapped_count"],
|
||||
"written_count": written_count,
|
||||
}
|
||||
|
||||
|
||||
def _items_from_ws_message(payload: Any, config: dict) -> Any:
|
||||
message_path = config.get("ws_message_path")
|
||||
items_path = config.get("ws_items_path")
|
||||
value = extract_path(payload, message_path) if message_path else payload
|
||||
return extract_path(value, items_path) if items_path else value
|
||||
|
||||
|
||||
async def _connect_websocket(endpoint: str, headers: dict[str, str]):
|
||||
import websockets
|
||||
|
||||
try:
|
||||
return await websockets.connect(endpoint, additional_headers=headers or None)
|
||||
except TypeError:
|
||||
return await websockets.connect(endpoint, extra_headers=headers or None)
|
||||
|
||||
|
||||
async def test_websocket_config(config: DataSourceConfig) -> dict[str, Any]:
|
||||
if not str(config.endpoint or "").startswith(("ws://", "wss://")):
|
||||
raise CustomDatasourceRuntimeError("WebSocket datasource endpoint must start with ws:// or wss://")
|
||||
|
||||
runtime_config = config.config or {}
|
||||
headers = build_request_headers(config.auth_type, config.auth_config or {}, config.headers or {})
|
||||
receive_timeout = float(runtime_config.get("receive_timeout_seconds") or runtime_config.get("timeout") or 10)
|
||||
async with await _connect_websocket(config.endpoint, headers) as websocket:
|
||||
subscribe_message = runtime_config.get("ws_subscribe_message")
|
||||
if isinstance(subscribe_message, (dict, list)):
|
||||
await websocket.send(json.dumps(subscribe_message))
|
||||
elif isinstance(subscribe_message, str) and subscribe_message.strip():
|
||||
await websocket.send(subscribe_message)
|
||||
raw_message = await asyncio.wait_for(websocket.recv(), timeout=receive_timeout)
|
||||
return {
|
||||
"success": True,
|
||||
"message_preview": raw_message[:1000] if isinstance(raw_message, str) else str(raw_message)[:1000],
|
||||
}
|
||||
|
||||
|
||||
async def run_mapped_websocket_config(
|
||||
db: AsyncSession,
|
||||
datasource: DataSourceConfig,
|
||||
*,
|
||||
debug_max_messages: int | None = None,
|
||||
use_config_debug_max_messages: bool = True,
|
||||
) -> dict[str, Any]:
|
||||
if not str(datasource.endpoint or "").startswith(("ws://", "wss://")):
|
||||
raise CustomDatasourceRuntimeError("WebSocket datasource endpoint must start with ws:// or wss://")
|
||||
|
||||
mapping = await load_active_mapping(db, datasource.id)
|
||||
runtime_config = datasource.config or {}
|
||||
max_messages = debug_max_messages
|
||||
if max_messages is None and use_config_debug_max_messages:
|
||||
max_messages = runtime_config.get("debug_max_messages")
|
||||
max_messages = int(max_messages) if max_messages else None
|
||||
receive_timeout = float(runtime_config.get("receive_timeout_seconds") or runtime_config.get("timeout") or 30)
|
||||
reconnect = bool(runtime_config.get("ws_reconnect", True))
|
||||
reconnect_delay = float(runtime_config.get("reconnect_delay_seconds") or 3)
|
||||
headers = build_request_headers(datasource.auth_type, datasource.auth_config or {}, datasource.headers or {})
|
||||
|
||||
messages_seen = 0
|
||||
mapped_count = 0
|
||||
failed_count = 0
|
||||
written_count = 0
|
||||
errors: list[dict[str, Any]] = []
|
||||
started_at = datetime.now(UTC)
|
||||
|
||||
while True:
|
||||
try:
|
||||
async with await _connect_websocket(datasource.endpoint, headers) as websocket:
|
||||
subscribe_message = runtime_config.get("ws_subscribe_message")
|
||||
if isinstance(subscribe_message, (dict, list)):
|
||||
await websocket.send(json.dumps(subscribe_message))
|
||||
elif isinstance(subscribe_message, str) and subscribe_message.strip():
|
||||
await websocket.send(subscribe_message)
|
||||
|
||||
while True:
|
||||
raw_message = await asyncio.wait_for(websocket.recv(), timeout=receive_timeout)
|
||||
messages_seen += 1
|
||||
try:
|
||||
payload = json.loads(raw_message)
|
||||
except json.JSONDecodeError as exc:
|
||||
failed_count += 1
|
||||
errors.append({"message": "invalid_json", "error": str(exc)})
|
||||
continue
|
||||
|
||||
extracted = _items_from_ws_message(payload, runtime_config)
|
||||
try:
|
||||
mapped = execute_mapping(extracted, mapping.mapping_json, mapping.target_schema)
|
||||
except (MappingError, ValueError) as exc:
|
||||
failed_count += 1
|
||||
errors.append({"message": "mapping_failed", "error": str(exc)})
|
||||
continue
|
||||
|
||||
mapped_count += mapped["mapped_count"]
|
||||
failed_count += mapped["failed_count"]
|
||||
if mapped["errors"]:
|
||||
errors.extend(mapped["errors"][:5])
|
||||
if mapped["records"]:
|
||||
written_count += await persist_mapped_records(
|
||||
db,
|
||||
datasource_name=datasource.name,
|
||||
datasource_config_id=datasource.id,
|
||||
target_schema=mapping.target_schema,
|
||||
records=mapped["records"],
|
||||
mapping_version=mapping.version,
|
||||
delivery_mode=runtime_config.get("delivery_mode") or "realtime_stream",
|
||||
transport="websocket",
|
||||
)
|
||||
|
||||
if max_messages and messages_seen >= max_messages:
|
||||
return {
|
||||
"status": "success",
|
||||
"datasource_config_id": datasource.id,
|
||||
"mapping_id": mapping.id,
|
||||
"mapping_version": mapping.version,
|
||||
"target_schema": mapping.target_schema,
|
||||
"messages_seen": messages_seen,
|
||||
"mapped_count": mapped_count,
|
||||
"failed_count": failed_count,
|
||||
"written_count": written_count,
|
||||
"errors": errors[:20],
|
||||
"execution_time_seconds": (datetime.now(UTC) - started_at).total_seconds(),
|
||||
}
|
||||
except asyncio.CancelledError:
|
||||
raise
|
||||
except Exception as exc:
|
||||
failed_count += 1
|
||||
errors.append({"message": "websocket_error", "error": f"{exc.__class__.__name__}: {exc}"})
|
||||
if not reconnect or max_messages:
|
||||
return {
|
||||
"status": "failed" if written_count == 0 else "partial",
|
||||
"datasource_config_id": datasource.id,
|
||||
"mapping_id": mapping.id,
|
||||
"mapping_version": mapping.version,
|
||||
"target_schema": mapping.target_schema,
|
||||
"messages_seen": messages_seen,
|
||||
"mapped_count": mapped_count,
|
||||
"failed_count": failed_count,
|
||||
"written_count": written_count,
|
||||
"errors": errors[:20],
|
||||
}
|
||||
await asyncio.sleep(reconnect_delay)
|
||||
|
||||
|
||||
async def run_custom_stream_by_id(config_id: int) -> dict[str, Any]:
|
||||
async with async_session_factory() as db:
|
||||
datasource = await db.get(DataSourceConfig, config_id)
|
||||
if not datasource:
|
||||
raise CustomDatasourceRuntimeError("Configuration not found")
|
||||
return await run_mapped_websocket_config(
|
||||
db,
|
||||
datasource,
|
||||
use_config_debug_max_messages=False,
|
||||
)
|
||||
|
||||
|
||||
def start_custom_stream(config_id: int) -> bool:
|
||||
existing = RUNNING_CUSTOM_STREAM_TASKS.get(config_id)
|
||||
if existing is not None and not existing.done():
|
||||
return False
|
||||
task = asyncio.create_task(run_custom_stream_by_id(config_id), name=f"custom-stream:{config_id}")
|
||||
RUNNING_CUSTOM_STREAM_TASKS[config_id] = task
|
||||
|
||||
def _cleanup(done_task: asyncio.Task[Any]) -> None:
|
||||
if RUNNING_CUSTOM_STREAM_TASKS.get(config_id) is done_task:
|
||||
RUNNING_CUSTOM_STREAM_TASKS.pop(config_id, None)
|
||||
|
||||
task.add_done_callback(_cleanup)
|
||||
return True
|
||||
|
||||
|
||||
async def stop_custom_stream(config_id: int) -> bool:
|
||||
task = RUNNING_CUSTOM_STREAM_TASKS.get(config_id)
|
||||
if task is None or task.done():
|
||||
RUNNING_CUSTOM_STREAM_TASKS.pop(config_id, None)
|
||||
return False
|
||||
task.cancel()
|
||||
try:
|
||||
await task
|
||||
except asyncio.CancelledError:
|
||||
return True
|
||||
return task.cancelled()
|
||||
|
||||
|
||||
def get_custom_stream_status(config_id: int) -> dict[str, Any]:
|
||||
task = RUNNING_CUSTOM_STREAM_TASKS.get(config_id)
|
||||
return {
|
||||
"config_id": config_id,
|
||||
"running": bool(task and not task.done()),
|
||||
"done": bool(task and task.done()),
|
||||
}
|
||||
@@ -290,25 +290,77 @@ async def persist_mapped_records(
|
||||
target_schema: str,
|
||||
records: list[dict[str, Any]],
|
||||
mapping_version: int,
|
||||
delivery_mode: str | None = None,
|
||||
transport: str | None = None,
|
||||
) -> int:
|
||||
"""Persist validated mapped records to the destination for a target schema."""
|
||||
if target_schema == "vessel_ais":
|
||||
from app.models.vessel import VesselPosition
|
||||
from app.core.time import to_iso8601_utc
|
||||
from app.core.websocket.broadcaster import broadcaster
|
||||
from app.services.vessel_ais_aggregation import (
|
||||
record_vessel_ais_observation,
|
||||
update_ais_source_health,
|
||||
)
|
||||
|
||||
now = datetime.now(UTC)
|
||||
latest_observed_at = now
|
||||
written_count = 0
|
||||
for record in records:
|
||||
db.add(
|
||||
VesselPosition(
|
||||
mmsi=record["mmsi"],
|
||||
lat=record["lat"],
|
||||
lon=record["lon"],
|
||||
sog=record.get("sog"),
|
||||
cog=record.get("cog"),
|
||||
heading=record.get("heading"),
|
||||
received_at=_parse_datetime(record.get("received_at")) or datetime.now(UTC),
|
||||
)
|
||||
observed_at = _parse_datetime(record.get("received_at")) or now
|
||||
observation = await record_vessel_ais_observation(
|
||||
db,
|
||||
source=datasource_name,
|
||||
normalized_payload=record,
|
||||
raw_payload=record,
|
||||
delivery_mode=delivery_mode or "polling",
|
||||
transport=transport or "http",
|
||||
message_type="PositionReport",
|
||||
observed_at=observed_at,
|
||||
collected_at=now,
|
||||
)
|
||||
if observation is not None:
|
||||
written_count += 1
|
||||
if observed_at > latest_observed_at:
|
||||
latest_observed_at = observed_at
|
||||
|
||||
await update_ais_source_health(
|
||||
db,
|
||||
source=datasource_name,
|
||||
connection_state="connected",
|
||||
observed_count=len(records),
|
||||
last_seen_at=latest_observed_at,
|
||||
last_success_at=now if records else None,
|
||||
lag_seconds=max((now - latest_observed_at).total_seconds(), 0),
|
||||
)
|
||||
await db.commit()
|
||||
return len(records)
|
||||
if records:
|
||||
await broadcaster.broadcast_custom(
|
||||
"vessels",
|
||||
{
|
||||
"action": "upsert",
|
||||
"source": datasource_name,
|
||||
"created": True,
|
||||
"vessels": [
|
||||
{
|
||||
"mmsi": record.get("mmsi"),
|
||||
"mmsi_display": str(record.get("mmsi")) if record.get("mmsi") is not None else None,
|
||||
"name": record.get("name"),
|
||||
"callsign": record.get("callsign"),
|
||||
"lat": record.get("lat"),
|
||||
"lon": record.get("lon"),
|
||||
"sog": record.get("sog"),
|
||||
"cog": record.get("cog"),
|
||||
"heading": record.get("heading"),
|
||||
"nav_status": record.get("nav_status"),
|
||||
"vessel_type": record.get("vessel_type"),
|
||||
"vessel_type_name": record.get("vessel_type_name"),
|
||||
"received_at": to_iso8601_utc(_parse_datetime(record.get("received_at"))),
|
||||
}
|
||||
for record in records
|
||||
],
|
||||
},
|
||||
)
|
||||
return written_count
|
||||
|
||||
from app.models.collected_data import CollectedData
|
||||
|
||||
|
||||
198
backend/app/services/vessel_aggregation_strategy.py
Normal file
198
backend/app/services/vessel_aggregation_strategy.py
Normal file
@@ -0,0 +1,198 @@
|
||||
"""Persistence + validation for the v4 vessel_ais aggregation strategy."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.models.system_setting import SystemSetting
|
||||
|
||||
VESSEL_AGGREGATION_STRATEGY_CATEGORY = "vessel_aggregation_strategy"
|
||||
|
||||
DYNAMIC_FIELDS: tuple[str, ...] = ("lat", "lon", "sog", "cog", "heading", "nav_status")
|
||||
STATIC_FIELDS: tuple[str, ...] = (
|
||||
"name",
|
||||
"callsign",
|
||||
"imo",
|
||||
"flag",
|
||||
"vessel_type",
|
||||
"vessel_type_name",
|
||||
"length",
|
||||
"width",
|
||||
"draught",
|
||||
)
|
||||
ALLOWED_FIELDS: frozenset[str] = frozenset(DYNAMIC_FIELDS + STATIC_FIELDS)
|
||||
ALLOWED_DYNAMIC_MODES: frozenset[str] = frozenset({"newest"})
|
||||
ALLOWED_STATIC_MODES: frozenset[str] = frozenset({"source_priority", "non_empty", "newest", "locked"})
|
||||
ALLOWED_LOCKED_DYNAMIC_MODES: frozenset[str] = frozenset({"newest", "source_priority", "locked"})
|
||||
|
||||
|
||||
DEFAULT_STRATEGY: dict[str, Any] = {
|
||||
"version": 1,
|
||||
"vessel_ais": {
|
||||
"source_priority": ["aisstream_vessels", "barentswatch_vessels"],
|
||||
"field_rules": {},
|
||||
"freshness": {
|
||||
"realtime_stream_seconds": 900,
|
||||
"polling_seconds": 3600,
|
||||
},
|
||||
"allow_dynamic_lock": False,
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
class StrategyValidationError(ValueError):
|
||||
"""Raised when a saved strategy payload is malformed."""
|
||||
|
||||
|
||||
def _coerce_str_list(value: Any, *, label: str) -> list[str]:
|
||||
if value is None:
|
||||
return []
|
||||
if not isinstance(value, list):
|
||||
raise StrategyValidationError(f"{label} must be a list of source names")
|
||||
out: list[str] = []
|
||||
for item in value:
|
||||
if not isinstance(item, str) or not item.strip():
|
||||
raise StrategyValidationError(f"{label} entries must be non-empty strings")
|
||||
out.append(item.strip())
|
||||
return out
|
||||
|
||||
|
||||
def validate_strategy(payload: dict[str, Any]) -> dict[str, Any]:
|
||||
"""Validate and normalize a strategy payload. Raise StrategyValidationError on issues."""
|
||||
|
||||
if not isinstance(payload, dict):
|
||||
raise StrategyValidationError("strategy payload must be an object")
|
||||
|
||||
vessel_ais = payload.get("vessel_ais")
|
||||
if not isinstance(vessel_ais, dict):
|
||||
raise StrategyValidationError("strategy.vessel_ais is required and must be an object")
|
||||
|
||||
allow_dynamic_lock = bool(vessel_ais.get("allow_dynamic_lock", False))
|
||||
source_priority = _coerce_str_list(
|
||||
vessel_ais.get("source_priority"),
|
||||
label="vessel_ais.source_priority",
|
||||
)
|
||||
|
||||
raw_rules = vessel_ais.get("field_rules") or {}
|
||||
if not isinstance(raw_rules, dict):
|
||||
raise StrategyValidationError("vessel_ais.field_rules must be an object")
|
||||
field_rules: dict[str, dict[str, Any]] = {}
|
||||
for field, rule in raw_rules.items():
|
||||
if field not in ALLOWED_FIELDS:
|
||||
raise StrategyValidationError(f"unknown vessel_ais field: {field}")
|
||||
if not isinstance(rule, dict):
|
||||
raise StrategyValidationError(f"field_rules.{field} must be an object")
|
||||
mode = str(rule.get("mode") or "").strip()
|
||||
if not mode:
|
||||
raise StrategyValidationError(f"field_rules.{field}.mode is required")
|
||||
is_dynamic = field in DYNAMIC_FIELDS
|
||||
if is_dynamic:
|
||||
allowed_modes = ALLOWED_LOCKED_DYNAMIC_MODES if allow_dynamic_lock else ALLOWED_DYNAMIC_MODES
|
||||
if mode not in allowed_modes:
|
||||
if not allow_dynamic_lock:
|
||||
raise StrategyValidationError(
|
||||
f"field_rules.{field}.mode='{mode}' requires allow_dynamic_lock=true"
|
||||
)
|
||||
raise StrategyValidationError(
|
||||
f"field_rules.{field}.mode must be one of {sorted(allowed_modes)}"
|
||||
)
|
||||
else:
|
||||
if mode not in ALLOWED_STATIC_MODES:
|
||||
raise StrategyValidationError(
|
||||
f"field_rules.{field}.mode must be one of {sorted(ALLOWED_STATIC_MODES)}"
|
||||
)
|
||||
normalized_rule: dict[str, Any] = {"mode": mode}
|
||||
rule_priority = rule.get("source_priority")
|
||||
if rule_priority is not None:
|
||||
normalized_rule["source_priority"] = _coerce_str_list(
|
||||
rule_priority,
|
||||
label=f"field_rules.{field}.source_priority",
|
||||
)
|
||||
if mode == "locked":
|
||||
locked_source = rule.get("locked_source")
|
||||
if not isinstance(locked_source, str) or not locked_source.strip():
|
||||
raise StrategyValidationError(
|
||||
f"field_rules.{field}.locked_source must be a non-empty string when mode=locked"
|
||||
)
|
||||
normalized_rule["locked_source"] = locked_source.strip()
|
||||
field_rules[field] = normalized_rule
|
||||
|
||||
raw_freshness = vessel_ais.get("freshness") or {}
|
||||
if not isinstance(raw_freshness, dict):
|
||||
raise StrategyValidationError("vessel_ais.freshness must be an object")
|
||||
freshness: dict[str, int] = {}
|
||||
for key in ("realtime_stream_seconds", "polling_seconds"):
|
||||
value = raw_freshness.get(key, DEFAULT_STRATEGY["vessel_ais"]["freshness"][key])
|
||||
try:
|
||||
seconds = int(value)
|
||||
except (TypeError, ValueError) as exc:
|
||||
raise StrategyValidationError(f"freshness.{key} must be an integer") from exc
|
||||
if seconds < 0:
|
||||
raise StrategyValidationError(f"freshness.{key} must be non-negative")
|
||||
freshness[key] = seconds
|
||||
|
||||
return {
|
||||
"version": int(payload.get("version") or 0) + 1,
|
||||
"vessel_ais": {
|
||||
"source_priority": source_priority,
|
||||
"field_rules": field_rules,
|
||||
"freshness": freshness,
|
||||
"allow_dynamic_lock": allow_dynamic_lock,
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
async def _select_setting(db: AsyncSession) -> SystemSetting | None:
|
||||
result = await db.execute(
|
||||
select(SystemSetting).where(SystemSetting.category == VESSEL_AGGREGATION_STRATEGY_CATEGORY)
|
||||
)
|
||||
return result.scalar_one_or_none()
|
||||
|
||||
|
||||
def _current_version(setting: SystemSetting | None) -> int:
|
||||
if setting is None:
|
||||
return 0
|
||||
payload = setting.payload or {}
|
||||
return int(payload.get("version") or 0)
|
||||
|
||||
|
||||
async def load_strategy(db: AsyncSession) -> dict[str, Any]:
|
||||
setting = await _select_setting(db)
|
||||
if setting is None or not isinstance(setting.payload, dict):
|
||||
return DEFAULT_STRATEGY
|
||||
payload = setting.payload
|
||||
if "vessel_ais" not in payload:
|
||||
return DEFAULT_STRATEGY
|
||||
return payload
|
||||
|
||||
|
||||
async def save_strategy(db: AsyncSession, payload: dict[str, Any]) -> dict[str, Any]:
|
||||
"""Validate + persist; bumps version automatically."""
|
||||
|
||||
existing = await _select_setting(db)
|
||||
incoming = dict(payload)
|
||||
incoming.setdefault("version", _current_version(existing))
|
||||
validated = validate_strategy(incoming)
|
||||
|
||||
if existing is None:
|
||||
existing = SystemSetting(category=VESSEL_AGGREGATION_STRATEGY_CATEGORY, payload=validated)
|
||||
db.add(existing)
|
||||
else:
|
||||
existing.payload = validated
|
||||
await db.commit()
|
||||
return validated
|
||||
|
||||
|
||||
async def reset_strategy(db: AsyncSession) -> dict[str, Any]:
|
||||
existing = await _select_setting(db)
|
||||
payload = {**DEFAULT_STRATEGY, "version": _current_version(existing) + 1}
|
||||
if existing is None:
|
||||
existing = SystemSetting(category=VESSEL_AGGREGATION_STRATEGY_CATEGORY, payload=payload)
|
||||
db.add(existing)
|
||||
else:
|
||||
existing.payload = payload
|
||||
await db.commit()
|
||||
return payload
|
||||
@@ -1,6 +1,6 @@
|
||||
"""AIS raw observation and aggregation support for vessel collectors."""
|
||||
|
||||
from datetime import UTC, datetime
|
||||
from datetime import UTC, datetime, timedelta
|
||||
from hashlib import sha256
|
||||
import json
|
||||
from typing import Any, Iterable
|
||||
@@ -9,9 +9,14 @@ from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.models.vessel import AISConflictRecord, AISRawObservation, AISSourceHealth
|
||||
from app.services.vessel_aggregation_strategy import (
|
||||
DEFAULT_STRATEGY,
|
||||
load_strategy,
|
||||
)
|
||||
from app.services.vessel_types import normalize_vessel_type_name
|
||||
|
||||
VESSEL_AIS_SCHEMA = "vessel_ais"
|
||||
DEFAULT_AGGREGATION_WINDOW_HOURS = 24
|
||||
BARENTSWATCH_DELIVERY_MODE = "polling"
|
||||
BARENTSWATCH_TRANSPORT = "http"
|
||||
AISSTREAM_DELIVERY_MODE = "realtime_stream"
|
||||
@@ -171,13 +176,43 @@ def _is_future_observation(observation: AISRawObservation, now: datetime) -> boo
|
||||
return observation.observed_at > now
|
||||
|
||||
|
||||
def _strategy_source_rank(
|
||||
source: str,
|
||||
strategy: dict[str, Any],
|
||||
) -> int:
|
||||
priority = (strategy.get("vessel_ais") or {}).get("source_priority") or []
|
||||
if source in priority:
|
||||
return len(priority) - priority.index(source)
|
||||
return 0
|
||||
|
||||
|
||||
def _is_stream_stale(
|
||||
observation: AISRawObservation,
|
||||
*,
|
||||
now: datetime,
|
||||
strategy: dict[str, Any],
|
||||
) -> bool:
|
||||
delivery_mode = str(observation.delivery_mode or "")
|
||||
freshness = (strategy.get("vessel_ais") or {}).get("freshness") or {}
|
||||
if delivery_mode == "realtime_stream":
|
||||
window = int(freshness.get("realtime_stream_seconds", 0) or 0)
|
||||
else:
|
||||
window = int(freshness.get("polling_seconds", 0) or 0)
|
||||
if window <= 0:
|
||||
return False
|
||||
return (now - observation.observed_at).total_seconds() > window
|
||||
|
||||
|
||||
def _select_position_observation(
|
||||
observations: list[AISRawObservation],
|
||||
*,
|
||||
now: datetime,
|
||||
strategy: dict[str, Any] | None = None,
|
||||
) -> tuple[AISRawObservation | None, list[str]]:
|
||||
strategy = strategy or DEFAULT_STRATEGY
|
||||
rejected_flags: list[str] = []
|
||||
candidates = []
|
||||
fresh_candidates: list[AISRawObservation] = []
|
||||
stale_candidates: list[AISRawObservation] = []
|
||||
for observation in observations:
|
||||
payload = observation.normalized_payload or {}
|
||||
if not _has_valid_position(payload):
|
||||
@@ -186,8 +221,13 @@ def _select_position_observation(
|
||||
if _is_future_observation(observation, now):
|
||||
rejected_flags.append("future_timestamp")
|
||||
continue
|
||||
candidates.append(observation)
|
||||
if _is_stream_stale(observation, now=now, strategy=strategy):
|
||||
stale_candidates.append(observation)
|
||||
rejected_flags.append("freshness_fallback")
|
||||
continue
|
||||
fresh_candidates.append(observation)
|
||||
|
||||
candidates = fresh_candidates or stale_candidates
|
||||
if not candidates:
|
||||
return None, sorted(set(rejected_flags))
|
||||
|
||||
@@ -195,6 +235,7 @@ def _select_position_observation(
|
||||
key=lambda item: (
|
||||
item.observed_at,
|
||||
_delivery_priority(item),
|
||||
_strategy_source_rank(item.source, strategy),
|
||||
item.collected_at,
|
||||
item.id or 0,
|
||||
),
|
||||
@@ -206,7 +247,9 @@ def _select_position_observation(
|
||||
def _select_static_field(
|
||||
observations: list[AISRawObservation],
|
||||
field: str,
|
||||
strategy: dict[str, Any] | None = None,
|
||||
) -> tuple[Any, str | None, str | None]:
|
||||
strategy = strategy or DEFAULT_STRATEGY
|
||||
candidates = []
|
||||
for observation in observations:
|
||||
value = _payload_value(observation.normalized_payload or {}, field)
|
||||
@@ -219,6 +262,39 @@ def _select_static_field(
|
||||
if not candidates:
|
||||
return None, None, None
|
||||
|
||||
field_rules = (strategy.get("vessel_ais") or {}).get("field_rules") or {}
|
||||
rule = field_rules.get(field) or {"mode": "source_priority"}
|
||||
mode = rule.get("mode")
|
||||
|
||||
if mode == "locked":
|
||||
locked_source = rule.get("locked_source")
|
||||
for observation, value in candidates:
|
||||
if observation.source == locked_source:
|
||||
return value, observation.source, "locked"
|
||||
|
||||
if mode in ("source_priority", "locked"):
|
||||
priority = rule.get("source_priority") or (strategy.get("vessel_ais") or {}).get("source_priority") or []
|
||||
ranked = sorted(
|
||||
candidates,
|
||||
key=lambda item: (
|
||||
priority.index(item[0].source) if item[0].source in priority else len(priority) + 1,
|
||||
-_delivery_priority(item[0]),
|
||||
-(item[0].observed_at.timestamp() if item[0].observed_at else 0),
|
||||
),
|
||||
)
|
||||
observation, value = ranked[0]
|
||||
return value, observation.source, "source_priority"
|
||||
|
||||
if mode == "newest":
|
||||
ranked = sorted(
|
||||
candidates,
|
||||
key=lambda item: (item[0].observed_at, _delivery_priority(item[0]), item[0].id or 0),
|
||||
reverse=True,
|
||||
)
|
||||
observation, value = ranked[0]
|
||||
return value, observation.source, "newest_observation"
|
||||
|
||||
# default / non_empty: prefer delivery mode priority, then newest
|
||||
candidates.sort(
|
||||
key=lambda item: (
|
||||
_delivery_priority(item[0]),
|
||||
@@ -261,8 +337,12 @@ def _build_aggregated_vessel(
|
||||
observations: list[AISRawObservation],
|
||||
*,
|
||||
now: datetime,
|
||||
strategy: dict[str, Any] | None = None,
|
||||
) -> dict[str, Any] | None:
|
||||
position_observation, rejected_flags = _select_position_observation(observations, now=now)
|
||||
strategy = strategy or DEFAULT_STRATEGY
|
||||
position_observation, rejected_flags = _select_position_observation(
|
||||
observations, now=now, strategy=strategy
|
||||
)
|
||||
if position_observation is None:
|
||||
return None
|
||||
|
||||
@@ -279,6 +359,7 @@ def _build_aggregated_vessel(
|
||||
"quality_flags": sorted(
|
||||
set((position_observation.quality_flags or []) + rejected_flags)
|
||||
),
|
||||
"aggregation_strategy_version": int(strategy.get("version") or 0),
|
||||
}
|
||||
|
||||
for field in DYNAMIC_FIELDS:
|
||||
@@ -289,7 +370,9 @@ def _build_aggregated_vessel(
|
||||
result["selected_reasons"][field] = "newest_observation"
|
||||
|
||||
for field in CONFLICT_FIELDS:
|
||||
selected_value, selected_source, reason = _select_static_field(observations, field)
|
||||
selected_value, selected_source, reason = _select_static_field(
|
||||
observations, field, strategy=strategy
|
||||
)
|
||||
if selected_value is None:
|
||||
continue
|
||||
result[field] = selected_value
|
||||
@@ -408,12 +491,16 @@ async def aggregate_vessel_observations(
|
||||
db: AsyncSession,
|
||||
observations: Iterable[AISRawObservation],
|
||||
*,
|
||||
write_conflicts: bool = True,
|
||||
write_conflicts: bool = False,
|
||||
strategy: dict[str, Any] | None = None,
|
||||
) -> list[dict[str, Any]]:
|
||||
strategy = strategy if strategy is not None else await _safe_load_strategy(db)
|
||||
now = datetime.now(UTC)
|
||||
vessels = []
|
||||
for entity_key, entity_observations in _group_observations(observations).items():
|
||||
aggregated = _build_aggregated_vessel(entity_key, entity_observations, now=now)
|
||||
aggregated = _build_aggregated_vessel(
|
||||
entity_key, entity_observations, now=now, strategy=strategy
|
||||
)
|
||||
if aggregated is None:
|
||||
continue
|
||||
if write_conflicts:
|
||||
@@ -431,15 +518,28 @@ async def aggregate_vessel_observations(
|
||||
return vessels
|
||||
|
||||
|
||||
async def _safe_load_strategy(db: AsyncSession) -> dict[str, Any]:
|
||||
"""Tolerate fake test sessions where load_strategy may misbehave."""
|
||||
try:
|
||||
return await load_strategy(db)
|
||||
except Exception:
|
||||
return DEFAULT_STRATEGY
|
||||
|
||||
|
||||
async def get_aggregated_vessels(
|
||||
db: AsyncSession,
|
||||
*,
|
||||
bbox: tuple[float, float, float, float] | None = None,
|
||||
limit: int | None = None,
|
||||
observed_since: datetime | None = None,
|
||||
) -> list[dict[str, Any]]:
|
||||
observed_since = observed_since or (
|
||||
datetime.now(UTC) - timedelta(hours=DEFAULT_AGGREGATION_WINDOW_HOURS)
|
||||
)
|
||||
stmt = (
|
||||
select(AISRawObservation)
|
||||
.where(AISRawObservation.target_schema == VESSEL_AIS_SCHEMA)
|
||||
.where(AISRawObservation.observed_at >= observed_since)
|
||||
.order_by(AISRawObservation.observed_at.desc(), AISRawObservation.id.desc())
|
||||
)
|
||||
if limit and limit > 0:
|
||||
@@ -545,6 +645,30 @@ async def update_ais_source_health(
|
||||
return health
|
||||
|
||||
|
||||
async def count_unique_raw_vessel_mmsi(
|
||||
db: AsyncSession,
|
||||
*,
|
||||
observed_since: datetime | None = None,
|
||||
) -> int:
|
||||
"""Count unique raw vessel MMSI values for HUD counts; never aggregates."""
|
||||
from sqlalchemy import func as sa_func
|
||||
|
||||
unique_mmsi_stmt = (
|
||||
select(AISRawObservation.entity_key)
|
||||
.where(AISRawObservation.target_schema == VESSEL_AIS_SCHEMA)
|
||||
.distinct()
|
||||
)
|
||||
if observed_since is not None:
|
||||
unique_mmsi_stmt = unique_mmsi_stmt.where(
|
||||
AISRawObservation.observed_at >= observed_since,
|
||||
)
|
||||
|
||||
result = await db.execute(
|
||||
select(sa_func.count()).select_from(unique_mmsi_stmt.subquery()),
|
||||
)
|
||||
return int(result.scalar() or 0)
|
||||
|
||||
|
||||
async def get_vessel_raw_observations(
|
||||
db: AsyncSession,
|
||||
mmsi: int,
|
||||
|
||||
109
backend/app/services/vessel_enrichment.py
Normal file
109
backend/app/services/vessel_enrichment.py
Normal file
@@ -0,0 +1,109 @@
|
||||
"""v5 vessel enrichment service.
|
||||
|
||||
Read-only side: `get_vessel_enrichment_bundle` is the only path the
|
||||
aggregation/detail endpoints use. It never reaches out to third parties; it
|
||||
just returns whatever the upsert side has already cached. Expired rows are
|
||||
filtered out so old data never leaks back into the live UI.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import UTC, datetime
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.models.vessel_enrichment import VesselMediaEnrichment, VesselProfileEnrichment
|
||||
|
||||
|
||||
def _coerce_datetime(value: Any) -> datetime | None:
|
||||
if value in (None, ""):
|
||||
return None
|
||||
if isinstance(value, datetime):
|
||||
return value if value.tzinfo else value.replace(tzinfo=UTC)
|
||||
if isinstance(value, (int, float)):
|
||||
ts = float(value)
|
||||
if ts > 10_000_000_000:
|
||||
ts /= 1000
|
||||
return datetime.fromtimestamp(ts, 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
|
||||
|
||||
|
||||
def _build_payload(record, *, now: datetime) -> dict[str, Any] | None:
|
||||
if record is None:
|
||||
return None
|
||||
expires_at = record.expires_at
|
||||
if isinstance(expires_at, datetime):
|
||||
if expires_at.tzinfo is None:
|
||||
expires_at = expires_at.replace(tzinfo=UTC)
|
||||
if expires_at < now:
|
||||
return None
|
||||
return record.to_dict()
|
||||
|
||||
|
||||
async def get_vessel_enrichment_bundle(db: AsyncSession, mmsi: int) -> dict[str, Any]:
|
||||
now = datetime.now(UTC)
|
||||
profile = await db.get(VesselProfileEnrichment, mmsi)
|
||||
media = await db.get(VesselMediaEnrichment, mmsi)
|
||||
return {
|
||||
"mmsi": mmsi,
|
||||
"profile": _build_payload(profile, now=now),
|
||||
"media": _build_payload(media, now=now),
|
||||
}
|
||||
|
||||
|
||||
async def upsert_vessel_profile_enrichment(
|
||||
db: AsyncSession,
|
||||
*,
|
||||
mmsi: int,
|
||||
payload: dict[str, Any],
|
||||
) -> dict[str, Any]:
|
||||
record = await db.get(VesselProfileEnrichment, mmsi)
|
||||
if record is None:
|
||||
record = VesselProfileEnrichment(mmsi=mmsi)
|
||||
db.add(record)
|
||||
return _apply_upsert(record, payload)
|
||||
|
||||
|
||||
async def upsert_vessel_media_enrichment(
|
||||
db: AsyncSession,
|
||||
*,
|
||||
mmsi: int,
|
||||
payload: dict[str, Any],
|
||||
) -> dict[str, Any]:
|
||||
record = await db.get(VesselMediaEnrichment, mmsi)
|
||||
if record is None:
|
||||
record = VesselMediaEnrichment(mmsi=mmsi)
|
||||
db.add(record)
|
||||
return _apply_upsert(record, payload)
|
||||
|
||||
|
||||
def _apply_upsert(record, payload: dict[str, Any]) -> dict[str, Any]:
|
||||
if not isinstance(payload, dict):
|
||||
raise ValueError("enrichment payload must be an object")
|
||||
body = payload.get("payload")
|
||||
if body is not None and not isinstance(body, dict):
|
||||
raise ValueError("payload.payload must be an object")
|
||||
if body is not None:
|
||||
record.payload = body
|
||||
if "source" in payload and isinstance(payload["source"], str) and payload["source"].strip():
|
||||
record.source = payload["source"].strip()
|
||||
fetched_at = _coerce_datetime(payload.get("fetched_at"))
|
||||
record.fetched_at = fetched_at or datetime.now(UTC)
|
||||
record.expires_at = _coerce_datetime(payload.get("expires_at"))
|
||||
confidence = payload.get("confidence")
|
||||
if confidence is not None:
|
||||
try:
|
||||
record.confidence = float(confidence)
|
||||
except (TypeError, ValueError):
|
||||
record.confidence = None
|
||||
if "reference_url" in payload:
|
||||
ref = payload.get("reference_url")
|
||||
record.reference_url = str(ref) if ref else None
|
||||
return record.to_dict()
|
||||
Reference in New Issue
Block a user