From 2d43263b9edaf56b8dd352e2f389c029c44d0ad6 Mon Sep 17 00:00:00 2001 From: linkong Date: Wed, 8 Apr 2026 09:33:48 +0800 Subject: [PATCH] fix: tighten datasource trigger guards --- backend/app/api/v1/datasources.py | 92 +++++++++++++++++++++++++++---- backend/app/services/scheduler.py | 5 ++ backend/tests/test_api.py | 9 +++ 3 files changed, 96 insertions(+), 10 deletions(-) diff --git a/backend/app/api/v1/datasources.py b/backend/app/api/v1/datasources.py index 1428e073..d7176a02 100644 --- a/backend/app/api/v1/datasources.py +++ b/backend/app/api/v1/datasources.py @@ -93,16 +93,20 @@ async def get_running_task(db: AsyncSession, datasource_id: int) -> Optional[Col if now - started_at <= timedelta(minutes=STALE_RUNNING_TASK_TIMEOUT_MINUTES): return task - existing_error = (task.error_message or "").strip() - stale_reason = ( - f"Marked failed automatically after stale running timeout " - f"({STALE_RUNNING_TASK_TIMEOUT_MINUTES}m)" - ) - task.status = "failed" - task.phase = "failed" - task.completed_at = now - task.error_message = f"{existing_error}\n{stale_reason}".strip() if existing_error else stale_reason - await db.commit() + datasource = await db.get(DataSource, datasource_id) + if datasource is not None: + await fail_and_rollback_stale_running_task(db, datasource, task) + else: + existing_error = (task.error_message or "").strip() + stale_reason = ( + f"Marked failed automatically after stale running timeout " + f"({STALE_RUNNING_TASK_TIMEOUT_MINUTES}m)" + ) + task.status = "failed" + task.phase = "failed" + task.completed_at = now + task.error_message = f"{existing_error}\n{stale_reason}".strip() if existing_error else stale_reason + await db.commit() return None @@ -170,6 +174,73 @@ async def rollback_orphaned_running_task( await db.commit() +async def fail_and_rollback_stale_running_task( + db: AsyncSession, + datasource: DataSource, + running_task: CollectionTask, +) -> None: + snapshot_result = await db.execute( + select(DataSnapshot) + .where( + DataSnapshot.datasource_id == datasource.id, + DataSnapshot.task_id == running_task.id, + ) + .order_by(DataSnapshot.id.desc()) + .limit(1) + ) + snapshot = snapshot_result.scalar_one_or_none() + + await db.execute(CollectedData.__table__.delete().where(CollectedData.task_id == running_task.id)) + + await db.execute( + text( + """ + UPDATE collected_data + SET is_current = FALSE + WHERE source = :source + """ + ), + {"source": datasource.source}, + ) + + if snapshot is not None: + snapshot.status = "failed" + snapshot.is_current = False + snapshot.completed_at = datetime.now(timezone.utc) + summary = dict(snapshot.summary or {}) + summary["rollback"] = True + summary["rollback_reason"] = "stale_running_task_timeout" + snapshot.summary = summary + + if snapshot.parent_snapshot_id is not None: + parent_snapshot = await db.get(DataSnapshot, snapshot.parent_snapshot_id) + if parent_snapshot: + parent_snapshot.is_current = True + await db.execute( + text( + """ + UPDATE collected_data + SET is_current = TRUE + WHERE snapshot_id = :snapshot_id + """ + ), + {"snapshot_id": snapshot.parent_snapshot_id}, + ) + + existing_error = (running_task.error_message or "").strip() + stale_reason = ( + f"Marked failed automatically after stale running timeout " + f"({STALE_RUNNING_TASK_TIMEOUT_MINUTES}m); incomplete writes rolled back" + ) + running_task.status = "failed" + running_task.phase = "failed" + running_task.completed_at = datetime.now(timezone.utc) + running_task.error_message = f"{existing_error}\n{stale_reason}".strip() if existing_error else stale_reason + datasource.last_status = "failed" + datasource.last_run_at = datetime.now(timezone.utc) + await db.commit() + + @router.get("") async def list_datasources( module: Optional[str] = None, @@ -504,6 +575,7 @@ async def clear_datasource_data( async def get_task_status( source_id: str, task_id: Optional[int] = None, + current_user: User = Depends(get_current_user), db: AsyncSession = Depends(get_db), ): datasource = await get_datasource_record(db, source_id) diff --git a/backend/app/services/scheduler.py b/backend/app/services/scheduler.py index 0baf24be..a16d84ff 100644 --- a/backend/app/services/scheduler.py +++ b/backend/app/services/scheduler.py @@ -274,6 +274,11 @@ def run_collector_now(collector_name: str) -> bool: logger.error("Collector not found: %s", collector_name) return False + existing_task = get_running_collector_task(collector_name) + if existing_task is not None and not existing_task.done(): + logger.warning("Collector %s is already running in-memory; skipping duplicate trigger", collector_name) + return False + try: task = asyncio.create_task(run_collector_task(collector_name), name=_collector_task_name(collector_name)) RUNNING_COLLECTOR_TASKS[collector_name] = task diff --git a/backend/tests/test_api.py b/backend/tests/test_api.py index 5faa479e..4d79ee92 100644 --- a/backend/tests/test_api.py +++ b/backend/tests/test_api.py @@ -90,6 +90,15 @@ async def test_alerts_without_auth(): assert response.status_code == 401 +@pytest.mark.asyncio +async def test_datasource_task_status_without_auth(): + """Test datasource task-status endpoint requires authentication""" + transport = ASGITransport(app=app) + async with AsyncClient(transport=transport, base_url="http://test") as client: + response = await client.get("/api/v1/datasources/1/task-status") + assert response.status_code == 401 + + @pytest.mark.asyncio async def test_alerts_endpoint_with_auth(auth_headers): """Test alerts endpoint with authentication"""