import asyncio
from types import SimpleNamespace

import pytest

from app.services.embedding_service import (
    EmbeddingDimensionError,
    EmbeddingError,
    EmbeddingService,
    OpenAIEmbedding,
    reset_embedding_service,
)


def _vector(dimension: int = 1536, fill: float = 0.1) -> list[float]:
    return [fill] * dimension


class _FakeEmbeddingsAPI:
    def __init__(self, vectors: list[list[float]] | None = None, dimension: int = 1536):
        self.vectors = vectors
        self.dimension = dimension
        self.calls: list[dict] = []

    async def create(self, *, model, input):
        self.calls.append({"model": model, "input": input})
        if self.vectors is not None:
            data = [
                SimpleNamespace(index=index, embedding=vector)
                for index, vector in enumerate(self.vectors)
            ]
        else:
            data = [
                SimpleNamespace(index=index, embedding=_vector(self.dimension))
                for index, _ in enumerate(input)
            ]
        return SimpleNamespace(
            data=data,
            usage=SimpleNamespace(total_tokens=42),
        )


class _FakeOpenAIClient:
    def __init__(self, api):
        self.embeddings = api


@pytest.fixture(autouse=True)
def _reset_service():
    reset_embedding_service()
    yield
    reset_embedding_service()


@pytest.mark.asyncio
async def test_openai_embed_sorts_by_response_index():
    api = _FakeEmbeddingsAPI(
        vectors=[
            _vector(fill=0.1),
            _vector(fill=0.2),
        ]
    )
    provider = OpenAIEmbedding(api_key="test-key", batch_size=10)
    provider.client = _FakeOpenAIClient(api)

    result = await provider.embed(["first", "second"])

    assert len(result.embeddings) == 2
    assert result.embeddings[0][0] == 0.1
    assert result.total_tokens == 42


@pytest.mark.asyncio
async def test_openai_embed_rejects_wrong_dimension():
    api = _FakeEmbeddingsAPI(vectors=[[0.1, 0.2]])
    provider = OpenAIEmbedding(api_key="test-key", expected_dimension=1536)
    provider.client = _FakeOpenAIClient(api)

    with pytest.raises(EmbeddingDimensionError):
        await provider.embed(["text"])


@pytest.mark.asyncio
async def test_openai_embed_rejects_empty_text():
    provider = OpenAIEmbedding(api_key="test-key")
    provider.client = _FakeOpenAIClient(_FakeEmbeddingsAPI())

    with pytest.raises(EmbeddingError, match="empty text"):
        await provider.embed(["valid", "   "])


@pytest.mark.asyncio
async def test_embedding_service_truncates_long_text(monkeypatch):
    api = _FakeEmbeddingsAPI()
    provider = OpenAIEmbedding(api_key="test-key")
    provider.client = _FakeOpenAIClient(api)
    service = EmbeddingService(provider=provider)

    long_text = "word " * 20_000
    vectors = await service.embed_texts([long_text])

    assert len(vectors) == 1
    assert len(api.calls) == 1
    assert len(api.calls[0]["input"][0]) < len(long_text)


@pytest.mark.asyncio
async def test_embedding_service_rejects_blank_text():
    provider = OpenAIEmbedding(api_key="test-key")
    provider.client = _FakeOpenAIClient(_FakeEmbeddingsAPI())
    service = EmbeddingService(provider=provider)

    with pytest.raises(EmbeddingError, match="empty"):
        await service.embed_texts(["   "])


@pytest.mark.asyncio
async def test_embed_batch_respects_count_cap():
    api = _FakeEmbeddingsAPI()
    provider = OpenAIEmbedding(
        api_key="test-key", batch_size=2, max_tokens_per_request=1_000_000
    )
    provider.client = _FakeOpenAIClient(api)

    texts = ["a", "b", "c", "d", "e"]
    vectors = await provider.embed_batch(texts)

    assert len(vectors) == len(texts)
    # batch_size=2 -> requests of 2, 2, 1
    assert [len(call["input"]) for call in api.calls] == [2, 2, 1]
    flat = [t for call in api.calls for t in call["input"]]
    assert flat == texts  # order preserved across requests


@pytest.mark.asyncio
async def test_embed_batch_splits_on_token_budget():
    # Regression: a run of chunks whose combined tokens exceed OpenAI's
    # per-request cap must be split across requests, not sent as one 400-ing call.
    api = _FakeEmbeddingsAPI()
    budget = 4  # tiny budget so short words force a split despite large batch_size
    provider = OpenAIEmbedding(
        api_key="test-key", batch_size=100, max_tokens_per_request=budget
    )
    provider.client = _FakeOpenAIClient(api)

    texts = ["one", "two", "three", "four", "five", "six"]
    vectors = await provider.embed_batch(texts)

    assert len(vectors) == len(texts)
    assert len(api.calls) >= 2  # never packed everything into one over-cap request
    for call in api.calls:
        total = sum(provider.tokenizer.count_tokens(t) for t in call["input"])
        # each request is within budget, or a lone chunk that itself exceeds it
        assert len(call["input"]) == 1 or total <= budget
    flat = [t for call in api.calls for t in call["input"]]
    assert flat == texts  # order preserved across the split
