from __future__ import annotations

import json
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock

import pytest
from openai import APITimeoutError, AuthenticationError, RateLimitError

from app.generation.llm_client import (
    LLMAuthError,
    LLMClient,
    LLMMalformedOutputError,
    LLMRateLimitError,
    LLMTimeoutError,
)


def _completion_response(text: str, *, prompt_tokens: int = 10, completion_tokens: int = 5):
    return SimpleNamespace(
        choices=[SimpleNamespace(message=SimpleNamespace(content=text))],
        usage=SimpleNamespace(prompt_tokens=prompt_tokens, completion_tokens=completion_tokens),
    )


@pytest.fixture
def mock_openai():
    client = MagicMock()
    client.chat = MagicMock()
    client.chat.completions = MagicMock()
    client.chat.completions.create = AsyncMock()
    return client


@pytest.mark.asyncio
async def test_complete_success(mock_openai):
    mock_openai.chat.completions.create.return_value = _completion_response("Hello world")
    llm = LLMClient(client=mock_openai, api_key="sk-test")

    result = await llm.complete([{"role": "user", "content": "Hi"}])

    assert result.text == "Hello world"
    assert result.prompt_tokens == 10
    assert result.completion_tokens == 5
    assert result.latency_ms >= 0


@pytest.mark.asyncio
async def test_complete_auth_error(mock_openai):
    mock_openai.chat.completions.create.side_effect = AuthenticationError(
        "invalid key", response=MagicMock(status_code=401), body=None
    )
    llm = LLMClient(client=mock_openai, api_key="sk-bad")

    with pytest.raises(LLMAuthError):
        await llm.complete([{"role": "user", "content": "Hi"}])


@pytest.mark.asyncio
async def test_complete_rate_limit(mock_openai):
    mock_openai.chat.completions.create.side_effect = RateLimitError(
        "rate limited", response=MagicMock(status_code=429), body=None
    )
    llm = LLMClient(client=mock_openai, api_key="sk-test", max_retries=1)

    with pytest.raises(LLMRateLimitError):
        await llm.complete([{"role": "user", "content": "Hi"}])


@pytest.mark.asyncio
async def test_complete_timeout(mock_openai):
    mock_openai.chat.completions.create.side_effect = APITimeoutError(request=MagicMock())
    llm = LLMClient(client=mock_openai, api_key="sk-test", max_retries=1)

    with pytest.raises(LLMTimeoutError):
        await llm.complete([{"role": "user", "content": "Hi"}])


@pytest.mark.asyncio
async def test_complete_json_parses(mock_openai):
    payload = {"intent": "code_lookup", "confidence": 0.9}
    mock_openai.chat.completions.create.return_value = _completion_response(json.dumps(payload))
    llm = LLMClient(client=mock_openai, api_key="sk-test")

    data = await llm.complete_json([{"role": "user", "content": "classify"}])
    assert data["intent"] == "code_lookup"


@pytest.mark.asyncio
async def test_complete_json_malformed(mock_openai):
    mock_openai.chat.completions.create.return_value = _completion_response("not json")
    llm = LLMClient(client=mock_openai, api_key="sk-test")

    with pytest.raises(LLMMalformedOutputError):
        await llm.complete_json([{"role": "user", "content": "classify"}])


@pytest.mark.asyncio
async def test_stream_yields_deltas(mock_openai):
    async def fake_stream():
        yield SimpleNamespace(choices=[SimpleNamespace(delta=SimpleNamespace(content="Hel"))])
        yield SimpleNamespace(choices=[SimpleNamespace(delta=SimpleNamespace(content="lo"))])

    mock_openai.chat.completions.create.return_value = fake_stream()
    llm = LLMClient(client=mock_openai, api_key="sk-test")

    parts = []
    async for token in llm.stream([{"role": "user", "content": "Hi"}]):
        parts.append(token)

    assert "".join(parts) == "Hello"
