from __future__ import annotations

import time

from app.core.config import settings
from app.core.logger import logger
from app.models.embedding_record import (
    embedding_record_id_for_analysis,
    embedding_record_id_for_chunk,
)
from app.models.enums import SourceType
from app.models.upstream.snapshot_graph import SnapshotGraphDocument, SnapshotGraphNode
from app.repositories import snapshot_graph_repo
from app.retrieval.hydrator import hydrate_vector_hits
from app.retrieval.models import GraphExpansionResult, HydratedHit, VectorHit

OUTGOING_EDGE_KINDS = frozenset({"has_chunk", "documents"})
INCOMING_EDGE_KINDS = frozenset({"documents", "related_context"})
CHUNK_NODE_KINDS = frozenset({"code_chunk", "doc_chunk"})
COMMIT_NODE_KINDS = frozenset({"commit"})


def apply_historical_commit_boost(
    hits: list[VectorHit],
    source_types: list[str] | None,
) -> list[VectorHit]:
    """Boost commit vector scores when historical retrieval includes commits."""
    if not source_types or SourceType.COMMIT.value not in source_types:
        return hits

    boost = settings.RETRIEVE_COMMIT_SCORE_BOOST
    boosted: list[VectorHit] = []
    for hit in hits:
        if hit.source_type != SourceType.COMMIT.value:
            boosted.append(hit)
            continue
        boosted.append(
            hit.model_copy(update={"score": min(1.0, hit.score * boost)})
        )
    return boosted


def _record_id_for_node(node: SnapshotGraphNode) -> str | None:
    if node.kind in CHUNK_NODE_KINDS:
        return embedding_record_id_for_chunk(node.id)
    if node.kind in COMMIT_NODE_KINDS:
        return embedding_record_id_for_analysis(node.id)
    return None


def _source_type_for_node(node: SnapshotGraphNode) -> str | None:
    if node.kind == "code_chunk":
        return SourceType.CODE.value
    if node.kind == "doc_chunk":
        return SourceType.DOCS.value
    if node.kind == "commit":
        return SourceType.COMMIT.value
    return None


def _chunk_type_for_source(source_type: str) -> str:
    if source_type == SourceType.CODE.value:
        return "symbol"
    if source_type == SourceType.DOCS.value:
        return "text"
    return "summary"


def _neighbor_node_ids(seed_node_id: str, graph: SnapshotGraphDocument) -> list[str]:
    neighbors: list[str] = []
    seen: set[str] = set()

    for edge in graph.edges:
        if edge.source_id == seed_node_id and edge.kind in OUTGOING_EDGE_KINDS:
            candidate = edge.target_id
        elif edge.target_id == seed_node_id and edge.kind in INCOMING_EDGE_KINDS:
            candidate = edge.source_id
        else:
            continue

        if candidate in seen:
            continue
        seen.add(candidate)
        neighbors.append(candidate)

    return neighbors


def _score_for_seed(
    graph_node_id: str,
    hydrated_by_graph_node: dict[str, HydratedHit],
    vector_by_graph_node: dict[str, VectorHit],
) -> float:
    if graph_node_id in hydrated_by_graph_node:
        return hydrated_by_graph_node[graph_node_id].score
    if graph_node_id in vector_by_graph_node:
        return vector_by_graph_node[graph_node_id].score
    return 0.0


async def expand_graph_context(
    vector_hits: list[VectorHit],
    hydrated_hits: list[HydratedHit],
    *,
    repo_id: str,
    snapshot_id: str,
) -> GraphExpansionResult:
    """Traverse snapshot_graphs one hop from vector/hydrated seeds and hydrate neighbors."""
    started = time.perf_counter()
    graph = await snapshot_graph_repo.get_merged(repo_id, snapshot_id)
    if graph is None:
        return GraphExpansionResult(latency_ms=0)

    nodes_by_id = {node.id: node for node in graph.nodes}
    existing_record_ids = {hit.record_id for hit in hydrated_hits}
    hydrated_by_graph_node = {
        hit.graph_node_id: hit for hit in hydrated_hits if hit.graph_node_id
    }
    vector_by_graph_node = {
        hit.graph_node_id: hit for hit in vector_hits if hit.graph_node_id
    }

    seed_node_ids: list[str] = []
    seen_seeds: set[str] = set()
    for hit in hydrated_hits:
        if hit.graph_node_id and hit.graph_node_id not in seen_seeds:
            seen_seeds.add(hit.graph_node_id)
            seed_node_ids.append(hit.graph_node_id)
    for hit in vector_hits:
        if hit.graph_node_id and hit.graph_node_id not in seen_seeds:
            seen_seeds.add(hit.graph_node_id)
            seed_node_ids.append(hit.graph_node_id)

    expansion_vectors: list[VectorHit] = []
    nodes_expanded = 0
    max_nodes = settings.RETRIEVE_GRAPH_EXPAND_MAX_NODES
    decay = settings.RETRIEVE_GRAPH_SCORE_DECAY

    for seed_node_id in seed_node_ids:
        if nodes_expanded >= max_nodes:
            break

        parent_score = _score_for_seed(
            seed_node_id,
            hydrated_by_graph_node,
            vector_by_graph_node,
        )
        if parent_score <= 0:
            continue

        for neighbor_id in _neighbor_node_ids(seed_node_id, graph):
            if nodes_expanded >= max_nodes:
                break

            node = nodes_by_id.get(neighbor_id)
            if node is None:
                continue

            record_id = _record_id_for_node(node)
            source_type = _source_type_for_node(node)
            if record_id is None or source_type is None:
                continue
            if record_id in existing_record_ids:
                continue

            existing_record_ids.add(record_id)
            nodes_expanded += 1
            expansion_vectors.append(
                VectorHit(
                    record_id=record_id,
                    score=round(parent_score * decay, 6),
                    source_type=source_type,
                    chunk_type=_chunk_type_for_source(source_type),
                    graph_node_id=neighbor_id,
                )
            )

    if not expansion_vectors:
        latency_ms = int((time.perf_counter() - started) * 1000)
        return GraphExpansionResult(
            nodes_expanded=nodes_expanded,
            records_added=0,
            latency_ms=latency_ms,
        )

    hydration = await hydrate_vector_hits(expansion_vectors)
    expanded_hits = [
        hit.model_copy(update={"graph_expanded": True}) for hit in hydration.hits
    ]

    latency_ms = int((time.perf_counter() - started) * 1000)
    logger.info(
        "Graph expansion completed repo_id={repo_id} snapshot_id={snapshot_id} "
        "nodes_expanded={nodes_expanded} records_added={records_added} latency_ms={latency_ms}",
        repo_id=repo_id,
        snapshot_id=snapshot_id,
        nodes_expanded=nodes_expanded,
        records_added=len(expanded_hits),
        latency_ms=latency_ms,
    )

    return GraphExpansionResult(
        additional_hits=expanded_hits,
        nodes_expanded=nodes_expanded,
        records_added=len(expanded_hits),
        latency_ms=latency_ms,
    )
