Add force rerun recovery and polish CLI startup
This commit is contained in:
@@ -3,7 +3,7 @@ from datetime import datetime, timedelta, timezone
|
||||
from typing import Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query
|
||||
from sqlalchemy import func, select
|
||||
from sqlalchemy import func, select, text
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.core.time import to_iso8601_utc
|
||||
@@ -11,10 +11,16 @@ from app.core.security import get_current_user
|
||||
from app.core.data_sources import get_data_sources_config
|
||||
from app.db.session import get_db
|
||||
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.models.user import User
|
||||
from app.services.scheduler import get_latest_task_id_for_datasource, run_collector_now, sync_datasource_job
|
||||
from app.services.scheduler import (
|
||||
cancel_running_collector_now,
|
||||
get_latest_task_id_for_datasource,
|
||||
run_collector_now,
|
||||
sync_datasource_job,
|
||||
)
|
||||
|
||||
router = APIRouter()
|
||||
STALE_RUNNING_TASK_TIMEOUT_MINUTES = 90
|
||||
@@ -100,6 +106,70 @@ async def get_running_task(db: AsyncSession, datasource_id: int) -> Optional[Col
|
||||
return None
|
||||
|
||||
|
||||
async def rollback_orphaned_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 = "cancelled"
|
||||
snapshot.is_current = False
|
||||
snapshot.completed_at = datetime.now(timezone.utc)
|
||||
summary = dict(snapshot.summary or {})
|
||||
summary["rollback"] = True
|
||||
summary["rollback_reason"] = "orphaned_running_task_after_backend_restart"
|
||||
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},
|
||||
)
|
||||
|
||||
running_task.status = "cancelled"
|
||||
running_task.phase = "cancelled"
|
||||
running_task.completed_at = datetime.now(timezone.utc)
|
||||
existing_error = (running_task.error_message or "").strip()
|
||||
cancel_reason = "Cancelled after backend restart because the running task handle was lost; incomplete writes rolled back"
|
||||
running_task.error_message = f"{existing_error}\n{cancel_reason}".strip() if existing_error else cancel_reason
|
||||
datasource.last_status = "cancelled"
|
||||
datasource.last_run_at = datetime.now(timezone.utc)
|
||||
await db.commit()
|
||||
|
||||
|
||||
@router.get("")
|
||||
async def list_datasources(
|
||||
module: Optional[str] = None,
|
||||
@@ -346,6 +416,7 @@ async def get_datasource_stats(
|
||||
@router.post("/{source_id}/trigger")
|
||||
async def trigger_datasource(
|
||||
source_id: str,
|
||||
force: bool = Query(False),
|
||||
current_user: User = Depends(get_current_user),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
@@ -356,6 +427,26 @@ async def trigger_datasource(
|
||||
if not datasource.is_active:
|
||||
raise HTTPException(status_code=400, detail="Data source is disabled")
|
||||
|
||||
running_task = await get_running_task(db, datasource.id)
|
||||
if running_task is not None and not force:
|
||||
raise HTTPException(
|
||||
status_code=409,
|
||||
detail={
|
||||
"reason": "running_task_in_progress",
|
||||
"message": "当前采集任务尚未完成,重新触发会丢失本次未完成进度。是否强制重新采集?",
|
||||
"task_id": running_task.id,
|
||||
"phase": running_task.phase,
|
||||
"progress": running_task.progress,
|
||||
"records_processed": running_task.records_processed,
|
||||
"total_records": running_task.total_records,
|
||||
},
|
||||
)
|
||||
|
||||
if running_task is not None and force:
|
||||
cancelled = await cancel_running_collector_now(datasource.source)
|
||||
if not cancelled:
|
||||
await rollback_orphaned_running_task(db, datasource, running_task)
|
||||
|
||||
previous_task_id = await get_latest_task_id_for_datasource(datasource.id)
|
||||
success = run_collector_now(datasource.source)
|
||||
if not success:
|
||||
@@ -375,6 +466,7 @@ async def trigger_datasource(
|
||||
"source_id": datasource.id,
|
||||
"task_id": task_id,
|
||||
"collector_name": datasource.source,
|
||||
"force": force,
|
||||
"message": f"Collector '{datasource.source}' has been triggered",
|
||||
}
|
||||
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
"""Base collector class for all data sources"""
|
||||
|
||||
import asyncio
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Dict, List, Any, Optional
|
||||
from datetime import UTC, datetime
|
||||
@@ -166,6 +167,59 @@ class BaseCollector(ABC):
|
||||
await db.commit()
|
||||
return snapshot.id
|
||||
|
||||
async def _rollback_incomplete_run(
|
||||
self,
|
||||
db: AsyncSession,
|
||||
*,
|
||||
task_id: int,
|
||||
snapshot_id: Optional[int],
|
||||
reason: str,
|
||||
) -> None:
|
||||
from app.models.collected_data import CollectedData
|
||||
from app.models.data_snapshot import DataSnapshot
|
||||
|
||||
await db.execute(CollectedData.__table__.delete().where(CollectedData.task_id == task_id))
|
||||
|
||||
parent_snapshot_id: Optional[int] = None
|
||||
if snapshot_id is not None:
|
||||
snapshot = await db.get(DataSnapshot, snapshot_id)
|
||||
if snapshot:
|
||||
parent_snapshot_id = snapshot.parent_snapshot_id
|
||||
snapshot.status = "cancelled"
|
||||
snapshot.is_current = False
|
||||
snapshot.completed_at = datetime.now(UTC)
|
||||
summary = dict(snapshot.summary or {})
|
||||
summary["rollback"] = True
|
||||
summary["rollback_reason"] = reason
|
||||
snapshot.summary = summary
|
||||
|
||||
await db.execute(
|
||||
text(
|
||||
"""
|
||||
UPDATE collected_data
|
||||
SET is_current = FALSE
|
||||
WHERE source = :source
|
||||
"""
|
||||
),
|
||||
{"source": self.name},
|
||||
)
|
||||
|
||||
if parent_snapshot_id is not None:
|
||||
parent_snapshot = await db.get(DataSnapshot, 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": parent_snapshot_id},
|
||||
)
|
||||
|
||||
async def run(self, db: AsyncSession) -> Dict[str, Any]:
|
||||
"""Full pipeline: fetch -> transform -> save"""
|
||||
from app.services.collectors.registry import collector_registry
|
||||
@@ -227,6 +281,21 @@ class BaseCollector(ABC):
|
||||
"records_processed": records_count,
|
||||
"execution_time_seconds": (datetime.now(UTC) - start_time).total_seconds(),
|
||||
}
|
||||
except asyncio.CancelledError:
|
||||
task.status = "cancelled"
|
||||
task.phase = "cancelled"
|
||||
task.error_message = "Collection cancelled by operator and rolled back"
|
||||
task.completed_at = datetime.now(UTC)
|
||||
if snapshot_id is not None:
|
||||
await self._rollback_incomplete_run(
|
||||
db,
|
||||
task_id=task_id,
|
||||
snapshot_id=snapshot_id,
|
||||
reason="cancelled_by_operator",
|
||||
)
|
||||
await db.commit()
|
||||
await self._publish_task_update(force=True)
|
||||
raise
|
||||
except Exception as e:
|
||||
task.status = "failed"
|
||||
task.phase = "failed"
|
||||
@@ -276,20 +345,34 @@ class BaseCollector(ABC):
|
||||
updated_count = 0
|
||||
unchanged_count = 0
|
||||
seen_entity_keys: set[str] = set()
|
||||
previous_current_keys: set[str] = set()
|
||||
progress_commit_interval = 1000
|
||||
|
||||
previous_current_result = await db.execute(
|
||||
select(CollectedData.entity_key).where(
|
||||
select(CollectedData)
|
||||
.where(
|
||||
CollectedData.source == self.name,
|
||||
CollectedData.is_current == True,
|
||||
)
|
||||
.order_by(CollectedData.entity_key.asc(), CollectedData.collected_at.desc().nullslast(), CollectedData.id.desc())
|
||||
)
|
||||
previous_current_keys = {row[0] for row in previous_current_result.fetchall() if row[0]}
|
||||
previous_current_records = previous_current_result.scalars().all()
|
||||
previous_current_keys = {record.entity_key for record in previous_current_records if record.entity_key}
|
||||
previous_current_map: dict[str, CollectedData] = {}
|
||||
stale_previous_records: list[CollectedData] = []
|
||||
|
||||
for existing_record in previous_current_records:
|
||||
entity_key = existing_record.entity_key
|
||||
if not entity_key:
|
||||
continue
|
||||
if entity_key not in previous_current_map:
|
||||
previous_current_map[entity_key] = existing_record
|
||||
continue
|
||||
stale_previous_records.append(existing_record)
|
||||
|
||||
for stale_record in stale_previous_records:
|
||||
stale_record.is_current = False
|
||||
|
||||
for i, item in enumerate(data):
|
||||
print(
|
||||
f"DEBUG: Saving item {i}: name={item.get('name')}, metadata={item.get('metadata', 'NOT FOUND')}"
|
||||
)
|
||||
raw_metadata = item.get("metadata", {})
|
||||
extra_data = build_dynamic_metadata(
|
||||
raw_metadata,
|
||||
@@ -318,20 +401,9 @@ class BaseCollector(ABC):
|
||||
previous_record = None
|
||||
|
||||
if entity_key and entity_key not in seen_entity_keys:
|
||||
result = await db.execute(
|
||||
select(CollectedData)
|
||||
.where(
|
||||
CollectedData.source == self.name,
|
||||
CollectedData.entity_key == entity_key,
|
||||
CollectedData.is_current == True,
|
||||
)
|
||||
.order_by(CollectedData.collected_at.desc().nullslast(), CollectedData.id.desc())
|
||||
)
|
||||
previous_records = result.scalars().all()
|
||||
if previous_records:
|
||||
previous_record = previous_records[0]
|
||||
for old_record in previous_records:
|
||||
old_record.is_current = False
|
||||
previous_record = previous_current_map.get(entity_key)
|
||||
if previous_record is not None:
|
||||
previous_record.is_current = False
|
||||
|
||||
record = CollectedData(
|
||||
snapshot_id=snapshot_id,
|
||||
@@ -375,7 +447,7 @@ class BaseCollector(ABC):
|
||||
seen_entity_keys.add(entity_key)
|
||||
records_added += 1
|
||||
|
||||
if i % 100 == 0:
|
||||
if (i + 1) % progress_commit_interval == 0:
|
||||
await self.update_progress(i + 1, commit=True)
|
||||
|
||||
if snapshot_id is not None:
|
||||
|
||||
@@ -19,6 +19,30 @@ logger = logging.getLogger(__name__)
|
||||
|
||||
scheduler = AsyncIOScheduler()
|
||||
RUNNING_TASK_GUARD_TIMEOUT_MINUTES = 90
|
||||
RUNNING_COLLECTOR_TASKS: dict[str, asyncio.Task[Any]] = {}
|
||||
|
||||
|
||||
def _collector_task_name(collector_name: str) -> str:
|
||||
return f"collector:{collector_name}"
|
||||
|
||||
|
||||
def get_running_collector_task(collector_name: str) -> asyncio.Task[Any] | None:
|
||||
task = RUNNING_COLLECTOR_TASKS.get(collector_name)
|
||||
if task is not None and not task.done():
|
||||
return task
|
||||
|
||||
if task is not None and task.done():
|
||||
RUNNING_COLLECTOR_TASKS.pop(collector_name, None)
|
||||
|
||||
target_name = _collector_task_name(collector_name)
|
||||
for candidate in asyncio.all_tasks():
|
||||
if candidate.done():
|
||||
continue
|
||||
if candidate.get_name() == target_name:
|
||||
RUNNING_COLLECTOR_TASKS[collector_name] = candidate
|
||||
return candidate
|
||||
|
||||
return None
|
||||
|
||||
|
||||
async def _update_next_run_at(datasource: DataSource, session) -> None:
|
||||
@@ -133,6 +157,12 @@ async def run_collector_task(collector_name: str):
|
||||
datasource.last_status = task_result.get("status")
|
||||
await _update_next_run_at(datasource, db)
|
||||
logger.info("Collector %s completed: %s", collector_name, task_result)
|
||||
except asyncio.CancelledError:
|
||||
datasource.last_run_at = datetime.now(UTC)
|
||||
datasource.last_status = "cancelled"
|
||||
await db.commit()
|
||||
logger.warning("Collector %s cancelled by operator", collector_name)
|
||||
raise
|
||||
except Exception as exc:
|
||||
datasource.last_run_at = datetime.now(UTC)
|
||||
datasource.last_status = "failed"
|
||||
@@ -245,9 +275,31 @@ def run_collector_now(collector_name: str) -> bool:
|
||||
return False
|
||||
|
||||
try:
|
||||
asyncio.create_task(run_collector_task(collector_name))
|
||||
task = asyncio.create_task(run_collector_task(collector_name), name=_collector_task_name(collector_name))
|
||||
RUNNING_COLLECTOR_TASKS[collector_name] = task
|
||||
|
||||
def _cleanup_task(done_task: asyncio.Task[Any]) -> None:
|
||||
current = RUNNING_COLLECTOR_TASKS.get(collector_name)
|
||||
if current is done_task:
|
||||
RUNNING_COLLECTOR_TASKS.pop(collector_name, None)
|
||||
|
||||
task.add_done_callback(_cleanup_task)
|
||||
logger.info("Triggered collector: %s", collector_name)
|
||||
return True
|
||||
except Exception as exc:
|
||||
logger.error("Failed to trigger collector %s: %s", collector_name, exc)
|
||||
return False
|
||||
return False
|
||||
|
||||
|
||||
async def cancel_running_collector_now(collector_name: str) -> bool:
|
||||
task = get_running_collector_task(collector_name)
|
||||
if task is None or task.done():
|
||||
RUNNING_COLLECTOR_TASKS.pop(collector_name, None)
|
||||
return False
|
||||
|
||||
task.cancel()
|
||||
try:
|
||||
await task
|
||||
except asyncio.CancelledError:
|
||||
return True
|
||||
return task.cancelled()
|
||||
|
||||
Reference in New Issue
Block a user