"""Unit tests for AdjustmentFileLoader.""" from pathlib import Path from unittest.mock import Mock, patch import pytest from src.constants import AdjustmentInputSchema from src.enums import FileType from src.errors import EmptyFileError, MissingHeadersError, RowCountExceededError from src.schemas import FileMetadata from src.services.file_loader import AdjustmentFileLoader class TestAdjustmentFileLoader: """Tests for AdjustmentFileLoader class.""" @pytest.fixture def mock_duck_conn(self): """Mock DuckDB connection.""" inner_conn = Mock() cursor = Mock() cursor.__enter__ = Mock(return_value=cursor) cursor.__exit__ = Mock(return_value=False) inner_conn.cursor.return_value = cursor # Mock execute().fetchall() for get_file_columns - return display names execute_result = Mock() execute_result.fetchall.return_value = [ (col.display_name,) for col in AdjustmentInputSchema.COLUMNS ] inner_conn.execute.return_value = execute_result # Create the wrapper mock that has _conn attribute wrapper = Mock() wrapper._conn = inner_conn wrapper.cursor.return_value = cursor return wrapper @pytest.fixture def loader(self, mock_duck_conn): """Create loader instance with mocked dependencies.""" return AdjustmentFileLoader(duck_conn=mock_duck_conn) def test_initialization(self, loader): """Test loader initialization.""" assert loader._duck_conn is not None @patch('src.services.file_loader.get_file_columns') def test_validate_file_columns_success(self, mock_get_columns, loader): """Test successful column validation.""" mock_metadata = FileMetadata( encoding='utf-8', file_path=Path('/tmp/test.csv'), file_type=FileType.CSV, gzipped=False, ) mock_get_columns.return_value = [ col.display_name for col in AdjustmentInputSchema.COLUMNS ] # Execute result = loader._validate_file_columns(mock_metadata) # Verify all columns mapped correctly assert isinstance(result, dict) assert len(result) == len(AdjustmentInputSchema.COLUMNS) @patch('src.services.file_loader.get_file_columns') def test_validate_file_columns_case_insensitive(self, mock_get_columns, loader): """Test case-insensitive column matching.""" mock_metadata = FileMetadata( encoding='utf-8', file_path=Path('/tmp/test.csv'), file_type=FileType.CSV, gzipped=False, ) # Return all required columns in different case with spaces mock_get_columns.return_value = [ col.display_name.upper() if i % 2 == 0 else ' ' + col.display_name + ' ' for i, col in enumerate(AdjustmentInputSchema.COLUMNS) ] # Should not raise - case insensitive matching should work result = loader._validate_file_columns(mock_metadata) assert isinstance(result, dict) # Should map all columns assert len(result) == len(AdjustmentInputSchema.COLUMNS) @patch('src.services.file_loader.get_file_columns') def test_validate_file_columns_missing_required(self, mock_get_columns, loader): """Test validation fails when required columns are missing.""" mock_metadata = FileMetadata( encoding='utf-8', file_path=Path('/tmp/test.csv'), file_type=FileType.CSV, gzipped=False, ) # Return one valid but optional column, missing all required ones # This will pass the "No valid columns" check but fail "Missing required columns" check optional_columns = [ col.display_name for col in AdjustmentInputSchema.COLUMNS if not col.required ] if optional_columns: mock_get_columns.return_value = [optional_columns[0]] else: # If all columns are required, provide just one required column mock_get_columns.return_value = [ AdjustmentInputSchema.COLUMNS[0].display_name ] # Should raise MissingHeadersError with pytest.raises(MissingHeadersError) as exc_info: loader._validate_file_columns(mock_metadata) # Check for either error message depending on what columns exist error_msg = str(exc_info.value) assert ( 'Missing required columns' in error_msg or 'No valid columns' in error_msg ) @patch('src.services.file_loader.get_file_columns') def test_validate_file_columns_no_valid_columns(self, mock_get_columns, loader): """Test validation fails when no valid columns found.""" mock_metadata = FileMetadata( encoding='utf-8', file_path=Path('/tmp/test.csv'), file_type=FileType.CSV, gzipped=False, ) # Return completely wrong columns mock_get_columns.return_value = ['InvalidCol1', 'InvalidCol2'] # Should raise MissingHeadersError with pytest.raises(MissingHeadersError) as exc_info: loader._validate_file_columns(mock_metadata) assert 'No valid columns' in str(exc_info.value) @patch('src.services.file_loader.count_table_rows') def test_validate_row_count_success(self, mock_count_rows, loader): """Test successful row count validation.""" mock_count_rows.return_value = 100 # Should not raise loader._validate_row_count() mock_count_rows.assert_called_once() @patch('src.services.file_loader.count_table_rows') def test_validate_row_count_empty_file(self, mock_count_rows, loader): """Test validation fails for empty file.""" mock_count_rows.return_value = 0 # Should raise EmptyFileError with pytest.raises(EmptyFileError) as exc_info: loader._validate_row_count() assert 'File cannot be empty' in str(exc_info.value) @patch('src.services.file_loader.count_table_rows') def test_validate_row_count_exceeded(self, mock_count_rows, mock_duck_conn): """Test validation fails when row count exceeds maximum.""" loader = AdjustmentFileLoader(duck_conn=mock_duck_conn, max_rows=1000) mock_count_rows.return_value = 1500 # Should raise RowCountExceededError with pytest.raises(RowCountExceededError) as exc_info: loader._validate_row_count() assert 'Row count exceeds maximum' in str(exc_info.value) assert '1500 > 1000' in str(exc_info.value) def test_normalize_header(self, loader): """Test header normalization.""" assert loader._normalize_header(' Account ID ') == 'account id' assert loader._normalize_header('AMOUNT') == 'amount' assert loader._normalize_header(' Currency ') == 'currency'