"""Cross-reference detection and resolution for extracted records.

This module provides functionality to detect cross-reference patterns
in extracted content and resolve them to target records in the database.
"""

import re
from dataclasses import dataclass, field
from typing import Any, Optional
from uuid import UUID

import structlog
from sqlalchemy.orm import Session

from src.db.models import CrossReference as CrossReferenceORM
from src.db.models import Document
from src.db.models import ExtractedRecord as ExtractedRecordORM

logger = structlog.get_logger(__name__)

# Batch size for database operations
BATCH_SIZE = 500


@dataclass
class CrossReferenceMatch:
    """Represents a detected cross-reference in content.

    Attributes:
        reference_text: The original reference text found in the content
        reference_type: Type of reference (sheet_ref, item_ref, page_ref, sheet_item_ref)
        field_name: The field/column where the reference was found
        sheet_target: Target sheet name if detected (for sheet_ref, sheet_item_ref)
        item_target: Target item code if detected (for item_ref, sheet_item_ref)
        page_target: Target page reference if detected (for page_ref)
        source_record_id: UUID of the source record containing the reference
        target_record_id: UUID of the resolved target record (if resolved)
        resolved: Whether the reference was successfully resolved
    """

    reference_text: str
    reference_type: str
    field_name: str
    sheet_target: Optional[str] = None
    item_target: Optional[str] = None
    page_target: Optional[str] = None
    source_record_id: Optional[UUID] = None
    target_record_id: Optional[UUID] = None
    resolved: bool = False

    def __post_init__(self):
        """Validate cross-reference match integrity after initialization."""
        valid_types = ("sheet_ref", "item_ref", "page_ref", "sheet_item_ref", "code_ref")
        if self.reference_type not in valid_types:
            raise ValueError(
                f"Invalid reference type: {self.reference_type}. "
                f"Must be one of: {valid_types}"
            )

        if not self.reference_text:
            raise ValueError("Reference text cannot be empty")

        if not self.field_name:
            raise ValueError("Field name cannot be empty")

    def to_dict(self) -> dict:
        """Convert CrossReferenceMatch to dictionary for JSON serialization."""
        result = {
            "reference_text": self.reference_text,
            "reference_type": self.reference_type,
            "field_name": self.field_name,
            "resolved": self.resolved,
        }

        if self.sheet_target:
            result["sheet_target"] = self.sheet_target
        if self.item_target:
            result["item_target"] = self.item_target
        if self.page_target:
            result["page_target"] = self.page_target
        if self.source_record_id:
            result["source_record_id"] = str(self.source_record_id)
        if self.target_record_id:
            result["target_record_id"] = str(self.target_record_id)

        return result


@dataclass
class CrossRefResolutionSummary:
    """Summary of cross-reference detection and resolution results.

    Attributes:
        document_id: UUID of the processed document
        references_found: Total number of cross-references detected
        references_resolved: Number of references successfully resolved
        unresolved_count: Number of references that could not be resolved
        unresolved_refs: List of unresolved reference texts
        warnings: List of warning messages during detection/resolution
    """

    document_id: UUID
    references_found: int
    references_resolved: int
    unresolved_count: int
    unresolved_refs: list[str] = field(default_factory=list)
    warnings: list[str] = field(default_factory=list)

    def to_dict(self) -> dict:
        """Convert CrossRefResolutionSummary to dictionary for JSON serialization."""
        return {
            "document_id": str(self.document_id),
            "references_found": self.references_found,
            "references_resolved": self.references_resolved,
            "unresolved_count": self.unresolved_count,
            "unresolved_refs": self.unresolved_refs,
            "warnings": self.warnings,
        }


# Cross-reference patterns for detection
# Each pattern is a tuple of (pattern_name, compiled_regex, reference_type)
CROSS_REF_PATTERNS = [
    # Japanese "別紙X" (Appendix X) patterns - very common in Japanese documents
    # Match: 別紙3_繰返し回数項目　参照, 別紙1参照, (別紙3参照), 別紙3 繰返し回数項目を参照
    (
        "appendix_ref_jp",
        re.compile(
            r"別紙\s*([0-9０-９]+)[_＿\-]?\s*([^\s　、。()（）参]*?)\s*(?:を)?参照",
            re.UNICODE,
        ),
        "sheet_ref",
    ),
    (
        "appendix_ref_jp_paren",
        re.compile(
            r"[（(]別紙\s*([0-9０-９]+)[_＿\-]?\s*([^)）]*?)[)）]\s*(?:を)?参照?",
            re.UNICODE,
        ),
        "sheet_ref",
    ),
    (
        "appendix_ref_jp_simple",
        re.compile(
            r"別紙\s*([0-9０-９]+)\s*参照",
            re.UNICODE,
        ),
        "sheet_ref",
    ),
    # "See Sheet X" or "参照 シート X" patterns - sheet references
    (
        "sheet_ref_en",
        re.compile(
            r"(?:See|see|SEE|Refer to|refer to)\s+(?:Sheet|sheet|SHEET)\s+['\"]?([^'\",\n]+)['\"]?",
            re.IGNORECASE,
        ),
        "sheet_ref",
    ),
    (
        "sheet_ref_jp",
        re.compile(r"(?:参照|参考)\s*(?:シート|Sheet)?\s*[「『]?([^」』\n]+)[」』]?"),
        "sheet_ref",
    ),
    # "Item Y" or "項目 Y" patterns - item references
    (
        "item_ref_en",
        re.compile(
            r"(?:Item|item|ITEM)\s*[:\s]?\s*([A-Za-z0-9_\-]+)",
            re.IGNORECASE,
        ),
        "item_ref",
    ),
    (
        "item_ref_jp",
        re.compile(r"(?:項目|アイテム)\s*[:\s]?\s*([A-Za-z0-9_\-]+)"),
        "item_ref",
    ),
    # "Ref: CODE" or "参照: CODE" patterns - code references
    (
        "code_ref_en",
        re.compile(
            r"(?:Ref|ref|REF|Reference|reference)[:\s]+([A-Za-z0-9_\-]+)",
            re.IGNORECASE,
        ),
        "code_ref",
    ),
    (
        "code_ref_jp",
        re.compile(r"(?:参照|Ref)[:\s]+([A-Za-z0-9_\-]+)"),
        "code_ref",
    ),
    # "Page X, Item Y" patterns - page and item references
    (
        "page_item_ref_en",
        re.compile(
            r"(?:Page|page|PAGE)\s+([0-9.\-]+)[,、]\s*(?:Item|item|ITEM)\s+([0-9]+)",
            re.IGNORECASE,
        ),
        "page_ref",
    ),
    (
        "page_item_ref_jp",
        re.compile(r"(?:ページ|Page)\s*([0-9.\-]+)[,、]\s*(?:項目|Item)\s*([0-9]+)"),
        "page_ref",
    ),
    # Combined "See Sheet X, Item Y" patterns
    (
        "sheet_item_ref",
        re.compile(
            r"(?:See|see|参照)\s+(?:Sheet|sheet|シート)\s+['\"]?([^'\",\n]+)['\"]?[,、]\s*(?:Item|item|項目)\s+([A-Za-z0-9_\-]+)",
            re.IGNORECASE,
        ),
        "sheet_item_ref",
    ),
]


class CrossRefDetector:
    """Detects and resolves cross-references in extracted records.

    This class scans extracted record content for cross-reference patterns,
    attempts to resolve them to target records, and stores the results
    in the database.
    """

    def __init__(self):
        """Initialize the CrossRefDetector."""
        pass

    def detect_references(
        self, content: dict[str, Any], field_name_override: Optional[str] = None
    ) -> list[CrossReferenceMatch]:
        """Detect cross-reference patterns in record content.

        Args:
            content: Dictionary mapping field names to values
            field_name_override: Optional field name to use for all matches

        Returns:
            List of CrossReferenceMatch objects for detected references
        """
        matches: list[CrossReferenceMatch] = []

        if not content:
            return matches

        for field_name, value in content.items():
            # Skip None values and non-string values
            if value is None:
                continue

            # Convert to string
            value_str = str(value).strip()
            if not value_str:
                continue

            # Check each pattern
            field_matches = self._detect_in_value(
                value_str, field_name_override or field_name
            )
            matches.extend(field_matches)

        return matches

    def _detect_in_value(
        self, value: str, field_name: str
    ) -> list[CrossReferenceMatch]:
        """Detect cross-references in a single string value.

        Args:
            value: String value to scan for references
            field_name: Name of the field containing the value

        Returns:
            List of CrossReferenceMatch objects
        """
        matches: list[CrossReferenceMatch] = []
        seen_texts: set[str] = set()  # Avoid duplicate matches for same text

        for pattern_name, pattern, ref_type in CROSS_REF_PATTERNS:
            for match in pattern.finditer(value):
                reference_text = match.group(0).strip()

                # Skip if we've already matched this exact text
                if reference_text in seen_texts:
                    continue
                seen_texts.add(reference_text)

                # Extract targets based on reference type and pattern
                sheet_target = None
                item_target = None
                page_target = None

                # Handle Japanese appendix patterns (別紙X)
                if pattern_name.startswith("appendix_ref_jp"):
                    # Convert full-width digits to half-width
                    appendix_num = match.group(1).strip()
                    appendix_num = self._normalize_digits(appendix_num)
                    
                    # Build sheet target name like "別紙X_..." or just "別紙X"
                    if match.lastindex and match.lastindex >= 2:
                        appendix_name = match.group(2).strip() if match.group(2) else ""
                        if appendix_name:
                            sheet_target = f"別紙{appendix_num}_{appendix_name}"
                        else:
                            sheet_target = f"別紙{appendix_num}"
                    else:
                        sheet_target = f"別紙{appendix_num}"
                elif ref_type == "sheet_ref":
                    sheet_target = match.group(1).strip()
                elif ref_type == "item_ref":
                    item_target = match.group(1).strip()
                elif ref_type == "code_ref":
                    item_target = match.group(1).strip()  # Treat code as item
                elif ref_type == "page_ref":
                    page_target = match.group(1).strip()
                    if match.lastindex and match.lastindex >= 2:
                        item_target = match.group(2).strip()
                elif ref_type == "sheet_item_ref":
                    sheet_target = match.group(1).strip()
                    if match.lastindex and match.lastindex >= 2:
                        item_target = match.group(2).strip()

                cross_ref = CrossReferenceMatch(
                    reference_text=reference_text,
                    reference_type=ref_type,
                    field_name=field_name,
                    sheet_target=sheet_target,
                    item_target=item_target,
                    page_target=page_target,
                )
                matches.append(cross_ref)

                logger.debug(
                    "cross_reference_detected",
                    pattern=pattern_name,
                    reference_type=ref_type,
                    reference_text=reference_text,
                    sheet_target=sheet_target,
                    field_name=field_name,
                )

        return matches

    def _normalize_digits(self, text: str) -> str:
        """Convert full-width digits to half-width.

        Args:
            text: String that may contain full-width digits

        Returns:
            String with half-width digits
        """
        full_width = "０１２３４５６７８９"
        half_width = "0123456789"
        trans_table = str.maketrans(full_width, half_width)
        return text.translate(trans_table)

    def resolve_reference(
        self,
        match: CrossReferenceMatch,
        document_id: UUID,
        session: Session,
        api_key_id: Optional[UUID] = None,
    ) -> Optional[UUID]:
        """Attempt to resolve a cross-reference to a target record.

        Args:
            match: CrossReferenceMatch to resolve
            document_id: UUID of the source document
            session: SQLAlchemy database session
            api_key_id: Optional API key ID for cross-document resolution

        Returns:
            UUID of target record if resolved, None otherwise
        """
        # Build query for target records
        query = session.query(ExtractedRecordORM)

        # If api_key_id provided, allow cross-document resolution
        if api_key_id:
            # Get all document IDs for this API key
            doc_ids = (
                session.query(Document.id)
                .filter(Document.api_key_id == api_key_id)
                .all()
            )
            doc_id_list = [d[0] for d in doc_ids]
            query = query.filter(ExtractedRecordORM.document_id.in_(doc_id_list))
        else:
            # Only search within the same document
            query = query.filter(ExtractedRecordORM.document_id == document_id)

        # Try different resolution strategies based on reference type
        target_record = None

        if match.sheet_target and match.item_target:
            # Combined sheet + item reference
            target_record = self._resolve_sheet_item(
                query, match.sheet_target, match.item_target
            )
        elif match.sheet_target:
            # Sheet-only reference - find first record in that sheet
            target_record = self._resolve_sheet_only(query, match.sheet_target)
        elif match.item_target:
            # Item/code reference - search in content
            target_record = self._resolve_item_code(query, match.item_target)

        if target_record:
            logger.debug(
                "cross_reference_resolved",
                reference_text=match.reference_text,
                target_record_id=str(target_record.id),
            )
            return target_record.id

        logger.warning(
            "cross_reference_unresolved",
            reference_text=match.reference_text,
            reference_type=match.reference_type,
            sheet_target=match.sheet_target,
            item_target=match.item_target,
        )
        return None

    def _resolve_sheet_item(
        self, query, sheet_target: str, item_target: str
    ) -> Optional[ExtractedRecordORM]:
        """Resolve a combined sheet + item reference."""
        # First filter by sheet name (case-insensitive)
        sheet_query = query.filter(
            ExtractedRecordORM.sheet_name.ilike(sheet_target)
        )

        # Then search for item in content
        for record in sheet_query.all():
            if record.content and self._content_contains_item(
                record.content, item_target
            ):
                return record

        return None

    def _resolve_sheet_only(
        self, query, sheet_target: str
    ) -> Optional[ExtractedRecordORM]:
        """Resolve a sheet-only reference to the first record in that sheet."""
        return (
            query.filter(ExtractedRecordORM.sheet_name.ilike(sheet_target))
            .order_by(ExtractedRecordORM.row_number)
            .first()
        )

    def _resolve_item_code(
        self, query, item_target: str
    ) -> Optional[ExtractedRecordORM]:
        """Resolve an item/code reference by searching record content."""
        # Search for exact match in content JSONB
        # Use JSONB operators to search for the item code
        for record in query.all():
            if record.content and self._content_contains_item(
                record.content, item_target
            ):
                return record

        return None

    def _content_contains_item(self, content: dict, item_code: str) -> bool:
        """Check if content dictionary contains the item code."""
        if not content:
            return False

        item_code_lower = item_code.lower()
        for key, value in content.items():
            if value is None:
                continue

            value_str = str(value).strip().lower()

            # Exact match
            if value_str == item_code_lower:
                return True

            # Check if item code is in the value
            if item_code_lower in value_str:
                return True

            # Check key names that typically contain item codes
            key_lower = key.lower()
            if any(
                keyword in key_lower
                for keyword in ("item", "code", "id", "番号", "コード")
            ):
                if value_str == item_code_lower or item_code_lower in value_str:
                    return True

        return False

    def process_document(
        self,
        document_id: UUID,
        session: Session,
        api_key_id: Optional[UUID] = None,
    ) -> CrossRefResolutionSummary:
        """Detect and resolve all cross-references in a document's records.

        Args:
            document_id: UUID of the document to process
            session: SQLAlchemy database session
            api_key_id: Optional API key ID for cross-document resolution

        Returns:
            CrossRefResolutionSummary with detection and resolution statistics
        """
        logger.info("cross_ref_detection_started", document_id=str(document_id))

        # Load all extracted records for this document
        records = (
            session.query(ExtractedRecordORM)
            .filter(ExtractedRecordORM.document_id == document_id)
            .all()
        )

        if not records:
            logger.info(
                "no_records_for_cross_ref_detection",
                document_id=str(document_id),
            )
            return CrossRefResolutionSummary(
                document_id=document_id,
                references_found=0,
                references_resolved=0,
                unresolved_count=0,
                unresolved_refs=[],
                warnings=["No extracted records found for document"],
            )

        # Detect and resolve cross-references
        all_matches: list[CrossReferenceMatch] = []
        unresolved_refs: list[str] = []

        for record in records:
            if not record.content:
                continue

            # Detect references in this record
            matches = self.detect_references(record.content)

            for match in matches:
                match.source_record_id = record.id

                # Try to resolve the reference
                target_id = self.resolve_reference(
                    match, document_id, session, api_key_id
                )

                if target_id:
                    match.target_record_id = target_id
                    match.resolved = True
                else:
                    match.resolved = False
                    if match.reference_text not in unresolved_refs:
                        unresolved_refs.append(match.reference_text)

                all_matches.append(match)

        # Store cross-references in database
        if all_matches:
            self._store_cross_references(all_matches, session)

        # Build summary
        resolved_count = sum(1 for m in all_matches if m.resolved)
        summary = CrossRefResolutionSummary(
            document_id=document_id,
            references_found=len(all_matches),
            references_resolved=resolved_count,
            unresolved_count=len(all_matches) - resolved_count,
            unresolved_refs=unresolved_refs,
            warnings=[],
        )

        logger.info(
            "cross_ref_detection_completed",
            document_id=str(document_id),
            references_found=summary.references_found,
            references_resolved=summary.references_resolved,
            unresolved_count=summary.unresolved_count,
        )

        return summary

    def _store_cross_references(
        self, matches: list[CrossReferenceMatch], session: Session
    ) -> int:
        """Store cross-references in the database.

        Args:
            matches: List of CrossReferenceMatch objects to store
            session: SQLAlchemy database session

        Returns:
            Number of cross-references stored
        """
        if not matches:
            return 0

        stored_count = 0

        # Process in batches
        for i in range(0, len(matches), BATCH_SIZE):
            batch = matches[i : i + BATCH_SIZE]

            try:
                for match in batch:
                    if not match.source_record_id:
                        continue

                    # Check if this cross-reference already exists
                    existing = (
                        session.query(CrossReferenceORM)
                        .filter(
                            CrossReferenceORM.source_record_id == match.source_record_id,
                            CrossReferenceORM.reference_text == match.reference_text,
                        )
                        .first()
                    )

                    if existing:
                        # Update existing record
                        existing.target_record_id = match.target_record_id
                        existing.reference_type = match.reference_type
                        existing.resolved = match.resolved
                    else:
                        # Create new record
                        cross_ref = CrossReferenceORM(
                            source_record_id=match.source_record_id,
                            target_record_id=match.target_record_id,
                            reference_text=match.reference_text,
                            reference_type=match.reference_type,
                            resolved=match.resolved,
                        )
                        session.add(cross_ref)

                    stored_count += 1

                session.commit()

                logger.debug(
                    "cross_ref_batch_stored",
                    batch_size=len(batch),
                    total_stored=stored_count,
                )

            except Exception as e:
                session.rollback()
                logger.error(
                    "cross_ref_storage_failed",
                    batch_start=i,
                    batch_size=len(batch),
                    error=str(e),
                )
                raise

        logger.info(
            "cross_references_stored",
            total_stored=stored_count,
        )

        return stored_count
