from __future__ import annotations

from typing import Any

from pydantic import BaseModel, Field, field_validator, model_validator

from app.models.enums import ChunkType, SourceType


class EmbeddingRecordMetadata(BaseModel):
    """Additional retrieval metadata stored on embedding_records."""

    language: str | None = None
    symbol_type: str | None = None
    impacted_symbols: list[str] = Field(default_factory=list)
    changed_files: list[str] = Field(default_factory=list)
    chunk_hash: str | None = None
    author_name: str | None = None
    author_email: str | None = None
    committer_name: str | None = None
    committer_email: str | None = None
    authored_at: str | None = None
    committed_at: str | None = None
    commit_message: str | None = None

    model_config = {"extra": "allow"}


class EmbeddingRecord(BaseModel):
    """
    Canonical embedding payload for MongoDB embedding_records and Qdrant payloads.

    See adpilot-common-composer.com/docs/indexing-pipeline.md §4.5.
    """

    record_id: str
    repo_id: str
    snapshot_id: str
    commit_sha: str | None = None
    source_type: SourceType
    chunk_type: ChunkType
    text: str
    file_path: str | None = None
    doc_path: str | None = None
    symbol_name: str | None = None
    section_title: str | None = None
    start_line: int | None = None
    end_line: int | None = None
    graph_node_id: str | None = None
    artifact_uri: str | None = None
    metadata: EmbeddingRecordMetadata | dict[str, Any] = Field(default_factory=dict)

    @field_validator("text")
    @classmethod
    def _text_not_blank(cls, value: str) -> str:
        if not value or not value.strip():
            raise ValueError("text must be non-empty")
        return value.strip()

    @field_validator("metadata", mode="before")
    @classmethod
    def _coerce_metadata(cls, value: Any) -> EmbeddingRecordMetadata | dict[str, Any]:
        if value is None:
            return EmbeddingRecordMetadata()
        if isinstance(value, EmbeddingRecordMetadata):
            return value
        if isinstance(value, dict):
            return EmbeddingRecordMetadata.model_validate(value)
        raise ValueError("metadata must be an object")

    @model_validator(mode="after")
    def _validate_source_shape(self) -> EmbeddingRecord:
        if self.source_type == SourceType.CODE:
            if self.chunk_type != ChunkType.SYMBOL:
                raise ValueError("code records must use chunk_type=symbol")
            if not self.file_path:
                raise ValueError("code records require file_path")
            if not self.symbol_name:
                raise ValueError("code records require symbol_name")
        elif self.source_type == SourceType.DOCS:
            if self.chunk_type not in (ChunkType.TEXT, ChunkType.SUMMARY):
                raise ValueError("docs records must use chunk_type=text or summary")
            if not self.doc_path:
                raise ValueError("docs records require doc_path")
            if not self.section_title:
                raise ValueError("docs records require section_title")
        elif self.source_type == SourceType.COMMIT:
            if self.chunk_type != ChunkType.SUMMARY:
                raise ValueError("commit records must use chunk_type=summary")
            if not self.commit_sha:
                raise ValueError("commit records require commit_sha")
        return self

    @classmethod
    def from_code_chunk(
        cls,
        *,
        record_id: str,
        repo_id: str,
        snapshot_id: str,
        commit_sha: str,
        text: str,
        file_path: str,
        symbol_name: str,
        start_line: int,
        end_line: int,
        graph_node_id: str,
        artifact_uri: str | None = None,
        language: str | None = None,
        symbol_type: str | None = None,
        chunk_hash: str | None = None,
    ) -> EmbeddingRecord:
        return cls(
            record_id=record_id,
            repo_id=repo_id,
            snapshot_id=snapshot_id,
            commit_sha=commit_sha,
            source_type=SourceType.CODE,
            chunk_type=ChunkType.SYMBOL,
            text=text,
            file_path=file_path,
            symbol_name=symbol_name,
            start_line=start_line,
            end_line=end_line,
            graph_node_id=graph_node_id,
            artifact_uri=artifact_uri,
            metadata=EmbeddingRecordMetadata(
                language=language,
                symbol_type=symbol_type,
                chunk_hash=chunk_hash,
            ),
        )

    @classmethod
    def from_doc_chunk(
        cls,
        *,
        record_id: str,
        repo_id: str,
        snapshot_id: str,
        commit_sha: str,
        text: str,
        doc_path: str,
        section_title: str,
        graph_node_id: str,
        section_level: int | None = None,
        chunk_hash: str | None = None,
    ) -> EmbeddingRecord:
        metadata_fields: dict[str, Any] = {}
        if chunk_hash is not None:
            metadata_fields["chunk_hash"] = chunk_hash
        if section_level is not None:
            metadata_fields["section_level"] = section_level
        metadata = EmbeddingRecordMetadata.model_validate(metadata_fields)
        return cls(
            record_id=record_id,
            repo_id=repo_id,
            snapshot_id=snapshot_id,
            commit_sha=commit_sha,
            source_type=SourceType.DOCS,
            chunk_type=ChunkType.TEXT,
            text=text,
            doc_path=doc_path,
            section_title=section_title,
            graph_node_id=graph_node_id,
            metadata=metadata,
        )

    @classmethod
    def from_repo_overview(
        cls,
        *,
        record_id: str,
        repo_id: str,
        snapshot_id: str,
        commit_sha: str,
        text: str,
        chunk_hash: str | None = None,
    ) -> EmbeddingRecord:
        metadata_fields: dict[str, Any] = {"overview": True}
        if chunk_hash is not None:
            metadata_fields["chunk_hash"] = chunk_hash
        return cls(
            record_id=record_id,
            repo_id=repo_id,
            snapshot_id=snapshot_id,
            commit_sha=commit_sha,
            source_type=SourceType.DOCS,
            chunk_type=ChunkType.SUMMARY,
            text=text,
            doc_path="REPO_OVERVIEW.md",
            section_title="Repository Overview",
            graph_node_id=f"overview_{snapshot_id}",
            metadata=EmbeddingRecordMetadata.model_validate(metadata_fields),
        )

    @classmethod
    def from_commit_analysis(
        cls,
        *,
        record_id: str,
        repo_id: str,
        snapshot_id: str,
        commit_sha: str,
        text: str,
        impacted_symbols: list[str] | None = None,
        changed_files: list[str] | None = None,
        author_name: str | None = None,
        author_email: str | None = None,
        committer_name: str | None = None,
        committer_email: str | None = None,
        authored_at: str | None = None,
        committed_at: str | None = None,
        commit_message: str | None = None,
    ) -> EmbeddingRecord:
        return cls(
            record_id=record_id,
            repo_id=repo_id,
            snapshot_id=snapshot_id,
            commit_sha=commit_sha,
            source_type=SourceType.COMMIT,
            chunk_type=ChunkType.SUMMARY,
            text=text,
            metadata=EmbeddingRecordMetadata(
                impacted_symbols=impacted_symbols or [],
                changed_files=changed_files or [],
                author_name=author_name,
                author_email=author_email,
                committer_name=committer_name,
                committer_email=committer_email,
                authored_at=authored_at,
                committed_at=committed_at,
                commit_message=commit_message,
            ),
        )


def validate_embedding_record(data: dict[str, Any]) -> EmbeddingRecord:
    """Validate a raw dict against the embedding record contract."""
    return EmbeddingRecord.model_validate(data)


def embedding_record_id_for_chunk(chunk_id: str) -> str:
    """Deterministic embedding record id for an upstream chunk."""
    normalized = chunk_id.strip()
    if normalized.startswith("emb_"):
        return normalized
    return f"emb_{normalized}"


def embedding_record_id_for_analysis(analysis_id: str) -> str:
    normalized = analysis_id.strip()
    if normalized.startswith("emb_"):
        return normalized
    return f"emb_{normalized}"
