import json import uuid from datetime import date, datetime from decimal import Decimal from typing import Any from fastapi import HTTPException, status from sqlalchemy import select, update from sqlalchemy.orm import Session from app.core.pagination import bounded_limit from app.core.time import utc_now from app.modules.approvals.constants import ( ApprovalActionValue, ApprovalErrorDetail, ApprovalPayloadKey, ApprovalStatus, approval_action, ) from app.modules.approvals.models import ApprovalRequest from app.modules.approvals.schemas import ApprovalCreate from app.modules.audit.constants import AuditAction, AuditRiskLevel, AuditSource from app.modules.audit.schemas import AuditLogCreate from app.modules.audit.service import AuditService class ApprovalService: """Create, decide, and validate approval tickets for guarded actions.""" def __init__(self, db: Session): self.db = db self.audit = AuditService(db) def create(self, payload: ApprovalCreate, applicant: str) -> ApprovalRequest: ticket = ApprovalRequest( ticket_id=f"APR-{uuid.uuid4().hex[:12].upper()}", domain=payload.domain, record_id=payload.record_id, action=payload.action, applicant=applicant, reason=payload.reason, payload=json.dumps(payload.payload, ensure_ascii=False, default=str), ) self.db.add(ticket) self.db.commit() self.db.refresh(ticket) self.audit.log( AuditLogCreate( actor=applicant, source=AuditSource.APPROVAL, action=AuditAction.APPROVAL_CREATE, target_type=payload.domain, target_id=payload.record_id, risk_level=AuditRiskLevel.MEDIUM, request_payload=payload.model_dump(), response_payload={ ApprovalPayloadKey.TICKET_ID: ticket.ticket_id, ApprovalPayloadKey.STATUS: ticket.status, }, ) ) return ticket def list(self, status_filter: str | None = None, limit: int = 100) -> list[ApprovalRequest]: stmt = ( select(ApprovalRequest) .order_by(ApprovalRequest.id.desc()) .limit(bounded_limit(limit)) ) if status_filter: stmt = stmt.where(ApprovalRequest.status == status_filter) return list(self.db.execute(stmt).scalars()) def get_by_ticket(self, ticket_id: str) -> ApprovalRequest: ticket = self.db.execute( select(ApprovalRequest).where(ApprovalRequest.ticket_id == ticket_id) ).scalar_one_or_none() if ticket is None: raise HTTPException( status_code=status.HTTP_404_NOT_FOUND, detail=ApprovalErrorDetail.NOT_FOUND, ) return ticket def decide( self, ticket_id: str, approver: str, approved: bool, comment: str | None, ) -> ApprovalRequest: ticket = self.get_by_ticket(ticket_id) if ticket.status != ApprovalStatus.PENDING: raise HTTPException( status_code=status.HTTP_409_CONFLICT, detail=ApprovalErrorDetail.ALREADY_DECIDED, ) if approved and approver == ticket.applicant: raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, detail=ApprovalErrorDetail.SELF_APPROVAL, ) ticket.status = ApprovalStatus.APPROVED if approved else ApprovalStatus.REJECTED ticket.approver = approver ticket.decision_comment = comment ticket.decided_at = utc_now() self.db.commit() self.db.refresh(ticket) self.audit.log( AuditLogCreate( actor=approver, source=AuditSource.APPROVAL, action=AuditAction.APPROVAL_APPROVE if approved else AuditAction.APPROVAL_REJECT, target_type=ticket.domain, target_id=ticket.record_id, risk_level=AuditRiskLevel.HIGH, request_payload={ ApprovalPayloadKey.TICKET_ID: ticket_id, ApprovalPayloadKey.COMMENT: comment, }, response_payload={ApprovalPayloadKey.STATUS: ticket.status}, ) ) return ticket def consume_for( self, ticket_id: str, domain: str, record_id: str | int | None, action: str, payload: dict[str, Any], actor: str, ) -> ApprovalRequest: ticket = self.get_by_ticket(ticket_id) if not self._is_ticket_scope_valid(ticket, domain, record_id, action): raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, detail=ApprovalErrorDetail.NOT_APPROVED, ) if not _payload_matches(ticket.payload, payload): raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, detail=ApprovalErrorDetail.PAYLOAD_MISMATCH, ) used_at = utc_now() values: dict[str, Any] = { ApprovalPayloadKey.STATUS: ApprovalStatus.USED, ApprovalPayloadKey.USED_BY: actor, ApprovalPayloadKey.USED_AT: used_at, } if record_id is not None and not ticket.record_id: values[ApprovalPayloadKey.RECORD_ID] = str(record_id) result = self.db.execute( update(ApprovalRequest) .where( ApprovalRequest.ticket_id == ticket_id, ApprovalRequest.status == ApprovalStatus.APPROVED, ) .values(**values) ) if result.rowcount != 1: raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, detail=ApprovalErrorDetail.NOT_APPROVED, ) for key, value in values.items(): setattr(ticket, key, value) return ticket def is_approved_for( self, ticket_id: str, domain: str, record_id: str | int | None, action: str, ) -> bool: ticket = self.get_by_ticket(ticket_id) return self._is_ticket_scope_valid(ticket, domain, record_id, action) @staticmethod def _is_ticket_scope_valid( ticket: ApprovalRequest, domain: str, record_id: str | int | None, action: str, ) -> bool: if ticket.status != ApprovalStatus.APPROVED: return False if ticket.domain != domain: return False if ticket.record_id and ( record_id is None or str(ticket.record_id) != str(record_id) ): return False return ticket.action in { action, ApprovalActionValue.UPDATE, approval_action(ApprovalActionValue.UPDATE, domain), } def _payload_matches(approved_payload: str | None, requested_payload: dict[str, Any]) -> bool: try: parsed_payload = json.loads(approved_payload or "{}") except json.JSONDecodeError: parsed_payload = {} return _canonical_payload(parsed_payload) == _canonical_payload(requested_payload) def _canonical_payload(value: Any) -> str: return json.dumps(_json_safe(value), ensure_ascii=False, sort_keys=True, default=str) def _json_safe(value: Any) -> Any: if isinstance(value, Decimal): return float(value) if isinstance(value, (datetime, date)): return value.isoformat() if isinstance(value, dict): return {str(key): _json_safe(item) for key, item in value.items()} if isinstance(value, list): return [_json_safe(item) for item in value] return value