import pytest

from app.intent.tier2_classifier import Tier2Classifier
from app.models.intent import QueryIntent


@pytest.mark.asyncio
async def test_tier2_parse_valid_payload():
    class FakeLLM:
        async def complete_json(self, messages, model=None):
            return {
                "intent": "code_lookup",
                "confidence": 0.91,
                "source_types": ["code"],
                "rewritten_query": "ProcessDelta definition location",
                "needs_graph_expansion": True,
            }

    from app.generation.llm_client import LLMClient

    llm = LLMClient(client=object())  # type: ignore[arg-type]
    llm.complete_json = FakeLLM().complete_json  # type: ignore[method-assign]

    result = await Tier2Classifier(llm_client=llm).classify("where is ProcessDelta?")
    assert result is not None
    assert result.type == QueryIntent.CODE_LOOKUP
    assert result.tier == "tier2"
    assert result.needs_graph_expansion is True


@pytest.mark.asyncio
async def test_tier2_invalid_payload_returns_none():
    class FakeLLM:
        async def complete_json(self, messages, model=None):
            return {"intent": "bogus", "confidence": 0.5}

    from app.generation.llm_client import LLMClient

    llm = LLMClient(client=object())  # type: ignore[arg-type]
    llm.complete_json = FakeLLM().complete_json  # type: ignore[method-assign]

    result = await Tier2Classifier(llm_client=llm).classify("ambiguous question")
    assert result is None
