fix: tighten datasource trigger guards

This commit is contained in:
linkong
2026-04-08 09:33:48 +08:00
parent 2da6ed166b
commit 2d43263b9e
3 changed files with 96 additions and 10 deletions

View File

@@ -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)

View File

@@ -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

View File

@@ -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"""