"""Unit tests for database models.""" import pytest from pydantic import ValidationError from src.schemas.models import AdjustmentBatch, FileUpload, StatementPeriod class TestFileUpload: """Tests for FileUpload model.""" def test_valid_file_upload(self): """Test FileUpload with valid data.""" file_upload = FileUpload( file_upload_id=123, original_file_name='test.csv', s3_bucket='my-bucket', s3_key='path/to/file.csv', upload_status='complete', upload_type='adjustments', created_by='user-123', ) assert file_upload.file_upload_id == 123 assert file_upload.original_file_name == 'test.csv' assert file_upload.s3_bucket == 'my-bucket' assert file_upload.s3_key == 'path/to/file.csv' assert file_upload.upload_status == 'complete' assert file_upload.upload_type == 'adjustments' assert file_upload.created_by == 'user-123' def test_s3_bucket_min_length(self): """Test s3_bucket minimum length validation.""" with pytest.raises(ValidationError) as exc_info: FileUpload( file_upload_id=123, s3_bucket='ab', # Too short (min 3) s3_key='file.csv', upload_status='complete', upload_type='adjustments', created_by='user', ) errors = exc_info.value.errors() assert any('s3_bucket' in str(error) for error in errors) def test_s3_bucket_max_length(self): """Test s3_bucket maximum length validation.""" with pytest.raises(ValidationError) as exc_info: FileUpload( file_upload_id=123, s3_bucket='a' * 64, # Too long (max 63) s3_key='file.csv', upload_status='complete', upload_type='adjustments', created_by='user', ) errors = exc_info.value.errors() assert any('s3_bucket' in str(error) for error in errors) def test_s3_bucket_with_slash_invalid(self): """Test s3_bucket with slash is invalid.""" with pytest.raises(ValidationError) as exc_info: FileUpload( file_upload_id=123, s3_bucket='my-bucket/with-slash', # Invalid - contains slash s3_key='file.csv', upload_status='complete', upload_type='adjustments', created_by='user', ) assert 'Invalid S3 bucket name' in str(exc_info.value) def test_s3_bucket_empty_string_invalid(self): """Test s3_bucket empty string is invalid.""" with pytest.raises(ValidationError) as exc_info: FileUpload( file_upload_id=123, s3_bucket='', # Empty s3_key='file.csv', upload_status='complete', upload_type='adjustments', created_by='user', ) errors = exc_info.value.errors() assert any('s3_bucket' in str(error) for error in errors) def test_s3_key_min_length(self): """Test s3_key minimum length validation.""" with pytest.raises(ValidationError) as exc_info: FileUpload( file_upload_id=123, s3_bucket='my-bucket', s3_key='', # Too short (min 1) upload_status='complete', upload_type='adjustments', created_by='user', ) errors = exc_info.value.errors() assert any('s3_key' in str(error) for error in errors) def test_s3_key_max_length(self): """Test s3_key maximum length validation.""" with pytest.raises(ValidationError) as exc_info: FileUpload( file_upload_id=123, s3_bucket='my-bucket', s3_key='a' * 1025, # Too long (max 1024) upload_status='complete', upload_type='adjustments', created_by='user', ) errors = exc_info.value.errors() assert any('s3_key' in str(error) for error in errors) def test_missing_required_fields(self): """Test missing required fields raise ValidationError.""" with pytest.raises(ValidationError) as exc_info: FileUpload(file_upload_id=123) errors = exc_info.value.errors() required_fields = [ 's3_bucket', 's3_key', 'upload_status', 'upload_type', 'created_by', ] for field in required_fields: assert any(field in str(error) for error in errors) class TestAdjustmentBatch: """Tests for AdjustmentBatch model.""" def test_valid_adjustment_batch(self): """Test AdjustmentBatch with valid data.""" batch = AdjustmentBatch( batch_id=789, batch_type='upload', batch_status='pending', source_file_upload_id=123, statement_period_id=456, ) assert batch.batch_id == 789 assert batch.batch_type == 'upload' assert batch.batch_status == 'pending' assert batch.source_file_upload_id == 123 assert batch.statement_period_id == 456 def test_adjustment_batch_with_null_source_file_upload_id(self): """Test AdjustmentBatch with None source_file_upload_id.""" batch = AdjustmentBatch( batch_id=789, batch_type='upload', batch_status='pending', source_file_upload_id=None, statement_period_id=456, ) assert batch.source_file_upload_id is None def test_missing_required_fields(self): """Test missing required fields raise ValidationError.""" with pytest.raises(ValidationError) as exc_info: AdjustmentBatch(batch_id=789) errors = exc_info.value.errors() required_fields = ['batch_type', 'batch_status', 'statement_period_id'] for field in required_fields: assert any(field in str(error) for error in errors) class TestStatementPeriod: """Tests for StatementPeriod model.""" def test_valid_statement_period(self): """Test StatementPeriod with valid data.""" period = StatementPeriod( statement_period_id=456, statement_period_name='2024-01', statement_period_status='current', statement_month=1, statement_year=2024, ) assert period.statement_period_id == 456 assert period.statement_period_name == '2024-01' assert period.statement_period_status == 'current' assert period.statement_month == 1 assert period.statement_year == 2024 def test_statement_period_with_null_month_year(self): """Test StatementPeriod with None month and year.""" period = StatementPeriod( statement_period_id=456, statement_period_name='2024-01', statement_period_status='current', statement_month=None, statement_year=None, ) assert period.statement_month is None assert period.statement_year is None def test_missing_required_fields(self): """Test missing required fields raise ValidationError.""" with pytest.raises(ValidationError) as exc_info: StatementPeriod(statement_period_id=456) errors = exc_info.value.errors() required_fields = ['statement_period_name', 'statement_period_status'] for field in required_fields: assert any(field in str(error) for error in errors)