"""Adjustment File Initialize Processor.""" from __future__ import annotations import uuid from lambdacommon.common_config import logger from pymysql.err import IntegrityError import config from src import features from src.connectors import ows_abacus_event, ows_abacus_state, ows_royalties from src.connectors.repository import Repository from src.enums import ( BatchStatus, EventType, TargetType, UploadStatus, UploadType, ) from src.errors import ( BatchStatementPeriodMismatchError, FileUploadNotFoundError, InvalidBatchStatusError, InvalidUploadStatusError, InvalidUploadTypeError, StatementPeriodAdjustmentFileCreateError, StatementPeriodAdjustmentFileStatesCreateError, StatementPeriodNotFoundError, TransientError, ) from src.schemas import ( AdjustmentBatch, AdjustmentFileInitializeEvent, AdjustmentFileInitializeResponse, AdjustmentFileInitializeResponseData, AdjustmentFileInitializeResponseDetail, AdjustmentFileInitializeResponseMetadata, FileUpload, StatementPeriod, ) class AdjustmentFileInitializeProcessor: """Adjustment file initialize processor. Creates a new batch record in worksheet_flowthrough_batch table with status 'pending'. """ def __init__( self, repository: Repository, ): """Initialize processor. Args: repository: Repository instance """ self._repository = repository def process( self, event: AdjustmentFileInitializeEvent ) -> AdjustmentFileInitializeResponse: """Create batch record in worksheet_flowthrough_batch table. Args: event: Validated event Returns: AdjustmentFileInitializeResponse: Response containing batch_id, s3_bucket, and s3_key Raises: FileUploadNotFoundError: If file_upload record not found (NOT retriable) InvalidUploadStatusError: If file_upload status is not 'complete' (NOT retriable) InvalidUploadTypeError: If file_upload type is not 'adjustments' (NOT retriable) StatementPeriodNotFoundError: If current statement period not found (NOT retriable) BatchStatementPeriodMismatchError: If batch belongs to different period (NOT retriable) InvalidBatchStatusError: If batch status is not pending (NOT retriable) StatementPeriodAdjustmentFileStatesCreateError: If cannot create adjustment file states (NOT retriable) StatementPeriodAdjustmentFileCreateError: If cannot create adjustment file (NOT retriable) TransientError: If database connection fails (RETRIABLE) Exception: For other unexpected errors (RETRIABLE) """ file_upload_id = event.detail.metadata.target_id correlation_id = event.detail.metadata.correlation_id or str(uuid.uuid4()) file_upload = self._get_and_validate_file_upload(file_upload_id) statement_period = self._get_current_statement_period() adjustment_file_id = None if features.is_abacus_flowthrough_automation_enabled( {'identity_id': config.ORCHARD_IDENTITY_ID} ): adjustment_file_id = self._create_statement_period_adjustment_file( file_upload.original_file_name, f's3://{file_upload.s3_bucket}/{file_upload.s3_key}', file_upload_id, statement_period.statement_period_id, file_upload.created_by, ) batch_id = self._get_or_create_batch_id( file_upload, statement_period.statement_period_id ) if ( features.is_abacus_flowthrough_automation_enabled( {'identity_id': config.ORCHARD_IDENTITY_ID} ) and adjustment_file_id ): ows_abacus_event.create_abacus_event( ows_abacus_event.ADJUSTMENT_FILE_UPLOAD_EVENT, adjustment_file_id, ows_abacus_event.ADJUSTMENT_FILE_UPLOAD_TARGET_TYPE, ) return AdjustmentFileInitializeResponse( detail_type=EventType.ADJUSTMENT_BATCH_INITIALIZED, detail=AdjustmentFileInitializeResponseDetail( metadata=AdjustmentFileInitializeResponseMetadata( correlation_id=correlation_id, target_id=batch_id, target_type=TargetType.WORKSHEET_ADJUSTMENT_BATCH, ), data=AdjustmentFileInitializeResponseData( s3_bucket=file_upload.s3_bucket, s3_key=file_upload.s3_key, ), ), ) def _get_and_validate_batch( self, file_upload: FileUpload, statement_period_id: int ) -> AdjustmentBatch | None: file_upload_id = file_upload.file_upload_id logger.info(f'Retrieving batch with source_file_upload_id={file_upload_id}') batch = self._repository.get_batch_by_file_upload(file_upload_id) # Batch not found if not batch: logger.info('Batch not found') return None logger.info(f'Checking statement_period_id = {statement_period_id}') if batch.statement_period_id != statement_period_id: raise BatchStatementPeriodMismatchError( f'Statement period mismatch for batch_id={batch.batch_id}. ' f'Expected: {statement_period_id}, Got: {batch.statement_period_id}. ' f'Batch may be for a previous statement period' ) logger.info(f'Checking batch_status = {BatchStatus.PENDING}') if batch.batch_status != BatchStatus.PENDING: raise InvalidBatchStatusError( f'Invalid batch status for batch_id={batch.batch_id}. ' f'Expected: {BatchStatus.PENDING}, Got: {batch.batch_status}. ' f'Batch may have already been processed' ) return batch def _get_and_validate_file_upload(self, file_upload_id: int) -> FileUpload: """Get file upload and validate status and type.""" logger.info('Getting file upload by ID') file_upload = self._repository.get_file_upload(file_upload_id) if not file_upload: raise FileUploadNotFoundError( f'File upload not found: file_upload_id={file_upload_id}' ) logger.info('File upload retrieved') logger.info(f'Checking upload_status = {UploadStatus.COMPLETE}') if file_upload.upload_status != UploadStatus.COMPLETE: raise InvalidUploadStatusError( f'Invalid upload status for file_upload_id={file_upload_id}. ' f'Expected: {UploadStatus.COMPLETE}, Got: {file_upload.upload_status}' ) logger.info(f'Checking upload_type = {UploadType.ADJUSTMENTS}') if file_upload.upload_type != UploadType.ADJUSTMENTS: raise InvalidUploadTypeError( f'Invalid upload type for file_upload_id={file_upload_id}. ' f'Expected: {UploadType.ADJUSTMENTS}, Got: {file_upload.upload_type}' ) return file_upload def _get_current_statement_period(self) -> StatementPeriod: """Get current statement period or raise error.""" logger.info('Getting current statement period') statement_period = self._repository.get_current_statement_period() if not statement_period: raise StatementPeriodNotFoundError( 'statement_period record not found ' "for statement_period_status='current'" ) logger.info(f'Statement period ID: {statement_period.statement_period_id}') return statement_period def _get_or_create_batch_id( self, file_upload: FileUpload, statement_period_id: int, ) -> int: logger.info('Checking if batch exists for file upload') batch = self._get_and_validate_batch(file_upload, statement_period_id) if batch is not None: logger.info(f'Batch exists. Batch ID: {batch.batch_id}') return batch.batch_id try: logger.info('Creating batch') batch_id = self._repository.create_batch( statement_period_id, file_upload.file_upload_id, file_upload.created_by, ) logger.info(f'Created batch. Batch ID: {batch_id}') return batch_id except IntegrityError: logger.warning('IntegrityError caught. Batch already exists') batch = self._get_and_validate_batch(file_upload, statement_period_id) if not batch: raise TransientError( 'Conflict creating batch. Batch creation ' 'failed but existing batch not found' ) logger.info(f'Batch exists. Batch ID: {batch.batch_id}') return batch.batch_id def _create_statement_period_adjustment_file( self, file_name: str, file_location: str, source_file_upload_id: int, statement_period_id: int, created_by: str, ) -> int: create_file_response = ows_royalties.create_statement_period_adjustment_file( statement_period_id, { 'file_name': file_name, 'valid_file_location': file_location, 'source_file_upload_id': source_file_upload_id, 'created_by': created_by, }, ) if create_file_response.status_code == 201: adjustment_file_id: int = create_file_response.json()[ 'statement_period_adjustment_file_id' ] create_abacus_state_response = ows_abacus_state.create_abacus_state( 'statement_period_adjustment_file', adjustment_file_id, 'upload_file', ) if create_abacus_state_response.status_code != 201: raise StatementPeriodAdjustmentFileStatesCreateError( f'Error creating statement period adjustment file states. ' f'file_name={file_name}, ' f'source_file_upload_id={source_file_upload_id}, statement_period_id={statement_period_id}. ' f'Response status code: {create_abacus_state_response.status_code}. ' f'Response body: {create_abacus_state_response.json()}' ) return adjustment_file_id else: raise StatementPeriodAdjustmentFileCreateError( f'Error creating statement period adjustment file. ' f'file_name={file_name}, ' f'source_file_upload_id={source_file_upload_id}, statement_period_id={statement_period_id}. ' f'Response status code: {create_file_response.status_code}. ' f'Response body: {create_file_response.json()}' )