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-frontend": { "command": ["./planet.sh", "restart", "-f"], "recovery_mode": "frontend", }, "restart-ai-provider": { "command": ["./planet.sh", "restart", "-a"], "recovery_mode": "ai-provider", }, "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"