import asyncio
from datetime import datetime

import pytest

from app.models.embedding_record import EmbeddingRecord
from app.models.enums import ChunkType, EmbedStatus, SourceType
from app.models.stored_embedding_record import StoredEmbeddingRecord
from app.models.upstream.snapshot_graph import SnapshotGraphDocument, SnapshotGraphEdge, SnapshotGraphNode
from app.retrieval import graph_expander
from app.retrieval.models import HydratedHit, HydrationSource, VectorHit


def _sample_graph() -> SnapshotGraphDocument:
    return SnapshotGraphDocument(
        repo_id="AD/example-repo",
        snapshot_id="snap_abc123",
        nodes=[
            SnapshotGraphNode(id="sym_Load", kind="symbol", name="Load"),
            SnapshotGraphNode(id="chk_doc_overview", kind="doc_chunk", name="Overview"),
            SnapshotGraphNode(id="chk_sym_Load", kind="code_chunk", name="Load"),
        ],
        edges=[
            SnapshotGraphEdge(
                id="e1",
                kind="documents",
                source_id="chk_doc_overview",
                target_id="sym_Load",
            ),
            SnapshotGraphEdge(
                id="e2",
                kind="has_chunk",
                source_id="sym_Load",
                target_id="chk_sym_Load",
            ),
        ],
    )


def test_apply_historical_commit_boost_only_when_commit_filter_enabled():
    hits = [
        VectorHit(
            record_id="emb_commit",
            score=0.7,
            source_type="commit",
            chunk_type="summary",
        ),
        VectorHit(
            record_id="emb_code",
            score=0.8,
            source_type="code",
            chunk_type="symbol",
        ),
    ]

    boosted = graph_expander.apply_historical_commit_boost(hits, ["commit", "code"])

    assert boosted[0].score == pytest.approx(0.805)
    assert boosted[1].score == 0.8


def test_apply_historical_commit_boost_skipped_without_commit_filter():
    hits = [
        VectorHit(
            record_id="emb_commit",
            score=0.7,
            source_type="commit",
            chunk_type="summary",
        )
    ]

    assert graph_expander.apply_historical_commit_boost(hits, ["code"]) == hits


@pytest.mark.asyncio
async def test_expand_graph_context_adds_linked_doc_chunk(monkeypatch):
    doc_record = StoredEmbeddingRecord(
        **EmbeddingRecord.from_doc_chunk(
            record_id="emb_chk_doc_overview",
            repo_id="AD/example-repo",
            snapshot_id="snap_abc123",
            commit_sha="a1b2c3d4e5f6",
            text="Commit analysis is triggered when parser events fire.",
            doc_path="docs/architecture.md",
            section_title="Overview",
            graph_node_id="chk_doc_overview",
        ).model_dump(),
        embed_status=EmbedStatus.EMBEDDED,
        created_at=datetime.utcnow(),
        updated_at=datetime.utcnow(),
    )

    async def _get_merged(repo_id, snapshot_id):
        return _sample_graph()

    async def _get_records_by_ids(record_ids):
        return {doc_record.record_id: doc_record}

    monkeypatch.setattr(graph_expander.snapshot_graph_repo, "get_merged", _get_merged)
    from app.retrieval import hydrator

    monkeypatch.setattr(hydrator.embedding_record_repo, "get_records_by_ids", _get_records_by_ids)

    vector_hits = [
        VectorHit(
            record_id="emb_chk_sym_Load",
            score=0.9,
            source_type=SourceType.CODE.value,
            chunk_type=ChunkType.SYMBOL.value,
            graph_node_id="sym_Load",
        )
    ]
    hydrated_hits = [
        HydratedHit(
            record_id="emb_chk_sym_Load",
            source_type=SourceType.CODE.value,
            score=0.9,
            chunk_type=ChunkType.SYMBOL.value,
            text="func Load() {}",
            file_path="internal/config/config.go",
            symbol_name="Load",
            start_line=1,
            end_line=2,
            graph_node_id="sym_Load",
            hydration_source=HydrationSource.EMBEDDING_RECORD,
        )
    ]

    result = await graph_expander.expand_graph_context(
        vector_hits,
        hydrated_hits,
        repo_id="AD/example-repo",
        snapshot_id="snap_abc123",
    )

    assert result.records_added == 1
    assert result.nodes_expanded == 1
    assert result.additional_hits[0].record_id == "emb_chk_doc_overview"
    assert result.additional_hits[0].graph_expanded is True
    assert result.additional_hits[0].score == pytest.approx(0.765)


@pytest.mark.asyncio
async def test_expand_graph_context_returns_empty_when_graph_missing(monkeypatch):
    async def _get_merged(repo_id, snapshot_id):
        return None

    monkeypatch.setattr(graph_expander.snapshot_graph_repo, "get_merged", _get_merged)

    result = await graph_expander.expand_graph_context(
        [],
        [],
        repo_id="AD/example-repo",
        snapshot_id="snap_abc123",
    )

    assert result.records_added == 0
    assert result.nodes_expanded == 0
