fix: tighten datasource trigger guards
This commit is contained in:
@@ -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):
|
if now - started_at <= timedelta(minutes=STALE_RUNNING_TASK_TIMEOUT_MINUTES):
|
||||||
return task
|
return task
|
||||||
|
|
||||||
existing_error = (task.error_message or "").strip()
|
datasource = await db.get(DataSource, datasource_id)
|
||||||
stale_reason = (
|
if datasource is not None:
|
||||||
f"Marked failed automatically after stale running timeout "
|
await fail_and_rollback_stale_running_task(db, datasource, task)
|
||||||
f"({STALE_RUNNING_TASK_TIMEOUT_MINUTES}m)"
|
else:
|
||||||
)
|
existing_error = (task.error_message or "").strip()
|
||||||
task.status = "failed"
|
stale_reason = (
|
||||||
task.phase = "failed"
|
f"Marked failed automatically after stale running timeout "
|
||||||
task.completed_at = now
|
f"({STALE_RUNNING_TASK_TIMEOUT_MINUTES}m)"
|
||||||
task.error_message = f"{existing_error}\n{stale_reason}".strip() if existing_error else stale_reason
|
)
|
||||||
await db.commit()
|
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
|
return None
|
||||||
|
|
||||||
|
|
||||||
@@ -170,6 +174,73 @@ async def rollback_orphaned_running_task(
|
|||||||
await db.commit()
|
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("")
|
@router.get("")
|
||||||
async def list_datasources(
|
async def list_datasources(
|
||||||
module: Optional[str] = None,
|
module: Optional[str] = None,
|
||||||
@@ -504,6 +575,7 @@ async def clear_datasource_data(
|
|||||||
async def get_task_status(
|
async def get_task_status(
|
||||||
source_id: str,
|
source_id: str,
|
||||||
task_id: Optional[int] = None,
|
task_id: Optional[int] = None,
|
||||||
|
current_user: User = Depends(get_current_user),
|
||||||
db: AsyncSession = Depends(get_db),
|
db: AsyncSession = Depends(get_db),
|
||||||
):
|
):
|
||||||
datasource = await get_datasource_record(db, source_id)
|
datasource = await get_datasource_record(db, source_id)
|
||||||
|
|||||||
@@ -274,6 +274,11 @@ def run_collector_now(collector_name: str) -> bool:
|
|||||||
logger.error("Collector not found: %s", collector_name)
|
logger.error("Collector not found: %s", collector_name)
|
||||||
return False
|
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:
|
try:
|
||||||
task = asyncio.create_task(run_collector_task(collector_name), name=_collector_task_name(collector_name))
|
task = asyncio.create_task(run_collector_task(collector_name), name=_collector_task_name(collector_name))
|
||||||
RUNNING_COLLECTOR_TASKS[collector_name] = task
|
RUNNING_COLLECTOR_TASKS[collector_name] = task
|
||||||
|
|||||||
@@ -90,6 +90,15 @@ async def test_alerts_without_auth():
|
|||||||
assert response.status_code == 401
|
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
|
@pytest.mark.asyncio
|
||||||
async def test_alerts_endpoint_with_auth(auth_headers):
|
async def test_alerts_endpoint_with_auth(auth_headers):
|
||||||
"""Test alerts endpoint with authentication"""
|
"""Test alerts endpoint with authentication"""
|
||||||
|
|||||||
Reference in New Issue
Block a user