# -*- coding: utf-8 -*- """Persistent task state store for meeting-style offline jobs.""" from __future__ import annotations import json import time from pathlib import Path from threading import RLock from typing import Any, Optional from app.core.config import settings tasks_db: dict[str, dict[str, Any]] = {} _tasks_lock = RLock() _last_cleanup_at = 0 def _task_state_dir() -> Path: return Path(settings.TASK_STATE_DIR) def _task_file_path(task_id: str) -> Path: return _task_state_dir() / f"{task_id}.json" def _task_result_file_path(task_id: str) -> Path: return _task_state_dir() / f"{task_id}.result.json" def _retention_seconds() -> int: return max(0, int(settings.TASK_RETENTION_HOURS)) * 3600 def _is_task_expired(task_payload: dict[str, Any], now_ts: Optional[int] = None) -> bool: retention_seconds = _retention_seconds() if retention_seconds <= 0: return False current_time = int(now_ts or time.time()) base_timestamp = int(task_payload.get("updated_at") or task_payload.get("created_at") or 0) if base_timestamp <= 0: return False return current_time - base_timestamp >= retention_seconds def _write_json_file(file_path: Path, payload: dict[str, Any]) -> None: file_path.parent.mkdir(parents=True, exist_ok=True) temporary_file_path = file_path.with_suffix(f"{file_path.suffix}.tmp") temporary_file_path.write_text( json.dumps(payload, ensure_ascii=False), encoding="utf-8", ) temporary_file_path.replace(file_path) def _write_task_file(task_id: str, payload: dict[str, Any]) -> None: _write_json_file(_task_file_path(task_id), payload) def _remove_task_file(task_id: str) -> None: task_file_path = _task_file_path(task_id) if task_file_path.exists(): task_file_path.unlink() def _write_task_result_file(task_id: str, payload: dict[str, Any]) -> None: _write_json_file(_task_result_file_path(task_id), payload) def _remove_task_result_file(task_id: str) -> None: task_result_file_path = _task_result_file_path(task_id) if task_result_file_path.exists(): task_result_file_path.unlink() def _load_task_payload_from_file(task_file_path: Path) -> Optional[tuple[str, dict[str, Any]]]: try: payload = json.loads(task_file_path.read_text(encoding="utf-8")) except Exception: return None if not isinstance(payload, dict): return None task_id = str(payload.get("task_id") or task_file_path.stem).strip() if not task_id: return None payload["task_id"] = task_id return task_id, payload def _load_task_result_payload(task_id: str) -> Optional[dict[str, Any]]: task_result_file_path = _task_result_file_path(task_id) if not task_result_file_path.exists(): return None try: payload = json.loads(task_result_file_path.read_text(encoding="utf-8")) except Exception: return None if not isinstance(payload, dict): return None return payload def _split_task_payload(task_id: str, payload: dict[str, Any], *, persist_legacy_result: bool = False) -> tuple[dict[str, Any], bool]: status_payload = dict(payload) has_result = "result" in status_payload result_payload = status_payload.pop("result", None) if has_result: if isinstance(result_payload, dict): _write_task_result_file(task_id, result_payload) elif result_payload is None: _remove_task_result_file(task_id) if persist_legacy_result: _write_task_file(task_id, status_payload) return status_payload, has_result def _status_file_paths() -> list[Path]: return [ task_file_path for task_file_path in _task_state_dir().glob("*.json") if not task_file_path.name.endswith(".result.json") ] def load_tasks_from_disk() -> None: with _tasks_lock: tasks_db.clear() task_directory = _task_state_dir() task_directory.mkdir(parents=True, exist_ok=True) current_time = int(time.time()) for task_file_path in _status_file_paths(): loaded_item = _load_task_payload_from_file(task_file_path) if loaded_item is None: try: task_file_path.unlink() except Exception: pass continue task_id, payload = loaded_item payload, _ = _split_task_payload(task_id, payload, persist_legacy_result=True) if _is_task_expired(payload, now_ts=current_time): try: task_file_path.unlink() except Exception: pass try: _remove_task_result_file(task_id) except Exception: pass continue tasks_db[task_id] = payload def cleanup_expired_tasks(force: bool = False) -> None: global _last_cleanup_at current_time = int(time.time()) if not force and current_time - _last_cleanup_at < 60: return with _tasks_lock: expired_task_ids = [ task_id for task_id, task_payload in tasks_db.items() if _is_task_expired(task_payload, now_ts=current_time) ] for task_id in expired_task_ids: tasks_db.pop(task_id, None) try: _remove_task_file(task_id) except Exception: pass try: _remove_task_result_file(task_id) except Exception: pass for task_file_path in _status_file_paths(): loaded_item = _load_task_payload_from_file(task_file_path) if loaded_item is None: try: task_file_path.unlink() except Exception: pass continue task_id, payload = loaded_item payload, _ = _split_task_payload(task_id, payload, persist_legacy_result=True) if _is_task_expired(payload, now_ts=current_time): try: task_file_path.unlink() except Exception: pass try: _remove_task_result_file(task_id) except Exception: pass _last_cleanup_at = current_time def recover_tasks_after_restart() -> None: load_tasks_from_disk() cleanup_expired_tasks(force=True) with _tasks_lock: for task_id, task_payload in list(tasks_db.items()): if str(task_payload.get("status") or "").strip() not in {"queued", "processing"}: continue interrupted_message = "服务已重启,原离线任务已中断,请重新提交。" task_payload.update( { "status": "failed", "stage": "failed", "message": interrupted_message, "error": interrupted_message, "percentage": 100, "updated_at": int(time.time()), } ) tasks_db[task_id] = task_payload _write_task_file(task_id, task_payload) def create_task_record(task_id: str, payload: dict[str, Any]) -> dict[str, Any]: current_time = int(time.time()) task_payload = dict(payload) task_payload["task_id"] = task_id task_payload.setdefault("created_at", current_time) task_payload["updated_at"] = current_time status_payload, has_result = _split_task_payload(task_id, task_payload) with _tasks_lock: tasks_db[task_id] = status_payload _write_task_file(task_id, status_payload) response_payload = dict(status_payload) if has_result: response_payload["result"] = _load_task_result_payload(task_id) return response_payload def get_task_record(task_id: str, *, include_result: bool = False) -> Optional[dict[str, Any]]: with _tasks_lock: task_file_path = _task_file_path(task_id) if task_file_path.exists(): loaded_item = _load_task_payload_from_file(task_file_path) if loaded_item is None: return None loaded_task_id, task_payload = loaded_item task_payload, _ = _split_task_payload( loaded_task_id, task_payload, persist_legacy_result=True, ) if loaded_task_id != task_id or _is_task_expired(task_payload): return None tasks_db[task_id] = task_payload else: task_payload = tasks_db.get(task_id) if task_payload is None: return None response_payload = dict(task_payload) if include_result: result_payload = _load_task_result_payload(task_id) if result_payload is not None: response_payload["result"] = result_payload return response_payload def update_task_record(task_id: str, payload: dict[str, Any]) -> dict[str, Any]: current_time = int(time.time()) with _tasks_lock: task_payload = dict(tasks_db.get(task_id) or {}) task_payload.update(payload) task_payload["task_id"] = task_id task_payload.setdefault("created_at", current_time) task_payload["updated_at"] = current_time status_payload, has_result = _split_task_payload(task_id, task_payload) tasks_db[task_id] = status_payload _write_task_file(task_id, status_payload) response_payload = dict(status_payload) if has_result: response_payload["result"] = _load_task_result_payload(task_id) return response_payload load_tasks_from_disk()