from __future__ import annotations

from enum import Enum
from typing import Literal

from pydantic import BaseModel, ConfigDict, Field, field_validator

from app.core.config import settings

from app.models.enums import SourceType

RETRIEVE_SOURCE_TYPES = frozenset({SourceType.CODE.value, SourceType.DOCS.value, SourceType.COMMIT.value})


class HydrationSource(str, Enum):
    EMBEDDING_RECORD = "embedding_record"
    CODE_CHUNK = "code_chunk"
    DOC_CHUNK = "doc_chunk"
    COMMIT_ANALYSIS = "commit_analysis"


class VectorHit(BaseModel):
    """Ranked vector search hit (payload metadata only; text comes from hydration)."""

    record_id: str
    score: float
    source_type: str
    chunk_type: str
    graph_node_id: str | None = None
    file_path: str | None = None
    doc_path: str | None = None
    symbol_name: str | None = None
    section_title: str | None = None
    commit_sha: str | None = None


class RetrieveFilters(BaseModel):
    source_types: list[Literal["code", "docs", "commit"]] | None = None

    @field_validator("source_types")
    @classmethod
    def _validate_source_types(
        cls, value: list[str] | None
    ) -> list[Literal["code", "docs", "commit"]] | None:
        if value is None:
            return None
        if not value:
            raise ValueError("source_types must not be empty when provided")
        invalid = [item for item in value if item not in RETRIEVE_SOURCE_TYPES]
        if invalid:
            raise ValueError(
                f"Invalid source_types: {invalid}. Allowed: code, docs, commit"
            )
        return value


class VectorSearchResult(BaseModel):
    """Output of query embedding + Qdrant search (pre-hydration)."""

    hits: list[VectorHit] = Field(default_factory=list)
    embedding_model: str
    embedding_dimension: int
    vector_latency_ms: int = 0


class HydratedHit(BaseModel):
    """Vector hit enriched with full text from Mongo."""

    record_id: str
    source_type: str
    score: float
    chunk_type: str
    text: str
    file_path: str | None = None
    doc_path: str | None = None
    section_title: str | None = None
    symbol_name: str | None = None
    symbol_type: str | None = None
    language: str | None = None
    section_level: int | None = None
    start_line: int | None = None
    end_line: int | None = None
    commit_sha: str | None = None
    graph_node_id: str | None = None
    impacted_symbols: list[str] = Field(default_factory=list)
    changed_files: list[str] = Field(default_factory=list)
    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
    hydration_source: HydrationSource
    graph_expanded: bool = False


class CodeSnippet(BaseModel):
    record_id: str
    file_path: str
    symbol_name: str
    symbol_type: str | None = None
    language: str | None = None
    start_line: int
    end_line: int
    text: str
    score: float
    graph_expanded: bool = False


class DocExcerpt(BaseModel):
    record_id: str
    doc_path: str
    section_title: str
    section_level: int | None = None
    text: str
    score: float
    graph_expanded: bool = False


class RelatedCommit(BaseModel):
    record_id: str
    commit_sha: str
    summary: str
    impacted_symbols: list[str] = Field(default_factory=list)
    changed_files: list[str] = Field(default_factory=list)
    score: float
    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


class HydrationResult(BaseModel):
    hits: list[HydratedHit] = Field(default_factory=list)
    retrieval_count: int = 0
    hydrated_count: int = 0
    skipped_count: int = 0


class GraphExpansionResult(BaseModel):
    additional_hits: list[HydratedHit] = Field(default_factory=list)
    nodes_expanded: int = 0
    records_added: int = 0
    latency_ms: int = 0


class GroupedRetrievalContext(BaseModel):
    hits: list[HydratedHit] = Field(default_factory=list)
    code_snippets: list[CodeSnippet] = Field(default_factory=list)
    doc_excerpts: list[DocExcerpt] = Field(default_factory=list)
    related_commits: list[RelatedCommit] = Field(default_factory=list)
    retrieval_count: int = 0
    hydrated_count: int = 0
    skipped_count: int = 0


class RetrieveRequest(BaseModel):
    model_config = ConfigDict(extra="forbid")

    repo_id: str
    snapshot_id: str
    query_text: str
    filters: RetrieveFilters | None = None
    top_k: int | None = Field(default=None, ge=1, le=settings.RETRIEVE_MAX_TOP_K)
    score_threshold: float | None = Field(default=None, ge=0.0, le=1.0)
    graph_expand: bool = False
    request_id: str | None = None

    @field_validator("repo_id", "snapshot_id", "query_text")
    @classmethod
    def _required_non_empty(cls, value: str) -> str:
        text = value.strip()
        if not text:
            raise ValueError("must be non-empty after trim")
        return text


class RetrieveMetadata(BaseModel):
    request_id: str
    latency_ms: int = 0
    vector_latency_ms: int = 0
    hydrate_latency_ms: int = 0
    graph_latency_ms: int = 0
    retrieval_count: int = 0
    hydrated_count: int = 0
    skipped_count: int = 0
    graph_nodes_expanded: int = 0
    graph_records_added: int = 0


class RetrieveResponse(BaseModel):
    embedding_model: str
    embedding_dimension: int
    hits: list[HydratedHit] = Field(default_factory=list)
    code_snippets: list[CodeSnippet] = Field(default_factory=list)
    doc_excerpts: list[DocExcerpt] = Field(default_factory=list)
    related_commits: list[RelatedCommit] = Field(default_factory=list)
    metadata: RetrieveMetadata


class RetrieveErrorBody(BaseModel):
    code: str
    message: str
    request_id: str


class RetrieveErrorResponse(BaseModel):
    error: RetrieveErrorBody
