from __future__ import annotations

import pytest

from app.intent.intent_analyzer import IntentAnalyzer
from app.intent.intent_pipeline import IntentPipeline
from app.intent.tier2_classifier import Tier2Classifier, clear_tier2_cache
from app.models.intent import QueryIntent


@pytest.fixture(autouse=True)
def _clear_cache():
    clear_tier2_cache()
    yield
    clear_tier2_cache()


@pytest.mark.asyncio
async def test_high_confidence_skips_tier2():
    pipeline = IntentPipeline(tier2=Tier2Classifier(llm_client=_failing_client()))
    result = await pipeline.resolve_intent("Where is ProcessDelta defined?")
    assert result.tier == "tier1"
    assert result.type == QueryIntent.CODE_LOOKUP


@pytest.mark.asyncio
async def test_low_confidence_runs_tier2(monkeypatch):
    monkeypatch.setattr("app.intent.intent_pipeline.settings.INTENT_CLASSIFIER", "tier1+tier2")
    monkeypatch.setattr("app.intent.intent_pipeline.settings.TIER1_CONFIDENCE_THRESHOLD", 0.99)

    class FakeLLM:
        async def complete_json(self, messages, model=None):
            return {
                "intent": "historical",
                "confidence": 0.92,
                "source_types": ["commit", "code"],
                "rewritten_query": "Redis consumer refactor history",
                "needs_graph_expansion": False,
            }

    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]

    pipeline = IntentPipeline(tier2=Tier2Classifier(llm_client=llm))
    result = await pipeline.resolve_intent("hello")
    assert result.tier == "tier2"
    assert result.type == QueryIntent.HISTORICAL
    assert result.rewritten_query == "Redis consumer refactor history"


@pytest.mark.asyncio
async def test_intent_override_skips_tier2(monkeypatch):
    monkeypatch.setattr("app.intent.intent_pipeline.settings.INTENT_CLASSIFIER", "tier1+tier2")
    pipeline = IntentPipeline(tier2=Tier2Classifier(llm_client=_failing_client()))
    result = await pipeline.resolve_intent("hello", intent_override="documentation")
    assert result.type == QueryIntent.DOCUMENTATION
    assert result.confidence == 1.0


@pytest.mark.asyncio
async def test_invalid_tier2_falls_back_to_tier1(monkeypatch):
    monkeypatch.setattr("app.intent.intent_pipeline.settings.INTENT_CLASSIFIER", "tier1+tier2")
    monkeypatch.setattr("app.intent.intent_pipeline.settings.TIER1_CONFIDENCE_THRESHOLD", 0.99)

    class BadLLM:
        async def complete_json(self, messages, model=None):
            return {"intent": "not_valid", "confidence": 0.5}

    from app.generation.llm_client import LLMClient

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

    pipeline = IntentPipeline(tier2=Tier2Classifier(llm_client=llm))
    result = await pipeline.resolve_intent("hello")
    assert result.tier == "tier1"
    assert result.type == QueryIntent.GENERAL


def _failing_client():
    class Fail:
        async def complete_json(self, *args, **kwargs):
            raise RuntimeError("should not be called")

    from app.generation.llm_client import LLMClient

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