import asyncio
from datetime import datetime

import pytest
from fastapi.testclient import TestClient

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.retrieval.errors import SnapshotNotFoundError
from app.retrieval.models import (
    CodeSnippet,
    DocExcerpt,
    HydratedHit,
    HydrationSource,
    RelatedCommit,
    RetrieveMetadata,
    RetrieveResponse,
)
from app.api.routes import internal_retrieve
from app.retrieval import retrieval_service


@pytest.fixture(autouse=True)
def _skip_startup_bootstrap(monkeypatch):
    monkeypatch.setattr("app.main.bootstrap_qdrant", lambda *_args, **_kwargs: "adpilot_embeddings")


@pytest.fixture
def client():
    from app.main import app

    return TestClient(app)


def _sample_response() -> RetrieveResponse:
    return RetrieveResponse(
        embedding_model="text-embedding-3-small",
        embedding_dimension=1536,
        hits=[
            HydratedHit(
                record_id="emb_chk_sym_config_go_Load",
                source_type=SourceType.CODE.value,
                score=0.89,
                chunk_type=ChunkType.SYMBOL.value,
                text="func Load() {}",
                file_path="internal/config/config.go",
                symbol_name="Load",
                symbol_type="function",
                language="go",
                start_line=24,
                end_line=58,
                hydration_source=HydrationSource.EMBEDDING_RECORD,
            )
        ],
        code_snippets=[
            CodeSnippet(
                record_id="emb_chk_sym_config_go_Load",
                file_path="internal/config/config.go",
                symbol_name="Load",
                symbol_type="function",
                language="go",
                start_line=24,
                end_line=58,
                text="func Load() {}",
                score=0.89,
            )
        ],
        doc_excerpts=[],
        related_commits=[],
        metadata=RetrieveMetadata(
            request_id="req_test123",
            latency_ms=120,
            vector_latency_ms=80,
            hydrate_latency_ms=40,
            retrieval_count=1,
            hydrated_count=1,
            skipped_count=0,
        ),
    )


def test_internal_retrieve_returns_hydrated_response(client, monkeypatch):
    async def _retrieve(_request):
        return _sample_response()

    monkeypatch.setattr(internal_retrieve, "retrieve", _retrieve)

    response = client.post(
        "/internal/retrieve",
        json={
            "repo_id": "AD/example-repo",
            "snapshot_id": "snap_abc123",
            "query_text": "where is Load defined?",
            "request_id": "req_test123",
        },
    )

    assert response.status_code == 200
    body = response.json()
    assert body["embedding_model"] == "text-embedding-3-small"
    assert len(body["code_snippets"]) == 1
    assert body["metadata"]["request_id"] == "req_test123"
    assert body["metadata"]["hydrated_count"] == 1


def test_internal_retrieve_empty_query_returns_validation_error(client):
    response = client.post(
        "/internal/retrieve",
        json={
            "repo_id": "AD/example-repo",
            "snapshot_id": "snap_abc123",
            "query_text": "   ",
        },
    )

    assert response.status_code == 422
    body = response.json()
    assert body["error"]["code"] == "validation_error"


def test_internal_retrieve_rejects_unknown_fields(client):
    response = client.post(
        "/internal/retrieve",
        json={
            "repo_id": "AD/example-repo",
            "snapshot_id": "snap_abc123",
            "query_text": "hello",
            "unexpected": True,
        },
    )

    assert response.status_code == 422
    assert response.json()["error"]["code"] == "validation_error"


def test_internal_retrieve_snapshot_not_found(client, monkeypatch):
    async def _retrieve(_request):
        raise SnapshotNotFoundError("AD/example-repo", "snap_missing")

    monkeypatch.setattr(internal_retrieve, "retrieve", _retrieve)

    response = client.post(
        "/internal/retrieve",
        json={
            "repo_id": "AD/example-repo",
            "snapshot_id": "snap_missing",
            "query_text": "hello",
            "request_id": "req_missing",
        },
    )

    assert response.status_code == 404
    body = response.json()
    assert body["error"]["code"] == "snapshot_not_found"
    assert body["error"]["request_id"] == "req_missing"


def test_retrieval_service_end_to_end_with_mocks(monkeypatch):
    stored = StoredEmbeddingRecord(
        **EmbeddingRecord.from_code_chunk(
            record_id="emb_chk_sym_config_go_Load",
            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",
        ).model_dump(),
        embed_status=EmbedStatus.EMBEDDED,
        created_at=datetime.utcnow(),
        updated_at=datetime.utcnow(),
    )

    from app.retrieval.models import RetrieveRequest, VectorHit, VectorSearchResult

    vector_result = VectorSearchResult(
        hits=[
            VectorHit(
                record_id=stored.record_id,
                score=0.91,
                source_type=SourceType.CODE.value,
                chunk_type=ChunkType.SYMBOL.value,
                file_path=stored.file_path,
                symbol_name=stored.symbol_name,
            )
        ],
        embedding_model="text-embedding-3-small",
        embedding_dimension=1536,
        vector_latency_ms=12,
    )

    async def _count_records_by_snapshot(repo_id, snapshot_id, **kwargs):
        return 1

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

    async def _search_by_query_text(*_args, **_kwargs):
        return vector_result

    monkeypatch.setattr(
        retrieval_service.embedding_record_repo,
        "count_records_by_snapshot",
        _count_records_by_snapshot,
    )
    monkeypatch.setattr(
        retrieval_service,
        "search_by_query_text",
        _search_by_query_text,
    )
    from app.retrieval import hydrator

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

    request = RetrieveRequest(
        repo_id="AD/example-repo",
        snapshot_id="snap_abc123",
        query_text="where is Load defined?",
        request_id="req_e2e",
    )
    response = asyncio.run(retrieval_service.retrieve(request))

    assert response.metadata.request_id == "req_e2e"
    assert len(response.code_snippets) == 1
    assert response.code_snippets[0].text == "func Load() {}"
    assert response.metadata.retrieval_count == 1
