175 lines
4.8 KiB
Python
175 lines
4.8 KiB
Python
from __future__ import annotations
|
|
|
|
import json
|
|
import secrets
|
|
from datetime import UTC, datetime, timedelta
|
|
from pathlib import Path
|
|
from typing import Any
|
|
|
|
from app.core.config import ROOT_DIR
|
|
from app.core.security import redis_client
|
|
|
|
SYSTEM_TASK_TTL_SECONDS = 24 * 60 * 60
|
|
SYSTEM_TASK_LOG_LIMIT = 100
|
|
SYSTEM_TASK_ACTIVE_KEY = "system:restart_task:active"
|
|
SYSTEM_TASK_STALE_SECONDS = 5 * 60
|
|
|
|
ALLOWED_ACTIONS: dict[str, dict[str, Any]] = {
|
|
"restart-backend": {
|
|
"command": ["./planet.sh", "restart", "-b"],
|
|
"recovery_mode": "backend",
|
|
},
|
|
"restart-database": {
|
|
"command": ["./planet.sh", "restart", "-d"],
|
|
"recovery_mode": "database",
|
|
},
|
|
"restart-system": {
|
|
"command": ["./planet.sh", "restart"],
|
|
"recovery_mode": "system",
|
|
},
|
|
}
|
|
|
|
|
|
def utc_now_iso() -> str:
|
|
return datetime.now(UTC).isoformat()
|
|
|
|
|
|
def normalize_user_role(role: Any) -> str:
|
|
return role.value if hasattr(role, "value") else str(role)
|
|
|
|
|
|
def require_super_admin(user_role: Any) -> bool:
|
|
return normalize_user_role(user_role) == "super_admin"
|
|
|
|
|
|
def build_task_id(prefix: str = "restart") -> str:
|
|
timestamp = datetime.now(UTC).strftime("%Y%m%d_%H%M%S")
|
|
return f"{prefix}_{timestamp}_{secrets.token_hex(3)}"
|
|
|
|
|
|
def get_task_key(task_id: str) -> str:
|
|
return f"system:restart_task:{task_id}"
|
|
|
|
|
|
def get_task_logs_key(task_id: str) -> str:
|
|
return f"{get_task_key(task_id)}:logs"
|
|
|
|
|
|
def get_allowed_command(action: str) -> list[str] | None:
|
|
config = ALLOWED_ACTIONS.get(action)
|
|
if config is None:
|
|
return None
|
|
return list(config["command"])
|
|
|
|
|
|
def get_action_recovery_mode(action: str) -> str | None:
|
|
config = ALLOWED_ACTIONS.get(action)
|
|
if config is None:
|
|
return None
|
|
return str(config["recovery_mode"])
|
|
|
|
|
|
def serialize_task(task_id: str) -> dict[str, Any] | None:
|
|
payload = redis_client.hgetall(get_task_key(task_id))
|
|
if not payload:
|
|
return None
|
|
|
|
if payload.get("requested_by"):
|
|
try:
|
|
payload["requested_by"] = json.loads(payload["requested_by"])
|
|
except json.JSONDecodeError:
|
|
payload["requested_by"] = {"username": payload["requested_by"]}
|
|
|
|
return payload
|
|
|
|
|
|
def append_task_log(task_id: str, line: str) -> None:
|
|
logs_key = get_task_logs_key(task_id)
|
|
redis_client.rpush(logs_key, line)
|
|
redis_client.ltrim(logs_key, -SYSTEM_TASK_LOG_LIMIT, -1)
|
|
redis_client.expire(logs_key, SYSTEM_TASK_TTL_SECONDS)
|
|
|
|
|
|
def upsert_task_state(
|
|
task_id: str,
|
|
*,
|
|
action: str | None = None,
|
|
status: str,
|
|
stage: str,
|
|
message: str,
|
|
requested_by: dict[str, Any] | None = None,
|
|
) -> dict[str, Any]:
|
|
existing = serialize_task(task_id) or {}
|
|
now = utc_now_iso()
|
|
|
|
payload: dict[str, Any] = {
|
|
"task_id": task_id,
|
|
"action": action or existing.get("action") or "",
|
|
"status": status,
|
|
"stage": stage,
|
|
"message": message,
|
|
"created_at": existing.get("created_at") or now,
|
|
"updated_at": now,
|
|
}
|
|
|
|
if requested_by is not None:
|
|
payload["requested_by"] = requested_by
|
|
elif existing.get("requested_by") is not None:
|
|
payload["requested_by"] = existing["requested_by"]
|
|
|
|
redis_payload = {
|
|
key: json.dumps(value, ensure_ascii=False) if key == "requested_by" else str(value)
|
|
for key, value in payload.items()
|
|
if value is not None
|
|
}
|
|
task_key = get_task_key(task_id)
|
|
redis_client.hset(task_key, mapping=redis_payload)
|
|
redis_client.expire(task_key, SYSTEM_TASK_TTL_SECONDS)
|
|
return payload
|
|
|
|
|
|
def get_task_logs(task_id: str) -> list[str]:
|
|
return [str(item) for item in redis_client.lrange(get_task_logs_key(task_id), 0, -1)]
|
|
|
|
|
|
def get_active_task_id() -> str | None:
|
|
value = redis_client.get(SYSTEM_TASK_ACTIVE_KEY)
|
|
return str(value) if value else None
|
|
|
|
|
|
def set_active_task_id(task_id: str) -> None:
|
|
redis_client.set(SYSTEM_TASK_ACTIVE_KEY, task_id, ex=SYSTEM_TASK_TTL_SECONDS)
|
|
|
|
|
|
def clear_active_task_id(task_id: str) -> None:
|
|
current = get_active_task_id()
|
|
if current == task_id:
|
|
redis_client.delete(SYSTEM_TASK_ACTIVE_KEY)
|
|
|
|
|
|
def parse_task_timestamp(value: str | None) -> datetime | None:
|
|
if not value:
|
|
return None
|
|
try:
|
|
return datetime.fromisoformat(value)
|
|
except ValueError:
|
|
return None
|
|
|
|
|
|
def is_task_stale(task: dict[str, Any], *, max_age_seconds: int = SYSTEM_TASK_STALE_SECONDS) -> bool:
|
|
if task.get("status") not in {"queued", "running"}:
|
|
return False
|
|
|
|
updated_at = parse_task_timestamp(str(task.get("updated_at") or ""))
|
|
if updated_at is None:
|
|
return False
|
|
|
|
if updated_at.tzinfo is None:
|
|
updated_at = updated_at.replace(tzinfo=UTC)
|
|
|
|
return datetime.now(UTC) - updated_at > timedelta(seconds=max_age_seconds)
|
|
|
|
|
|
def get_runner_script_path() -> Path:
|
|
return ROOT_DIR / "backend" / "scripts" / "system_restart_runner.py"
|