"""Statement Period Adjustment File model tests.""" from unittest.mock import patch import pytest from sqlalchemy import bindparam from abacus_file_upload.tests.utils.factories import FileUploadFactory from royalties.constants.constants import STATEMENT_PERIOD_ADJUSTMENT_FILE_USER_ACTIONS from royalties.models.statement_period_adjustment_file import ( StatementPeriodAdjustmentFile, ) from royalties.tests.utils.factories import ( StatementPeriodAdjustmentFileFactory, StatementPeriodFactory, ) def test_create(): """Test creating a statement_period_adjustment_file.""" assert len(StatementPeriodAdjustmentFile.query.all()) == 0 statement_period = StatementPeriodFactory.create() StatementPeriodAdjustmentFile.create( statement_period_id=statement_period.statement_period_id, file_name='test_file_name.xlsx', ) StatementPeriodAdjustmentFile.create( statement_period_id=statement_period.statement_period_id, file_name='test_file_name.xlsx', ) result = StatementPeriodAdjustmentFile.query.all() assert len(result) == 2 assert result[0].statement_period_adjustment_file_id == 1 assert result[1].statement_period_adjustment_file_id == 2 assert statement_period.statement_period_adjustment_file == result assert all([record.statement_period == statement_period for record in result]) def test_filter_deleted_records(mock_statement_period_adjustment_files): """Test to get non deleted records.""" result = StatementPeriodAdjustmentFile.filter_deleted_records().all() statement_period_adjustment_file_ids = [ res.statement_period_adjustment_file_id for res in result ] assert len(result) == 8 assert statement_period_adjustment_file_ids == [ 90, 91, 92, 93, 1001, 1002, 1003, 1004, ] @patch( 'royalties.models.statement_period_adjustment_file.is_abacus_auto_generate_adjustments_flowthrough_ff_enabled' ) def test_get_statement_period_adjustment_files( mock_is_ff_enabled, mock_statement_period_adjustment_files ): """Test to get a list of statement period adjustment files.""" mock_is_ff_enabled.return_value = False assert len(StatementPeriodAdjustmentFile.query.all()) == 10 items, total_count = ( StatementPeriodAdjustmentFile.get_statement_period_adjustment_files( limit=10, offset=0, sort_by='status', sort_order='desc' ) ) assert total_count == 3 assert items[0]['statement_period_adjustment_file_id'] == 91 assert items[1]['statement_period_adjustment_file_id'] == 92 assert items[2]['statement_period_adjustment_file_id'] == 93 assert items[0]['status'] == 'not_approved' assert items[1]['status'] == 'approved' assert items[2]['status'] == 'applied' def test_get_statement_period_adjustment_files_by_id( mock_statement_period_adjustment_files, ): """Test to get a list of statement period adjustment files by id.""" # filter by statement_period_adjustment_file_id items, total_count = ( StatementPeriodAdjustmentFile.get_statement_period_adjustment_files( limit=10, offset=0, sort_by='status', sort_order='desc', statement_period_adjustment_file_id=91, ) ) assert total_count == 1 assert items[0]['statement_period_adjustment_file_id'] == 91 assert items[0]['status'] == 'not_approved' def test_get_statement_period_adjustment_files_by_id_and_status( mock_statement_period_adjustment_files, ): """Test to get a list of statement period adjustment files by id and status.""" # filter by statement_period_adjustment_file_id # but status of 'upload_file' state is in 'init' items, total_count = ( StatementPeriodAdjustmentFile.get_statement_period_adjustment_files( limit=10, offset=0, sort_by='status', sort_order='desc', statement_period_adjustment_file_id=90, ) ) assert total_count == 0 def test_get_statement_period_adjustment_files_by_status( mock_statement_period_adjustment_files, ): """Test to get a list of statement period adjustment files by status.""" items, total_count = ( StatementPeriodAdjustmentFile.get_statement_period_adjustment_files( limit=10, offset=0, sort_by='status', sort_order='asc', status='applied, approved', ) ) assert total_count == 2 assert items[0]['statement_period_adjustment_file_id'] == 93 assert items[0]['status'] == 'applied' assert items[1]['statement_period_adjustment_file_id'] == 92 assert items[1]['status'] == 'approved' def test_get_statement_period_adjustment_files_by_id_multiple_status( mock_statement_period_adjustment_files, ): """Test to get a list of statement period adjustment files by id and multiple statuses.""" items, total_count = ( StatementPeriodAdjustmentFile.get_statement_period_adjustment_files( limit=10, offset=0, sort_by='status', sort_order='desc', statement_period_adjustment_file_id=91, status='not_approved, applied', ) ) assert total_count == 1 assert items[0]['statement_period_adjustment_file_id'] == 91 assert items[0]['status'] == 'not_approved' def test_get_statement_period_adjustment_files_by_name( mock_statement_period_adjustment_files, ): """Test to get a list of statement period adjustment files by file_name.""" items, total_count = ( StatementPeriodAdjustmentFile.get_statement_period_adjustment_files( limit=10, offset=0, sort_by='status', sort_order='desc', file_name='2024' ) ) assert total_count == 2 assert items[0]['statement_period_adjustment_file_id'] == 92 assert items[0]['file_name'] == 'adjustment-file-1-2024.xlsx' assert items[0]['status'] == 'approved' assert items[1]['statement_period_adjustment_file_id'] == 93 assert items[1]['file_name'] == 'adjustment-file-2-2024.xlsx' assert items[1]['status'] == 'applied' def test_get_statement_period_adjustment_files_by_statement_period( mock_statement_period_adjustment_files, ): """Test to get a list of statement period adjustment files by statement_period_id.""" items, total_count = ( StatementPeriodAdjustmentFile.get_statement_period_adjustment_files( limit=10, offset=0, sort_by='status', sort_order='desc', file_name='adjustment-file-1-2024.xlsx', statement_period_id=999, ) ) assert total_count == 1 assert items[0]['file_name'] == 'adjustment-file-1-2024.xlsx' assert items[0]['statement_period_id'] == 999 def test_get_statement_period_adjustment_files_by_created_by( mock_statement_period_adjustment_files, ): """Test to get a list of statement period adjustment files by created_by.""" items, total_count = ( StatementPeriodAdjustmentFile.get_statement_period_adjustment_files( limit=10, offset=0, sort_by='status', sort_order='desc', created_by='e5ca8bc3-7e52-4793-8775-50d11282504c', ) ) assert total_count == 2 assert items[0]['statement_period_adjustment_file_id'] == 91 assert items[0]['created_by'] == 'e5ca8bc3-7e52-4793-8775-50d11282504c' assert items[0]['status'] == 'not_approved' assert items[1]['statement_period_adjustment_file_id'] == 92 assert items[1]['created_by'] == 'e5ca8bc3-7e52-4793-8775-50d11282504c' assert items[1]['status'] == 'approved' def test_get_statement_period_adjustment_file_users( mock_statement_period_adjustment_files, ): """Test to get a list of users who uploaded files.""" items, total_count = ( StatementPeriodAdjustmentFile.get_statement_period_adjustment_file_users( STATEMENT_PERIOD_ADJUSTMENT_FILE_USER_ACTIONS.UPLOADED_FILE ) ) identityIds = [ 'd5ca8ac3-7e51-4793-8775-50d11282504c', 'e5ca8bc3-7e52-4793-8775-50d11282504c', ] assert total_count == 2 assert all([item[0] in identityIds for item in items]) def test_get_by_source_file_key_success( mock_statement_period_adjustment_files, ): """Test to get a statement period adjustment file by source file key.""" file_upload = FileUploadFactory.create(file_key='file-key') file = StatementPeriodAdjustmentFileFactory.create( source_file_upload_id=file_upload.file_upload_id ) item = StatementPeriodAdjustmentFile.get_by_source_file_key('file-key') assert ( item.statement_period_adjustment_file_id == file.statement_period_adjustment_file_id ) def test_get_by_source_file_key_not_existing( mock_statement_period_adjustment_files, ): """Test to get a statement period adjustment file by source file key when not existing.""" file_upload = FileUploadFactory.create(file_key='file-key') file = StatementPeriodAdjustmentFileFactory.create( source_file_upload_id=file_upload.file_upload_id ) item = StatementPeriodAdjustmentFile.get_by_source_file_key('not-existing-file-key') assert item is None @patch( 'royalties.models.statement_period_adjustment_file.is_abacus_auto_generate_adjustments_flowthrough_ff_enabled' ) def test_get_statement_period_adjustment_files_ff_enabled( mock_is_ff_enabled, mock_statement_period_adjustment_files ): """Test to get a list of statement period adjustment files when ff is enabled.""" mock_is_ff_enabled.return_value = True assert len(StatementPeriodAdjustmentFile.query.all()) == 10 items, total_count = ( StatementPeriodAdjustmentFile.get_statement_period_adjustment_files( limit=10, offset=0, sort_by='statement_period_adjustment_file_id', sort_order='desc', ) ) assert total_count == 7 assert items[0]['statement_period_adjustment_file_id'] == 1004 assert items[1]['statement_period_adjustment_file_id'] == 1003 assert items[2]['statement_period_adjustment_file_id'] == 1002 assert items[3]['statement_period_adjustment_file_id'] == 1001 assert items[4]['statement_period_adjustment_file_id'] == 93 assert items[5]['statement_period_adjustment_file_id'] == 92 assert items[6]['statement_period_adjustment_file_id'] == 91 assert items[0]['status'] == 'no_records' assert items[1]['status'] == 'generating' assert items[2]['status'] == 'failed_to_generate' assert items[3]['status'] == 'not_approved' assert items[4]['status'] == 'applied' assert items[5]['status'] == 'approved' assert items[6]['status'] == 'not_approved' @pytest.mark.db('mysql') @patch('royalties.models.statement_period_adjustment_file.func') def test_get_in_progress_auto_generated_adjustments( mock_utc_func, mock_statement_period_adjustment_files, ): """Test to get in-progress auto-generated adjustment file.""" mock_utc_func.utc_timestamp.return_value = bindparam( 'mock_now', value='2024-02-01 05:16:00' ) statement_period_id = 998 identity_id = 'effff8ac3-7e51-4793-8775-50d11282504c' items = StatementPeriodAdjustmentFile.get_in_progress_auto_generated_adjustments( statement_period_id, identity_id ) assert items[0]['statement_period_adjustment_file_id'] == 1003 assert items[0]['status'] == 'generating' @pytest.mark.db('mysql') def test_get_in_progress_adjustments_returns_none_when_no_match( mock_statement_period_adjustment_files, ): """Test returns None when no auto-generated adjustment files are in generating/error state.""" identity_id = 'e5ca8bc3-7e52-4793-8775-50d11282504c' statement_period_id = 999 items = StatementPeriodAdjustmentFile.get_in_progress_auto_generated_adjustments( statement_period_id, identity_id ) assert len(items) == 0 @patch( 'royalties.models.statement_period_adjustment_file.is_abacus_auto_generate_adjustments_flowthrough_ff_enabled' ) def test_get_failed_statement_period_adjustment_files( mock_is_ff_enabled, mock_statement_period_adjustment_files ): """Test to get adjustment files that are failed to generate.""" items, total_count = ( StatementPeriodAdjustmentFile.get_statement_period_adjustment_files( limit=10, offset=0, sort_by='status', sort_order='desc', status='failed_to_generate', ) ) assert total_count == 1 assert items[0]['statement_period_adjustment_file_id'] == 1002 assert items[0]['status'] == 'failed_to_generate'