"""Source metadata service for querying and validating record metadata."""

from dataclasses import asdict, dataclass
from datetime import datetime
from typing import Any
from uuid import UUID

import structlog
from sqlalchemy import func, select
from sqlalchemy.ext.asyncio import AsyncSession

from src.db.models import Document as DocumentModel
from src.db.models import ExtractedRecord as ExtractedRecordModel

logger = structlog.get_logger()


@dataclass
class SourceCitation:
    """Structured source citation for a record.

    Provides both human-readable and machine-readable formats for
    citing the exact location of a record in source documents.
    """

    record_id: UUID
    document_id: UUID
    filename: str
    sheet_name: str
    row_number: int
    col_range: str | None
    headers: list[str] | None
    created_at: datetime

    def to_human_readable(self) -> str:
        """Format as human-readable citation.

        Returns:
            Citation like: "Found in distribution_spec.xlsx, sheet 'Item Mapping', row 42"

        Example:
            >>> citation = SourceCitation(...)
            >>> citation.to_human_readable()
            "Found in distribution_spec.xlsx, sheet 'Item Mapping', row 42, columns A:F"
        """
        citation = f"Found in {self.filename}, sheet '{self.sheet_name}', row {self.row_number}"

        if self.col_range:
            citation += f", columns {self.col_range}"

        if self.headers:
            headers_str = ", ".join(self.headers)
            citation += f" (Headers: {headers_str})"

        return citation

    def to_json(self) -> dict[str, Any]:
        """Return structured JSON representation.

        Returns:
            Dictionary with all source fields, suitable for API responses.

        Example:
            >>> citation = SourceCitation(...)
            >>> citation.to_json()
            {
                "record_id": "123e4567-e89b-12d3-a456-426614174000",
                "filename": "distribution_spec.xlsx",
                "sheet_name": "Item Mapping",
                ...
            }
        """
        return {
            "record_id": str(self.record_id),
            "document_id": str(self.document_id),
            "filename": self.filename,
            "sheet_name": self.sheet_name,
            "row_number": self.row_number,
            "col_range": self.col_range,
            "headers": self.headers,
            "created_at": self.created_at.isoformat() if self.created_at else None,
        }


@dataclass
class ValidationReport:
    """Source metadata validation report for a document.

    Tracks completeness of source metadata across all records
    in a document, identifying any missing required fields.
    """

    document_id: UUID
    total_records: int
    complete_records: int
    records_with_col_range: int
    records_with_headers: int
    missing_required_fields: list[dict[str, Any]]

    @property
    def is_valid(self) -> bool:
        """Check if all records have required source metadata.

        Returns:
            True if no records have missing required fields.
        """
        return len(self.missing_required_fields) == 0

    @property
    def completeness_percentage(self) -> float:
        """Calculate percentage of records with complete metadata.

        Returns:
            Percentage (0-100) of complete records.
        """
        if self.total_records == 0:
            return 100.0
        return (self.complete_records / self.total_records) * 100

    def to_dict(self) -> dict[str, Any]:
        """Return structured dictionary for logging/API responses.

        Returns:
            Dictionary with validation statistics and issues.
        """
        return {
            "document_id": str(self.document_id),
            "total_records": self.total_records,
            "complete_records": self.complete_records,
            "completeness_percentage": round(self.completeness_percentage, 2),
            "records_with_col_range": self.records_with_col_range,
            "records_with_headers": self.records_with_headers,
            "is_valid": self.is_valid,
            "issues_count": len(self.missing_required_fields),
            "issues": self.missing_required_fields,
        }


class MetadataService:
    """Service for querying and validating source metadata.

    Provides helper methods to:
    - Query source citations for records
    - Find records by location (document, sheet, row range)
    - Validate source metadata completeness

    This service reads metadata stored by StructuredStore (Story 4.1).
    It does not modify or store data, only queries and validates.
    """

    def __init__(self, db_session: AsyncSession):
        """Initialize MetadataService.

        Args:
            db_session: Async SQLAlchemy session for database queries.
        """
        self.db = db_session
        self.logger = logger.bind(component="metadata_service")

    async def get_source_citation(self, record_id: UUID) -> SourceCitation | None:
        """Get source citation for a record.

        Queries extracted_records table and joins with documents table
        to get complete source metadata including filename.

        Args:
            record_id: UUID of the extracted record.

        Returns:
            SourceCitation object with complete metadata, or None if not found.

        Raises:
            Exception: If database query fails.

        Example:
            >>> service = MetadataService(session)
            >>> citation = await service.get_source_citation(record_id)
            >>> print(citation.to_human_readable())
            "Found in distribution_spec.xlsx, sheet 'Item Mapping', row 42"
        """
        start_time = datetime.now()

        try:
            # Query with JOIN to get filename from documents table
            stmt = (
                select(ExtractedRecordModel, DocumentModel.filename)
                .join(
                    DocumentModel,
                    ExtractedRecordModel.document_id == DocumentModel.id,
                )
                .where(ExtractedRecordModel.id == record_id)
            )

            result = await self.db.execute(stmt)
            row = result.first()

            if not row:
                self.logger.warning(
                    "record_not_found",
                    record_id=str(record_id),
                )
                return None

            record, filename = row

            # Build SourceCitation
            citation = SourceCitation(
                record_id=record.id,
                document_id=record.document_id,
                filename=filename,
                sheet_name=record.sheet_name,
                row_number=record.row_number,
                col_range=record.col_range,
                headers=record.headers,
                created_at=record.created_at,
            )

            duration_ms = int((datetime.now() - start_time).total_seconds() * 1000)
            self.logger.debug(
                "source_citation_retrieved",
                record_id=str(record_id),
                duration_ms=duration_ms,
            )

            return citation

        except Exception as e:
            duration_ms = int((datetime.now() - start_time).total_seconds() * 1000)
            self.logger.error(
                "source_citation_failed",
                record_id=str(record_id),
                error=str(e),
                error_type=type(e).__name__,
                duration_ms=duration_ms,
            )
            raise

    async def get_records_by_location(
        self,
        document_id: UUID,
        sheet_name: str | None = None,
        row_range: tuple[int, int] | None = None,
    ) -> list[ExtractedRecordModel]:
        """Query records by source location.

        Filters records by document_id, optional sheet_name, and optional row range.

        Args:
            document_id: UUID of the source document.
            sheet_name: Optional sheet name to filter by.
            row_range: Optional tuple (start_row, end_row) to filter by row number.

        Returns:
            List of ExtractedRecord models matching the filters.

        Raises:
            Exception: If database query fails.

        Example:
            >>> service = MetadataService(session)
            >>> # Get all records from a document
            >>> records = await service.get_records_by_location(doc_id)
            >>> # Get records from specific sheet
            >>> records = await service.get_records_by_location(doc_id, "Sheet1")
            >>> # Get records from sheet in row range
            >>> records = await service.get_records_by_location(doc_id, "Sheet1", (10, 50))
        """
        start_time = datetime.now()

        try:
            # Base query
            stmt = select(ExtractedRecordModel).where(
                ExtractedRecordModel.document_id == document_id
            )

            # Add sheet_name filter if provided
            if sheet_name:
                stmt = stmt.where(ExtractedRecordModel.sheet_name == sheet_name)

            # Add row range filter if provided
            if row_range:
                start_row, end_row = row_range
                stmt = stmt.where(
                    ExtractedRecordModel.row_number >= start_row,
                    ExtractedRecordModel.row_number <= end_row,
                )

            # Order by sheet name and row number for consistent results
            stmt = stmt.order_by(
                ExtractedRecordModel.sheet_name,
                ExtractedRecordModel.row_number,
            )

            result = await self.db.execute(stmt)
            records = result.scalars().all()

            duration_ms = int((datetime.now() - start_time).total_seconds() * 1000)
            self.logger.debug(
                "records_by_location_retrieved",
                document_id=str(document_id),
                sheet_name=sheet_name,
                row_range=row_range,
                record_count=len(records),
                duration_ms=duration_ms,
            )

            return list(records)

        except Exception as e:
            duration_ms = int((datetime.now() - start_time).total_seconds() * 1000)
            self.logger.error(
                "records_by_location_failed",
                document_id=str(document_id),
                sheet_name=sheet_name,
                error=str(e),
                error_type=type(e).__name__,
                duration_ms=duration_ms,
            )
            raise

    async def validate_source_completeness(
        self, document_id: UUID
    ) -> ValidationReport:
        """Validate that all records have complete source metadata.

        Checks all records for a document to ensure required fields are present.
        Counts optional fields and identifies any incomplete records.

        Args:
            document_id: UUID of the document to validate.

        Returns:
            ValidationReport with completeness statistics and any issues.

        Raises:
            Exception: If database query fails.

        Example:
            >>> service = MetadataService(session)
            >>> report = await service.validate_source_completeness(doc_id)
            >>> if report.is_valid:
            ...     print(f"All {report.total_records} records are complete!")
            >>> else:
            ...     print(f"Found {len(report.missing_required_fields)} incomplete records")
        """
        start_time = datetime.now()

        try:
            # Query all records for the document
            stmt = select(ExtractedRecordModel).where(
                ExtractedRecordModel.document_id == document_id
            )

            result = await self.db.execute(stmt)
            records = result.scalars().all()

            # Analyze metadata completeness
            missing_fields: list[dict[str, Any]] = []
            records_with_col_range = 0
            records_with_headers = 0

            for record in records:
                # Check required fields (document_id always present due to foreign key)
                record_missing: list[str] = []

                if not record.sheet_name:
                    record_missing.append("sheet_name")

                if record.row_number is None:  # Could be 0, so check None explicitly
                    record_missing.append("row_number")

                if record_missing:
                    missing_fields.append(
                        {
                            "record_id": str(record.id),
                            "sheet_name": record.sheet_name,
                            "row_number": record.row_number,
                            "missing_fields": record_missing,
                        }
                    )

                # Count optional fields
                if record.col_range:
                    records_with_col_range += 1

                if record.headers:
                    records_with_headers += 1

            # Build validation report
            complete_records = len(records) - len(missing_fields)

            report = ValidationReport(
                document_id=document_id,
                total_records=len(records),
                complete_records=complete_records,
                records_with_col_range=records_with_col_range,
                records_with_headers=records_with_headers,
                missing_required_fields=missing_fields,
            )

            duration_ms = int((datetime.now() - start_time).total_seconds() * 1000)
            self.logger.info(
                "source_validation_completed",
                document_id=str(document_id),
                is_valid=report.is_valid,
                total_records=report.total_records,
                complete_records=report.complete_records,
                completeness_percentage=report.completeness_percentage,
                issues_count=len(report.missing_required_fields),
                duration_ms=duration_ms,
            )

            return report

        except Exception as e:
            duration_ms = int((datetime.now() - start_time).total_seconds() * 1000)
            self.logger.error(
                "source_validation_failed",
                document_id=str(document_id),
                error=str(e),
                error_type=type(e).__name__,
                duration_ms=duration_ms,
            )
            raise

    async def get_citation_batch(
        self, record_ids: list[UUID]
    ) -> dict[UUID, SourceCitation]:
        """Get source citations for multiple records in a single query.

        More efficient than calling get_source_citation() multiple times.

        Args:
            record_ids: List of record UUIDs to get citations for.

        Returns:
            Dictionary mapping record_id to SourceCitation.
            Missing records are omitted from the result.

        Raises:
            Exception: If database query fails.

        Example:
            >>> service = MetadataService(session)
            >>> citations = await service.get_citation_batch([id1, id2, id3])
            >>> for record_id, citation in citations.items():
            ...     print(citation.to_human_readable())
        """
        if not record_ids:
            return {}

        start_time = datetime.now()

        try:
            # Query with JOIN for all record_ids at once
            stmt = (
                select(ExtractedRecordModel, DocumentModel.filename)
                .join(
                    DocumentModel,
                    ExtractedRecordModel.document_id == DocumentModel.id,
                )
                .where(ExtractedRecordModel.id.in_(record_ids))
            )

            result = await self.db.execute(stmt)
            rows = result.all()

            # Build dictionary of citations
            citations: dict[UUID, SourceCitation] = {}
            for record, filename in rows:
                citations[record.id] = SourceCitation(
                    record_id=record.id,
                    document_id=record.document_id,
                    filename=filename,
                    sheet_name=record.sheet_name,
                    row_number=record.row_number,
                    col_range=record.col_range,
                    headers=record.headers,
                    created_at=record.created_at,
                )

            duration_ms = int((datetime.now() - start_time).total_seconds() * 1000)
            self.logger.debug(
                "citation_batch_retrieved",
                requested_count=len(record_ids),
                found_count=len(citations),
                duration_ms=duration_ms,
            )

            return citations

        except Exception as e:
            duration_ms = int((datetime.now() - start_time).total_seconds() * 1000)
            self.logger.error(
                "citation_batch_failed",
                requested_count=len(record_ids),
                error=str(e),
                error_type=type(e).__name__,
                duration_ms=duration_ms,
            )
            raise
