from __future__ import annotations import asyncio import hashlib import json import logging import os import time from dataclasses import dataclass, field from enum import Enum, auto from typing import Any logger = logging.getLogger(__name__) class ApprovalStatus(Enum): PENDING = auto() APPROVED = auto() REJECTED = auto() EXPIRED = auto() CANCELLED = auto() @dataclass class ApprovalRequest: request_id: str action: str description: str requester_id: str requester_name: str params: dict[str, Any] = field(default_factory=dict) status: ApprovalStatus = ApprovalStatus.PENDING created_at: float = field(default_factory=time.time) expires_at: float = 300.0 approver_id: str = "" approved_at: float = 0.0 def __post_init__(self): self.expires_at = self.created_at + 300.0 def to_dict(self) -> dict: return { "request_id": self.request_id, "action": self.action, "description": self.description, "requester_id": self.requester_id, "requester_name": self.requester_name, "params": self.params, "status": self.status.name, "created_at": self.created_at, "expires_at": self.expires_at, "approver_id": self.approver_id, "approved_at": self.approved_at, } class ExecApprovalManager: def __init__( self, store_dir: str | None = None, timeout_s: float = 300.0, max_pending_per_user: int = 5, ): self._store_dir = store_dir or os.path.join(os.path.dirname(__file__), "..", "approval_data") self._timeout_s = timeout_s self._max_pending_per_user = max_pending_per_user self._requests: dict[str, ApprovalRequest] = {} self._lock = asyncio.Lock() self._subscribers: dict[str, asyncio.Event] = {} self._load() def _load(self) -> None: path = os.path.join(self._store_dir, "approval_state.json") try: if os.path.exists(path): with open(path, encoding="utf-8") as f: data = json.load(f) for item in data: req = ApprovalRequest( request_id=item["request_id"], action=item["action"], description=item["description"], requester_id=item["requester_id"], requester_name=item["requester_name"], params=item.get("params", {}), status=ApprovalStatus[item["status"]], created_at=item["created_at"], expires_at=item.get("expires_at", item["created_at"] + 300), ) self._requests[req.request_id] = req logger.info("ExecApprovalManager: loaded %d requests", len(self._requests)) except (OSError, json.JSONDecodeError): logger.exception("ExecApprovalManager: failed to load state") async def _save(self) -> None: os.makedirs(self._store_dir, exist_ok=True) path = os.path.join(self._store_dir, "approval_state.json") tmp = path + ".tmp" data = [r.to_dict() for r in self._requests.values()] try: with open(tmp, "w", encoding="utf-8") as f: json.dump(data, f, ensure_ascii=False) os.replace(tmp, path) except OSError: logger.exception("ExecApprovalManager: failed to save") async def request_approval( self, action: str, description: str, requester_id: str, requester_name: str = "", params: dict[str, Any] | None = None, ) -> ApprovalRequest | None: async with self._lock: pending = sum( 1 for r in self._requests.values() if r.requester_id == requester_id and r.status == ApprovalStatus.PENDING ) if pending >= self._max_pending_per_user: logger.warning( "ExecApprovalManager: user %s has %d pending requests", requester_id, pending, ) return None req_id = hashlib.sha256(f"{requester_id}:{action}:{time.time()}".encode()).hexdigest()[:12] req = ApprovalRequest( request_id=req_id, action=action, description=description, requester_id=requester_id, requester_name=requester_name, params=params or {}, ) self._requests[req_id] = req await self._save() logger.info("ExecApprovalManager: created request %s for %s", req_id, action) return req async def approve(self, request_id: str, approver_id: str) -> ApprovalRequest | None: async with self._lock: req = self._requests.get(request_id) if req is None: return None if req.status != ApprovalStatus.PENDING: return req req.status = ApprovalStatus.APPROVED req.approver_id = approver_id req.approved_at = time.time() await self._save() event = self._subscribers.pop(request_id, None) if event: event.set() logger.info("ExecApprovalManager: approved %s by %s", request_id, approver_id) return req async def reject(self, request_id: str, approver_id: str, reason: str = "") -> ApprovalRequest | None: async with self._lock: req = self._requests.get(request_id) if req is None: return None if req.status != ApprovalStatus.PENDING: return req req.status = ApprovalStatus.REJECTED req.approver_id = approver_id req.approved_at = time.time() await self._save() event = self._subscribers.pop(request_id, None) if event: event.set() logger.info( "ExecApprovalManager: rejected %s by %s reason=%s", request_id, approver_id, reason, ) return req async def wait_for_approval(self, request_id: str, timeout_s: float | None = None) -> ApprovalRequest: event = asyncio.Event() self._subscribers[request_id] = event timeout = timeout_s or self._timeout_s try: await asyncio.wait_for(event.wait(), timeout=timeout) except TimeoutError: async with self._lock: req = self._requests.get(request_id) if req and req.status == ApprovalStatus.PENDING: req.status = ApprovalStatus.EXPIRED await self._save() logger.warning("ExecApprovalManager: request %s expired", request_id) finally: self._subscribers.pop(request_id, None) return self._requests.get( request_id, ApprovalRequest( request_id=request_id, action="unknown", description="", requester_id="", status=ApprovalStatus.EXPIRED, ), ) async def cancel(self, request_id: str, requester_id: str) -> bool: async with self._lock: req = self._requests.get(request_id) if req is None or req.requester_id != requester_id: return False if req.status != ApprovalStatus.PENDING: return False req.status = ApprovalStatus.CANCELLED await self._save() event = self._subscribers.pop(request_id, None) if event: event.set() return True async def list_pending(self, requester_id: str = "") -> list[ApprovalRequest]: async with self._lock: self._cleanup_expired() return [ r for r in self._requests.values() if r.status == ApprovalStatus.PENDING and (not requester_id or r.requester_id == requester_id) ] def _cleanup_expired(self) -> None: now = time.time() for req_id, req in list(self._requests.items()): if req.status == ApprovalStatus.PENDING and now >= req.expires_at: req.status = ApprovalStatus.EXPIRED logger.info("ExecApprovalManager: cleaned up expired request %s", req_id)