"""Kafka-ready datasource job queue backed by PostgreSQL for v1.""" from __future__ import annotations import asyncio from datetime import UTC, datetime, timedelta from uuid import uuid4 from typing import Any from sqlalchemy import bindparam, select, text from sqlalchemy.ext.asyncio import AsyncSession from app.core.cache import cache from app.core.config import settings from app.core.enums import JobStatus, JobType, RollbackPolicy from app.core.logging import get_logger from app.core.time import to_iso8601_utc from app.core.websocket.broadcaster import broadcaster from app.db.session import async_session_factory from app.models.collected_data import CollectedData from app.models.data_snapshot import DataSnapshot from app.models.datasource import DataSource from app.models.task import CollectionTask from app.services.collectors.registry import collector_registry from app.services.datasource_connectivity import ( build_builtin_connectivity_checksum, get_builtin_effective_candidate, save_connectivity_success, ) from app.services.earth_layer_adapters import ( clear_derived_datasource_data, get_earth_refresh_strategy_for_change, get_earth_update_layers_for_source, ) from app.services.earth_layer_cache import invalidate_earth_layer_cache_for_source from app.services.scheduler import sync_datasource_job logger = get_logger(__name__) JOB_TYPE_COLLECT = JobType.COLLECT.value JOB_TYPE_CLEAR_DATA = JobType.CLEAR_DATA.value JOB_TYPE_CLEAR_CACHE = JobType.CLEAR_CACHE.value JOB_TYPE_EARTH_REFRESH = JobType.EARTH_REFRESH.value JOB_STATUS_QUEUED = JobStatus.QUEUED.value JOB_STATUS_RUNNING = JobStatus.RUNNING.value JOB_STATUS_CANCELLING = JobStatus.CANCELLING.value JOB_STATUS_SUCCESS = JobStatus.SUCCESS.value JOB_STATUS_FAILED = JobStatus.FAILED.value JOB_STATUS_CANCELLED = JobStatus.CANCELLED.value ACTIVE_JOB_STATUSES = (JOB_STATUS_QUEUED, JOB_STATUS_RUNNING, JOB_STATUS_CANCELLING) TERMINAL_JOB_STATUSES = (JOB_STATUS_SUCCESS, JOB_STATUS_FAILED, JOB_STATUS_CANCELLED) DATA_WRITE_JOB_TYPES = (JOB_TYPE_COLLECT, JOB_TYPE_CLEAR_DATA, JOB_TYPE_CLEAR_CACHE) SOURCE_LOCK_JOB_STATUSES = (JOB_STATUS_RUNNING, JOB_STATUS_CANCELLING) QUEUE_POLL_SECONDS = 0.35 JOB_STALE_LOCK_MINUTES = 90 ORPHAN_CANCELLING_GRACE_SECONDS = 30 JOB_RECOVERY_SWEEP_SECONDS = 15 DATA_DELETE_BATCH_SIZE = 50_000 DEFAULT_WORKER_CONCURRENCY = 2 RUNNING_DATA_JOB_TASKS: dict[int, asyncio.Task[Any]] = {} def _utcnow() -> datetime: return datetime.now(UTC) def _job_worker_id() -> str: return f"{settings.PROJECT_NAME}:data-job-worker:{uuid4().hex[:8]}" def is_terminal_job_status(status: str | None) -> bool: return status in TERMINAL_JOB_STATUSES async def enqueue_datasource_job( db: AsyncSession, datasource: DataSource, task_type: str, *, payload: dict[str, Any] | None = None, rollback_policy: str = RollbackPolicy.KEEP_COMMITTED_BATCHES.value, dedupe_key: str | None = None, ) -> CollectionTask: if dedupe_key: existing = await _get_active_job_by_dedupe_key(db, dedupe_key) if existing is not None: return existing task = CollectionTask( datasource_id=datasource.id, source=datasource.source, task_type=task_type, status=JOB_STATUS_QUEUED, phase="queued", phase_message="任务已进入队列", payload=payload or {}, rollback_policy=rollback_policy, dedupe_key=dedupe_key, ) db.add(task) await db.commit() await db.refresh(task) await _broadcast_task_update(task) return task async def enqueue_earth_refresh_job( db: AsyncSession, *, source: str, payload: dict[str, Any] | None = None, ) -> CollectionTask | None: layers = list((payload or {}).get("layers") or get_earth_update_layers_for_source(source)) if not layers: return None datasource = await _get_or_create_virtual_datasource(db, source) refresh_payload = { "source": source, "layers": layers, "refresh_strategy": (payload or {}).get("refresh_strategy") or get_earth_refresh_strategy_for_change((payload or {}).get("table"), source) or "clear_then_reload", **(payload or {}), } return await enqueue_datasource_job( db, datasource, JOB_TYPE_EARTH_REFRESH, payload=refresh_payload, dedupe_key=f"earth_refresh:{source}", ) async def enqueue_earth_refresh_from_update(payload: dict[str, Any]) -> None: source = str(payload.get("source") or "").strip() if not source: return async with async_session_factory() as db: await enqueue_earth_refresh_job(db, source=source, payload=payload) async def request_cancel_datasource_task( db: AsyncSession, task: CollectionTask, *, reason: str = "cancelled_by_operator", ) -> CollectionTask: if is_terminal_job_status(task.status): return task running_task = RUNNING_DATA_JOB_TASKS.get(task.id) if task.status == JOB_STATUS_QUEUED or ( running_task is None and ( task.status == JOB_STATUS_CANCELLING or (task.status == JOB_STATUS_RUNNING and task.task_type != JOB_TYPE_COLLECT) ) ): return await _cancel_task_without_runner(db, task, reason=reason) task.status = JOB_STATUS_CANCELLING task.phase = JOB_STATUS_CANCELLING task.phase_message = "正在停止任务" task.requested_cancel_at = _utcnow() task.cancel_reason = reason await db.commit() await db.refresh(task) if running_task is not None and not running_task.done(): running_task.cancel() await _broadcast_task_update(task) return task async def _cancel_task_without_runner( db: AsyncSession, task: CollectionTask, *, reason: str, ) -> CollectionTask: if task.task_type == JOB_TYPE_COLLECT: await db.execute(CollectedData.__table__.delete().where(CollectedData.task_id == task.id)) snapshot_result = await db.execute(select(DataSnapshot).where(DataSnapshot.task_id == task.id)) for snapshot in snapshot_result.scalars().all(): snapshot.status = JOB_STATUS_CANCELLED snapshot.completed_at = _utcnow() snapshot.error_message = "Cancelled after operator stop request; no active worker handle remained" datasource = await db.get(DataSource, task.datasource_id) if datasource is not None: datasource.last_status = JOB_STATUS_CANCELLED task.status = JOB_STATUS_CANCELLED task.phase = JOB_STATUS_CANCELLED task.phase_message = "任务已停止" task.completed_at = _utcnow() task.requested_cancel_at = task.requested_cancel_at or _utcnow() task.cancel_reason = reason await db.commit() await db.refresh(task) await _broadcast_task_update(task) return task async def get_active_datasource_job( db: AsyncSession, datasource_id: int, *, task_types: tuple[str, ...] = DATA_WRITE_JOB_TYPES, ) -> CollectionTask | None: result = await db.execute( select(CollectionTask) .where(CollectionTask.datasource_id == datasource_id) .where(CollectionTask.task_type.in_(task_types)) .where(CollectionTask.status.in_(ACTIVE_JOB_STATUSES)) .order_by(CollectionTask.created_at.desc().nullslast(), CollectionTask.id.desc()) .limit(1) ) return result.scalar_one_or_none() async def _get_active_job_by_dedupe_key(db: AsyncSession, dedupe_key: str) -> CollectionTask | None: result = await db.execute( select(CollectionTask) .where(CollectionTask.dedupe_key == dedupe_key) .where(CollectionTask.status.in_(ACTIVE_JOB_STATUSES)) .order_by(CollectionTask.created_at.desc().nullslast(), CollectionTask.id.desc()) .limit(1) ) return result.scalar_one_or_none() async def _get_or_create_virtual_datasource(db: AsyncSession, source: str) -> DataSource: result = await db.execute(select(DataSource).where(DataSource.source == source)) datasource = result.scalar_one_or_none() if datasource is not None: return datasource datasource = DataSource( name=f"Earth refresh: {source}", source=source, module="SYS", collector_class="EarthRefreshJob", is_active=True, ) db.add(datasource) await db.commit() await db.refresh(datasource) return datasource async def _broadcast_task_update(task: CollectionTask) -> None: await broadcaster.broadcast_datasource_task_update( { "datasource_id": task.datasource_id, "collector_name": task.source, "task_id": task.id, "task_type": task.task_type, "status": task.status, "phase": task.phase, "phase_progress": task.phase_progress, "phase_message": task.phase_message, "phase_current": task.phase_current, "phase_total": task.phase_total, "phase_unit": task.phase_unit, "progress": task.progress, "records_processed": task.records_processed, "total_records": task.total_records, "started_at": to_iso8601_utc(task.started_at), "completed_at": to_iso8601_utc(task.completed_at), "requested_cancel_at": to_iso8601_utc(task.requested_cancel_at), "error_message": task.error_message, } ) class DataJobWorker: def __init__(self, *, concurrency: int = DEFAULT_WORKER_CONCURRENCY) -> None: self.worker_id = _job_worker_id() self.concurrency = max(1, concurrency) self._task: asyncio.Task[None] | None = None self._stop_event: asyncio.Event | None = None self._running: set[asyncio.Task[Any]] = set() self._last_recovery_sweep_at: datetime | None = None def start(self) -> None: if self._task and not self._task.done(): return self._stop_event = asyncio.Event() self._task = asyncio.create_task(self._run(), name="data-job-worker") async def stop(self) -> None: if self._stop_event: self._stop_event.set() for task in list(self._running): task.cancel() if self._task: await asyncio.gather(self._task, return_exceptions=True) if self._running: await asyncio.gather(*self._running, return_exceptions=True) async def _run(self) -> None: assert self._stop_event is not None await self._recover_stale_running_jobs() while not self._stop_event.is_set(): self._running = {task for task in self._running if not task.done()} if ( self._last_recovery_sweep_at is None or (_utcnow() - self._last_recovery_sweep_at).total_seconds() >= JOB_RECOVERY_SWEEP_SECONDS ): await self._recover_stale_running_jobs() if len(self._running) >= self.concurrency: await asyncio.sleep(QUEUE_POLL_SECONDS) continue task_id = await self._claim_next_job() if task_id is None: await asyncio.sleep(QUEUE_POLL_SECONDS) continue runner = asyncio.create_task(self._run_claimed_job(task_id), name=f"data-job:{task_id}") self._running.add(runner) async def _recover_stale_running_jobs(self) -> None: self._last_recovery_sweep_at = _utcnow() cutoff = _utcnow() - timedelta(minutes=JOB_STALE_LOCK_MINUTES) orphan_cancelling_cutoff = _utcnow() - timedelta(seconds=ORPHAN_CANCELLING_GRACE_SECONDS) async with async_session_factory() as db: result = await db.execute( select(CollectionTask) .where(CollectionTask.status.in_((JOB_STATUS_RUNNING, JOB_STATUS_CANCELLING))) .where(CollectionTask.locked_at.is_not(None)) .where(CollectionTask.locked_at < cutoff) ) stale_jobs = list(result.scalars().all()) for job in stale_jobs: job.status = JOB_STATUS_FAILED job.phase = JOB_STATUS_FAILED job.completed_at = _utcnow() job.error_message = "Marked failed after stale data job lock timeout" if stale_jobs: await db.commit() orphan_result = await db.execute( select(CollectionTask) .where(CollectionTask.status == JOB_STATUS_CANCELLING) .where(CollectionTask.locked_at.is_(None)) .where(CollectionTask.requested_cancel_at.is_not(None)) .where(CollectionTask.requested_cancel_at < orphan_cancelling_cutoff) ) for job in orphan_result.scalars().all(): if job.id in RUNNING_DATA_JOB_TASKS: continue await _cancel_task_without_runner( db, job, reason=job.cancel_reason or "cancelled_after_orphaned_runner", ) async def _claim_next_job(self) -> int | None: async with async_session_factory() as db: row = await db.execute( text( """ SELECT queued.id FROM collection_tasks AS queued WHERE queued.status = :queued_status AND NOT EXISTS ( SELECT 1 FROM collection_tasks AS active WHERE active.source = queued.source AND active.id <> queued.id AND active.status IN :active_statuses ) ORDER BY queued.created_at ASC NULLS FIRST, queued.id ASC LIMIT 1 FOR UPDATE SKIP LOCKED """ ).bindparams(bindparam("active_statuses", expanding=True)), { "queued_status": JOB_STATUS_QUEUED, "active_statuses": SOURCE_LOCK_JOB_STATUSES, }, ) task_id = row.scalar_one_or_none() if task_id is None: return None task = await db.get(CollectionTask, int(task_id)) if task is None: return None task.status = JOB_STATUS_RUNNING task.phase = "starting" task.phase_message = "任务开始执行" task.started_at = task.started_at or _utcnow() task.worker_id = self.worker_id task.locked_at = _utcnow() await db.commit() await _broadcast_task_update(task) return int(task_id) async def _run_claimed_job(self, task_id: int) -> None: current_task = asyncio.current_task() if current_task is not None: RUNNING_DATA_JOB_TASKS[task_id] = current_task try: async with async_session_factory() as db: task = await db.get(CollectionTask, task_id) if task is None: return await self._execute_job(db, task) except asyncio.CancelledError: async with async_session_factory() as db: task = await db.get(CollectionTask, task_id) if task is not None and not is_terminal_job_status(task.status): task.status = JOB_STATUS_CANCELLED task.phase = JOB_STATUS_CANCELLED task.phase_message = "任务已停止" task.completed_at = _utcnow() await db.commit() await _broadcast_task_update(task) raise except Exception as exc: logger.exception_event( "Data job failed", event="data_jobs.job_failed", context={"task_id": task_id, "error": str(exc)}, ) async with async_session_factory() as db: task = await db.get(CollectionTask, task_id) if task is not None: task.status = JOB_STATUS_FAILED task.phase = JOB_STATUS_FAILED task.phase_message = str(exc) task.error_message = str(exc) task.completed_at = _utcnow() await db.commit() await _broadcast_task_update(task) finally: RUNNING_DATA_JOB_TASKS.pop(task_id, None) async def _execute_job(self, db: AsyncSession, task: CollectionTask) -> None: if task.task_type == JOB_TYPE_COLLECT: await _run_collect_job(db, task) elif task.task_type == JOB_TYPE_CLEAR_DATA: await _run_clear_data_job(db, task) elif task.task_type == JOB_TYPE_CLEAR_CACHE: await _run_clear_cache_job(db, task) elif task.task_type == JOB_TYPE_EARTH_REFRESH: await _run_earth_refresh_job(db, task) else: raise RuntimeError(f"Unsupported data job type: {task.task_type}") async def _run_collect_job(db: AsyncSession, task: CollectionTask) -> None: datasource = await db.get(DataSource, task.datasource_id) if datasource is None: raise RuntimeError("Data source not found") collector = collector_registry.get(datasource.source) if collector is None: raise RuntimeError(f"Collector '{datasource.source}' not found") if not datasource.is_active: raise RuntimeError("Data source is disabled") collector._datasource_id = datasource.id collector._current_task = task collector._db_session = db result = await collector.run(db) datasource.last_run_at = _utcnow() datasource.last_status = result.get("status") if datasource.last_status == JOB_STATUS_SUCCESS: effective_candidate = await get_builtin_effective_candidate(db, datasource.source) checksum, _credential_context = await build_builtin_connectivity_checksum( datasource.source, effective_candidate["endpoint"], effective_candidate["auth_type"], effective_candidate["headers"], effective_candidate["config"], db, ) await save_connectivity_success( db, datasource.source, checksum, {"status_code": None}, connected_by="collection", ) await db.commit() await sync_datasource_job(datasource.id) async def _run_clear_data_job(db: AsyncSession, task: CollectionTask) -> None: source = str(task.source or (task.payload or {}).get("source") or "").strip() if not source: raise RuntimeError("Clear data job has no source") task.phase = "clearing_data" task.phase_message = "正在删除数据库数据" await db.commit() await _broadcast_task_update(task) deleted_count = await _delete_table_rows_by_source( db, task, table_name="collected_data", source_column="source", source=source, ) derived_deleted_counts = await _clear_derived_datasource_data_in_batches( db, task, source, progress_offset=deleted_count, ) derived_deleted_count = sum(derived_deleted_counts.values()) if any(key.startswith("ais_") for key in derived_deleted_counts): await db.execute(text("ANALYZE ais_raw_observations")) task.records_processed = deleted_count + derived_deleted_count task.total_records = task.records_processed task.progress = 100.0 task.phase_progress = 100.0 task.phase_current = task.records_processed task.phase_total = task.records_processed task.phase_unit = "records" task.payload = { **(task.payload or {}), "deleted_count": deleted_count, "derived_deleted_count": derived_deleted_count, "derived_deleted_counts": derived_deleted_counts, } task.status = JOB_STATUS_SUCCESS task.phase = "completed" task.phase_message = "数据库数据已清理" task.completed_at = _utcnow() datasource = await db.get(DataSource, task.datasource_id) if datasource is not None: datasource.last_status = JOB_STATUS_SUCCESS datasource.last_run_at = task.completed_at await db.execute( DataSnapshot.__table__.update() .where(DataSnapshot.source == source) .values(is_current=False) ) await db.commit() await _broadcast_task_update(task) async def _delete_table_rows_by_source( db: AsyncSession, task: CollectionTask, *, table_name: str, source_column: str, source: str, progress_offset: int = 0, ) -> int: deleted = 0 while True: result = await db.execute( text( f""" WITH doomed AS ( SELECT ctid FROM {table_name} WHERE {source_column} = :source LIMIT :batch_size ), deleted_rows AS ( DELETE FROM {table_name} USING doomed WHERE {table_name}.ctid = doomed.ctid RETURNING 1 ) SELECT COUNT(*) FROM deleted_rows """ ), {"source": source, "batch_size": DATA_DELETE_BATCH_SIZE}, ) batch_deleted = max(int(result.scalar_one() or 0), 0) if batch_deleted <= 0: break deleted += batch_deleted task.records_processed = progress_offset + deleted task.phase_current = task.records_processed task.phase_unit = "records" task.phase_message = f"正在删除数据:{task.records_processed} 条" await db.commit() await _broadcast_task_update(task) return deleted async def _clear_derived_datasource_data_in_batches( db: AsyncSession, task: CollectionTask, source: str, progress_offset: int = 0, ) -> dict[str, int]: deleted_counts: dict[str, int] = {} if source in {"barentswatch_vessels", "aisstream_vessels"}: deleted_counts["ais_conflict_records"] = await _delete_table_rows_by_source( db, task, table_name="ais_conflict_records", source_column="selected_source", source=source, progress_offset=progress_offset + sum(deleted_counts.values()), ) deleted_counts["ais_source_health"] = await _delete_table_rows_by_source( db, task, table_name="ais_source_health", source_column="source", source=source, progress_offset=progress_offset + sum(deleted_counts.values()), ) deleted_counts["ais_raw_observations"] = await _delete_table_rows_by_source( db, task, table_name="ais_raw_observations", source_column="source", source=source, progress_offset=progress_offset + sum(deleted_counts.values()), ) return deleted_counts return await clear_derived_datasource_data(db, source) async def _run_clear_cache_job(db: AsyncSession, task: CollectionTask) -> None: source = str(task.source or (task.payload or {}).get("source") or "").strip() if not source: raise RuntimeError("Clear cache job has no source") earth_deleted_count = invalidate_earth_layer_cache_for_source(source) dashboard_deleted_count = int(cache.delete("dashboard:stats")) + int(cache.delete("dashboard:summary")) deleted_count = earth_deleted_count + dashboard_deleted_count task.records_processed = deleted_count task.total_records = deleted_count task.progress = 100.0 task.phase_progress = 100.0 task.phase = "completed" task.phase_message = "缓存已清理" task.payload = { **(task.payload or {}), "earth_layer_deleted_count": earth_deleted_count, "dashboard_deleted_count": dashboard_deleted_count, } task.status = JOB_STATUS_SUCCESS task.completed_at = _utcnow() await db.commit() await _broadcast_task_update(task) await enqueue_earth_refresh_job(db, source=source, payload={"operation": "CACHE_INVALIDATED"}) async def _run_earth_refresh_job(db: AsyncSession, task: CollectionTask) -> None: payload = task.payload or {} source = str(payload.get("source") or task.source or "").strip() layers = list(payload.get("layers") or get_earth_update_layers_for_source(source)) if not source or not layers: task.status = JOB_STATUS_SUCCESS task.phase = "completed" task.phase_message = "没有需要刷新的 Earth 图层" task.completed_at = _utcnow() await db.commit() await _broadcast_task_update(task) return deleted_cache_entries = invalidate_earth_layer_cache_for_source(source) update_payload = { "event": "earth.layer.changed", "action": "database_changed", "source": source, "table": payload.get("table"), "data_type": source, "layers": layers, "refresh_strategy": payload.get("refresh_strategy") or "clear_then_reload", "records_processed": payload.get("records_processed", 0), "operations": payload.get("operations") or [payload.get("operation") or "CHANGE"], "operation": payload.get("operation"), "cache_entries_invalidated": deleted_cache_entries, "timestamp": to_iso8601_utc(_utcnow()), } if payload.get("entity") == "interactable": update_payload.update( { "entity": "interactable", "action": payload.get("action") or "changed", "ids": payload.get("ids") or payload.get("entity_keys") or [], "item": payload.get("item"), } ) await broadcaster.broadcast_earth_update(update_payload) task.records_processed = int(payload.get("records_processed") or 0) task.progress = 100.0 task.phase_progress = 100.0 task.phase = "completed" task.phase_message = "Earth 图层刷新通知已发送" task.status = JOB_STATUS_SUCCESS task.completed_at = _utcnow() task.payload = {**payload, "cache_entries_invalidated": deleted_cache_entries} await db.commit() await _broadcast_task_update(task) _worker = DataJobWorker() def start_data_job_worker() -> None: _worker.start() async def stop_data_job_worker() -> None: await _worker.stop()