from __future__ import annotations

import time
import uuid

from pymongo.errors import PyMongoError

from app.core.logger import logger
from app.repositories import embedding_record_repo
from app.retrieval.errors import (
    EmbeddingRetrievalError,
    MongoRetrievalError,
    QdrantRetrievalError,
    SnapshotNotFoundError,
)
from app.retrieval.graph_expander import apply_historical_commit_boost, expand_graph_context
from app.retrieval.grouper import group_hydrated_hits
from app.retrieval.hydrator import hydrate_vector_hits
from app.retrieval.models import RetrieveMetadata, RetrieveRequest, RetrieveResponse
from app.retrieval.vector_search import search_by_query_text
from app.services.embedding_service import EmbeddingError
from app.qdrant.vector_store import QdrantVectorStoreError


def _resolve_request_id(request: RetrieveRequest) -> str:
    if request.request_id and request.request_id.strip():
        return request.request_id.strip()
    return f"req_{uuid.uuid4().hex[:12]}"


async def retrieve(request: RetrieveRequest) -> RetrieveResponse:
    """
    Orchestrate vector search, Mongo hydration, and grouped response shaping.

    Graph expansion is applied when ``graph_expand`` is true (TW-102).
    """
    request_id = _resolve_request_id(request)
    started = time.perf_counter()

    indexed_count = await embedding_record_repo.count_records_by_snapshot(
        request.repo_id,
        request.snapshot_id,
    )
    if indexed_count == 0:
        raise SnapshotNotFoundError(request.repo_id, request.snapshot_id)

    source_types = (
        list(request.filters.source_types)
        if request.filters and request.filters.source_types
        else None
    )

    vector_started = time.perf_counter()
    try:
        vector_result = await search_by_query_text(
            request.query_text,
            repo_id=request.repo_id,
            snapshot_id=request.snapshot_id,
            source_types=source_types,
            top_k=request.top_k,
            score_threshold=request.score_threshold,
        )
    except EmbeddingError as exc:
        raise EmbeddingRetrievalError(str(exc)) from exc
    except QdrantVectorStoreError as exc:
        raise QdrantRetrievalError(str(exc)) from exc

    vector_result = vector_result.model_copy(
        update={
            "hits": apply_historical_commit_boost(vector_result.hits, source_types),
        }
    )
    vector_latency_ms = int((time.perf_counter() - vector_started) * 1000)

    hydrate_started = time.perf_counter()
    try:
        hydration = await hydrate_vector_hits(vector_result.hits)
    except PyMongoError as exc:
        raise MongoRetrievalError(str(exc)) from exc

    hydrate_latency_ms = int((time.perf_counter() - hydrate_started) * 1000)

    graph_latency_ms = 0
    graph_nodes_expanded = 0
    graph_records_added = 0
    combined_hits = list(hydration.hits)

    if request.graph_expand:
        try:
            expansion = await expand_graph_context(
                vector_result.hits,
                hydration.hits,
                repo_id=request.repo_id,
                snapshot_id=request.snapshot_id,
            )
        except PyMongoError as exc:
            raise MongoRetrievalError(str(exc)) from exc

        graph_latency_ms = expansion.latency_ms
        graph_nodes_expanded = expansion.nodes_expanded
        graph_records_added = expansion.records_added
        combined_hits.extend(expansion.additional_hits)

    hydration = hydration.model_copy(
        update={
            "hits": combined_hits,
            "hydrated_count": len(combined_hits),
        }
    )
    grouped = group_hydrated_hits(hydration)

    latency_ms = int((time.perf_counter() - started) * 1000)

    logger.info(
        "Retrieve completed request_id={request_id} repo_id={repo_id} snapshot_id={snapshot_id} "
        "graph_expand={graph_expand} retrieval_count={retrieval_count} "
        "hydrated_count={hydrated_count} skipped_count={skipped_count} "
        "latency_ms={latency_ms} vector_latency_ms={vector_latency_ms} "
        "hydrate_latency_ms={hydrate_latency_ms} graph_latency_ms={graph_latency_ms}",
        request_id=request_id,
        repo_id=request.repo_id,
        snapshot_id=request.snapshot_id,
        graph_expand=request.graph_expand,
        retrieval_count=grouped.retrieval_count,
        hydrated_count=grouped.hydrated_count,
        skipped_count=grouped.skipped_count,
        latency_ms=latency_ms,
        vector_latency_ms=vector_latency_ms,
        hydrate_latency_ms=hydrate_latency_ms,
        graph_latency_ms=graph_latency_ms,
    )

    return RetrieveResponse(
        embedding_model=vector_result.embedding_model,
        embedding_dimension=vector_result.embedding_dimension,
        hits=grouped.hits,
        code_snippets=grouped.code_snippets,
        doc_excerpts=grouped.doc_excerpts,
        related_commits=grouped.related_commits,
        metadata=RetrieveMetadata(
            request_id=request_id,
            latency_ms=latency_ms,
            vector_latency_ms=vector_latency_ms,
            hydrate_latency_ms=hydrate_latency_ms,
            graph_latency_ms=graph_latency_ms,
            retrieval_count=grouped.retrieval_count,
            hydrated_count=grouped.hydrated_count,
            skipped_count=grouped.skipped_count,
            graph_nodes_expanded=graph_nodes_expanded,
            graph_records_added=graph_records_added,
        ),
    )
