import json

import httpx
import pytest

from app.core.config import settings
from app.intent.intent_analyzer import IntentAnalyzer
from app.models.intent import IntentResult, QueryIntent
from app.models.search import SearchOptions
from app.models.internal_retrieve import SourceType
from app.retrieval.retrieval_client import (
    RetrievalClient,
    RetrievalClientError,
    build_retrieve_request,
    resolve_retrieval_score_threshold,
)


@pytest.fixture
def intent():
    return IntentAnalyzer().analyze("Where is ProcessDelta defined?")


def test_build_retrieve_request_maps_intent(intent):
    request = build_retrieve_request(
        repo_id="ad/example.com",
        snapshot_id="snap_abc",
        question="Where is ProcessDelta defined?",
        intent=intent,
        request_id="req_test123",
    )
    assert request.repo_id == "ad/example.com"
    assert request.snapshot_id == "snap_abc"
    assert request.query_text == "Where is ProcessDelta defined?"
    assert request.filters is not None
    assert request.filters.source_types == [SourceType.CODE]
    assert request.top_k == 25
    assert request.graph_expand is True
    assert request.request_id == "req_test123"


def test_resolve_score_threshold_uses_historical_default():
    intent = IntentResult(
        type=QueryIntent.HISTORICAL,
        confidence=0.9,
        source_types_used=[SourceType.COMMIT, SourceType.CODE],
    )
    assert resolve_retrieval_score_threshold(intent) == settings.QUERY_HISTORICAL_SCORE_THRESHOLD


def test_resolve_score_threshold_honors_explicit_override():
    intent = IntentResult(type=QueryIntent.HISTORICAL, confidence=0.9)
    options = SearchOptions(score_threshold=0.8)
    assert resolve_retrieval_score_threshold(intent, options) == 0.8


def test_build_retrieve_request_uses_historical_score_threshold():
    intent = IntentResult(
        type=QueryIntent.HISTORICAL,
        confidence=0.9,
        source_types_used=[SourceType.COMMIT, SourceType.CODE],
    )
    request = build_retrieve_request(
        repo_id="ad/example.com",
        snapshot_id="snap_abc",
        question="What files changed in the latest commit?",
        intent=intent,
    )
    assert request.score_threshold == settings.QUERY_HISTORICAL_SCORE_THRESHOLD
    assert request.filters is not None
    assert SourceType.COMMIT in request.filters.source_types


@pytest.mark.asyncio
async def test_retrieval_client_mock_mode(monkeypatch):
    monkeypatch.setattr("app.core.config.settings.RETRIEVAL_MOCK_SCENARIO", "empty")
    client = RetrievalClient(mock=True)
    from app.models.internal_retrieve import RetrieveRequest

    response = await client.retrieve(
        RetrieveRequest(
            repo_id="ad/example.com",
            snapshot_id="snap_abc",
            query_text="test query",
        )
    )
    assert response.embedding_model == "text-embedding-3-small"
    assert response.embedding_dimension == 1536
    assert response.hits == []


@pytest.mark.asyncio
async def test_retrieval_client_http_success(intent):
    mock_response = {
        "embedding_model": "text-embedding-3-small",
        "embedding_dimension": 1536,
        "hits": [],
        "code_snippets": [],
        "doc_excerpts": [],
        "related_commits": [],
        "metadata": {"request_id": "req_test123", "latency_ms": 100},
    }

    def handler(request: httpx.Request) -> httpx.Response:
        assert request.url.path == "/internal/retrieve"
        body = json.loads(request.content)
        assert body["repo_id"] == "ad/example.com"
        assert body["filters"]["source_types"] == ["code"]
        return httpx.Response(200, json=mock_response)

    transport = httpx.MockTransport(handler)
    async with httpx.AsyncClient(transport=transport, base_url="http://embedding:6004") as http_client:
        client = RetrievalClient(base_url="http://embedding:6004", client=http_client)
        request = build_retrieve_request(
            repo_id="ad/example.com",
            snapshot_id="snap_abc",
            question="Where is ProcessDelta defined?",
            intent=intent,
            request_id="req_test123",
        )
        response = await client.retrieve(request)
        assert response.embedding_model == "text-embedding-3-small"


@pytest.mark.asyncio
async def test_retrieval_client_http_error():
    def handler(_request: httpx.Request) -> httpx.Response:
        return httpx.Response(503, json={"detail": "service unavailable"})

    transport = httpx.MockTransport(handler)
    async with httpx.AsyncClient(transport=transport, base_url="http://embedding:6004") as http_client:
        client = RetrievalClient(base_url="http://embedding:6004", client=http_client)
        from app.models.internal_retrieve import RetrieveRequest

        with pytest.raises(RetrievalClientError, match="503"):
            await client.retrieve(
                RetrieveRequest(
                    repo_id="ad/example.com",
                    snapshot_id="snap_abc",
                    query_text="test",
                    request_id="req_err",
                )
            )
