import asyncio

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.code_chunk import CodeChunkDocument
from app.models.upstream.commit_analysis import CommitAnalysisDocument
from app.models.upstream.doc_chunk import DocChunkDocument
from app.retrieval import hydrator
from app.retrieval.grouper import group_hydrated_hits
from app.retrieval.models import HydrationSource, VectorHit


def _vector_hit(
    record_id: str,
    *,
    source_type: str = "code",
    score: float = 0.9,
) -> VectorHit:
    return VectorHit(
        record_id=record_id,
        score=score,
        source_type=source_type,
        chunk_type="symbol" if source_type == "code" else "text",
        file_path="internal/config/config.go" if source_type == "code" else None,
        symbol_name="Load" if source_type == "code" else None,
        doc_path="docs/architecture.md" if source_type == "docs" else None,
        section_title="Overview" if source_type == "docs" else None,
        commit_sha="abc123" if source_type == "commit" else None,
    )


def _stored_code_record(record_id: str = "emb_chk_sym_config_go_Load") -> StoredEmbeddingRecord:
    record = EmbeddingRecord.from_code_chunk(
        record_id=record_id,
        repo_id="AD/example-repo",
        snapshot_id="snap_abc123",
        commit_sha="a1b2c3d4e5f6",
        text="func Load() {}",
        file_path="internal/config/config.go",
        symbol_name="Load",
        start_line=24,
        end_line=58,
        graph_node_id="sym_config_go_Load",
        language="go",
        symbol_type="function",
    )
    return StoredEmbeddingRecord(
        **record.model_dump(),
        embed_status=EmbedStatus.EMBEDDED,
        upstream_chunk_id="chk_sym_config_go_Load",
    )


def test_hydrate_uses_embedding_record_primary(monkeypatch):
    stored = _stored_code_record()

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

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

    result = asyncio.run(hydrator.hydrate_vector_hits([_vector_hit(stored.record_id)]))

    assert result.hydrated_count == 1
    assert result.skipped_count == 0
    assert result.hits[0].text == "func Load() {}"
    assert result.hits[0].hydration_source == HydrationSource.EMBEDDING_RECORD
    assert result.hits[0].language == "go"


def test_hydrate_falls_back_to_code_chunk(monkeypatch):
    record_id = "emb_chk_sym_config_go_Load"
    chunk = CodeChunkDocument(
        chunk_id="chk_sym_config_go_Load",
        repo_id="AD/example-repo",
        snapshot_id="snap_abc123",
        commit_sha="a1b2c3d4e5f6",
        file_path="internal/config/config.go",
        symbol_name="Load",
        text="func Load() { /* fallback */ }",
        start_line=24,
        end_line=58,
        graph_node_id="sym_config_go_Load",
        language="go",
        symbol_type="function",
    )

    async def _get_records_by_ids(record_ids):
        return {}

    async def _get_by_chunk_ids(chunk_ids):
        return {chunk.chunk_id: chunk}

    monkeypatch.setattr(hydrator.embedding_record_repo, "get_records_by_ids", _get_records_by_ids)
    monkeypatch.setattr(hydrator.code_chunk_repo, "get_by_chunk_ids", _get_by_chunk_ids)

    result = asyncio.run(hydrator.hydrate_vector_hits([_vector_hit(record_id)]))

    assert result.hydrated_count == 1
    assert result.hits[0].hydration_source == HydrationSource.CODE_CHUNK
    assert "fallback" in result.hits[0].text


def test_hydrate_skips_missing_records(monkeypatch):
    async def _get_records_by_ids(record_ids):
        return {}

    async def _get_by_chunk_ids(chunk_ids):
        return {}

    monkeypatch.setattr(hydrator.embedding_record_repo, "get_records_by_ids", _get_records_by_ids)
    monkeypatch.setattr(hydrator.code_chunk_repo, "get_by_chunk_ids", _get_by_chunk_ids)

    result = asyncio.run(
        hydrator.hydrate_vector_hits([_vector_hit("emb_missing_record")])
    )

    assert result.hydrated_count == 0
    assert result.skipped_count == 1


def test_hydrate_doc_and_commit_fallbacks(monkeypatch):
    doc_hit = _vector_hit("emb_chk_doc_overview", source_type="docs", score=0.8)
    commit_hit = _vector_hit("emb_analysis_b2c3d4", source_type="commit", score=0.7)
    doc_chunk = DocChunkDocument(
        chunk_id="chk_doc_overview",
        repo_id="AD/example-repo",
        doc_path="docs/architecture.md",
        section_title="Overview",
        chunk_text="Commit analysis is triggered when...",
        metadata={"section_level": 2},
    )
    analysis = CommitAnalysisDocument(
        analysis_id="analysis_b2c3d4",
        repo_id="AD/example-repo",
        commit_sha="b2c3d4e5f6a7",
        summary="Introduced Redis consumer refactor.",
        impacted_symbols=["ProcessDelta"],
        changed_files=["internal/service.go"],
    )

    async def _get_records_by_ids(record_ids):
        return {}

    async def _get_doc_chunks(chunk_ids):
        return {doc_chunk.chunk_id: doc_chunk}

    async def _get_analyses(analysis_ids):
        return {analysis.analysis_id: analysis}

    monkeypatch.setattr(hydrator.embedding_record_repo, "get_records_by_ids", _get_records_by_ids)
    monkeypatch.setattr(hydrator.doc_chunk_repo, "get_by_chunk_ids", _get_doc_chunks)
    monkeypatch.setattr(hydrator.commit_analysis_repo, "get_by_analysis_ids", _get_analyses)

    result = asyncio.run(hydrator.hydrate_vector_hits([doc_hit, commit_hit]))

    assert result.hydrated_count == 2
    assert {hit.source_type for hit in result.hits} == {"docs", "commit"}


def test_grouper_splits_hits_into_arrays():
    from app.retrieval.models import HydratedHit, HydrationResult

    hydration = HydrationResult(
        hits=[
            HydratedHit(
                record_id="emb_code",
                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",
                symbol_type="function",
                language="go",
                start_line=1,
                end_line=2,
                hydration_source=HydrationSource.EMBEDDING_RECORD,
            ),
            HydratedHit(
                record_id="emb_docs",
                source_type=SourceType.DOCS.value,
                score=0.8,
                chunk_type=ChunkType.TEXT.value,
                text="Overview text",
                doc_path="docs/architecture.md",
                section_title="Overview",
                section_level=2,
                hydration_source=HydrationSource.EMBEDDING_RECORD,
            ),
            HydratedHit(
                record_id="emb_commit",
                source_type=SourceType.COMMIT.value,
                score=0.7,
                chunk_type=ChunkType.SUMMARY.value,
                text="Refactored consumer",
                commit_sha="abc123",
                impacted_symbols=["Load"],
                changed_files=["config.go"],
                hydration_source=HydrationSource.COMMIT_ANALYSIS,
            ),
        ],
        retrieval_count=3,
        hydrated_count=3,
        skipped_count=0,
    )

    grouped = group_hydrated_hits(hydration)

    assert len(grouped.code_snippets) == 1
    assert len(grouped.doc_excerpts) == 1
    assert len(grouped.related_commits) == 1
    assert grouped.code_snippets[0].symbol_name == "Load"
    assert grouped.doc_excerpts[0].section_level == 2
    assert grouped.related_commits[0].impacted_symbols == ["Load"]


def test_grouper_dedupes_by_record_id_keeps_higher_score():
    from app.retrieval.models import HydratedHit, HydrationResult

    hydration = HydrationResult(
        hits=[
            HydratedHit(
                record_id="emb_dup",
                source_type=SourceType.CODE.value,
                score=0.7,
                chunk_type=ChunkType.SYMBOL.value,
                text="lower",
                file_path="a.go",
                symbol_name="A",
                start_line=1,
                end_line=2,
                hydration_source=HydrationSource.EMBEDDING_RECORD,
            ),
            HydratedHit(
                record_id="emb_dup",
                source_type=SourceType.CODE.value,
                score=0.95,
                chunk_type=ChunkType.SYMBOL.value,
                text="higher",
                file_path="a.go",
                symbol_name="A",
                start_line=1,
                end_line=2,
                hydration_source=HydrationSource.EMBEDDING_RECORD,
            ),
        ],
        retrieval_count=2,
        hydrated_count=2,
        skipped_count=0,
    )

    grouped = group_hydrated_hits(hydration)

    assert len(grouped.hits) == 1
    assert grouped.hits[0].score == 0.95
    assert grouped.hits[0].text == "higher"
