281 lines
9.6 KiB
Python
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()
|