"""Unit tests for Repository.""" from unittest.mock import MagicMock, Mock import pymysql import pytest from src.connectors.repository import Repository from src.errors import TransientError from src.schemas import FileUpload class TestRepositoryInit: """Tests for Repository initialization.""" def test_init_stores_connection(self): """Test repository stores database connection.""" mock_conn = Mock() repo = Repository(mock_conn) assert repo.conn is mock_conn class TestRepositoryGetFileUpload: """Tests for Repository.get_file_upload method.""" def test_get_file_upload_success(self): """Test successful file_upload retrieval.""" # Setup mock cursor and connection mock_cursor = MagicMock() mock_cursor.fetchone.return_value = { 'file_upload_id': 123, 'original_file_name': 'tets.csv', 's3_bucket': 'test-bucket', 's3_key': 'uploads/test.csv', 'upload_status': 'complete', 'upload_type': 'adjustments', 'created_by': 'user-456', } mock_conn = MagicMock() mock_conn.cursor.return_value.__enter__.return_value = mock_cursor # Execute repo = Repository(mock_conn) result = repo.get_file_upload(123) # Verify assert isinstance(result, FileUpload) assert result.file_upload_id == 123 assert result.original_file_name == 'tets.csv' assert result.s3_bucket == 'test-bucket' assert result.s3_key == 'uploads/test.csv' assert result.upload_status == 'complete' assert result.created_by == 'user-456' # Verify SQL query mock_cursor.execute.assert_called_once() sql = mock_cursor.execute.call_args[0][0] assert 'SELECT' in sql assert 'fu.file_upload_id' in sql assert 'fu.original_file_name' in sql assert 'fu.s3_bucket' in sql assert 'fu.s3_key' in sql assert 'fu.upload_status' in sql assert 'fu.created_by' in sql assert 'fuc.upload_type' in sql assert 'FROM file_upload AS fu' in sql assert 'INNER JOIN file_upload_config AS fuc' in sql assert 'ON fu.file_upload_config_id = fuc.file_upload_config_id' in sql assert 'WHERE fu.file_upload_id = %s' in sql assert 'fu.deleted_at IS NULL' in sql def test_get_file_upload_not_found(self): """Test get_file_upload returns None when record not found.""" mock_cursor = MagicMock() mock_cursor.fetchone.return_value = None mock_conn = MagicMock() mock_conn.cursor.return_value.__enter__.return_value = mock_cursor repo = Repository(mock_conn) result = repo.get_file_upload(999) assert result is None def test_get_file_upload_filters_deleted(self): """Test get_file_upload filters out soft-deleted records.""" mock_cursor = MagicMock() mock_cursor.fetchone.return_value = None mock_conn = MagicMock() mock_conn.cursor.return_value.__enter__.return_value = mock_cursor repo = Repository(mock_conn) repo.get_file_upload(123) # Verify deleted_at filter in query sql = mock_cursor.execute.call_args[0][0] assert 'deleted_at IS NULL' in sql def test_get_file_upload_handles_transient_errors(self): """Test get_file_upload converts pymysql connection errors to TransientError.""" mock_conn = Mock() mock_conn.cursor.side_effect = pymysql.err.OperationalError( 2003, "Can't connect to MySQL server" ) repo = Repository(mock_conn) with pytest.raises(TransientError) as exc_info: repo.get_file_upload(123) assert 'Database connection failed' in str(exc_info.value) def test_get_file_upload_handles_mysql_gone_away(self): """Test get_file_upload handles 'MySQL server has gone away' error.""" mock_conn = Mock() mock_conn.cursor.side_effect = pymysql.err.OperationalError( 2006, 'MySQL server has gone away' ) repo = Repository(mock_conn) with pytest.raises(TransientError) as exc_info: repo.get_file_upload(123) assert 'Database connection failed' in str(exc_info.value) def test_get_file_upload_handles_connection_lost(self): """Test get_file_upload handles lost connection error.""" mock_conn = Mock() mock_conn.cursor.side_effect = pymysql.err.OperationalError( 2013, 'Lost connection to MySQL server during query' ) repo = Repository(mock_conn) with pytest.raises(TransientError) as exc_info: repo.get_file_upload(123) assert 'Database connection failed' in str(exc_info.value) def test_get_file_upload_non_transient_error_propagates(self): """Test non-transient pymysql errors propagate unchanged.""" mock_conn = Mock() mock_conn.cursor.side_effect = pymysql.err.OperationalError( 1045, 'Access denied for user' ) repo = Repository(mock_conn) with pytest.raises(pymysql.err.OperationalError) as exc_info: repo.get_file_upload(123) assert exc_info.value.args[0] == 1045 class TestRepositoryGetBatchByFileUpload: """Tests for Repository.get_batch_by_file_upload method.""" def test_get_batch_success(self): """Test successful batch retrieval.""" mock_cursor = MagicMock() mock_cursor.fetchone.return_value = { 'batch_id': 789, 'batch_type': 'upload', 'batch_status': 'pending', 'source_file_upload_id': 123, 'statement_period_id': 456, } mock_conn = MagicMock() mock_conn.cursor.return_value.__enter__.return_value = mock_cursor repo = Repository(mock_conn) result = repo.get_batch_by_file_upload(123) assert result.batch_id == 789 assert result.batch_type == 'upload' assert result.batch_status == 'pending' # Verify SQL query mock_cursor.execute.assert_called_once() sql = mock_cursor.execute.call_args[0][0] assert 'SELECT' in sql assert 'worksheet_flowthrough_batch_id AS batch_id' in sql assert 'FROM worksheet_flowthrough_batch' in sql assert 'WHERE source_file_upload_id = %s' in sql assert 'deleted_at IS NULL' in sql def test_get_batch_not_found(self): """Test get_batch returns None when batch not found.""" mock_cursor = MagicMock() mock_cursor.fetchone.return_value = None mock_conn = MagicMock() mock_conn.cursor.return_value.__enter__.return_value = mock_cursor repo = Repository(mock_conn) result = repo.get_batch_by_file_upload(999) assert result is None def test_get_batch_filters_deleted(self): """Test get_batch filters out soft-deleted batches.""" mock_cursor = MagicMock() mock_cursor.fetchone.return_value = None mock_conn = MagicMock() mock_conn.cursor.return_value.__enter__.return_value = mock_cursor repo = Repository(mock_conn) repo.get_batch_by_file_upload(123) # Verify deleted_at filter in query sql = mock_cursor.execute.call_args[0][0] assert 'deleted_at IS NULL' in sql def test_get_batch_handles_transient_errors(self): """Test get_batch converts pymysql connection errors to TransientError.""" mock_conn = Mock() mock_conn.cursor.side_effect = pymysql.err.OperationalError( 2003, "Can't connect to MySQL server" ) repo = Repository(mock_conn) with pytest.raises(TransientError) as exc_info: repo.get_batch_by_file_upload(123) assert 'Database connection failed' in str(exc_info.value) class TestRepositoryGetCurrentStatementPeriod: """Tests for Repository.get_current_statement_period method.""" def test_get_current_statement_period_success(self): """Test successful statement period retrieval.""" mock_cursor = MagicMock() mock_cursor.fetchone.return_value = { 'statement_period_id': 456, 'statement_period_name': '2024-01', 'statement_period_status': 'current', 'statement_month': 1, 'statement_year': 2024, } mock_conn = MagicMock() mock_conn.cursor.return_value.__enter__.return_value = mock_cursor repo = Repository(mock_conn) result = repo.get_current_statement_period() assert result.statement_period_id == 456 assert result.statement_period_name == '2024-01' assert result.statement_period_status == 'current' assert result.statement_month == 1 assert result.statement_year == 2024 # Verify SQL query mock_cursor.execute.assert_called_once() sql = mock_cursor.execute.call_args[0][0] assert 'SELECT' in sql assert 'statement_period_id' in sql assert 'statement_period_name' in sql assert 'statement_period_status' in sql assert 'statement_month' in sql assert 'statement_year' in sql assert 'FROM statement_period' in sql assert 'WHERE statement_period_status = %s' in sql def test_get_current_statement_period_not_found(self): """Test get_current_statement_period returns None when not found.""" mock_cursor = MagicMock() mock_cursor.fetchone.return_value = None mock_conn = MagicMock() mock_conn.cursor.return_value.__enter__.return_value = mock_cursor repo = Repository(mock_conn) result = repo.get_current_statement_period() assert result is None def test_get_current_statement_period_handles_transient_errors(self): """Test get_current_statement_period converts connection errors to TransientError.""" mock_conn = Mock() mock_conn.cursor.side_effect = pymysql.err.OperationalError( 2006, 'MySQL server has gone away' ) repo = Repository(mock_conn) with pytest.raises(TransientError) as exc_info: repo.get_current_statement_period() assert 'Database connection failed' in str(exc_info.value) class TestRepositoryCreateBatch: """Tests for Repository.create_batch method.""" def test_create_batch_success(self): """Test successful batch creation.""" mock_cursor = MagicMock() mock_cursor.lastrowid = 789 mock_conn = MagicMock() mock_conn.cursor.return_value.__enter__.return_value = mock_cursor repo = Repository(mock_conn) result = repo.create_batch(456, 123, 'user-123') assert result == 789 # Verify SQL query mock_cursor.execute.assert_called_once() sql = mock_cursor.execute.call_args[0][0] params = mock_cursor.execute.call_args[0][1] assert 'INSERT INTO worksheet_flowthrough_batch' in sql assert 'statement_period_id' in sql assert 'source_file_upload_id' in sql assert 'batch_type' in sql assert 'batch_status' in sql assert 'created_by' in sql assert 'last_modified_by' in sql # Check that values are parameterized assert 'VALUES (%s, %s, %s, %s, %s, %s)' in sql # Verify params include constants and user_id assert params == (456, 123, 'upload', 'pending', 'user-123', 'user-123') def test_create_batch_sets_status_pending(self): """Test create_batch sets status to 'pending'.""" mock_cursor = MagicMock() mock_cursor.lastrowid = 456 mock_conn = MagicMock() mock_conn.cursor.return_value.__enter__.return_value = mock_cursor repo = Repository(mock_conn) repo.create_batch(456, 123, 'user') params = mock_cursor.execute.call_args[0][1] assert params[3] == 'pending' def test_create_batch_uses_user_for_audit_fields(self): """Test create_batch uses user_id for created_by and last_modified_by.""" mock_cursor = MagicMock() mock_cursor.lastrowid = 123 mock_conn = MagicMock() mock_conn.cursor.return_value.__enter__.return_value = mock_cursor repo = Repository(mock_conn) repo.create_batch(456, 123, 'audit-user-789') params = mock_cursor.execute.call_args[0][1] # created_by and last_modified_by should both be the user assert params[4] == 'audit-user-789' # created_by assert params[5] == 'audit-user-789' # last_modified_by def test_create_batch_does_not_commit(self): """Test create_batch does not commit transaction (caller's responsibility).""" mock_cursor = MagicMock() mock_cursor.lastrowid = 999 mock_conn = MagicMock() mock_conn.cursor.return_value.__enter__.return_value = mock_cursor repo = Repository(mock_conn) repo.create_batch(456, 123, 'user') # Verify commit was NOT called mock_conn.commit.assert_not_called() def test_create_batch_handles_transient_errors(self): """Test create_batch converts pymysql connection errors to TransientError.""" mock_conn = Mock() mock_conn.cursor.side_effect = pymysql.err.OperationalError( 2013, 'Lost connection to MySQL server during query' ) repo = Repository(mock_conn) with pytest.raises(TransientError) as exc_info: repo.create_batch(456, 123, 'user') assert 'Database connection failed' in str(exc_info.value) def test_create_batch_non_transient_error_propagates(self): """Test non-transient pymysql errors propagate unchanged.""" mock_cursor = MagicMock() mock_cursor.execute.side_effect = pymysql.err.IntegrityError( 1062, 'Duplicate entry' ) mock_conn = MagicMock() mock_conn.cursor.return_value.__enter__.return_value = mock_cursor repo = Repository(mock_conn) with pytest.raises(pymysql.err.IntegrityError): repo.create_batch(456, 123, 'user') class TestRepositoryEdgeCases: """Tests for Repository edge cases.""" def test_get_file_upload_sql_injection_attempt(self): """Test get_file_upload with SQL injection attempt in file_upload_id.""" mock_cursor = MagicMock() mock_cursor.fetchone.return_value = { 'file_upload_id': 123, 'original_file_name': 'test.csv', 's3_bucket': 'test-bucket', 's3_key': 'test-key', 'upload_status': 'complete', 'upload_type': 'adjustments', 'created_by': 'user-123', } mock_conn = MagicMock() mock_conn.cursor.return_value.__enter__.return_value = mock_cursor repo = Repository(mock_conn) # Attempt SQL injection - should be safely handled by parameterized query # The file_upload_id will be passed as an integer, but test the query uses params result = repo.get_file_upload(123) # Verify parameterized query was used (not string concatenation) mock_cursor.execute.assert_called_once() sql, params = mock_cursor.execute.call_args[0] assert '%s' in sql # Verify parameterized placeholder assert params == (123,) # Verify parameter passed separately assert result is not None def test_create_batch_with_large_ids(self): """Test create_batch with large statement_period_id and file_upload_id.""" mock_cursor = MagicMock() mock_cursor.lastrowid = 999 mock_conn = MagicMock() mock_conn.cursor.return_value.__enter__.return_value = mock_cursor repo = Repository(mock_conn) # Test with large integer IDs large_period_id = 2147483647 # Max 32-bit int large_file_id = 2147483647 result = repo.create_batch(large_period_id, large_file_id, 'user-123') assert result == 999 # Verify the IDs were passed correctly params = mock_cursor.execute.call_args[0][1] assert params[0] == large_period_id assert params[1] == large_file_id class TestFileUploadRecord: """Tests for FileUploadRecord structure.""" def test_file_upload_record_structure(self): """Test FileUploadRecord model validation.""" record = FileUpload( file_upload_id=123, original_file_name='test.csv', s3_bucket='test-bucket', s3_key='test-key', upload_status='complete', upload_type='adjustments', created_by='user-123', ) assert record.file_upload_id == 123 assert record.original_file_name == 'test.csv' assert record.s3_bucket == 'test-bucket' assert record.s3_key == 'test-key' assert record.upload_status == 'complete' assert record.created_by == 'user-123'