test/app/core/task_store.py

281 lines
9.6 KiB
Python

# -*- 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()