"""Unit tests for MetadataService."""

from datetime import datetime
from unittest.mock import AsyncMock, MagicMock
from uuid import uuid4

import pytest

from src.db.models import Document as DocumentModel
from src.db.models import ExtractedRecord as ExtractedRecordModel
from src.knowledge.metadata_service import (
    MetadataService,
    SourceCitation,
    ValidationReport,
)


@pytest.fixture
def mock_db_session():
    """Create a mock async database session."""
    session = AsyncMock()
    return session


@pytest.fixture
def sample_record_data():
    """Sample record data for testing."""
    record_id = uuid4()
    document_id = uuid4()

    return {
        "record_id": record_id,
        "document_id": document_id,
        "filename": "test_document.xlsx",
        "sheet_name": "Sheet1",
        "row_number": 42,
        "col_range": "A:F",
        "headers": ["Item Code", "Description", "Status"],
        "created_at": datetime(2025, 11, 29, 12, 0, 0),
    }


class TestSourceCitation:
    """Test SourceCitation dataclass."""

    def test_to_human_readable_with_all_fields(self, sample_record_data):
        """Test human-readable formatting with all fields present."""
        citation = SourceCitation(
            record_id=sample_record_data["record_id"],
            document_id=sample_record_data["document_id"],
            filename=sample_record_data["filename"],
            sheet_name=sample_record_data["sheet_name"],
            row_number=sample_record_data["row_number"],
            col_range=sample_record_data["col_range"],
            headers=sample_record_data["headers"],
            created_at=sample_record_data["created_at"],
        )

        result = citation.to_human_readable()

        assert "test_document.xlsx" in result
        assert "Sheet1" in result
        assert "row 42" in result
        assert "columns A:F" in result
        assert "Item Code, Description, Status" in result

    def test_to_human_readable_without_col_range(self, sample_record_data):
        """Test human-readable formatting without col_range."""
        citation = SourceCitation(
            record_id=sample_record_data["record_id"],
            document_id=sample_record_data["document_id"],
            filename=sample_record_data["filename"],
            sheet_name=sample_record_data["sheet_name"],
            row_number=sample_record_data["row_number"],
            col_range=None,
            headers=sample_record_data["headers"],
            created_at=sample_record_data["created_at"],
        )

        result = citation.to_human_readable()

        assert "test_document.xlsx" in result
        assert "columns" not in result
        assert "Headers:" in result

    def test_to_human_readable_minimal(self, sample_record_data):
        """Test human-readable formatting with only required fields."""
        citation = SourceCitation(
            record_id=sample_record_data["record_id"],
            document_id=sample_record_data["document_id"],
            filename=sample_record_data["filename"],
            sheet_name=sample_record_data["sheet_name"],
            row_number=sample_record_data["row_number"],
            col_range=None,
            headers=None,
            created_at=sample_record_data["created_at"],
        )

        result = citation.to_human_readable()

        assert result == "Found in test_document.xlsx, sheet 'Sheet1', row 42"

    def test_to_json(self, sample_record_data):
        """Test JSON formatting."""
        citation = SourceCitation(
            record_id=sample_record_data["record_id"],
            document_id=sample_record_data["document_id"],
            filename=sample_record_data["filename"],
            sheet_name=sample_record_data["sheet_name"],
            row_number=sample_record_data["row_number"],
            col_range=sample_record_data["col_range"],
            headers=sample_record_data["headers"],
            created_at=sample_record_data["created_at"],
        )

        result = citation.to_json()

        assert result["record_id"] == str(sample_record_data["record_id"])
        assert result["document_id"] == str(sample_record_data["document_id"])
        assert result["filename"] == "test_document.xlsx"
        assert result["sheet_name"] == "Sheet1"
        assert result["row_number"] == 42
        assert result["col_range"] == "A:F"
        assert result["headers"] == ["Item Code", "Description", "Status"]
        assert result["created_at"] == "2025-11-29T12:00:00"


class TestValidationReport:
    """Test ValidationReport dataclass."""

    def test_is_valid_when_no_issues(self):
        """Test is_valid property returns True when no issues."""
        report = ValidationReport(
            document_id=uuid4(),
            total_records=10,
            complete_records=10,
            records_with_col_range=8,
            records_with_headers=10,
            missing_required_fields=[],
        )

        assert report.is_valid is True

    def test_is_valid_when_has_issues(self):
        """Test is_valid property returns False when issues exist."""
        report = ValidationReport(
            document_id=uuid4(),
            total_records=10,
            complete_records=8,
            records_with_col_range=8,
            records_with_headers=10,
            missing_required_fields=[
                {"record_id": str(uuid4()), "missing_fields": ["sheet_name"]}
            ],
        )

        assert report.is_valid is False

    def test_completeness_percentage(self):
        """Test completeness percentage calculation."""
        report = ValidationReport(
            document_id=uuid4(),
            total_records=10,
            complete_records=8,
            records_with_col_range=5,
            records_with_headers=9,
            missing_required_fields=[],
        )

        assert report.completeness_percentage == 80.0

    def test_completeness_percentage_zero_records(self):
        """Test completeness percentage when no records."""
        report = ValidationReport(
            document_id=uuid4(),
            total_records=0,
            complete_records=0,
            records_with_col_range=0,
            records_with_headers=0,
            missing_required_fields=[],
        )

        assert report.completeness_percentage == 100.0

    def test_to_dict(self):
        """Test dictionary conversion."""
        document_id = uuid4()
        report = ValidationReport(
            document_id=document_id,
            total_records=10,
            complete_records=9,
            records_with_col_range=8,
            records_with_headers=10,
            missing_required_fields=[
                {"record_id": str(uuid4()), "missing_fields": ["row_number"]}
            ],
        )

        result = report.to_dict()

        assert result["document_id"] == str(document_id)
        assert result["total_records"] == 10
        assert result["complete_records"] == 9
        assert result["completeness_percentage"] == 90.0
        assert result["records_with_col_range"] == 8
        assert result["records_with_headers"] == 10
        assert result["is_valid"] is False
        assert result["issues_count"] == 1
        assert len(result["issues"]) == 1


class TestMetadataService:
    """Test MetadataService class."""

    @pytest.mark.asyncio
    async def test_get_source_citation_success(
        self, mock_db_session, sample_record_data
    ):
        """Test successful source citation retrieval."""
        # Create mock record and document
        mock_record = MagicMock(spec=ExtractedRecordModel)
        mock_record.id = sample_record_data["record_id"]
        mock_record.document_id = sample_record_data["document_id"]
        mock_record.sheet_name = sample_record_data["sheet_name"]
        mock_record.row_number = sample_record_data["row_number"]
        mock_record.col_range = sample_record_data["col_range"]
        mock_record.headers = sample_record_data["headers"]
        mock_record.created_at = sample_record_data["created_at"]

        # Mock database response
        mock_result = MagicMock()
        mock_result.first.return_value = (
            mock_record,
            sample_record_data["filename"],
        )
        mock_db_session.execute = AsyncMock(return_value=mock_result)

        # Execute
        service = MetadataService(mock_db_session)
        citation = await service.get_source_citation(sample_record_data["record_id"])

        # Assert
        assert citation is not None
        assert citation.record_id == sample_record_data["record_id"]
        assert citation.filename == sample_record_data["filename"]
        assert citation.sheet_name == sample_record_data["sheet_name"]
        assert citation.row_number == sample_record_data["row_number"]

    @pytest.mark.asyncio
    async def test_get_source_citation_not_found(self, mock_db_session):
        """Test source citation when record not found."""
        # Mock database response (no result)
        mock_result = MagicMock()
        mock_result.first.return_value = None
        mock_db_session.execute = AsyncMock(return_value=mock_result)

        # Execute
        service = MetadataService(mock_db_session)
        citation = await service.get_source_citation(uuid4())

        # Assert
        assert citation is None

    @pytest.mark.asyncio
    async def test_get_records_by_location_document_only(self, mock_db_session):
        """Test getting records by document_id only."""
        # Create mock records
        mock_records = [
            MagicMock(spec=ExtractedRecordModel),
            MagicMock(spec=ExtractedRecordModel),
        ]

        # Mock database response
        mock_scalars = MagicMock()
        mock_scalars.all.return_value = mock_records
        mock_result = MagicMock()
        mock_result.scalars.return_value = mock_scalars
        mock_db_session.execute = AsyncMock(return_value=mock_result)

        # Execute
        service = MetadataService(mock_db_session)
        document_id = uuid4()
        records = await service.get_records_by_location(document_id)

        # Assert
        assert len(records) == 2
        assert records == mock_records

    @pytest.mark.asyncio
    async def test_get_records_by_location_with_sheet(self, mock_db_session):
        """Test getting records filtered by sheet_name."""
        # Create mock records
        mock_records = [MagicMock(spec=ExtractedRecordModel)]

        # Mock database response
        mock_scalars = MagicMock()
        mock_scalars.all.return_value = mock_records
        mock_result = MagicMock()
        mock_result.scalars.return_value = mock_scalars
        mock_db_session.execute = AsyncMock(return_value=mock_result)

        # Execute
        service = MetadataService(mock_db_session)
        document_id = uuid4()
        records = await service.get_records_by_location(document_id, sheet_name="Sheet1")

        # Assert
        assert len(records) == 1

    @pytest.mark.asyncio
    async def test_get_records_by_location_with_row_range(self, mock_db_session):
        """Test getting records filtered by row range."""
        # Create mock records
        mock_records = [
            MagicMock(spec=ExtractedRecordModel),
            MagicMock(spec=ExtractedRecordModel),
        ]

        # Mock database response
        mock_scalars = MagicMock()
        mock_scalars.all.return_value = mock_records
        mock_result = MagicMock()
        mock_result.scalars.return_value = mock_scalars
        mock_db_session.execute = AsyncMock(return_value=mock_result)

        # Execute
        service = MetadataService(mock_db_session)
        document_id = uuid4()
        records = await service.get_records_by_location(
            document_id, sheet_name="Sheet1", row_range=(10, 50)
        )

        # Assert
        assert len(records) == 2

    @pytest.mark.asyncio
    async def test_validate_source_completeness_all_valid(self, mock_db_session):
        """Test validation with all complete records."""
        # Create mock complete records
        mock_records = []
        for i in range(10):
            record = MagicMock(spec=ExtractedRecordModel)
            record.id = uuid4()
            record.sheet_name = f"Sheet{i}"
            record.row_number = i + 1
            record.col_range = "A:D"
            record.headers = ["Col1", "Col2"]
            mock_records.append(record)

        # Mock database response
        mock_scalars = MagicMock()
        mock_scalars.all.return_value = mock_records
        mock_result = MagicMock()
        mock_result.scalars.return_value = mock_scalars
        mock_db_session.execute = AsyncMock(return_value=mock_result)

        # Execute
        service = MetadataService(mock_db_session)
        document_id = uuid4()
        report = await service.validate_source_completeness(document_id)

        # Assert
        assert report.total_records == 10
        assert report.complete_records == 10
        assert report.is_valid is True
        assert len(report.missing_required_fields) == 0
        assert report.records_with_col_range == 10
        assert report.records_with_headers == 10

    @pytest.mark.asyncio
    async def test_validate_source_completeness_missing_sheet_name(
        self, mock_db_session
    ):
        """Test validation with missing sheet_name."""
        # Create mock records with one missing sheet_name
        mock_records = [
            MagicMock(
                spec=ExtractedRecordModel,
                id=uuid4(),
                sheet_name="Sheet1",
                row_number=1,
                col_range="A:D",
                headers=["Col1"],
            ),
            MagicMock(
                spec=ExtractedRecordModel,
                id=uuid4(),
                sheet_name=None,  # Missing
                row_number=2,
                col_range="A:D",
                headers=["Col1"],
            ),
        ]

        # Mock database response
        mock_scalars = MagicMock()
        mock_scalars.all.return_value = mock_records
        mock_result = MagicMock()
        mock_result.scalars.return_value = mock_scalars
        mock_db_session.execute = AsyncMock(return_value=mock_result)

        # Execute
        service = MetadataService(mock_db_session)
        document_id = uuid4()
        report = await service.validate_source_completeness(document_id)

        # Assert
        assert report.total_records == 2
        assert report.complete_records == 1
        assert report.is_valid is False
        assert len(report.missing_required_fields) == 1
        assert "sheet_name" in report.missing_required_fields[0]["missing_fields"]

    @pytest.mark.asyncio
    async def test_validate_source_completeness_missing_row_number(
        self, mock_db_session
    ):
        """Test validation with missing row_number."""
        # Create mock records with one missing row_number
        mock_records = [
            MagicMock(
                spec=ExtractedRecordModel,
                id=uuid4(),
                sheet_name="Sheet1",
                row_number=1,
                col_range="A:D",
                headers=None,
            ),
            MagicMock(
                spec=ExtractedRecordModel,
                id=uuid4(),
                sheet_name="Sheet1",
                row_number=None,  # Missing
                col_range=None,
                headers=None,
            ),
        ]

        # Mock database response
        mock_scalars = MagicMock()
        mock_scalars.all.return_value = mock_records
        mock_result = MagicMock()
        mock_result.scalars.return_value = mock_scalars
        mock_db_session.execute = AsyncMock(return_value=mock_result)

        # Execute
        service = MetadataService(mock_db_session)
        document_id = uuid4()
        report = await service.validate_source_completeness(document_id)

        # Assert
        assert report.total_records == 2
        assert report.complete_records == 1
        assert report.is_valid is False
        assert len(report.missing_required_fields) == 1

    @pytest.mark.asyncio
    async def test_validate_source_completeness_optional_fields(
        self, mock_db_session
    ):
        """Test validation counts optional fields correctly."""
        # Create mock records with varying optional fields
        mock_records = [
            MagicMock(
                spec=ExtractedRecordModel,
                id=uuid4(),
                sheet_name="Sheet1",
                row_number=1,
                col_range="A:D",  # Has col_range
                headers=["Col1"],  # Has headers
            ),
            MagicMock(
                spec=ExtractedRecordModel,
                id=uuid4(),
                sheet_name="Sheet1",
                row_number=2,
                col_range=None,  # No col_range
                headers=["Col1"],  # Has headers
            ),
            MagicMock(
                spec=ExtractedRecordModel,
                id=uuid4(),
                sheet_name="Sheet1",
                row_number=3,
                col_range="A:D",  # Has col_range
                headers=None,  # No headers
            ),
        ]

        # Mock database response
        mock_scalars = MagicMock()
        mock_scalars.all.return_value = mock_records
        mock_result = MagicMock()
        mock_result.scalars.return_value = mock_scalars
        mock_db_session.execute = AsyncMock(return_value=mock_result)

        # Execute
        service = MetadataService(mock_db_session)
        document_id = uuid4()
        report = await service.validate_source_completeness(document_id)

        # Assert
        assert report.total_records == 3
        assert report.complete_records == 3
        assert report.is_valid is True
        assert report.records_with_col_range == 2
        assert report.records_with_headers == 2

    @pytest.mark.asyncio
    async def test_get_citation_batch_success(self, mock_db_session, sample_record_data):
        """Test batch citation retrieval."""
        record_id_1 = uuid4()
        record_id_2 = uuid4()

        # Create mock records
        mock_record_1 = MagicMock(
            spec=ExtractedRecordModel,
            id=record_id_1,
            document_id=sample_record_data["document_id"],
            sheet_name="Sheet1",
            row_number=1,
            col_range="A:D",
            headers=["Col1"],
            created_at=datetime.now(),
        )
        mock_record_2 = MagicMock(
            spec=ExtractedRecordModel,
            id=record_id_2,
            document_id=sample_record_data["document_id"],
            sheet_name="Sheet2",
            row_number=2,
            col_range="A:E",
            headers=["Col2"],
            created_at=datetime.now(),
        )

        # Mock database response
        mock_result = MagicMock()
        mock_result.all.return_value = [
            (mock_record_1, "test1.xlsx"),
            (mock_record_2, "test2.xlsx"),
        ]
        mock_db_session.execute = AsyncMock(return_value=mock_result)

        # Execute
        service = MetadataService(mock_db_session)
        citations = await service.get_citation_batch([record_id_1, record_id_2])

        # Assert
        assert len(citations) == 2
        assert record_id_1 in citations
        assert record_id_2 in citations
        assert citations[record_id_1].filename == "test1.xlsx"
        assert citations[record_id_2].filename == "test2.xlsx"

    @pytest.mark.asyncio
    async def test_get_citation_batch_empty_list(self, mock_db_session):
        """Test batch citation with empty list."""
        service = MetadataService(mock_db_session)
        citations = await service.get_citation_batch([])

        # Assert
        assert citations == {}
        # Verify no database query was made
        mock_db_session.execute.assert_not_called()
