release: bump version to 0.48.0
This commit is contained in:
@@ -12,6 +12,7 @@ from app.api.v1 import (
|
||||
settings,
|
||||
collected_data,
|
||||
visualization,
|
||||
vessel_aggregation,
|
||||
bgp,
|
||||
news,
|
||||
system_control,
|
||||
@@ -34,6 +35,11 @@ api_router.include_router(alerts.router, prefix="/alerts", tags=["alerts"])
|
||||
api_router.include_router(settings.router, prefix="/settings", tags=["settings"])
|
||||
api_router.include_router(system_control.router, prefix="/system", tags=["system"])
|
||||
api_router.include_router(visualization.router, prefix="/visualization", tags=["visualization"])
|
||||
api_router.include_router(
|
||||
vessel_aggregation.router,
|
||||
prefix="/vessel-aggregation",
|
||||
tags=["vessel-aggregation"],
|
||||
)
|
||||
api_router.include_router(bgp.router, prefix="/bgp", tags=["bgp"])
|
||||
api_router.include_router(tv.router, prefix="/tv", tags=["tv"])
|
||||
api_router.include_router(news.router, prefix="/news", tags=["news"])
|
||||
|
||||
@@ -5,8 +5,8 @@ from datetime import datetime
|
||||
import base64
|
||||
import json
|
||||
import re
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
from sqlalchemy import select, func
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query
|
||||
from sqlalchemy import delete, select, func
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from pydantic import BaseModel, Field
|
||||
import httpx
|
||||
@@ -17,6 +17,8 @@ from app.db.session import get_db
|
||||
from app.models.user import User
|
||||
from app.models.datasource_config import DataSourceConfig
|
||||
from app.models.datasource_mapping import DataSourceMappingTemplate
|
||||
from app.models.collected_data import CollectedData
|
||||
from app.models.vessel import AISRawObservation, AISSourceHealth
|
||||
from app.core.security import get_current_user
|
||||
from app.core.cache import cache
|
||||
from app.core.time import to_iso8601_utc
|
||||
@@ -26,10 +28,19 @@ from app.services.datasource_mapping import (
|
||||
MappingError,
|
||||
build_heuristic_mapping,
|
||||
execute_mapping,
|
||||
persist_mapped_records,
|
||||
redact_for_llm,
|
||||
stable_payload_hash,
|
||||
)
|
||||
from app.services.custom_datasource_runtime import (
|
||||
CustomDatasourceRuntimeError,
|
||||
fetch_rest_payload,
|
||||
get_custom_stream_status,
|
||||
run_mapped_rest_config,
|
||||
run_mapped_websocket_config,
|
||||
start_custom_stream,
|
||||
stop_custom_stream,
|
||||
test_websocket_config,
|
||||
)
|
||||
from app.services.datasource_connectivity import (
|
||||
get_builtin_connection_status,
|
||||
save_connectivity_success,
|
||||
@@ -43,7 +54,7 @@ router = APIRouter()
|
||||
class DataSourceConfigCreate(BaseModel):
|
||||
name: str = Field(..., min_length=1, max_length=100)
|
||||
description: Optional[str] = None
|
||||
source_type: str = Field(..., description="http, api, database")
|
||||
source_type: str = Field(..., description="rest, websocket, http, api, database")
|
||||
endpoint: str = Field(..., max_length=500)
|
||||
auth_type: str = Field(default="none", description="none, bearer, api_key, basic")
|
||||
auth_config: dict = Field(default={})
|
||||
@@ -219,6 +230,8 @@ def _build_query_params(auth_type: str, auth_config: dict, config: dict) -> dict
|
||||
|
||||
|
||||
async def fetch_custom_sample_from_config(config: DataSourceConfig, limit_bytes: int) -> Any:
|
||||
if str(config.source_type or "").lower() in {"websocket", "ws"}:
|
||||
raise HTTPException(status_code=400, detail="WebSocket sources must use connection test or run-mapped stream.")
|
||||
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"}:
|
||||
@@ -488,6 +501,8 @@ async def update_config(
|
||||
@router.delete("/configs/{config_id}")
|
||||
async def delete_config(
|
||||
config_id: int,
|
||||
delete_mappings: bool = Query(False),
|
||||
delete_source_data: bool = Query(False),
|
||||
current_user: User = Depends(get_current_user),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
@@ -498,12 +513,59 @@ async def delete_config(
|
||||
if not config:
|
||||
raise HTTPException(status_code=404, detail="Configuration not found")
|
||||
|
||||
deleted_mappings = 0
|
||||
deleted_records = {
|
||||
"collected_data": 0,
|
||||
"ais_raw_observations": 0,
|
||||
"ais_source_health": 0,
|
||||
}
|
||||
|
||||
if delete_source_data:
|
||||
collected_result = await db.execute(
|
||||
delete(CollectedData).where(CollectedData.source == config.name)
|
||||
)
|
||||
raw_result = await db.execute(
|
||||
delete(AISRawObservation).where(AISRawObservation.source == config.name)
|
||||
)
|
||||
health_result = await db.execute(
|
||||
delete(AISSourceHealth).where(AISSourceHealth.source == config.name)
|
||||
)
|
||||
deleted_records = {
|
||||
"collected_data": collected_result.rowcount or 0,
|
||||
"ais_raw_observations": raw_result.rowcount or 0,
|
||||
"ais_source_health": health_result.rowcount or 0,
|
||||
}
|
||||
|
||||
if delete_mappings or delete_source_data:
|
||||
mapping_result = await db.execute(
|
||||
delete(DataSourceMappingTemplate).where(
|
||||
DataSourceMappingTemplate.datasource_config_id == config_id
|
||||
)
|
||||
)
|
||||
deleted_mappings = mapping_result.rowcount or 0
|
||||
|
||||
await db.delete(config)
|
||||
await db.commit()
|
||||
|
||||
cache.delete_pattern("datasource_configs:*")
|
||||
|
||||
return {"message": "Configuration deleted successfully"}
|
||||
if delete_source_data and (config.config or {}).get("target_schema") == "vessel_ais":
|
||||
from app.core.websocket.broadcaster import broadcaster
|
||||
|
||||
await broadcaster.broadcast_custom(
|
||||
"vessels",
|
||||
{
|
||||
"action": "reload",
|
||||
"source": config.name,
|
||||
"reason": "custom_source_deleted",
|
||||
},
|
||||
)
|
||||
|
||||
return {
|
||||
"message": "Configuration deleted successfully",
|
||||
"deleted_mappings": deleted_mappings,
|
||||
"deleted_records": deleted_records,
|
||||
}
|
||||
|
||||
|
||||
@router.post("/configs/{config_id}/test")
|
||||
@@ -520,6 +582,8 @@ async def test_config(
|
||||
raise HTTPException(status_code=404, detail="Configuration not found")
|
||||
|
||||
try:
|
||||
if str(config.source_type or "").lower() in {"websocket", "ws"}:
|
||||
return await test_websocket_config(config)
|
||||
result = await test_endpoint(
|
||||
endpoint=config.endpoint,
|
||||
auth_type=config.auth_type,
|
||||
@@ -550,6 +614,18 @@ async def test_new_config(
|
||||
):
|
||||
"""Test a new data source configuration without saving"""
|
||||
try:
|
||||
if str(config_data.source_type or "").lower() in {"websocket", "ws"}:
|
||||
config = DataSourceConfig(
|
||||
name=config_data.name,
|
||||
description=config_data.description,
|
||||
source_type=config_data.source_type,
|
||||
endpoint=config_data.endpoint,
|
||||
auth_type=config_data.auth_type,
|
||||
auth_config=config_data.auth_config,
|
||||
headers=config_data.headers,
|
||||
config=config_data.config,
|
||||
)
|
||||
return await test_websocket_config(config)
|
||||
result = await test_endpoint(
|
||||
endpoint=config_data.endpoint,
|
||||
auth_type=config_data.auth_type,
|
||||
@@ -875,6 +951,8 @@ async def update_datasource_mapping(
|
||||
@router.post("/{config_id}/run-mapped")
|
||||
async def run_mapped_datasource(
|
||||
config_id: int,
|
||||
background: bool = Query(False, description="For WebSocket sources, start a background stream task."),
|
||||
debug_max_messages: int | None = Query(None, ge=1),
|
||||
current_user: User = Depends(get_current_user),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
@@ -883,20 +961,24 @@ async def run_mapped_datasource(
|
||||
if not datasource:
|
||||
raise HTTPException(status_code=404, detail="Configuration not found")
|
||||
|
||||
result = await db.execute(
|
||||
select(DataSourceMappingTemplate)
|
||||
.where(DataSourceMappingTemplate.datasource_config_id == config_id)
|
||||
.where(DataSourceMappingTemplate.is_active.is_(True))
|
||||
.order_by(DataSourceMappingTemplate.version.desc())
|
||||
.limit(1)
|
||||
)
|
||||
mapping = result.scalar_one_or_none()
|
||||
if not mapping:
|
||||
raise HTTPException(status_code=404, detail="No active mapping template found")
|
||||
|
||||
try:
|
||||
sample = await fetch_custom_sample_from_config(datasource, 5_000_000)
|
||||
mapped = execute_mapping(sample, mapping.mapping_json, mapping.target_schema)
|
||||
if str(datasource.source_type or "").lower() in {"websocket", "ws"}:
|
||||
if background and debug_max_messages is None:
|
||||
started = start_custom_stream(config_id)
|
||||
if not started:
|
||||
raise HTTPException(status_code=409, detail="Custom WebSocket source is already running")
|
||||
return {
|
||||
"status": "started",
|
||||
"datasource_config_id": config_id,
|
||||
"stream": get_custom_stream_status(config_id),
|
||||
}
|
||||
return await run_mapped_websocket_config(
|
||||
db,
|
||||
datasource,
|
||||
debug_max_messages=debug_max_messages,
|
||||
)
|
||||
|
||||
return await run_mapped_rest_config(db, datasource)
|
||||
except httpx.HTTPStatusError as exc:
|
||||
raise HTTPException(
|
||||
status_code=exc.response.status_code,
|
||||
@@ -904,36 +986,26 @@ async def run_mapped_datasource(
|
||||
) from exc
|
||||
except httpx.HTTPError as exc:
|
||||
raise HTTPException(status_code=502, detail=f"Datasource request failed: {exc}") from exc
|
||||
except (MappingError, ValueError) as exc:
|
||||
except (CustomDatasourceRuntimeError, MappingError, ValueError) as exc:
|
||||
raise HTTPException(status_code=400, detail=f"Mapping failed: {exc}") from exc
|
||||
|
||||
if mapped["failed_count"] > 0:
|
||||
return {
|
||||
"status": "failed",
|
||||
"datasource_config_id": config_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],
|
||||
}
|
||||
|
||||
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,
|
||||
)
|
||||
@router.post("/{config_id}/stop-mapped")
|
||||
async def stop_mapped_datasource(
|
||||
config_id: int,
|
||||
current_user: User = Depends(get_current_user),
|
||||
):
|
||||
stopped = await stop_custom_stream(config_id)
|
||||
return {
|
||||
"status": "success",
|
||||
"status": "stopped" if stopped else "not_running",
|
||||
"datasource_config_id": config_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,
|
||||
"stream": get_custom_stream_status(config_id),
|
||||
}
|
||||
|
||||
|
||||
@router.get("/{config_id}/stream-status")
|
||||
async def get_mapped_stream_status(
|
||||
config_id: int,
|
||||
current_user: User = Depends(get_current_user),
|
||||
):
|
||||
return get_custom_stream_status(config_id)
|
||||
|
||||
132
backend/app/api/v1/vessel_aggregation.py
Normal file
132
backend/app/api/v1/vessel_aggregation.py
Normal file
@@ -0,0 +1,132 @@
|
||||
"""v4 strategy + v5 conflict-promotion + enrichment APIs for vessel_ais."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.core.security import get_current_user
|
||||
from app.db.session import get_db
|
||||
from app.models.user import User
|
||||
from app.models.vessel import AISConflictRecord
|
||||
from app.services.vessel_aggregation_strategy import (
|
||||
StrategyValidationError,
|
||||
load_strategy,
|
||||
reset_strategy,
|
||||
save_strategy,
|
||||
)
|
||||
from app.services.vessel_enrichment import (
|
||||
get_vessel_enrichment_bundle,
|
||||
upsert_vessel_media_enrichment,
|
||||
upsert_vessel_profile_enrichment,
|
||||
)
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
@router.get("/strategy")
|
||||
async def get_aggregation_strategy(db: AsyncSession = Depends(get_db)):
|
||||
return await load_strategy(db)
|
||||
|
||||
|
||||
@router.put("/strategy")
|
||||
async def put_aggregation_strategy(
|
||||
payload: dict[str, Any],
|
||||
current_user: User = Depends(get_current_user),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
try:
|
||||
return await save_strategy(db, payload)
|
||||
except StrategyValidationError as exc:
|
||||
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||||
|
||||
|
||||
@router.delete("/strategy")
|
||||
async def reset_aggregation_strategy(
|
||||
current_user: User = Depends(get_current_user),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
return await reset_strategy(db)
|
||||
|
||||
|
||||
@router.post("/conflicts/{mmsi}/{field}/promote-to-rule")
|
||||
async def promote_conflict_to_rule(
|
||||
mmsi: int,
|
||||
field: str,
|
||||
current_user: User = Depends(get_current_user),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
"""Lift the current conflict resolution into a persistent strategy rule."""
|
||||
|
||||
result = await db.execute(
|
||||
select(AISConflictRecord)
|
||||
.where(AISConflictRecord.target_schema == "vessel_ais")
|
||||
.where(AISConflictRecord.entity_key == str(mmsi))
|
||||
.where(AISConflictRecord.field == field)
|
||||
.order_by(AISConflictRecord.updated_at.desc(), AISConflictRecord.id.desc())
|
||||
.limit(1)
|
||||
)
|
||||
record = result.scalar_one_or_none()
|
||||
if record is None or not record.selected_source:
|
||||
raise HTTPException(status_code=404, detail="Conflict record with selected_source not found")
|
||||
|
||||
strategy = await load_strategy(db)
|
||||
vessel_ais = dict(strategy.get("vessel_ais") or {})
|
||||
field_rules = dict(vessel_ais.get("field_rules") or {})
|
||||
field_rules[field] = {"mode": "source_priority", "source_priority": [record.selected_source]}
|
||||
vessel_ais["field_rules"] = field_rules
|
||||
|
||||
incoming = {"version": int(strategy.get("version") or 0), "vessel_ais": vessel_ais}
|
||||
try:
|
||||
return await save_strategy(db, incoming)
|
||||
except StrategyValidationError as exc:
|
||||
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||||
|
||||
|
||||
@router.delete("/conflicts/{mmsi}/{field}/promote-to-rule")
|
||||
async def revert_conflict_rule(
|
||||
mmsi: int,
|
||||
field: str,
|
||||
current_user: User = Depends(get_current_user),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
strategy = await load_strategy(db)
|
||||
vessel_ais = dict(strategy.get("vessel_ais") or {})
|
||||
field_rules = dict(vessel_ais.get("field_rules") or {})
|
||||
if field in field_rules:
|
||||
del field_rules[field]
|
||||
vessel_ais["field_rules"] = field_rules
|
||||
|
||||
incoming = {"version": int(strategy.get("version") or 0), "vessel_ais": vessel_ais}
|
||||
try:
|
||||
return await save_strategy(db, incoming)
|
||||
except StrategyValidationError as exc:
|
||||
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||||
|
||||
|
||||
@router.get("/enrichment/{mmsi}")
|
||||
async def get_vessel_enrichment(mmsi: int, db: AsyncSession = Depends(get_db)):
|
||||
return await get_vessel_enrichment_bundle(db, mmsi)
|
||||
|
||||
|
||||
@router.put("/enrichment/{mmsi}/profile")
|
||||
async def put_vessel_profile_enrichment(
|
||||
mmsi: int,
|
||||
payload: dict[str, Any],
|
||||
current_user: User = Depends(get_current_user),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
return await upsert_vessel_profile_enrichment(db, mmsi=mmsi, payload=payload)
|
||||
|
||||
|
||||
@router.put("/enrichment/{mmsi}/media")
|
||||
async def put_vessel_media_enrichment(
|
||||
mmsi: int,
|
||||
payload: dict[str, Any],
|
||||
current_user: User = Depends(get_current_user),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
return await upsert_vessel_media_enrichment(db, mmsi=mmsi, payload=payload)
|
||||
@@ -6,6 +6,7 @@ Returns GeoJSON format compatible with Three.js, CesiumJS, and Unreal Cesium.
|
||||
|
||||
from datetime import UTC, datetime, timedelta
|
||||
import math
|
||||
import re
|
||||
import httpx
|
||||
from fastapi import APIRouter, HTTPException, Depends, Query, Response
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
@@ -19,14 +20,16 @@ from app.core.time import to_iso8601_utc
|
||||
from app.db.session import get_db
|
||||
from app.models.bgp_anomaly import BGPAnomaly
|
||||
from app.models.bgp_incident import BGPIncident
|
||||
from app.models.bgp_observation import BGPObservation
|
||||
from app.models.collected_data import CollectedData
|
||||
from app.models.vessel import VesselPosition, VesselStatic
|
||||
from app.models.vessel import AISSourceHealth, VesselPosition, VesselStatic
|
||||
from app.services.bgp_collectors import build_bgp_collector_coverage
|
||||
from app.services.cable_graph import build_graph_from_data, CableGraph, haversine_distance
|
||||
from app.services.collectors.bgp_common import RIPE_RIS_COLLECTOR_COORDS
|
||||
from app.services.persistent_logs import record_system_log
|
||||
from app.services.vessel_ais_aggregation import (
|
||||
build_field_conflict_candidates,
|
||||
count_unique_raw_vessel_mmsi,
|
||||
get_aggregated_vessel,
|
||||
get_aggregated_vessel_track,
|
||||
get_aggregated_vessels,
|
||||
@@ -40,6 +43,7 @@ logger = get_logger(__name__, service="api")
|
||||
TERRAIN_TILE_URL_TEMPLATE = (
|
||||
"https://s3.amazonaws.com/elevation-tiles-prod/terrarium/{z}/{x}/{y}.png"
|
||||
)
|
||||
VESSEL_NAME_FALLBACK_PATTERN = re.compile(r"^mmsi\s*\d+$", re.IGNORECASE)
|
||||
|
||||
|
||||
# ============== Converter Functions ==============
|
||||
@@ -281,6 +285,120 @@ async def _load_current_collected_data(
|
||||
return list(result.scalars().all())
|
||||
|
||||
|
||||
async def _latest_task_id_for_source(
|
||||
db: AsyncSession,
|
||||
source: str,
|
||||
*,
|
||||
exclude_unknown_name: bool = False,
|
||||
) -> int | None:
|
||||
stmt = (
|
||||
select(
|
||||
CollectedData.task_id,
|
||||
func.max(CollectedData.collected_at).label("latest_collected_at"),
|
||||
func.max(CollectedData.id).label("latest_id"),
|
||||
)
|
||||
.where(CollectedData.source == source)
|
||||
.where(CollectedData.task_id.isnot(None))
|
||||
.group_by(CollectedData.task_id)
|
||||
.order_by(func.max(CollectedData.collected_at).desc(), func.max(CollectedData.id).desc())
|
||||
.limit(1)
|
||||
)
|
||||
if exclude_unknown_name:
|
||||
stmt = stmt.where(CollectedData.name != "Unknown")
|
||||
|
||||
result = await db.execute(stmt)
|
||||
row = result.first()
|
||||
return int(row.task_id) if row and row.task_id is not None else None
|
||||
|
||||
|
||||
async def _load_current_or_latest_task_data(
|
||||
db: AsyncSession,
|
||||
source: str,
|
||||
*,
|
||||
exclude_unknown_name: bool = False,
|
||||
limit: Optional[int] = None,
|
||||
) -> List[CollectedData]:
|
||||
records = await _load_current_collected_data(
|
||||
db,
|
||||
source,
|
||||
exclude_unknown_name=exclude_unknown_name,
|
||||
limit=limit,
|
||||
)
|
||||
if records:
|
||||
return records
|
||||
|
||||
latest_task_id = await _latest_task_id_for_source(
|
||||
db,
|
||||
source,
|
||||
exclude_unknown_name=exclude_unknown_name,
|
||||
)
|
||||
if latest_task_id is None:
|
||||
return []
|
||||
|
||||
stmt = (
|
||||
select(CollectedData)
|
||||
.where(CollectedData.source == source)
|
||||
.where(CollectedData.task_id == latest_task_id)
|
||||
.order_by(CollectedData.id.desc())
|
||||
)
|
||||
if exclude_unknown_name:
|
||||
stmt = stmt.where(CollectedData.name != "Unknown")
|
||||
if limit is not None:
|
||||
stmt = stmt.limit(limit)
|
||||
|
||||
result = await db.execute(stmt)
|
||||
return list(result.scalars().all())
|
||||
|
||||
|
||||
async def _count_current_or_latest_task_data(
|
||||
db: AsyncSession,
|
||||
source: str,
|
||||
*,
|
||||
exclude_unknown_name: bool = False,
|
||||
) -> int:
|
||||
current_stmt = (
|
||||
select(func.count(CollectedData.id))
|
||||
.where(CollectedData.source == source)
|
||||
.where(CollectedData.is_current.is_(True))
|
||||
)
|
||||
if exclude_unknown_name:
|
||||
current_stmt = current_stmt.where(CollectedData.name != "Unknown")
|
||||
|
||||
current_result = await db.execute(current_stmt)
|
||||
current_scalar = current_result.scalar()
|
||||
if current_scalar is None and hasattr(current_result, "scalars"):
|
||||
current_rows = current_result.scalars().all()
|
||||
current_count = sum(
|
||||
1
|
||||
for row in current_rows
|
||||
if getattr(row, "source", None) == source
|
||||
and (not exclude_unknown_name or getattr(row, "name", None) != "Unknown")
|
||||
)
|
||||
else:
|
||||
current_count = int(current_scalar or 0)
|
||||
if current_count > 0:
|
||||
return current_count
|
||||
|
||||
latest_task_id = await _latest_task_id_for_source(
|
||||
db,
|
||||
source,
|
||||
exclude_unknown_name=exclude_unknown_name,
|
||||
)
|
||||
if latest_task_id is None:
|
||||
return 0
|
||||
|
||||
latest_stmt = (
|
||||
select(func.count(CollectedData.id))
|
||||
.where(CollectedData.source == source)
|
||||
.where(CollectedData.task_id == latest_task_id)
|
||||
)
|
||||
if exclude_unknown_name:
|
||||
latest_stmt = latest_stmt.where(CollectedData.name != "Unknown")
|
||||
|
||||
latest_result = await db.execute(latest_stmt)
|
||||
return int(latest_result.scalar() or 0)
|
||||
|
||||
|
||||
async def _load_current_collected_data_by_sources(
|
||||
db: AsyncSession,
|
||||
sources: List[str],
|
||||
@@ -636,14 +754,21 @@ VESSEL_TYPE_FILTERS = {
|
||||
|
||||
def convert_vessels_to_geojson(rows: List[Any]) -> Dict[str, Any]:
|
||||
features = []
|
||||
seen_mmsi: set[int] = set()
|
||||
for position, static in rows:
|
||||
if position.lat is None or position.lon is None:
|
||||
continue
|
||||
if position.mmsi in seen_mmsi:
|
||||
continue
|
||||
seen_mmsi.add(position.mmsi)
|
||||
props = {
|
||||
"mmsi": position.mmsi,
|
||||
"mmsi_display": str(position.mmsi),
|
||||
"name": getattr(static, "name", None) or f"MMSI {position.mmsi}",
|
||||
"name_is_fallback": _is_vessel_name_fallback(getattr(static, "name", None), position.mmsi),
|
||||
"callsign": getattr(static, "callsign", None),
|
||||
"imo": getattr(static, "imo", None),
|
||||
"imo_display": str(getattr(static, "imo")) if getattr(static, "imo", None) else None,
|
||||
"vessel_type": getattr(static, "vessel_type", None),
|
||||
"vessel_type_name": getattr(static, "vessel_type_name", None) or "Other",
|
||||
"flag": getattr(static, "flag", None),
|
||||
@@ -685,9 +810,12 @@ def convert_aggregated_vessels_to_geojson(vessels: List[dict[str, Any]]) -> Dict
|
||||
}
|
||||
props = {
|
||||
"mmsi": vessel["mmsi"],
|
||||
"mmsi_display": str(vessel["mmsi"]),
|
||||
"name": vessel.get("name") or f"MMSI {vessel['mmsi']}",
|
||||
"name_is_fallback": _is_vessel_name_fallback(vessel.get("name"), vessel["mmsi"]),
|
||||
"callsign": vessel.get("callsign"),
|
||||
"imo": vessel.get("imo"),
|
||||
"imo_display": str(vessel.get("imo")) if vessel.get("imo") else None,
|
||||
"vessel_type": vessel.get("vessel_type"),
|
||||
"vessel_type_name": vessel.get("vessel_type_name") or "Other",
|
||||
"flag": vessel.get("flag"),
|
||||
@@ -704,6 +832,7 @@ def convert_aggregated_vessels_to_geojson(vessels: List[dict[str, Any]]) -> Dict
|
||||
"source_summary": source_summary,
|
||||
"quality_flags": vessel.get("quality_flags") or [],
|
||||
"conflict_count": vessel.get("conflict_count", 0),
|
||||
"aggregation_strategy_version": vessel.get("aggregation_strategy_version", 0),
|
||||
"data_type": "vessel",
|
||||
}
|
||||
features.append(
|
||||
@@ -737,6 +866,24 @@ def _parse_bbox(value: Optional[str]) -> tuple[float, float, float, float] | Non
|
||||
return lon_min, lat_min, lon_max, lat_max
|
||||
|
||||
|
||||
def _is_vessel_name_fallback(name: Any, mmsi: Any) -> bool:
|
||||
text = str(name or "").strip()
|
||||
mmsi_text = str(mmsi or "").strip()
|
||||
if not text:
|
||||
return True
|
||||
if mmsi_text and text == mmsi_text:
|
||||
return True
|
||||
return bool(VESSEL_NAME_FALLBACK_PATTERN.match(text))
|
||||
|
||||
|
||||
def _requested_vessel_types(value: Optional[str]) -> set[str]:
|
||||
return {
|
||||
item.strip().lower()
|
||||
for item in (value or "").split(",")
|
||||
if item.strip()
|
||||
}
|
||||
|
||||
|
||||
def _matches_vessel_type(props: dict[str, Any], requested_types: set[str]) -> bool:
|
||||
if not requested_types:
|
||||
return True
|
||||
@@ -747,6 +894,88 @@ def _matches_vessel_type(props: dict[str, Any], requested_types: set[str]) -> bo
|
||||
return False
|
||||
|
||||
|
||||
def _feature_mmsi_key(feature: dict[str, Any]) -> str | None:
|
||||
props = feature.get("properties", {})
|
||||
mmsi = props.get("mmsi") or feature.get("id")
|
||||
if mmsi in (None, ""):
|
||||
return None
|
||||
return str(mmsi)
|
||||
|
||||
|
||||
def _feature_in_bbox(feature: dict[str, Any], bbox: tuple[float, float, float, float] | None) -> bool:
|
||||
if bbox is None:
|
||||
return True
|
||||
coordinates = feature.get("geometry", {}).get("coordinates") or []
|
||||
if len(coordinates) < 2:
|
||||
return False
|
||||
try:
|
||||
lon = float(coordinates[0])
|
||||
lat = float(coordinates[1])
|
||||
except (TypeError, ValueError):
|
||||
return False
|
||||
lon_min, lat_min, lon_max, lat_max = bbox
|
||||
return lon_min <= lon <= lon_max and lat_min <= lat <= lat_max
|
||||
|
||||
|
||||
def _filter_vessel_features(
|
||||
features: list[dict[str, Any]],
|
||||
*,
|
||||
bbox: tuple[float, float, float, float] | None,
|
||||
requested_types: set[str],
|
||||
) -> list[dict[str, Any]]:
|
||||
return [
|
||||
feature
|
||||
for feature in features
|
||||
if _feature_in_bbox(feature, bbox)
|
||||
and _matches_vessel_type(feature.get("properties", {}), requested_types)
|
||||
]
|
||||
|
||||
|
||||
def _merge_vessel_features(
|
||||
raw_features: list[dict[str, Any]],
|
||||
legacy_features: list[dict[str, Any]],
|
||||
) -> tuple[list[dict[str, Any]], dict[str, Any]]:
|
||||
"""Prefer aggregated raw observations as the canonical source of truth.
|
||||
|
||||
Legacy `vessel_position` rows only fill MMSIs that the unified pipeline does
|
||||
not yet know about, so a vessel never appears twice when both BarentsWatch
|
||||
and AISStream observe it. Once the legacy table drains, this branch becomes
|
||||
a no-op.
|
||||
"""
|
||||
|
||||
merged: list[dict[str, Any]] = []
|
||||
seen: set[str] = set()
|
||||
raw_keys: set[str] = set()
|
||||
legacy_keys: set[str] = set()
|
||||
|
||||
for feature in raw_features:
|
||||
key = _feature_mmsi_key(feature)
|
||||
if key is None or key in seen:
|
||||
continue
|
||||
seen.add(key)
|
||||
raw_keys.add(key)
|
||||
merged.append(feature)
|
||||
|
||||
legacy_added = 0
|
||||
for feature in legacy_features:
|
||||
key = _feature_mmsi_key(feature)
|
||||
if key is None:
|
||||
continue
|
||||
legacy_keys.add(key)
|
||||
if key in seen:
|
||||
continue
|
||||
seen.add(key)
|
||||
legacy_added += 1
|
||||
merged.append(feature)
|
||||
|
||||
return merged, {
|
||||
"raw_unique_mmsi": len(raw_keys),
|
||||
"legacy_unique_mmsi": len(legacy_keys),
|
||||
"legacy_backfilled_mmsi": legacy_added,
|
||||
"final_unique_mmsi": len(seen),
|
||||
}
|
||||
|
||||
|
||||
def _build_vessel_stats(features: List[dict[str, Any]]) -> dict[str, Any]:
|
||||
by_type: dict[str, int] = {}
|
||||
underway = 0
|
||||
@@ -1347,7 +1576,7 @@ async def get_satellites_geojson(
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
"""获取卫星 TLE GeoJSON 数据"""
|
||||
records = await _load_current_collected_data(
|
||||
records = await _load_current_or_latest_task_data(
|
||||
db,
|
||||
"celestrak_tle",
|
||||
exclude_unknown_name=True,
|
||||
@@ -1476,27 +1705,30 @@ async def get_vessels_geojson(
|
||||
):
|
||||
"""Return latest vessel positions as GeoJSON points."""
|
||||
parsed_bbox = _parse_bbox(bbox)
|
||||
aggregated_vessels = await get_aggregated_vessels(db, bbox=parsed_bbox, limit=limit)
|
||||
if aggregated_vessels:
|
||||
geojson = convert_aggregated_vessels_to_geojson(aggregated_vessels)
|
||||
requested_types = {
|
||||
item.strip().lower()
|
||||
for item in (type or "").split(",")
|
||||
if item.strip()
|
||||
}
|
||||
if requested_types:
|
||||
geojson["features"] = [
|
||||
feature
|
||||
for feature in geojson.get("features", [])
|
||||
if _matches_vessel_type(feature.get("properties", {}), requested_types)
|
||||
]
|
||||
requested_types = _requested_vessel_types(type)
|
||||
merged_features, diagnostics = await _load_merged_vessel_features(db)
|
||||
features = _filter_vessel_features(
|
||||
merged_features,
|
||||
bbox=parsed_bbox,
|
||||
requested_types=requested_types,
|
||||
)
|
||||
if limit and limit > 0:
|
||||
features = features[:limit]
|
||||
return {
|
||||
"type": "FeatureCollection",
|
||||
"features": features,
|
||||
"count": len(features),
|
||||
"stats": _build_vessel_stats(features),
|
||||
"diagnostics": {
|
||||
**diagnostics,
|
||||
"filtered_count": len(features),
|
||||
},
|
||||
}
|
||||
|
||||
features = geojson.get("features", [])
|
||||
return {
|
||||
**geojson,
|
||||
"count": len(features),
|
||||
"stats": _build_vessel_stats(features),
|
||||
}
|
||||
|
||||
async def _load_merged_vessel_features(db: AsyncSession) -> tuple[list[dict[str, Any]], dict[str, Any]]:
|
||||
aggregated_vessels = await get_aggregated_vessels(db)
|
||||
raw_geojson = convert_aggregated_vessels_to_geojson(aggregated_vessels)
|
||||
|
||||
latest_times = (
|
||||
select(
|
||||
@@ -1516,50 +1748,122 @@ async def get_vessels_geojson(
|
||||
.outerjoin(VesselStatic, VesselStatic.mmsi == VesselPosition.mmsi)
|
||||
.order_by(VesselPosition.received_at.desc())
|
||||
)
|
||||
if limit and limit > 0:
|
||||
stmt = stmt.limit(limit)
|
||||
|
||||
if parsed_bbox is not None:
|
||||
lon_min, lat_min, lon_max, lat_max = parsed_bbox
|
||||
stmt = stmt.where(
|
||||
VesselPosition.lon >= lon_min,
|
||||
VesselPosition.lon <= lon_max,
|
||||
VesselPosition.lat >= lat_min,
|
||||
VesselPosition.lat <= lat_max,
|
||||
)
|
||||
|
||||
result = await db.execute(stmt)
|
||||
rows = list(result.all())
|
||||
geojson = convert_vessels_to_geojson(rows)
|
||||
requested_types = {
|
||||
item.strip().lower()
|
||||
for item in (type or "").split(",")
|
||||
if item.strip()
|
||||
legacy_geojson = convert_vessels_to_geojson(rows)
|
||||
merged_features, diagnostics = _merge_vessel_features(
|
||||
raw_geojson.get("features", []),
|
||||
legacy_geojson.get("features", []),
|
||||
)
|
||||
return merged_features, {
|
||||
**diagnostics,
|
||||
"raw_feature_count": len(raw_geojson.get("features", [])),
|
||||
"legacy_feature_count": len(legacy_geojson.get("features", [])),
|
||||
}
|
||||
if requested_types:
|
||||
geojson["features"] = [
|
||||
feature
|
||||
for feature in geojson.get("features", [])
|
||||
if _matches_vessel_type(feature.get("properties", {}), requested_types)
|
||||
]
|
||||
|
||||
features = geojson.get("features", [])
|
||||
|
||||
@router.get("/vessels/custom-supplements")
|
||||
async def get_vessel_custom_supplements(db: AsyncSession = Depends(get_db)):
|
||||
"""Group custom vessel_ais sources by their declared merge target for diagnostics."""
|
||||
|
||||
from app.models.datasource_config import DataSourceConfig
|
||||
|
||||
result = await db.execute(
|
||||
select(DataSourceConfig.name, DataSourceConfig.config, DataSourceConfig.is_active)
|
||||
.where(DataSourceConfig.config["target_schema"].as_string() == "vessel_ais")
|
||||
)
|
||||
grouped: dict[str, dict[str, Any]] = {}
|
||||
for name, config, is_active in result.all():
|
||||
config = config or {}
|
||||
merge_target = str(config.get("merge_target_source") or "barentswatch_vessels")
|
||||
bucket = grouped.setdefault(merge_target, {"merge_target": merge_target, "sources": []})
|
||||
bucket["sources"].append({"name": name, "is_active": bool(is_active)})
|
||||
return {"groups": list(grouped.values())}
|
||||
|
||||
|
||||
@router.get("/vessels/name-fallbacks")
|
||||
async def get_vessel_name_fallbacks(
|
||||
limit: int = Query(500, ge=0, description="Maximum fallback-name vessels to return. 0 means no limit."),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
"""Return vessels whose display name still falls back to MMSI."""
|
||||
aggregated_vessels = await get_aggregated_vessels(db)
|
||||
raw_geojson = convert_aggregated_vessels_to_geojson(aggregated_vessels)
|
||||
|
||||
latest_times = (
|
||||
select(
|
||||
VesselPosition.mmsi.label("mmsi"),
|
||||
func.max(VesselPosition.received_at).label("received_at"),
|
||||
)
|
||||
.group_by(VesselPosition.mmsi)
|
||||
.subquery()
|
||||
)
|
||||
result = await db.execute(
|
||||
select(VesselPosition, VesselStatic)
|
||||
.join(
|
||||
latest_times,
|
||||
(VesselPosition.mmsi == latest_times.c.mmsi)
|
||||
& (VesselPosition.received_at == latest_times.c.received_at),
|
||||
)
|
||||
.outerjoin(VesselStatic, VesselStatic.mmsi == VesselPosition.mmsi)
|
||||
.order_by(VesselPosition.received_at.desc())
|
||||
)
|
||||
legacy_geojson = convert_vessels_to_geojson(list(result.all()))
|
||||
features, diagnostics = _merge_vessel_features(
|
||||
raw_geojson.get("features", []),
|
||||
legacy_geojson.get("features", []),
|
||||
)
|
||||
|
||||
fallback_items = []
|
||||
for feature in features:
|
||||
props = feature.get("properties", {})
|
||||
mmsi = props.get("mmsi")
|
||||
name = props.get("name")
|
||||
if not _is_vessel_name_fallback(name, mmsi):
|
||||
continue
|
||||
source_summary = props.get("source_summary") or {}
|
||||
fallback_items.append(
|
||||
{
|
||||
"mmsi": str(mmsi),
|
||||
"display_name": name or f"MMSI {mmsi}",
|
||||
"reason": "missing_real_name",
|
||||
"received_at": props.get("received_at"),
|
||||
"sources": sorted(source_summary.keys()),
|
||||
"source_summary": source_summary,
|
||||
"message_types": sorted(
|
||||
{
|
||||
message_type
|
||||
for summary in source_summary.values()
|
||||
for message_type in (summary.get("message_types") or [])
|
||||
}
|
||||
),
|
||||
"field_sources": props.get("field_sources") or {},
|
||||
}
|
||||
)
|
||||
|
||||
if limit and limit > 0:
|
||||
fallback_items = fallback_items[:limit]
|
||||
return {
|
||||
**geojson,
|
||||
"count": len(features),
|
||||
"stats": _build_vessel_stats(features),
|
||||
"count": len(fallback_items),
|
||||
"items": fallback_items,
|
||||
"diagnostics": diagnostics,
|
||||
}
|
||||
|
||||
|
||||
@router.get("/vessels/{mmsi}")
|
||||
async def get_vessel_detail(mmsi: int, db: AsyncSession = Depends(get_db)):
|
||||
from app.services.vessel_enrichment import get_vessel_enrichment_bundle
|
||||
|
||||
aggregated = await get_aggregated_vessel(db, mmsi)
|
||||
enrichment = await get_vessel_enrichment_bundle(db, mmsi)
|
||||
if aggregated is not None:
|
||||
return {
|
||||
**aggregated,
|
||||
"received_at": to_iso8601_utc(aggregated.get("received_at")),
|
||||
"latitude": aggregated["lat"],
|
||||
"longitude": aggregated["lon"],
|
||||
"enrichment": enrichment,
|
||||
}
|
||||
|
||||
latest_position_stmt = (
|
||||
@@ -1578,6 +1882,7 @@ async def get_vessel_detail(mmsi: int, db: AsyncSession = Depends(get_db)):
|
||||
**(geojson["features"][0]["properties"]),
|
||||
"latitude": position.lat,
|
||||
"longitude": position.lon,
|
||||
"enrichment": enrichment,
|
||||
}
|
||||
|
||||
|
||||
@@ -1737,31 +2042,16 @@ async def get_bgp_collectors_geojson(db: AsyncSession = Depends(get_db)):
|
||||
@router.get("/geo/summary")
|
||||
async def get_visualization_geo_summary(db: AsyncSession = Depends(get_db)):
|
||||
"""Return lightweight Earth HUD counts without loading layer GeoJSON payloads."""
|
||||
records_by_source = await _load_current_collected_data_by_sources(
|
||||
cable_count = await _count_current_or_latest_task_data(db, "arcgis_cables")
|
||||
landing_point_count = await _count_current_or_latest_task_data(db, "arcgis_landing_points")
|
||||
satellite_count = await _count_current_or_latest_task_data(
|
||||
db,
|
||||
[
|
||||
"arcgis_cables",
|
||||
"arcgis_landing_points",
|
||||
"celestrak_tle",
|
||||
"top500",
|
||||
"epoch_ai_gpu",
|
||||
],
|
||||
"celestrak_tle",
|
||||
exclude_unknown_name=True,
|
||||
)
|
||||
|
||||
cables = convert_cable_to_geojson(records_by_source.get("arcgis_cables", []))
|
||||
landing_points = convert_landing_point_to_geojson(
|
||||
records_by_source.get("arcgis_landing_points", []),
|
||||
)
|
||||
satellites = convert_satellite_to_geojson(
|
||||
_filter_known_records(records_by_source.get("celestrak_tle", [])),
|
||||
)
|
||||
compute_centers = convert_compute_centers_to_geojson(
|
||||
_filter_known_records(
|
||||
records_by_source.get("top500", [])
|
||||
+ records_by_source.get("epoch_ai_gpu", []),
|
||||
),
|
||||
)
|
||||
compute_features = compute_centers.get("features", [])
|
||||
supercomputer_count = await _count_current_or_latest_task_data(db, "top500")
|
||||
gpu_cluster_count = await _count_current_or_latest_task_data(db, "epoch_ai_gpu")
|
||||
compute_center_count = supercomputer_count + gpu_cluster_count
|
||||
|
||||
active_incident_result = await db.execute(
|
||||
select(func.count(BGPIncident.id)).where(BGPIncident.status == "active"),
|
||||
@@ -1771,35 +2061,56 @@ async def get_visualization_geo_summary(db: AsyncSession = Depends(get_db)):
|
||||
)
|
||||
active_incident_count = int(active_incident_result.scalar() or 0)
|
||||
active_anomaly_count = int(active_anomaly_result.scalar() or 0)
|
||||
bgp_collectors = await build_bgp_collector_coverage(
|
||||
bgp_collector_result = await db.execute(
|
||||
select(func.count(func.distinct(BGPObservation.collector)))
|
||||
.where(BGPObservation.collector.isnot(None))
|
||||
.where(func.length(func.btrim(BGPObservation.collector)) > 0)
|
||||
.where(BGPObservation.source.in_(("ris_live_bgp", "bgpstream_bgp")))
|
||||
)
|
||||
bgp_collector_scalar = bgp_collector_result.scalar()
|
||||
if bgp_collector_scalar is None:
|
||||
bgp_collectors = await build_bgp_collector_coverage(
|
||||
db,
|
||||
source_filter=("ris_live_bgp", "bgpstream_bgp"),
|
||||
)
|
||||
bgp_collector_count = len(
|
||||
[item for item in bgp_collectors if item.get("collector")]
|
||||
)
|
||||
else:
|
||||
bgp_collector_count = int(bgp_collector_scalar or 0)
|
||||
raw_unique_window_hours = 24
|
||||
raw_unique_mmsi = await count_unique_raw_vessel_mmsi(
|
||||
db,
|
||||
source_filter=("ris_live_bgp", "bgpstream_bgp"),
|
||||
observed_since=datetime.now(UTC) - timedelta(hours=raw_unique_window_hours),
|
||||
)
|
||||
vessel_count_result = await db.execute(
|
||||
select(func.count(func.distinct(VesselPosition.mmsi))),
|
||||
legacy_unique_result = await db.execute(
|
||||
select(func.count(func.distinct(VesselPosition.mmsi)))
|
||||
)
|
||||
vessel_count = int(vessel_count_result.scalar() or 0)
|
||||
legacy_unique_mmsi = int(legacy_unique_result.scalar() or 0)
|
||||
vessel_count = max(raw_unique_mmsi, legacy_unique_mmsi)
|
||||
aisstream_health = await db.get(AISSourceHealth, "aisstream_vessels")
|
||||
|
||||
return {
|
||||
"generated_at": to_iso8601_utc(datetime.now(UTC)),
|
||||
"stats": {
|
||||
"cable_count": len(cables.get("features", [])),
|
||||
"landing_point_count": len(landing_points.get("features", [])),
|
||||
"satellite_count": len(satellites.get("features", [])),
|
||||
"compute_center_count": len(compute_features),
|
||||
"cable_count": cable_count,
|
||||
"landing_point_count": landing_point_count,
|
||||
"satellite_count": satellite_count,
|
||||
"compute_center_count": compute_center_count,
|
||||
"vessel_count": vessel_count,
|
||||
"supercomputer_count": sum(
|
||||
1 for feature in compute_features
|
||||
if feature.get("properties", {}).get("site_type") == "supercomputer"
|
||||
),
|
||||
"gpu_cluster_count": sum(
|
||||
1 for feature in compute_features
|
||||
if feature.get("properties", {}).get("site_type") == "gpu_cluster"
|
||||
),
|
||||
"vessel_raw_unique_mmsi": raw_unique_mmsi,
|
||||
"vessel_raw_unique_window_hours": raw_unique_window_hours,
|
||||
"vessel_legacy_unique_mmsi": legacy_unique_mmsi,
|
||||
"aisstream_connection_state": aisstream_health.connection_state if aisstream_health else None,
|
||||
"aisstream_last_seen_at": to_iso8601_utc(aisstream_health.last_seen_at) if aisstream_health else None,
|
||||
"aisstream_message_rate": aisstream_health.message_rate if aisstream_health else None,
|
||||
"aisstream_lag_seconds": aisstream_health.lag_seconds if aisstream_health else None,
|
||||
"supercomputer_count": supercomputer_count,
|
||||
"gpu_cluster_count": gpu_cluster_count,
|
||||
"bgp_event_count": active_incident_count or active_anomaly_count,
|
||||
"bgp_incident_count": active_incident_count,
|
||||
"bgp_anomaly_count": active_anomaly_count,
|
||||
"bgp_collector_count": len([item for item in bgp_collectors if item.get("collector")]),
|
||||
"bgp_collector_count": bgp_collector_count,
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
@@ -40,16 +40,16 @@ async def authenticate_token(token: str) -> Optional[dict]:
|
||||
@router.websocket("/ws")
|
||||
async def websocket_endpoint(
|
||||
websocket: WebSocket,
|
||||
token: str = Query(...),
|
||||
token: str | None = Query(None),
|
||||
):
|
||||
"""WebSocket endpoint for real-time data"""
|
||||
logger.info_event(
|
||||
"WebSocket connection attempt",
|
||||
event="auth.websocket.connection_attempt",
|
||||
context={"token_preview": f"{token[:8]}..."},
|
||||
context={"token_preview": f"{token[:8]}..." if token else "anonymous"},
|
||||
)
|
||||
payload = await authenticate_token(token)
|
||||
if payload is None:
|
||||
payload = await authenticate_token(token) if token else None
|
||||
if token and payload is None:
|
||||
logger.warning_event(
|
||||
"WebSocket authentication failed, closing connection",
|
||||
event="auth.websocket.connection_rejected",
|
||||
@@ -57,7 +57,17 @@ async def websocket_endpoint(
|
||||
await websocket.close(code=4001)
|
||||
return
|
||||
|
||||
user_id = str(payload.get("sub"))
|
||||
is_anonymous = payload is None
|
||||
user_id = str(payload.get("sub")) if payload else f"anonymous:{id(websocket)}"
|
||||
supported_channels = ["vessels"] if is_anonymous else [
|
||||
"gpu_clusters",
|
||||
"submarine_cables",
|
||||
"ixp_nodes",
|
||||
"alerts",
|
||||
"dashboard",
|
||||
"datasource_tasks",
|
||||
"vessels",
|
||||
]
|
||||
await manager.connect(websocket, user_id)
|
||||
|
||||
try:
|
||||
@@ -68,14 +78,7 @@ async def websocket_endpoint(
|
||||
"connection_id": f"conn_{user_id}",
|
||||
"server_version": settings.VERSION,
|
||||
"heartbeat_interval": 30,
|
||||
"supported_channels": [
|
||||
"gpu_clusters",
|
||||
"submarine_cables",
|
||||
"ixp_nodes",
|
||||
"alerts",
|
||||
"dashboard",
|
||||
"datasource_tasks",
|
||||
],
|
||||
"supported_channels": supported_channels,
|
||||
},
|
||||
}
|
||||
)
|
||||
@@ -93,12 +96,24 @@ async def websocket_endpoint(
|
||||
)
|
||||
elif data.get("type") == "subscribe":
|
||||
channels = data.get("data", {}).get("channels", [])
|
||||
if is_anonymous:
|
||||
channels = [channel for channel in channels if channel in supported_channels]
|
||||
manager.subscribe(websocket, channels)
|
||||
await websocket.send_json(
|
||||
{
|
||||
"type": "subscription_confirmed",
|
||||
"data": {"action": "subscribe", "channels": channels},
|
||||
}
|
||||
)
|
||||
elif data.get("type") == "unsubscribe":
|
||||
channels = data.get("data", {}).get("channels", [])
|
||||
manager.unsubscribe(websocket, channels)
|
||||
await websocket.send_json(
|
||||
{
|
||||
"type": "subscription_confirmed",
|
||||
"data": {"action": "unsubscribe", "channels": channels},
|
||||
}
|
||||
)
|
||||
elif data.get("type") == "control_frame":
|
||||
await websocket.send_json(
|
||||
{"type": "control_acknowledged", "data": {"received": True}}
|
||||
|
||||
@@ -16,8 +16,11 @@ class VesselAISRecord(BaseModel):
|
||||
sog: float | None = None
|
||||
cog: float | None = Field(default=None, ge=0, le=360)
|
||||
heading: int | None = Field(default=None, ge=0, le=511)
|
||||
nav_status: int | None = None
|
||||
name: str | None = None
|
||||
callsign: str | None = None
|
||||
vessel_type: str | int | None = None
|
||||
vessel_type_name: str | None = None
|
||||
received_at: datetime | None = None
|
||||
|
||||
|
||||
@@ -104,8 +107,11 @@ TARGET_SCHEMAS: dict[str, TargetSchema] = {
|
||||
TargetField("sog", "float", False, "对地航速,单位节", 12.4),
|
||||
TargetField("cog", "float", False, "对地航向,0-360 度", 184.5),
|
||||
TargetField("heading", "integer", False, "船首向,0-511", 186),
|
||||
TargetField("nav_status", "integer", False, "导航状态码", 0),
|
||||
TargetField("name", "string", False, "船名", "OSLO EXPRESS"),
|
||||
TargetField("vessel_type", "string", False, "船型", "cargo"),
|
||||
TargetField("callsign", "string", False, "呼号", "LAAB"),
|
||||
TargetField("vessel_type", "string", False, "船型代码", 70),
|
||||
TargetField("vessel_type_name", "string", False, "船型名称", "Cargo"),
|
||||
TargetField("received_at", "datetime", False, "数据接收时间", "2026-04-28T00:00:00Z"),
|
||||
),
|
||||
),
|
||||
|
||||
@@ -75,7 +75,7 @@ class DataBroadcaster:
|
||||
"timestamp": to_iso8601_utc(datetime.now(UTC)),
|
||||
"payload": data,
|
||||
},
|
||||
channel=channel if channel in manager.active_connections else "all",
|
||||
channel=channel,
|
||||
)
|
||||
|
||||
async def broadcast_datasource_task_update(self, data: Dict[str, Any]):
|
||||
|
||||
@@ -1,9 +1,6 @@
|
||||
"""WebSocket Connection Manager"""
|
||||
|
||||
import json
|
||||
import asyncio
|
||||
from typing import Dict, Set, Optional
|
||||
from datetime import datetime
|
||||
from fastapi import WebSocket
|
||||
import redis.asyncio as redis
|
||||
|
||||
@@ -15,6 +12,8 @@ class ConnectionManager:
|
||||
|
||||
def __init__(self):
|
||||
self.active_connections: Dict[str, Set[WebSocket]] = {} # user_id -> connections
|
||||
self.channel_subscriptions: Dict[str, Set[WebSocket]] = {}
|
||||
self.websocket_channels: Dict[WebSocket, Set[str]] = {}
|
||||
self.redis_client: Optional[redis.Redis] = None
|
||||
|
||||
async def connect(self, websocket: WebSocket, user_id: str):
|
||||
@@ -40,6 +39,39 @@ class ConnectionManager:
|
||||
self.active_connections[user_id].discard(websocket)
|
||||
if not self.active_connections[user_id]:
|
||||
del self.active_connections[user_id]
|
||||
self.unsubscribe_all(websocket)
|
||||
|
||||
def subscribe(self, websocket: WebSocket, channels: list[str]):
|
||||
normalized_channels = {
|
||||
str(channel).strip()
|
||||
for channel in channels
|
||||
if str(channel).strip()
|
||||
}
|
||||
if not normalized_channels:
|
||||
return
|
||||
|
||||
socket_channels = self.websocket_channels.setdefault(websocket, set())
|
||||
for channel in normalized_channels:
|
||||
self.channel_subscriptions.setdefault(channel, set()).add(websocket)
|
||||
socket_channels.add(channel)
|
||||
|
||||
def unsubscribe(self, websocket: WebSocket, channels: list[str]):
|
||||
for channel in {str(channel).strip() for channel in channels if str(channel).strip()}:
|
||||
subscribers = self.channel_subscriptions.get(channel)
|
||||
if subscribers is not None:
|
||||
subscribers.discard(websocket)
|
||||
if not subscribers:
|
||||
del self.channel_subscriptions[channel]
|
||||
socket_channels = self.websocket_channels.get(websocket)
|
||||
if socket_channels is not None:
|
||||
socket_channels.discard(channel)
|
||||
if not socket_channels:
|
||||
del self.websocket_channels[websocket]
|
||||
|
||||
def unsubscribe_all(self, websocket: WebSocket):
|
||||
channels = list(self.websocket_channels.get(websocket, set()))
|
||||
if channels:
|
||||
self.unsubscribe(websocket, channels)
|
||||
|
||||
async def send_personal_message(self, message: dict, user_id: str):
|
||||
if user_id in self.active_connections:
|
||||
@@ -54,13 +86,19 @@ class ConnectionManager:
|
||||
for user_id in self.active_connections:
|
||||
await self.send_personal_message(message, user_id)
|
||||
else:
|
||||
await self.send_personal_message(message, channel)
|
||||
for connection in list(self.channel_subscriptions.get(channel, set())):
|
||||
try:
|
||||
await connection.send_json(message)
|
||||
except Exception:
|
||||
self.unsubscribe_all(connection)
|
||||
|
||||
async def close_all(self):
|
||||
for user_id in self.active_connections:
|
||||
for connection in self.active_connections[user_id]:
|
||||
await connection.close()
|
||||
self.active_connections.clear()
|
||||
self.channel_subscriptions.clear()
|
||||
self.websocket_channels.clear()
|
||||
|
||||
|
||||
manager = ConnectionManager()
|
||||
|
||||
@@ -111,6 +111,7 @@ async def init_db():
|
||||
import app.models.playground_message # noqa: F401
|
||||
import app.models.system_log # noqa: F401
|
||||
import app.models.vessel # noqa: F401
|
||||
import app.models.vessel_enrichment # noqa: F401
|
||||
import app.models.datasource_mapping # noqa: F401
|
||||
|
||||
logger.warning_event(
|
||||
@@ -163,6 +164,30 @@ async def init_db():
|
||||
"""
|
||||
)
|
||||
)
|
||||
await conn.execute(
|
||||
text(
|
||||
"""
|
||||
CREATE INDEX IF NOT EXISTS idx_collected_data_source_current_id
|
||||
ON collected_data (source, is_current, id)
|
||||
"""
|
||||
)
|
||||
)
|
||||
await conn.execute(
|
||||
text(
|
||||
"""
|
||||
CREATE INDEX IF NOT EXISTS idx_collected_data_source_task_id
|
||||
ON collected_data (source, task_id, id)
|
||||
"""
|
||||
)
|
||||
)
|
||||
await conn.execute(
|
||||
text(
|
||||
"""
|
||||
CREATE INDEX IF NOT EXISTS idx_ais_raw_schema_observed_entity
|
||||
ON ais_raw_observations (target_schema, observed_at, entity_key)
|
||||
"""
|
||||
)
|
||||
)
|
||||
await conn.execute(
|
||||
text(
|
||||
"""
|
||||
|
||||
@@ -48,6 +48,8 @@ class CollectedData(Base):
|
||||
# Indexes for common queries
|
||||
__table_args__ = (
|
||||
Index("idx_collected_data_source_collected", "source", "collected_at"),
|
||||
Index("idx_collected_data_source_current_id", "source", "is_current", "id"),
|
||||
Index("idx_collected_data_source_task_id", "source", "task_id", "id"),
|
||||
Index("idx_collected_data_source_type", "source", "data_type"),
|
||||
Index("idx_collected_data_source_source_id", "source", "source_id"),
|
||||
)
|
||||
|
||||
@@ -97,6 +97,7 @@ class AISRawObservation(Base):
|
||||
|
||||
__table_args__ = (
|
||||
Index("idx_ais_raw_entity_observed", "target_schema", "entity_key", "observed_at"),
|
||||
Index("idx_ais_raw_schema_observed_entity", "target_schema", "observed_at", "entity_key"),
|
||||
Index("idx_ais_raw_source_entity", "source", "entity_key"),
|
||||
)
|
||||
|
||||
|
||||
63
backend/app/models/vessel_enrichment.py
Normal file
63
backend/app/models/vessel_enrichment.py
Normal file
@@ -0,0 +1,63 @@
|
||||
"""Vessel enrichment cache tables (v5).
|
||||
|
||||
Profile and media enrichment are stored separately so cache TTLs can differ
|
||||
and so the conflict-resolution + display layers can read either independently.
|
||||
"""
|
||||
|
||||
from sqlalchemy import BigInteger, Column, DateTime, Float, JSON, String
|
||||
from sqlalchemy.sql import func
|
||||
|
||||
from app.core.time import to_iso8601_utc
|
||||
from app.db.session import Base
|
||||
|
||||
|
||||
class VesselProfileEnrichment(Base):
|
||||
"""Cached static vessel profile (type, flag, dimensions, operator, etc.)."""
|
||||
|
||||
__tablename__ = "vessel_profile_enrichment"
|
||||
|
||||
mmsi = Column(BigInteger, primary_key=True)
|
||||
source = Column(String(100), nullable=False, default="system")
|
||||
payload = Column(JSON, nullable=False, default=dict)
|
||||
fetched_at = Column(DateTime(timezone=True), nullable=False, server_default=func.now())
|
||||
expires_at = Column(DateTime(timezone=True), nullable=True)
|
||||
confidence = Column(Float, nullable=True)
|
||||
reference_url = Column(String(500), nullable=True)
|
||||
updated_at = Column(DateTime(timezone=True), nullable=False, server_default=func.now(), onupdate=func.now())
|
||||
|
||||
def to_dict(self) -> dict:
|
||||
return {
|
||||
"mmsi": self.mmsi,
|
||||
"source": self.source,
|
||||
"payload": self.payload or {},
|
||||
"fetched_at": to_iso8601_utc(self.fetched_at),
|
||||
"expires_at": to_iso8601_utc(self.expires_at),
|
||||
"confidence": self.confidence,
|
||||
"reference_url": self.reference_url,
|
||||
}
|
||||
|
||||
|
||||
class VesselMediaEnrichment(Base):
|
||||
"""Cached vessel imagery / external detail references."""
|
||||
|
||||
__tablename__ = "vessel_media_enrichment"
|
||||
|
||||
mmsi = Column(BigInteger, primary_key=True)
|
||||
source = Column(String(100), nullable=False, default="system")
|
||||
payload = Column(JSON, nullable=False, default=dict)
|
||||
fetched_at = Column(DateTime(timezone=True), nullable=False, server_default=func.now())
|
||||
expires_at = Column(DateTime(timezone=True), nullable=True)
|
||||
confidence = Column(Float, nullable=True)
|
||||
reference_url = Column(String(500), nullable=True)
|
||||
updated_at = Column(DateTime(timezone=True), nullable=False, server_default=func.now(), onupdate=func.now())
|
||||
|
||||
def to_dict(self) -> dict:
|
||||
return {
|
||||
"mmsi": self.mmsi,
|
||||
"source": self.source,
|
||||
"payload": self.payload or {},
|
||||
"fetched_at": to_iso8601_utc(self.fetched_at),
|
||||
"expires_at": to_iso8601_utc(self.expires_at),
|
||||
"confidence": self.confidence,
|
||||
"reference_url": self.reference_url,
|
||||
}
|
||||
@@ -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()
|
||||
149
backend/tests/test_custom_datasource_runtime_live.py
Normal file
149
backend/tests/test_custom_datasource_runtime_live.py
Normal file
@@ -0,0 +1,149 @@
|
||||
"""End-to-end integration test for the custom WebSocket datasource runner.
|
||||
|
||||
Boots an in-process WebSocket server that mimics the bun mock AIS server
|
||||
(`scripts/mock-ais-ws-server.ts`) and runs the real
|
||||
`run_mapped_websocket_config` against it. Catches regressions where the
|
||||
runner stops connecting, fails to extract the configured message path,
|
||||
or quietly drops mapped records before broadcasting.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
from contextlib import asynccontextmanager
|
||||
from datetime import UTC, datetime
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
import pytest
|
||||
import websockets
|
||||
|
||||
from app.models.datasource_config import DataSourceConfig
|
||||
from app.services import custom_datasource_runtime
|
||||
from app.services.custom_datasource_runtime import run_mapped_websocket_config
|
||||
|
||||
|
||||
def _make_payload(seq: int) -> str:
|
||||
return json.dumps(
|
||||
{
|
||||
"type": "vessel",
|
||||
"sequence": seq,
|
||||
"data": {
|
||||
"mmsi": str(999_000_000 + seq),
|
||||
"name": f"MOCK VESSEL {seq:03d}",
|
||||
"lat": 36.20 + seq * 0.001,
|
||||
"lon": 14.20 + seq * 0.001,
|
||||
"sog": 12.0,
|
||||
"cog": 90.0,
|
||||
"heading": 90,
|
||||
"vessel_type": 70,
|
||||
"vessel_type_name": "Cargo",
|
||||
"received_at": datetime.now(UTC).isoformat(),
|
||||
},
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def _mock_ais_server(emit_count: int):
|
||||
received_subscribe: list[str] = []
|
||||
|
||||
async def handler(ws):
|
||||
try:
|
||||
try:
|
||||
msg = await asyncio.wait_for(ws.recv(), timeout=0.5)
|
||||
received_subscribe.append(msg)
|
||||
except (asyncio.TimeoutError, websockets.ConnectionClosed):
|
||||
pass
|
||||
for seq in range(1, emit_count + 1):
|
||||
await ws.send(_make_payload(seq))
|
||||
await asyncio.sleep(0.01)
|
||||
# keep the socket open briefly so the runner observes the messages
|
||||
await asyncio.sleep(0.05)
|
||||
except websockets.ConnectionClosed:
|
||||
return
|
||||
|
||||
async with websockets.serve(handler, "127.0.0.1", 0) as server:
|
||||
port = next(iter(server.sockets)).getsockname()[1]
|
||||
yield port, received_subscribe
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_websocket_runner_streams_from_live_mock(monkeypatch):
|
||||
mapping = SimpleNamespace(
|
||||
id=11,
|
||||
version=3,
|
||||
target_schema="vessel_ais",
|
||||
mapping_json={
|
||||
"source": {"items_path": "$"},
|
||||
"fields": {
|
||||
"mmsi": {"path": "$.mmsi", "type": "integer"},
|
||||
"lat": {"path": "$.lat", "type": "float"},
|
||||
"lon": {"path": "$.lon", "type": "float"},
|
||||
"name": {"path": "$.name", "type": "string"},
|
||||
"vessel_type": {"path": "$.vessel_type", "type": "integer", "default": None},
|
||||
"vessel_type_name": {"path": "$.vessel_type_name", "type": "string", "default": None},
|
||||
"sog": {"path": "$.sog", "type": "float", "default": None},
|
||||
"cog": {"path": "$.cog", "type": "float", "default": None},
|
||||
"heading": {"path": "$.heading", "type": "integer", "default": None},
|
||||
"received_at": {"path": "$.received_at", "type": "datetime"},
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
class FakeResult:
|
||||
def scalar_one_or_none(self):
|
||||
return mapping
|
||||
|
||||
class FakeDB:
|
||||
async def execute(self, _stmt):
|
||||
return FakeResult()
|
||||
|
||||
persist = AsyncMock(return_value=1)
|
||||
monkeypatch.setattr(custom_datasource_runtime, "persist_mapped_records", persist)
|
||||
|
||||
async with _mock_ais_server(emit_count=3) as (port, received_subscribe):
|
||||
result = await run_mapped_websocket_config(
|
||||
FakeDB(),
|
||||
DataSourceConfig(
|
||||
id=99,
|
||||
name="mock_ais_ws",
|
||||
source_type="websocket",
|
||||
endpoint=f"ws://127.0.0.1:{port}",
|
||||
auth_type="none",
|
||||
headers={},
|
||||
config={
|
||||
"ws_message_path": "$.data",
|
||||
"ws_subscribe_message": {
|
||||
"type": "subscribe",
|
||||
"anchor": {"lat": 36.2, "lon": 14.2},
|
||||
"spread_km": 50,
|
||||
"rate_hz": 1,
|
||||
},
|
||||
"debug_max_messages": 2,
|
||||
"delivery_mode": "realtime_stream",
|
||||
"ws_reconnect": False,
|
||||
},
|
||||
),
|
||||
use_config_debug_max_messages=True,
|
||||
)
|
||||
|
||||
assert result["status"] == "success"
|
||||
assert result["messages_seen"] == 2
|
||||
assert result["written_count"] == 2
|
||||
assert result["mapped_count"] == 2
|
||||
assert result["target_schema"] == "vessel_ais"
|
||||
# subscribe message must reach the server unchanged
|
||||
assert received_subscribe, "runner did not forward ws_subscribe_message"
|
||||
parsed = json.loads(received_subscribe[0])
|
||||
assert parsed["type"] == "subscribe"
|
||||
assert parsed["anchor"] == {"lat": 36.2, "lon": 14.2}
|
||||
assert parsed["rate_hz"] == 1
|
||||
# mapped records carry the real MMSIs from the mock stream
|
||||
persisted_records = []
|
||||
for call in persist.await_args_list:
|
||||
persisted_records.extend(call.kwargs["records"])
|
||||
assert {record["mmsi"] for record in persisted_records} == {999_000_001, 999_000_002}
|
||||
assert all(record["vessel_type"] == 70 for record in persisted_records)
|
||||
assert all(record["vessel_type_name"] == "Cargo" for record in persisted_records)
|
||||
@@ -1,13 +1,18 @@
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
from unittest.mock import AsyncMock
|
||||
from httpx import ASGITransport, AsyncClient
|
||||
|
||||
from app.api.v1.datasource_config import get_ai_provider_client
|
||||
from app.core.websocket import broadcaster as broadcaster_module
|
||||
from app.core.security import get_current_user
|
||||
from app.core.target_schema_registry import get_target_schema, list_target_schemas
|
||||
from app.main import app
|
||||
from app.models.user import User
|
||||
from app.models.datasource_config import DataSourceConfig
|
||||
from app.services import custom_datasource_runtime
|
||||
from app.services.custom_datasource_runtime import run_mapped_websocket_config
|
||||
from app.services.datasource_mapping import execute_mapping, persist_mapped_records, redact_for_llm
|
||||
|
||||
|
||||
@@ -106,6 +111,130 @@ async def test_persist_mapped_records_writes_generic_records():
|
||||
assert db.added[0].extra_data["mapping_version"] == 3
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_persist_mapped_vessel_records_writes_raw_and_broadcasts(monkeypatch):
|
||||
record_observation = AsyncMock(return_value=object())
|
||||
update_health = AsyncMock()
|
||||
broadcast_custom = AsyncMock()
|
||||
monkeypatch.setattr(
|
||||
"app.services.vessel_ais_aggregation.record_vessel_ais_observation",
|
||||
record_observation,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"app.services.vessel_ais_aggregation.update_ais_source_health",
|
||||
update_health,
|
||||
)
|
||||
monkeypatch.setattr(broadcaster_module, "broadcast_custom", broadcast_custom)
|
||||
|
||||
class FakeDB:
|
||||
def __init__(self):
|
||||
self.committed = False
|
||||
|
||||
async def commit(self):
|
||||
self.committed = True
|
||||
|
||||
db = FakeDB()
|
||||
|
||||
count = await persist_mapped_records(
|
||||
db,
|
||||
datasource_name="mock_ais_ws",
|
||||
datasource_config_id=42,
|
||||
target_schema="vessel_ais",
|
||||
records=[
|
||||
{
|
||||
"mmsi": 999000001,
|
||||
"lat": 31.2,
|
||||
"lon": 121.4,
|
||||
"name": "MOCK VESSEL 001",
|
||||
"received_at": "2026-05-01T00:00:00Z",
|
||||
}
|
||||
],
|
||||
mapping_version=1,
|
||||
delivery_mode="realtime_stream",
|
||||
transport="websocket",
|
||||
)
|
||||
|
||||
assert count == 1
|
||||
assert db.committed is True
|
||||
record_observation.assert_awaited_once()
|
||||
assert record_observation.await_args.kwargs["source"] == "mock_ais_ws"
|
||||
assert record_observation.await_args.kwargs["delivery_mode"] == "realtime_stream"
|
||||
assert record_observation.await_args.kwargs["transport"] == "websocket"
|
||||
update_health.assert_awaited_once()
|
||||
broadcast_custom.assert_awaited_once()
|
||||
assert broadcast_custom.await_args.args[0] == "vessels"
|
||||
assert broadcast_custom.await_args.args[1]["vessels"][0]["mmsi_display"] == "999000001"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_custom_websocket_runner_maps_and_persists_vessel_records(monkeypatch):
|
||||
mapping = SimpleNamespace(
|
||||
id=7,
|
||||
version=2,
|
||||
target_schema="vessel_ais",
|
||||
mapping_json={
|
||||
"source": {"items_path": "$"},
|
||||
"fields": {
|
||||
"mmsi": {"path": "$.mmsi", "type": "integer"},
|
||||
"lat": {"path": "$.lat", "type": "float"},
|
||||
"lon": {"path": "$.lon", "type": "float"},
|
||||
"name": {"path": "$.name", "type": "string"},
|
||||
"received_at": {"path": "$.received_at", "type": "datetime"},
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
class FakeResult:
|
||||
def scalar_one_or_none(self):
|
||||
return mapping
|
||||
|
||||
class FakeDB:
|
||||
async def execute(self, _stmt):
|
||||
return FakeResult()
|
||||
|
||||
class FakeWebSocket:
|
||||
async def __aenter__(self):
|
||||
return self
|
||||
|
||||
async def __aexit__(self, *_args):
|
||||
return None
|
||||
|
||||
async def send(self, _message):
|
||||
return None
|
||||
|
||||
async def recv(self):
|
||||
return (
|
||||
'{"type":"vessel","data":{"mmsi":"999000001","name":"MOCK VESSEL 001",'
|
||||
'"lat":31.2,"lon":121.4,"received_at":"2026-05-01T00:00:00Z"}}'
|
||||
)
|
||||
|
||||
persist = AsyncMock(return_value=1)
|
||||
monkeypatch.setattr(custom_datasource_runtime, "_connect_websocket", AsyncMock(return_value=FakeWebSocket()))
|
||||
monkeypatch.setattr(custom_datasource_runtime, "persist_mapped_records", persist)
|
||||
|
||||
result = await run_mapped_websocket_config(
|
||||
FakeDB(),
|
||||
DataSourceConfig(
|
||||
id=42,
|
||||
name="mock_ais_ws",
|
||||
source_type="websocket",
|
||||
endpoint="ws://localhost:8787/ais",
|
||||
auth_type="none",
|
||||
headers={},
|
||||
config={"ws_message_path": "$.data", "debug_max_messages": 1},
|
||||
),
|
||||
)
|
||||
|
||||
assert result["status"] == "success"
|
||||
assert result["messages_seen"] == 1
|
||||
assert result["written_count"] == 1
|
||||
persist.assert_awaited_once()
|
||||
assert persist.await_args.kwargs["datasource_name"] == "mock_ais_ws"
|
||||
assert persist.await_args.kwargs["records"][0]["mmsi"] == 999000001
|
||||
assert persist.await_args.kwargs["delivery_mode"] == "realtime_stream"
|
||||
assert persist.await_args.kwargs["transport"] == "websocket"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mapping_preview_api_uses_deterministic_engine():
|
||||
def override_get_current_user():
|
||||
|
||||
161
backend/tests/test_vessel_aggregation_strategy.py
Normal file
161
backend/tests/test_vessel_aggregation_strategy.py
Normal file
@@ -0,0 +1,161 @@
|
||||
"""Tests for the v4 vessel_ais aggregation strategy."""
|
||||
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
import pytest
|
||||
|
||||
from app.models.vessel import AISRawObservation
|
||||
from app.services.vessel_aggregation_strategy import (
|
||||
DEFAULT_STRATEGY,
|
||||
StrategyValidationError,
|
||||
validate_strategy,
|
||||
)
|
||||
from app.services.vessel_ais_aggregation import aggregate_vessel_observations
|
||||
|
||||
|
||||
def _obs(*, source: str, mmsi: int, observed_at: datetime, **payload) -> AISRawObservation:
|
||||
payload = {"mmsi": mmsi, "lat": 50.0, "lon": 10.0, **payload}
|
||||
delivery_mode = "realtime_stream" if source == "aisstream_vessels" else "polling"
|
||||
transport = "websocket" if source == "aisstream_vessels" else "http"
|
||||
return AISRawObservation(
|
||||
target_schema="vessel_ais",
|
||||
source=source,
|
||||
entity_key=str(mmsi),
|
||||
delivery_mode=delivery_mode,
|
||||
transport=transport,
|
||||
message_type="PositionReport",
|
||||
observation_hash=f"{source}:{mmsi}:{observed_at.isoformat()}",
|
||||
observed_at=observed_at,
|
||||
collected_at=observed_at,
|
||||
normalized_payload=payload,
|
||||
raw_payload=payload,
|
||||
quality_flags=[],
|
||||
)
|
||||
|
||||
|
||||
def test_validate_rejects_unknown_field():
|
||||
with pytest.raises(StrategyValidationError, match="unknown vessel_ais field"):
|
||||
validate_strategy({"vessel_ais": {"field_rules": {"definitely_not_a_field": {"mode": "newest"}}}})
|
||||
|
||||
|
||||
def test_validate_rejects_dynamic_lock_without_flag():
|
||||
with pytest.raises(StrategyValidationError, match="allow_dynamic_lock"):
|
||||
validate_strategy(
|
||||
{
|
||||
"vessel_ais": {
|
||||
"field_rules": {"lat": {"mode": "source_priority"}},
|
||||
"allow_dynamic_lock": False,
|
||||
}
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def test_validate_allows_dynamic_lock_with_flag():
|
||||
normalized = validate_strategy(
|
||||
{
|
||||
"version": 0,
|
||||
"vessel_ais": {
|
||||
"field_rules": {"lat": {"mode": "source_priority", "source_priority": ["barentswatch_vessels"]}},
|
||||
"allow_dynamic_lock": True,
|
||||
},
|
||||
}
|
||||
)
|
||||
assert normalized["vessel_ais"]["field_rules"]["lat"]["mode"] == "source_priority"
|
||||
assert normalized["version"] == 1
|
||||
|
||||
|
||||
def test_validate_increments_version():
|
||||
first = validate_strategy({"version": 5, "vessel_ais": {}})
|
||||
assert first["version"] == 6
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_strategy_field_rule_promotes_specific_source(monkeypatch):
|
||||
now = datetime(2026, 5, 4, 12, 0, tzinfo=timezone.utc)
|
||||
|
||||
obs_a = _obs(
|
||||
source="aisstream_vessels",
|
||||
mmsi=257123000,
|
||||
observed_at=now,
|
||||
name="AISSTREAM ONE",
|
||||
vessel_type_name="Cargo",
|
||||
)
|
||||
obs_b = _obs(
|
||||
source="barentswatch_vessels",
|
||||
mmsi=257123000,
|
||||
observed_at=now - timedelta(seconds=1),
|
||||
name="BARENTSWATCH ONE",
|
||||
vessel_type_name="Cargo",
|
||||
)
|
||||
|
||||
strategy = {
|
||||
"version": 7,
|
||||
"vessel_ais": {
|
||||
"source_priority": [],
|
||||
"field_rules": {
|
||||
"name": {"mode": "source_priority", "source_priority": ["barentswatch_vessels", "aisstream_vessels"]},
|
||||
},
|
||||
"freshness": {"realtime_stream_seconds": 0, "polling_seconds": 0},
|
||||
"allow_dynamic_lock": False,
|
||||
},
|
||||
}
|
||||
|
||||
db = AsyncMock()
|
||||
vessels = await aggregate_vessel_observations(
|
||||
db,
|
||||
[obs_a, obs_b],
|
||||
write_conflicts=False,
|
||||
strategy=strategy,
|
||||
)
|
||||
assert len(vessels) == 1
|
||||
vessel = vessels[0]
|
||||
assert vessel["name"] == "BARENTSWATCH ONE"
|
||||
assert vessel["field_sources"]["name"] == "barentswatch_vessels"
|
||||
assert vessel["selected_reasons"]["name"] == "source_priority"
|
||||
assert vessel["aggregation_strategy_version"] == 7
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_strategy_freshness_falls_back_to_polling_when_realtime_stale():
|
||||
now = datetime(2026, 5, 4, 12, 0, tzinfo=timezone.utc)
|
||||
|
||||
stale_realtime = _obs(
|
||||
source="aisstream_vessels",
|
||||
mmsi=257123000,
|
||||
observed_at=now - timedelta(hours=1),
|
||||
lat=58.0,
|
||||
lon=10.0,
|
||||
)
|
||||
fresh_polling = _obs(
|
||||
source="barentswatch_vessels",
|
||||
mmsi=257123000,
|
||||
observed_at=now - timedelta(seconds=30),
|
||||
lat=60.0,
|
||||
lon=11.0,
|
||||
)
|
||||
|
||||
strategy = {
|
||||
"version": 1,
|
||||
"vessel_ais": {
|
||||
"source_priority": ["aisstream_vessels", "barentswatch_vessels"],
|
||||
"field_rules": {},
|
||||
"freshness": {"realtime_stream_seconds": 900, "polling_seconds": 7200},
|
||||
"allow_dynamic_lock": False,
|
||||
},
|
||||
}
|
||||
|
||||
db = AsyncMock()
|
||||
vessels = await aggregate_vessel_observations(
|
||||
db,
|
||||
[stale_realtime, fresh_polling],
|
||||
write_conflicts=False,
|
||||
strategy=strategy,
|
||||
)
|
||||
assert vessels[0]["field_sources"]["lat"] == "barentswatch_vessels"
|
||||
assert vessels[0]["lat"] == 60.0
|
||||
|
||||
|
||||
def test_default_strategy_is_stable():
|
||||
assert DEFAULT_STRATEGY["vessel_ais"]["allow_dynamic_lock"] is False
|
||||
assert "freshness" in DEFAULT_STRATEGY["vessel_ais"]
|
||||
155
backend/tests/test_vessel_enrichment.py
Normal file
155
backend/tests/test_vessel_enrichment.py
Normal file
@@ -0,0 +1,155 @@
|
||||
"""Tests for v5 enrichment + conflict promote-to-rule."""
|
||||
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
import pytest
|
||||
|
||||
from app.models.vessel import AISConflictRecord, AISRawObservation
|
||||
from app.models.vessel_enrichment import VesselMediaEnrichment, VesselProfileEnrichment
|
||||
from app.services.vessel_ais_aggregation import aggregate_vessel_observations
|
||||
from app.services.vessel_enrichment import (
|
||||
_apply_upsert,
|
||||
get_vessel_enrichment_bundle,
|
||||
)
|
||||
|
||||
|
||||
class _StoreSession:
|
||||
"""Minimal AsyncSession stand-in that tracks mmsi-keyed enrichment + a strategy."""
|
||||
|
||||
def __init__(self, *, profile=None, media=None, conflicts=None):
|
||||
self.profile = profile
|
||||
self.media = media
|
||||
self.conflicts = list(conflicts or [])
|
||||
self.added: list = []
|
||||
self.committed = False
|
||||
|
||||
async def get(self, model, key):
|
||||
if model is VesselProfileEnrichment:
|
||||
return self.profile if self.profile and self.profile.mmsi == key else None
|
||||
if model is VesselMediaEnrichment:
|
||||
return self.media if self.media and self.media.mmsi == key else None
|
||||
return None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_enrichment_bundle_filters_expired_records():
|
||||
now = datetime.now(timezone.utc)
|
||||
fresh = VesselProfileEnrichment(
|
||||
mmsi=257123000,
|
||||
source="local_cache",
|
||||
payload={"vessel_subtype": "Container"},
|
||||
fetched_at=now - timedelta(hours=1),
|
||||
expires_at=now + timedelta(days=7),
|
||||
confidence=0.9,
|
||||
)
|
||||
expired_media = VesselMediaEnrichment(
|
||||
mmsi=257123000,
|
||||
source="vesselfinder",
|
||||
payload={"images": ["https://example.com/a.jpg"]},
|
||||
fetched_at=now - timedelta(days=30),
|
||||
expires_at=now - timedelta(days=1),
|
||||
)
|
||||
db = _StoreSession(profile=fresh, media=expired_media)
|
||||
|
||||
bundle = await get_vessel_enrichment_bundle(db, 257123000)
|
||||
|
||||
assert bundle["profile"]["payload"]["vessel_subtype"] == "Container"
|
||||
assert bundle["media"] is None
|
||||
|
||||
|
||||
def test_apply_upsert_preserves_payload_and_metadata():
|
||||
record = VesselProfileEnrichment(mmsi=257123000)
|
||||
out = _apply_upsert(
|
||||
record,
|
||||
{
|
||||
"source": "vesselfinder",
|
||||
"payload": {"vessel_subtype": "Container", "operator": "Maersk"},
|
||||
"expires_at": "2026-12-31T00:00:00Z",
|
||||
"confidence": 0.85,
|
||||
"reference_url": "https://www.vesselfinder.com/vessels/257123000",
|
||||
},
|
||||
)
|
||||
assert out["payload"]["operator"] == "Maersk"
|
||||
assert out["confidence"] == 0.85
|
||||
assert record.reference_url == "https://www.vesselfinder.com/vessels/257123000"
|
||||
assert record.expires_at is not None
|
||||
assert record.expires_at.year == 2026
|
||||
|
||||
|
||||
def _obs(*, source: str, mmsi: int, observed_at, **payload) -> AISRawObservation:
|
||||
payload = {"mmsi": mmsi, "lat": 60.0, "lon": 5.0, **payload}
|
||||
delivery_mode = "realtime_stream" if source == "aisstream_vessels" else "polling"
|
||||
transport = "websocket" if source == "aisstream_vessels" else "http"
|
||||
return AISRawObservation(
|
||||
target_schema="vessel_ais",
|
||||
source=source,
|
||||
entity_key=str(mmsi),
|
||||
delivery_mode=delivery_mode,
|
||||
transport=transport,
|
||||
message_type="PositionReport",
|
||||
observation_hash=f"{source}:{mmsi}:{observed_at.isoformat()}",
|
||||
observed_at=observed_at,
|
||||
collected_at=observed_at,
|
||||
normalized_payload=payload,
|
||||
raw_payload=payload,
|
||||
quality_flags=[],
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_promoted_rule_wins_during_aggregation():
|
||||
"""Simulate the strategy that conflict-promote-to-rule writes."""
|
||||
now = datetime.now(timezone.utc)
|
||||
obs_a = _obs(
|
||||
source="aisstream_vessels",
|
||||
mmsi=257111000,
|
||||
observed_at=now,
|
||||
name="STREAM NAME",
|
||||
vessel_type_name="Cargo",
|
||||
)
|
||||
obs_b = _obs(
|
||||
source="barentswatch_vessels",
|
||||
mmsi=257111000,
|
||||
observed_at=now - timedelta(seconds=1),
|
||||
name="REST NAME",
|
||||
vessel_type_name="Cargo",
|
||||
)
|
||||
promoted_strategy = {
|
||||
"version": 99,
|
||||
"vessel_ais": {
|
||||
"source_priority": [],
|
||||
"field_rules": {
|
||||
"name": {"mode": "source_priority", "source_priority": ["barentswatch_vessels"]}
|
||||
},
|
||||
"freshness": {"realtime_stream_seconds": 0, "polling_seconds": 0},
|
||||
"allow_dynamic_lock": False,
|
||||
},
|
||||
}
|
||||
|
||||
db = AsyncMock()
|
||||
vessels = await aggregate_vessel_observations(
|
||||
db,
|
||||
[obs_a, obs_b],
|
||||
write_conflicts=False,
|
||||
strategy=promoted_strategy,
|
||||
)
|
||||
assert vessels[0]["name"] == "REST NAME"
|
||||
assert vessels[0]["selected_reasons"]["name"] == "source_priority"
|
||||
assert vessels[0]["aggregation_strategy_version"] == 99
|
||||
|
||||
|
||||
def test_conflict_record_holds_selected_source():
|
||||
"""Sanity: the promote-to-rule API reads selected_source from this column."""
|
||||
record = AISConflictRecord(
|
||||
target_schema="vessel_ais",
|
||||
entity_key="257111000",
|
||||
field="name",
|
||||
candidates={"a": "X", "b": "Y"},
|
||||
selected_source="barentswatch_vessels",
|
||||
selected_value="Y",
|
||||
selected_reason="delivery_mode_priority",
|
||||
)
|
||||
serialized = record.to_dict()
|
||||
assert serialized["selected_source"] == "barentswatch_vessels"
|
||||
assert serialized["field"] == "name"
|
||||
@@ -4,6 +4,7 @@ from unittest.mock import AsyncMock
|
||||
import pytest
|
||||
from httpx import ASGITransport, AsyncClient
|
||||
|
||||
from app.api.v1 import visualization
|
||||
from app.api.v1.visualization import convert_vessels_to_geojson
|
||||
from app.db.session import get_db
|
||||
from app.main import app
|
||||
@@ -200,11 +201,12 @@ async def test_aggregate_vessel_observations_prefers_realtime_and_records_confli
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_vessel_collector_writes_raw_observations_without_changing_position_save(monkeypatch):
|
||||
async def test_vessel_collector_writes_raw_observations_only(monkeypatch):
|
||||
collector = VesselAISCollector()
|
||||
collector.update_progress = AsyncMock()
|
||||
record_observation = AsyncMock()
|
||||
update_health = AsyncMock()
|
||||
broadcast_custom = AsyncMock()
|
||||
monkeypatch.setattr(
|
||||
"app.services.collectors.vessel_ais.record_vessel_ais_observation",
|
||||
record_observation,
|
||||
@@ -213,6 +215,10 @@ async def test_vessel_collector_writes_raw_observations_without_changing_positio
|
||||
"app.services.collectors.vessel_ais.update_ais_source_health",
|
||||
update_health,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"app.services.collectors.vessel_ais.broadcaster.broadcast_custom",
|
||||
broadcast_custom,
|
||||
)
|
||||
|
||||
class _Session:
|
||||
def __init__(self):
|
||||
@@ -249,12 +255,17 @@ async def test_vessel_collector_writes_raw_observations_without_changing_positio
|
||||
|
||||
assert saved == 1
|
||||
assert db.committed is True
|
||||
assert any(isinstance(item, VesselStatic) for item in db.added)
|
||||
assert any(isinstance(item, VesselPosition) for item in db.added)
|
||||
# BarentsWatch must funnel through the unified AIS pipeline only — no legacy writes.
|
||||
assert not any(isinstance(item, VesselStatic) for item in db.added)
|
||||
assert not any(isinstance(item, VesselPosition) for item in db.added)
|
||||
record_observation.assert_awaited_once()
|
||||
assert record_observation.await_args.kwargs["source"] == "barentswatch_vessels"
|
||||
assert record_observation.await_args.kwargs["normalized_payload"]["mmsi"] == 257123000
|
||||
update_health.assert_awaited_once()
|
||||
broadcast_custom.assert_awaited_once()
|
||||
assert broadcast_custom.await_args.args[0] == "vessels"
|
||||
assert broadcast_custom.await_args.args[1]["action"] == "upsert"
|
||||
assert broadcast_custom.await_args.args[1]["vessels"][0]["mmsi_display"] == "257123000"
|
||||
|
||||
|
||||
def test_aisstream_collector_normalizes_position_report():
|
||||
@@ -365,6 +376,48 @@ async def test_aisstream_collector_writes_only_raw_observations(monkeypatch):
|
||||
update_health.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aisstream_stream_record_broadcasts_vessel_delta(monkeypatch):
|
||||
collector = AISStreamCollector()
|
||||
record_observation = AsyncMock(return_value=object())
|
||||
update_health = AsyncMock()
|
||||
broadcast_custom = AsyncMock()
|
||||
monkeypatch.setattr(
|
||||
"app.services.collectors.aisstream.record_vessel_ais_observation",
|
||||
record_observation,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"app.services.collectors.aisstream.update_ais_source_health",
|
||||
update_health,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"app.services.collectors.aisstream.broadcaster.broadcast_custom",
|
||||
broadcast_custom,
|
||||
)
|
||||
|
||||
class _Session:
|
||||
async def commit(self):
|
||||
pass
|
||||
|
||||
created = await collector._save_stream_record(
|
||||
_Session(),
|
||||
{
|
||||
"mmsi": 257123000,
|
||||
"lat": 59.91,
|
||||
"lon": 10.73,
|
||||
"cog": 214,
|
||||
"received_at": datetime(2026, 4, 30, 12, 0, tzinfo=timezone.utc),
|
||||
},
|
||||
)
|
||||
|
||||
assert created is True
|
||||
record_observation.assert_awaited_once()
|
||||
broadcast_custom.assert_awaited_once()
|
||||
assert broadcast_custom.await_args.args[0] == "vessels"
|
||||
assert broadcast_custom.await_args.args[1]["action"] == "upsert"
|
||||
assert broadcast_custom.await_args.args[1]["vessels"][0]["mmsi_display"] == "257123000"
|
||||
|
||||
|
||||
def test_barentswatch_reads_credentials_from_zshrc(tmp_path):
|
||||
zshrc = tmp_path / ".zshrc"
|
||||
zshrc.write_text(
|
||||
@@ -436,6 +489,39 @@ def test_convert_vessels_to_geojson():
|
||||
assert payload["features"][0]["properties"]["vessel_type_name"] == "Cargo"
|
||||
|
||||
|
||||
def test_convert_vessels_to_geojson_dedupes_mmsi_rows():
|
||||
first = VesselPosition(
|
||||
mmsi=257123000,
|
||||
lat=59.91,
|
||||
lon=10.73,
|
||||
received_at=datetime(2026, 4, 28, 1, 0, tzinfo=timezone.utc),
|
||||
)
|
||||
duplicate = VesselPosition(
|
||||
mmsi=257123000,
|
||||
lat=60.01,
|
||||
lon=10.83,
|
||||
received_at=datetime(2026, 4, 28, 1, 0, tzinfo=timezone.utc),
|
||||
)
|
||||
other = VesselPosition(
|
||||
mmsi=257456000,
|
||||
lat=60.3,
|
||||
lon=5.3,
|
||||
received_at=datetime(2026, 4, 28, 0, 59, tzinfo=timezone.utc),
|
||||
)
|
||||
|
||||
payload = convert_vessels_to_geojson(
|
||||
[
|
||||
(first, VesselStatic(mmsi=257123000, name="OSLO TRADER")),
|
||||
(duplicate, VesselStatic(mmsi=257123000, name="OSLO TRADER DUP")),
|
||||
(other, VesselStatic(mmsi=257456000, name="BERGEN FERRY")),
|
||||
]
|
||||
)
|
||||
|
||||
mmsis = [feature["properties"]["mmsi"] for feature in payload["features"]]
|
||||
assert mmsis == [257123000, 257456000]
|
||||
assert payload["features"][0]["geometry"]["coordinates"] == [10.73, 59.91]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_vessels_geojson_endpoint_filters_type_and_bbox():
|
||||
now = datetime(2026, 4, 28, 1, 0, tzinfo=timezone.utc)
|
||||
@@ -477,3 +563,113 @@ async def test_vessels_geojson_endpoint_filters_type_and_bbox():
|
||||
assert data["stats"]["by_type"]["Cargo"] == 1
|
||||
finally:
|
||||
app.dependency_overrides.clear()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_vessels_geojson_merges_raw_and_legacy_sources(monkeypatch):
|
||||
now = datetime(2026, 4, 28, 1, 0, tzinfo=timezone.utc)
|
||||
monkeypatch.setattr(
|
||||
visualization,
|
||||
"get_aggregated_vessels",
|
||||
AsyncMock(
|
||||
return_value=[
|
||||
{
|
||||
"mmsi": 1,
|
||||
"lat": 59.9,
|
||||
"lon": 10.7,
|
||||
"received_at": now,
|
||||
"name": "AISSTREAM SHIP",
|
||||
"vessel_type_name": "Cargo",
|
||||
"source_summary": {"aisstream_vessels": {"message_types": ["PositionReport"]}},
|
||||
}
|
||||
]
|
||||
),
|
||||
)
|
||||
rows = [
|
||||
(
|
||||
VesselPosition(mmsi=1, lat=60.0, lon=10.8, received_at=now),
|
||||
VesselStatic(mmsi=1, name="LEGACY DUP", vessel_type_name="Cargo"),
|
||||
),
|
||||
(
|
||||
VesselPosition(mmsi=2, lat=60.3, lon=5.3, received_at=now),
|
||||
VesselStatic(mmsi=2, name="BARENTSWATCH ONLY", vessel_type_name="Passenger"),
|
||||
),
|
||||
]
|
||||
|
||||
class _Result:
|
||||
def all(self):
|
||||
return rows
|
||||
|
||||
class _FakeSession:
|
||||
async def execute(self, _query):
|
||||
return _Result()
|
||||
|
||||
async def override_get_db():
|
||||
yield _FakeSession()
|
||||
|
||||
app.dependency_overrides[get_db] = override_get_db
|
||||
transport = ASGITransport(app=app)
|
||||
try:
|
||||
async with AsyncClient(transport=transport, base_url="http://test") as client:
|
||||
response = await client.get("/api/v1/visualization/geo/vessels")
|
||||
|
||||
assert response.status_code == 200
|
||||
data = response.json()
|
||||
names = {feature["properties"]["mmsi"]: feature["properties"]["name"] for feature in data["features"]}
|
||||
assert data["count"] == 2
|
||||
assert names == {1: "AISSTREAM SHIP", 2: "BARENTSWATCH ONLY"}
|
||||
assert data["diagnostics"]["legacy_backfilled_mmsi"] == 1
|
||||
finally:
|
||||
app.dependency_overrides.clear()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_vessel_name_fallbacks_reports_mmsi_display_names(monkeypatch):
|
||||
now = datetime(2026, 4, 28, 1, 0, tzinfo=timezone.utc)
|
||||
monkeypatch.setattr(
|
||||
visualization,
|
||||
"get_aggregated_vessels",
|
||||
AsyncMock(
|
||||
return_value=[
|
||||
{
|
||||
"mmsi": 257123000,
|
||||
"lat": 59.9,
|
||||
"lon": 10.7,
|
||||
"received_at": now,
|
||||
"name": "MMSI 257123000",
|
||||
"vessel_type_name": "Other",
|
||||
"source_summary": {
|
||||
"aisstream_vessels": {
|
||||
"latest_observed_at": now,
|
||||
"message_types": ["PositionReport"],
|
||||
}
|
||||
},
|
||||
}
|
||||
]
|
||||
),
|
||||
)
|
||||
|
||||
class _Result:
|
||||
def all(self):
|
||||
return []
|
||||
|
||||
class _FakeSession:
|
||||
async def execute(self, _query):
|
||||
return _Result()
|
||||
|
||||
async def override_get_db():
|
||||
yield _FakeSession()
|
||||
|
||||
app.dependency_overrides[get_db] = override_get_db
|
||||
transport = ASGITransport(app=app)
|
||||
try:
|
||||
async with AsyncClient(transport=transport, base_url="http://test") as client:
|
||||
response = await client.get("/api/v1/visualization/vessels/name-fallbacks")
|
||||
|
||||
assert response.status_code == 200
|
||||
data = response.json()
|
||||
assert data["count"] == 1
|
||||
assert data["items"][0]["mmsi"] == "257123000"
|
||||
assert data["items"][0]["message_types"] == ["PositionReport"]
|
||||
finally:
|
||||
app.dependency_overrides.clear()
|
||||
|
||||
@@ -292,6 +292,9 @@ async def test_visualization_geo_summary_returns_counts(monkeypatch):
|
||||
def scalar(self):
|
||||
return self._scalar_value
|
||||
|
||||
def all(self):
|
||||
return list(self._rows)
|
||||
|
||||
def scalars(self):
|
||||
class _Scalars:
|
||||
def __init__(self, rows):
|
||||
@@ -304,13 +307,18 @@ async def test_visualization_geo_summary_returns_counts(monkeypatch):
|
||||
|
||||
class _FakeSession:
|
||||
async def execute(self, query):
|
||||
query_text = str(query)
|
||||
query_text = str(query).lower()
|
||||
if "bgp_incidents" in query_text:
|
||||
return _ScalarResult(scalar_value=2)
|
||||
if "bgp_anomalies" in query_text:
|
||||
return _ScalarResult(scalar_value=3)
|
||||
if "ais_raw_observations" in query_text or "vessel_position" in query_text:
|
||||
return _ScalarResult(rows=[])
|
||||
return _ScalarResult(rows=records)
|
||||
|
||||
async def get(self, *_args, **_kwargs):
|
||||
return None
|
||||
|
||||
async def override_get_db():
|
||||
yield _FakeSession()
|
||||
|
||||
|
||||
46
backend/tests/test_websocket_manager.py
Normal file
46
backend/tests/test_websocket_manager.py
Normal file
@@ -0,0 +1,46 @@
|
||||
import pytest
|
||||
|
||||
from app.core.websocket.manager import ConnectionManager
|
||||
|
||||
|
||||
class FakeWebSocket:
|
||||
def __init__(self):
|
||||
self.accepted = False
|
||||
self.sent = []
|
||||
self.closed = False
|
||||
|
||||
async def accept(self):
|
||||
self.accepted = True
|
||||
|
||||
async def send_json(self, message):
|
||||
self.sent.append(message)
|
||||
|
||||
async def close(self):
|
||||
self.closed = True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_channel_subscribers_receive_channel_broadcasts():
|
||||
manager = ConnectionManager()
|
||||
socket = FakeWebSocket()
|
||||
|
||||
await manager.connect(socket, "user-1")
|
||||
manager.subscribe(socket, ["dashboard"])
|
||||
await manager.broadcast({"type": "data_frame", "channel": "dashboard"}, channel="dashboard")
|
||||
|
||||
assert socket.accepted is True
|
||||
assert socket.sent == [{"type": "data_frame", "channel": "dashboard"}]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_disconnect_removes_channel_subscriptions():
|
||||
manager = ConnectionManager()
|
||||
socket = FakeWebSocket()
|
||||
|
||||
await manager.connect(socket, "user-1")
|
||||
manager.subscribe(socket, ["dashboard"])
|
||||
manager.disconnect(socket, "user-1")
|
||||
await manager.broadcast({"type": "data_frame", "channel": "dashboard"}, channel="dashboard")
|
||||
|
||||
assert socket.sent == []
|
||||
assert "dashboard" not in manager.channel_subscriptions
|
||||
Reference in New Issue
Block a user