"""Unit tests for Lambda handler.""" from unittest.mock import MagicMock, Mock, patch import pytest from src import app from src.enums import EventType from src.errors import InvalidFileTypeError, PermanentError, TransientError from src.schemas import ( AdjustmentFilePrepareResponse, AdjustmentFilePrepareResponseData, AdjustmentFilePrepareResponseDetail, AdjustmentFilePrepareResponseMetadata, ) class TestHandler: """Tests for Lambda handler.""" @pytest.fixture def event(self): """Sample event.""" return { 'detail-type': 'adjustment_batch.initialized', 'detail': { 'metadata': { 'target_type': 'worksheet_flowthrough_batch', 'target_id': 123, 'correlation_id': 'corr-123', }, 'data': { 's3_bucket': 'test-bucket', 's3_key': 'test/file.xlsx', }, }, } @patch('src.app.os.remove') @patch('src.app.create_temp_file') @patch('src.app.add_snowflake_secret') @patch('src.app.DuckDBConnectionFactory') @patch('src.app.MySQLConnectionFactory') @patch('src.app.SnowflakeConnectionFactory') @patch('src.app.RoyaltyAccountingClient') @patch('src.app.ResourceManager.get_s3_connection') @patch('src.app.S3Downloader') @patch('src.app.DuckDBValidator') @patch('src.app.AdjustmentFileLoader') @patch('src.app.SnowflakeGateway') @patch('src.app.ReferenceLoader') @patch('src.app.AdjustmentFilePrepareProcessor') def test_handler_success( self, mock_processor_cls, mock_ref_service_cls, mock_gateway_cls, mock_loader_cls, mock_validator_cls, mock_downloader_cls, mock_get_s3, mock_repo_cls, mock_snow_factory, mock_db_factory, mock_duck_factory, mock_snowflake_secret, mock_create_temp, mock_remove, event, ): """Test successful handler execution.""" # Setup mocks mock_context = Mock() mock_context.aws_request_id = 'test-request-id' mock_context.invoked_function_arn = ( 'arn:aws:lambda:us-east-1:123456789012:function:test' ) mock_context.memory_limit_in_mb = '512' # Setup file mocks mock_create_temp.return_value = '/tmp/test.duckdb' # Setup MySQL factory mock mock_mysql_factory_instance = Mock() mock_mysql_conn = Mock() mock_mysql_cm = MagicMock() mock_mysql_cm.__enter__.return_value = mock_mysql_conn mock_mysql_cm.__exit__.return_value = None mock_mysql_factory_instance.connection.return_value = mock_mysql_cm mock_db_factory.return_value = mock_mysql_factory_instance # Setup DuckDB factory mock mock_duck_factory_instance = Mock() mock_duck_conn = Mock() mock_duck_cursor = MagicMock() mock_duck_cursor_cm = MagicMock() mock_duck_cursor_cm.__enter__.return_value = mock_duck_cursor mock_duck_cursor_cm.__exit__.return_value = None mock_duck_conn.cursor.return_value = mock_duck_cursor_cm mock_duck_cm = MagicMock() mock_duck_cm.__enter__.return_value = mock_duck_conn mock_duck_cm.__exit__.return_value = None mock_duck_factory_instance.connection.return_value = mock_duck_cm mock_duck_factory.return_value = mock_duck_factory_instance # Setup Snowflake factory mock mock_snow_factory_instance = Mock() mock_snow_conn = Mock() mock_snow_cm = MagicMock() mock_snow_cm.__enter__.return_value = mock_snow_conn mock_snow_cm.__exit__.return_value = None mock_snow_factory_instance.connection.return_value = mock_snow_cm mock_snow_factory.return_value = mock_snow_factory_instance mock_processor = Mock() mock_processor_cls.return_value = mock_processor response_model = AdjustmentFilePrepareResponse( detail_type=EventType.ADJUSTMENT_BATCH_PREPARED, detail=AdjustmentFilePrepareResponseDetail( metadata=AdjustmentFilePrepareResponseMetadata( target_type='worksheet_flowthrough_batch', target_id=123, correlation_id='corr-123', ), data=AdjustmentFilePrepareResponseData( valid_row_count=100, invalid_row_count=5, total_file_amount_multicurrency=1000.0, total_rounded_amount_multicurrency=1000.0, s3_bucket='test-bucket', s3_key='staging/123/prepared.csv.gz', ), ), ) mock_processor.process.return_value = response_model # Execute result = app.handler(event, mock_context) # Verify assert result['detail_type'] == 'adjustment_batch.prepared' assert result['detail']['metadata']['target_id'] == 123 assert result['detail']['data']['valid_row_count'] == 100 assert result['detail']['data']['invalid_row_count'] == 5 # Verify factories were instantiated assert mock_db_factory.call_count == 1 assert mock_duck_factory.call_count == 1 assert mock_ref_service_cls.call_count == 1 # Verify commit was called on the MySQL connection mock_mysql_conn.commit.assert_called_once() @patch('src.app.os.remove') @patch('src.app.create_temp_file') @patch('src.app.add_snowflake_secret') @patch('src.app.DuckDBConnectionFactory') @patch('src.app.MySQLConnectionFactory') @patch('src.app.SnowflakeConnectionFactory') @patch('src.app.RoyaltyAccountingClient') @patch('src.app.ResourceManager.get_s3_connection') @patch('src.app.S3Downloader') @patch('src.app.DuckDBValidator') @patch('src.app.AdjustmentFileLoader') @patch('src.app.SnowflakeGateway') @patch('src.app.ReferenceLoader') @patch('src.app.AdjustmentFilePrepareProcessor') def test_handler_failure_updates_status( self, mock_processor_cls, mock_ref_service_cls, mock_gateway_cls, mock_loader_cls, mock_validator_cls, mock_downloader_cls, mock_get_s3, mock_repo_cls, mock_snow_factory, mock_db_factory, mock_duck_factory, mock_snowflake_secret, mock_create_temp, mock_remove, event, ): """Test handler updates status to FAILED on exception.""" # Setup mocks mock_context = Mock() mock_context.aws_request_id = 'test-request-id' mock_context.invoked_function_arn = ( 'arn:aws:lambda:us-east-1:123456789012:function:test' ) mock_context.memory_limit_in_mb = '512' # Setup file mocks mock_create_temp.return_value = '/tmp/test.duckdb' # Setup MySQL factory mock mock_mysql_factory_instance = Mock() mock_mysql_conn = Mock() mock_mysql_cm = MagicMock() mock_mysql_cm.__enter__.return_value = mock_mysql_conn mock_mysql_cm.__exit__.return_value = None mock_mysql_factory_instance.connection.return_value = mock_mysql_cm mock_db_factory.return_value = mock_mysql_factory_instance # Setup DuckDB factory mock mock_duck_factory_instance = Mock() mock_duck_conn = Mock() mock_duck_cursor = MagicMock() mock_duck_cursor_cm = MagicMock() mock_duck_cursor_cm.__enter__.return_value = mock_duck_cursor mock_duck_cursor_cm.__exit__.return_value = None mock_duck_conn.cursor.return_value = mock_duck_cursor_cm mock_duck_cm = MagicMock() mock_duck_cm.__enter__.return_value = mock_duck_conn mock_duck_cm.__exit__.return_value = None mock_duck_factory_instance.connection.return_value = mock_duck_cm mock_duck_factory.return_value = mock_duck_factory_instance # Setup Snowflake factory mock mock_snow_factory_instance = Mock() mock_snow_conn = Mock() mock_snow_cm = MagicMock() mock_snow_cm.__enter__.return_value = mock_snow_conn mock_snow_cm.__exit__.return_value = None mock_snow_factory_instance.connection.return_value = mock_snow_cm mock_snow_factory.return_value = mock_snow_factory_instance mock_processor = Mock() mock_processor_cls.return_value = mock_processor mock_processor.process.side_effect = Exception('Processing failed') mock_repo = Mock() mock_repo_cls.return_value = mock_repo mock_repo.conn = mock_mysql_conn # Execute with pytest.raises(Exception): app.handler(event, mock_context) # Verify rollback called mock_mysql_conn.rollback.assert_called_once() @patch('src.app.os.remove') @patch('src.app.create_temp_file') @patch('src.app.add_snowflake_secret') @patch('src.app.DuckDBConnectionFactory') @patch('src.app.SnowflakeConnectionFactory') @patch('src.app.S3Downloader') @patch('src.app.DuckDBValidator') @patch('src.app.AdjustmentFileLoader') @patch('src.app.SnowflakeGateway') @patch('src.app.ReferenceLoader') @patch('src.app.MySQLConnectionFactory') @patch('src.app.RoyaltyAccountingClient') @patch('src.app.ResourceManager.get_s3_connection') @patch('src.app.AdjustmentFilePrepareProcessor') def test_handler_transient_error( self, mock_processor_cls, mock_get_s3, mock_repo_cls, mock_db_factory, mock_ref_service_cls, mock_gateway_cls, mock_loader_cls, mock_validator_cls, mock_downloader_cls, mock_snow_factory, mock_duck_factory, mock_snowflake_secret, mock_create_temp, mock_remove, event, ): """Test handler re-raises TransientError.""" # Setup mocks mock_context = Mock() mock_context.aws_request_id = 'test-request-id' mock_context.invoked_function_arn = ( 'arn:aws:lambda:us-east-1:123456789012:function:test' ) mock_context.memory_limit_in_mb = '512' # Setup file mocks mock_create_temp.return_value = '/tmp/test.duckdb' # Setup MySQL factory mock mock_factory_instance = Mock() mock_conn = Mock() mock_cm = MagicMock() mock_cm.__enter__.return_value = mock_conn mock_cm.__exit__.return_value = None mock_factory_instance.connection.return_value = mock_cm mock_db_factory.return_value = mock_factory_instance mock_processor = Mock() mock_processor_cls.return_value = mock_processor mock_processor.process.side_effect = TransientError('Temporary glitch') # Execute & Verify with pytest.raises(TransientError) as exc: app.handler(event, mock_context) assert str(exc.value) == 'Temporary glitch' @patch('src.app.os.remove') @patch('src.app.create_temp_file') @patch('src.app.add_snowflake_secret') @patch('src.app.DuckDBConnectionFactory') @patch('src.app.SnowflakeConnectionFactory') @patch('src.app.S3Downloader') @patch('src.app.DuckDBValidator') @patch('src.app.AdjustmentFileLoader') @patch('src.app.SnowflakeGateway') @patch('src.app.ReferenceLoader') @patch('src.app.MySQLConnectionFactory') @patch('src.app.RoyaltyAccountingClient') @patch('src.app.ResourceManager.get_s3_connection') @patch('src.app.AdjustmentFilePrepareProcessor') def test_handler_permanent_error( self, mock_processor_cls, mock_get_s3, mock_repo_cls, mock_db_factory, mock_ref_service_cls, mock_gateway_cls, mock_loader_cls, mock_validator_cls, mock_downloader_cls, mock_snow_factory, mock_duck_factory, mock_snowflake_secret, mock_create_temp, mock_remove, event, ): """Test handler updates DB and re-raises PermanentError.""" # Setup mocks mock_context = Mock() mock_context.aws_request_id = 'test-request-id' mock_context.invoked_function_arn = ( 'arn:aws:lambda:us-east-1:123456789012:function:test' ) mock_context.memory_limit_in_mb = '512' # Setup file mocks mock_create_temp.return_value = '/tmp/test.duckdb' # Setup MySQL factory mock mock_factory_instance = Mock() mock_conn = Mock() mock_cm = MagicMock() mock_cm.__enter__.return_value = mock_conn mock_cm.__exit__.return_value = None mock_factory_instance.connection.return_value = mock_cm mock_db_factory.return_value = mock_factory_instance mock_processor = Mock() mock_processor_cls.return_value = mock_processor mock_processor.process.side_effect = InvalidFileTypeError('Bad file type') mock_repository_instance = Mock() # Create an instance mock for the Repository mock_repo_cls.return_value = ( mock_repository_instance # Configure the class mock to return this instance ) # Execute & Verify with pytest.raises(PermanentError) as exc: app.handler(event, mock_context) assert str(exc.value) == 'Bad file type' def test_handler_validation_error(self): """Test handler raises PermanentError for invalid input.""" mock_context = Mock() mock_context.aws_request_id = 'test-request-id' mock_context.invoked_function_arn = ( 'arn:aws:lambda:us-east-1:123456789012:function:test' ) mock_context.memory_limit_in_mb = '512' event = {'invalid': 'input'} with pytest.raises(PermanentError): app.handler(event, mock_context) @patch('src.app.os.remove') @patch('src.app.create_temp_file') @patch('src.app.add_snowflake_secret') @patch('src.app.DuckDBConnectionFactory') @patch('src.app.SnowflakeConnectionFactory') @patch('src.app.S3Downloader') @patch('src.app.DuckDBValidator') @patch('src.app.AdjustmentFileLoader') @patch('src.app.SnowflakeGateway') @patch('src.app.ReferenceLoader') @patch('src.app.MySQLConnectionFactory') @patch('src.app.RoyaltyAccountingClient') @patch('src.app.ResourceManager.get_s3_connection') @patch('src.app.AdjustmentFilePrepareProcessor') def test_handler_rollback_failure( self, mock_processor_cls, mock_get_s3, mock_repo_cls, mock_db_factory, mock_ref_service_cls, mock_gateway_cls, mock_loader_cls, mock_validator_cls, mock_downloader_cls, mock_snow_factory, mock_duck_factory, mock_snowflake_secret, mock_create_temp, mock_remove, event, ): """Test handler handles rollback failure gracefully.""" # Setup mocks mock_context = Mock() mock_context.aws_request_id = 'test-request-id' mock_context.invoked_function_arn = ( 'arn:aws:lambda:us-east-1:123456789012:function:test' ) mock_context.memory_limit_in_mb = '512' # Setup file mocks mock_create_temp.return_value = '/tmp/test.duckdb' # Setup MySQL factory mock mock_mysql_factory_instance = Mock() mock_mysql_conn = Mock() mock_mysql_cm = MagicMock() mock_mysql_cm.__enter__.return_value = mock_mysql_conn mock_mysql_cm.__exit__.return_value = None mock_mysql_factory_instance.connection.return_value = mock_mysql_cm mock_db_factory.return_value = mock_mysql_factory_instance # Setup DuckDB factory mock mock_duck_factory_instance = Mock() mock_duck_conn = Mock() mock_duck_cursor = MagicMock() mock_duck_cursor_cm = MagicMock() mock_duck_cursor_cm.__enter__.return_value = mock_duck_cursor mock_duck_cursor_cm.__exit__.return_value = None mock_duck_conn.cursor.return_value = mock_duck_cursor_cm mock_duck_cm = MagicMock() mock_duck_cm.__enter__.return_value = mock_duck_conn mock_duck_cm.__exit__.return_value = None mock_duck_factory_instance.connection.return_value = mock_duck_cm mock_duck_factory.return_value = mock_duck_factory_instance # Setup Snowflake factory mock mock_snow_factory_instance = Mock() mock_snow_conn = Mock() mock_snow_cm = MagicMock() mock_snow_cm.__enter__.return_value = mock_snow_conn mock_snow_cm.__exit__.return_value = None mock_snow_factory_instance.connection.return_value = mock_snow_cm mock_snow_factory.return_value = mock_snow_factory_instance mock_processor = Mock() mock_processor_cls.return_value = mock_processor processing_error = Exception('Processing failed') mock_processor.process.side_effect = processing_error # Make rollback fail rollback_error = Exception('Rollback failed') mock_mysql_conn.rollback.side_effect = rollback_error # Execute - should raise original error with pytest.raises(Exception) as exc_info: app.handler(event, mock_context) # Verify the exception is raised assert exc_info.value is processing_error @patch('src.app.add_snowflake_secret') @patch('src.app.os.remove') @patch('src.app.create_temp_file') @patch('src.app.DuckDBConnectionFactory') @patch('src.app.SnowflakeConnectionFactory') @patch('src.app.S3Downloader') @patch('src.app.DuckDBValidator') @patch('src.app.AdjustmentFileLoader') @patch('src.app.SnowflakeGateway') @patch('src.app.ReferenceLoader') @patch('src.app.MySQLConnectionFactory') @patch('src.app.RoyaltyAccountingClient') @patch('src.app.ResourceManager.get_s3_connection') @patch('src.app.AdjustmentFilePrepareProcessor') def test_handler_cleanup_oserror( self, mock_processor_cls, mock_get_s3, mock_repo_cls, mock_db_factory, mock_ref_service_cls, mock_gateway_cls, mock_loader_cls, mock_validator_cls, mock_downloader_cls, mock_snow_factory, mock_duck_factory, mock_create_temp, mock_remove, mock_snowflake_secret, event, ): """Test handler handles OS error during cleanup gracefully.""" # Setup mocks mock_context = Mock() mock_context.aws_request_id = 'test-request-id' mock_context.invoked_function_arn = ( 'arn:aws:lambda:us-east-1:123456789012:function:test' ) mock_context.memory_limit_in_mb = '512' # Setup MySQL factory mock mock_mysql_factory_instance = Mock() mock_mysql_conn = Mock() mock_mysql_cm = MagicMock() mock_mysql_cm.__enter__.return_value = mock_mysql_conn mock_mysql_cm.__exit__.return_value = None mock_mysql_factory_instance.connection.return_value = mock_mysql_cm mock_db_factory.return_value = mock_mysql_factory_instance # Setup DuckDB factory mock mock_duck_factory_instance = Mock() mock_duck_conn = Mock() mock_duck_cursor = MagicMock() mock_duck_cursor_cm = MagicMock() mock_duck_cursor_cm.__enter__.return_value = mock_duck_cursor mock_duck_cursor_cm.__exit__.return_value = None mock_duck_conn.cursor.return_value = mock_duck_cursor_cm mock_duck_cm = MagicMock() mock_duck_cm.__enter__.return_value = mock_duck_conn mock_duck_cm.__exit__.return_value = None mock_duck_factory_instance.connection.return_value = mock_duck_cm mock_duck_factory.return_value = mock_duck_factory_instance duck_path = '/tmp/test.duckdb' mock_create_temp.return_value = duck_path mock_processor = Mock() mock_processor_cls.return_value = mock_processor response_model = Mock() response_model.model_dump.return_value = {'result': 'success'} mock_processor.process.return_value = response_model # Make os.remove raise OSError mock_remove.side_effect = OSError('Permission denied') # Execute - should succeed despite cleanup error result = app.handler(event, mock_context) # Verify execution succeeded assert result == {'result': 'success'} # Verify cleanup was attempted mock_remove.assert_called_once_with(duck_path)