from __future__ import annotations

from unittest.mock import AsyncMock

import pytest

from app.models.intent import IntentResult, QueryIntent
from app.models.internal_retrieve import (
    HydratedHit,
    HydrationSource,
    RetrieveMetadata,
    RetrieveResponse,
    SourceType,
    ChunkType,
)
from app.models.search import SearchRequest
from app.services.rag_orchestrator import RagOrchestrator, _needs_weak_retrieval_fallback


def _hit(record_id: str, score: float, source_type: SourceType = SourceType.CODE) -> HydratedHit:
    return HydratedHit(
        record_id=record_id,
        source_type=source_type,
        score=score,
        chunk_type=ChunkType.SYMBOL,
        text="sample",
        hydration_source=HydrationSource.EMBEDDING_RECORD,
    )


def _response(*hits: HydratedHit) -> RetrieveResponse:
    return RetrieveResponse(
        embedding_model="test-model",
        embedding_dimension=1536,
        hits=list(hits),
        metadata=RetrieveMetadata(retrieval_count=len(hits)),
    )


@pytest.mark.parametrize(
    ("intent", "hits", "top_score", "expected"),
    [
        (QueryIntent.ARCHITECTURE, [], None, True),
        (QueryIntent.ARCHITECTURE, [_hit("a", 0.15)], 0.15, True),
        (QueryIntent.ARCHITECTURE, [_hit("a", 0.25)], 0.25, False),
        (QueryIntent.GENERAL, [_hit("a", 0.1)], 0.1, True),
        (QueryIntent.CODE_LOOKUP, [], None, False),
    ],
)
def test_needs_weak_retrieval_fallback(intent, hits, top_score, expected):
    response = _response(*hits)
    intent_result = IntentResult(type=intent, confidence=0.9)
    assert _needs_weak_retrieval_fallback(intent_result, response, top_score) is expected


@pytest.mark.asyncio
async def test_retrieve_context_applies_docs_only_fallback(monkeypatch):
    weak = _response(_hit("code_1", 0.12))
    docs = _response(_hit("doc_1", 0.35, SourceType.DOCS))
    broad = _response(_hit("doc_2", 0.4, SourceType.DOCS))

    client = AsyncMock()
    client.retrieve = AsyncMock(side_effect=[weak, docs, broad])

    orchestrator = RagOrchestrator(retrieval_client=client)
    orchestrator.ensure_repo_resolved = AsyncMock(return_value="ad/example")
    orchestrator.resolve_snapshot_id = AsyncMock(return_value="snap_1")
    orchestrator.resolve_intent = AsyncMock(
        return_value=IntentResult(type=QueryIntent.ARCHITECTURE, confidence=0.9)
    )

    body = SearchRequest(repo_id="ad/example", question="What are the main features?")
    context = await orchestrator.retrieve_context(body)

    assert context.retrieval_fallback_used is True
    assert context.retrieval_fallback_pass == "docs_only"
    assert context.retrieve_response.hits[0].record_id == "doc_1"
    assert client.retrieve.await_count == 2


@pytest.mark.asyncio
async def test_retrieve_context_applies_broad_fallback_when_docs_weak(monkeypatch):
    weak = _response()
    docs = _response(_hit("doc_1", 0.05, SourceType.DOCS))
    broad = _response(_hit("summary_1", 0.22, SourceType.DOCS))

    client = AsyncMock()
    client.retrieve = AsyncMock(side_effect=[weak, docs, broad])

    orchestrator = RagOrchestrator(retrieval_client=client)
    orchestrator.ensure_repo_resolved = AsyncMock(return_value="ad/example")
    orchestrator.resolve_snapshot_id = AsyncMock(return_value="snap_1")
    orchestrator.resolve_intent = AsyncMock(
        return_value=IntentResult(type=QueryIntent.GENERAL, confidence=0.8)
    )

    body = SearchRequest(repo_id="ad/example", question="What does this repo do?")
    context = await orchestrator.retrieve_context(body)

    assert context.retrieval_fallback_used is True
    assert context.retrieval_fallback_pass == "broad"
    assert context.retrieve_response.hits[0].record_id == "summary_1"
    assert client.retrieve.await_count == 3
