from typing import Any from fastapi import HTTPException, status from sqlalchemy import func, select from sqlalchemy.orm import Session from app.core.config import get_settings from app.core.constants import ActorValue from app.core.pagination import bounded_limit from app.core.time import utc_now from app.modules.approvals.constants import approval_action from app.modules.approvals.service import ApprovalService from app.modules.audit.constants import AuditAction, AuditRiskLevel, AuditSource from app.modules.audit.schemas import AuditLogCreate from app.modules.audit.service import AuditService from app.modules.business.service import serialize_model from app.modules.events.constants import EventAggregateType, EventSource, EventType from app.modules.events.service import EventService from app.modules.writebacks.adapters import get_writeback_adapter from app.modules.writebacks.constants import ( WRITEBACK_CODE_PREFIX, WRITEBACK_DISABLED_MESSAGE, WritebackActionValue, WritebackErrorDetail, WritebackPayloadKey, WritebackStatus, ) from app.modules.writebacks.models import OfficialWritebackRun class WritebackService: """Create and submit approved official-system writeback runs.""" def __init__(self, db: Session): self.db = db self.audit = AuditService(db) def create_run( self, domain: str, record_id: str | None, action: str, payload: dict[str, Any], actor: str = ActorValue.API, idempotency_key: str | None = None, ) -> dict[str, Any]: if idempotency_key: existing = self.db.execute( select(OfficialWritebackRun).where( OfficialWritebackRun.idempotency_key == idempotency_key ) ).scalar_one_or_none() if existing is not None: return serialize_model(existing) record = OfficialWritebackRun( code=f"{WRITEBACK_CODE_PREFIX}-{utc_now():%Y%m%d%H%M%S%f}", domain=domain, record_id=record_id, action=action, actor=actor, request_payload=payload, idempotency_key=idempotency_key, ) self.db.add(record) self.db.commit() self.db.refresh(record) EventService(self.db).emit( event_type=EventType.WRITEBACK_REQUESTED, source=EventSource.WRITEBACK, aggregate_type=EventAggregateType.WRITEBACK_RUN, aggregate_id=record.code, actor=actor, payload=self._event_payload(record), idempotency_key=f"{record.code}:requested", dispatch=True, ) self.audit.log( AuditLogCreate( actor=actor, source=AuditSource.API, action=AuditAction.WRITEBACK_CREATE, target_type=domain, target_id=record_id, risk_level=AuditRiskLevel.HIGH, request_payload=payload, response_payload={WritebackPayloadKey.CODE: record.code}, ) ) return serialize_model(record) def submit_run( self, code: str, approval_ticket_id: str | None, actor: str = ActorValue.API, ) -> dict[str, Any]: if not approval_ticket_id: raise HTTPException( status_code=status.HTTP_409_CONFLICT, detail=WritebackErrorDetail.APPROVAL_TICKET_REQUIRED, ) record = self._get_run(code) settings = get_settings() approval_action_name = approval_action(WritebackActionValue.WRITEBACK, record.domain) if settings.official_writeback_enabled: ApprovalService(self.db).consume_for( approval_ticket_id, record.domain, record.record_id, approval_action_name, record.request_payload or {}, actor, ) elif not ApprovalService(self.db).is_approved_for( approval_ticket_id, record.domain, record.record_id, approval_action_name, ): raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, detail=WritebackErrorDetail.APPROVAL_TICKET_REQUIRED, ) record.approval_ticket_id = approval_ticket_id record.submitted_at = utc_now() record.status = WritebackStatus.PENDING_APPROVAL self.db.commit() self.db.refresh(record) adapter = get_writeback_adapter(settings) try: result = adapter.submit(record) except Exception as exc: record.status = WritebackStatus.FAILED record.error_message = str(exc) self.db.commit() self.db.refresh(record) self._emit_submitted(record, actor) self._audit_submit(record, actor) return serialize_model(record) if result.get(WritebackPayloadKey.STATUS) == WritebackStatus.DISABLED: record.status = WritebackStatus.DISABLED record.error_message = WRITEBACK_DISABLED_MESSAGE else: record.status = WritebackStatus.SENT record.provider_response = result.get(WritebackPayloadKey.PROVIDER_RESPONSE) or result record.sent_at = utc_now() record.error_message = None self.db.commit() self.db.refresh(record) self._emit_submitted(record, actor) self._audit_submit(record, actor) return serialize_model(record) def list_runs( self, status_filter: str | None = None, limit: int = 100, ) -> list[dict[str, Any]]: stmt = ( select(OfficialWritebackRun) .order_by(OfficialWritebackRun.id.desc()) .limit(bounded_limit(limit)) ) if status_filter: stmt = stmt.where(OfficialWritebackRun.status == status_filter) return [serialize_model(item) for item in self.db.execute(stmt).scalars()] def get_run(self, code: str) -> dict[str, Any]: return serialize_model(self._get_run(code)) def count_by_status(self) -> dict[str, int]: rows = self.db.execute( select(OfficialWritebackRun.status, func.count()).group_by( OfficialWritebackRun.status ) ).all() return {str(status_value): int(count) for status_value, count in rows} def _get_run(self, code: str) -> OfficialWritebackRun: record = self.db.execute( select(OfficialWritebackRun).where(OfficialWritebackRun.code == code) ).scalar_one_or_none() if record is None: raise HTTPException( status_code=status.HTTP_404_NOT_FOUND, detail=WritebackErrorDetail.RUN_NOT_FOUND, ) return record def _emit_submitted(self, record: OfficialWritebackRun, actor: str) -> None: EventService(self.db).emit( event_type=EventType.WRITEBACK_SUBMITTED, source=EventSource.WRITEBACK, aggregate_type=EventAggregateType.WRITEBACK_RUN, aggregate_id=record.code, actor=actor, payload=self._event_payload(record), idempotency_key=f"{record.code}:submitted:{record.status}", dispatch=True, ) def _audit_submit(self, record: OfficialWritebackRun, actor: str) -> None: self.audit.log( AuditLogCreate( actor=actor, source=AuditSource.API, action=AuditAction.WRITEBACK_SUBMIT, target_type=record.domain, target_id=record.record_id, risk_level=AuditRiskLevel.HIGH, request_payload={ WritebackPayloadKey.CODE: record.code, WritebackPayloadKey.APPROVAL_TICKET_ID: record.approval_ticket_id, }, response_payload=self._event_payload(record), ) ) @staticmethod def _event_payload(record: OfficialWritebackRun) -> dict[str, Any]: return { WritebackPayloadKey.CODE: record.code, WritebackPayloadKey.STATUS: record.status, WritebackPayloadKey.DOMAIN: record.domain, WritebackPayloadKey.RECORD_ID: record.record_id, WritebackPayloadKey.ACTION: record.action, WritebackPayloadKey.ERROR_MESSAGE: record.error_message, }