from __future__ import annotations

import json
import re

import pytest

from app.core.config import settings
from app.models.intent import IntentResult, QueryIntent
from app.models.internal_retrieve import CodeSnippet
from app.models.query import QueryRequest
from app.models.search import SearchMetadata, SearchResponse
from app.services.query_service import QueryService
from tests.unit.test_query_service import FakeOrchestrator


@pytest.fixture(autouse=True)
def disable_query_logs(monkeypatch):
    monkeypatch.setattr(settings, "QUERY_LOGS_ENABLED", False)


@pytest.mark.asyncio
async def test_stream_query_emits_start_token_done():
    search_response = SearchResponse(
        request_id="req_stream",
        repo_id="ad/example.com",
        snapshot_id="snap_1",
        question="Where is Foo?",
        intent=IntentResult(
            type=QueryIntent.CODE_LOOKUP,
            confidence=0.9,
            source_types_used=[],
        ),
        code_snippets=[
            CodeSnippet(
                record_id="emb_1",
                file_path="foo.go",
                symbol_name="Foo",
                text="func Foo() {}",
                score=0.9,
            )
        ],
        metadata=SearchMetadata(
            latency_ms=20,
            retrieval_count=1,
            embedding_model="text-embedding-3-small",
        ),
    )

    class FakeLLM:
        async def stream(self, messages, model=None):
            yield "Foo "
            yield "answer [1]"

    from app.generation.llm_client import LLMClient

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

    service = QueryService(orchestrator=FakeOrchestrator(search_response), llm_client=llm)
    events = []
    async for frame in service.stream_query(
        QueryRequest(repo_id="ad/example.com", question="Where is Foo?")
    ):
        events.append(frame)

    joined = "".join(events)
    assert "event: start" in joined
    assert "event: sources" in joined
    assert "event: token" in joined
    assert "event: done" in joined
    assert "Foo" in joined
    done_match = re.search(r"event: done\ndata: (.+?)\n\n", joined, re.DOTALL)
    assert done_match is not None
    done_payload = json.loads(done_match.group(1))
    assert done_payload["answer"] == "Foo answer [1]"


@pytest.mark.asyncio
async def test_stream_query_no_context_skips_llm(monkeypatch):
    monkeypatch.setattr(settings, "GENERAL_LLM_FALLBACK_ENABLED", False)
    search_response = SearchResponse(
        request_id="req_empty",
        repo_id="ad/example.com",
        snapshot_id="snap_1",
        question="unknown?",
        intent=IntentResult(
            type=QueryIntent.GENERAL,
            confidence=0.5,
            source_types_used=[],
        ),
        metadata=SearchMetadata(latency_ms=5, retrieval_count=0),
    )

    class FakeLLM:
        async def stream(self, messages, model=None):
            raise AssertionError("LLM should not be called")

    from app.generation.llm_client import LLMClient

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

    service = QueryService(orchestrator=FakeOrchestrator(search_response), llm_client=llm)
    events = []
    async for frame in service.stream_query(
        QueryRequest(repo_id="ad/example.com", question="unknown?")
    ):
        events.append(frame)

    joined = "".join(events)
    assert "event: start" in joined
    assert "event: token" in joined
    assert "event: done" in joined
    assert "event: sources" not in joined


@pytest.mark.asyncio
async def test_stream_query_empty_context_general_fallback():
    search_response = SearchResponse(
        request_id="req_empty",
        repo_id="ad/example.com",
        snapshot_id="snap_1",
        question="Write Python to sort a list",
        intent=IntentResult(
            type=QueryIntent.GENERAL,
            confidence=0.5,
            source_types_used=[],
        ),
        metadata=SearchMetadata(latency_ms=5, retrieval_count=0),
    )

    class FakeLLM:
        async def stream(self, messages, model=None):
            yield "sorted_list = "
            yield "sorted(items)"

    from app.generation.llm_client import LLMClient

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

    service = QueryService(orchestrator=FakeOrchestrator(search_response), llm_client=llm)
    events = []
    async for frame in service.stream_query(
        QueryRequest(repo_id="ad/example.com", question="Write Python to sort a list")
    ):
        events.append(frame)

    joined = "".join(events)
    assert '"answer_mode":"general"' in joined.replace(" ", "")
    assert "event: sources" not in joined
    assert "sorted_list" in joined
    done_match = re.search(r"event: done\ndata: (.+?)\n\n", joined, re.DOTALL)
    assert done_match is not None
    done_payload = json.loads(done_match.group(1))
    assert done_payload["answer_mode"] == "general"


@pytest.mark.asyncio
async def test_stream_query_refuses_unsafe_output_in_buffer_mode():
    search_response = SearchResponse(
        request_id="req_stream_unsafe",
        repo_id="ad/example.com",
        snapshot_id="snap_1",
        question="How do I exploit this?",
        intent=IntentResult(
            type=QueryIntent.CODE_LOOKUP,
            confidence=0.9,
            source_types_used=[],
        ),
        code_snippets=[
            CodeSnippet(
                record_id="emb_1",
                file_path="auth.go",
                symbol_name="Auth",
                text="func Auth() {}",
                score=0.9,
            )
        ],
        metadata=SearchMetadata(
            latency_ms=20,
            retrieval_count=1,
            embedding_model="text-embedding-3-small",
        ),
    )

    class FakeLLM:
        async def stream(self, messages, model=None):
            yield "Run this curl command to exploit the endpoint and bypass auth."

    from app.generation.llm_client import LLMClient
    from app.security.guardrails import SAFE_REFUSAL_MESSAGE

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

    service = QueryService(orchestrator=FakeOrchestrator(search_response), llm_client=llm)
    events = []
    async for frame in service.stream_query(
        QueryRequest(repo_id="ad/example.com", question="How do I exploit this?")
    ):
        events.append(frame)

    joined = "".join(events)
    assert SAFE_REFUSAL_MESSAGE in joined
    assert "exploit the endpoint" not in joined
