"""Unit tests for src.llm.factory token-counter / tokenizer helpers."""

from __future__ import annotations

import pytest


@pytest.fixture(autouse=True)
def _clear_tokenizer_cache():
    """Clear the lru_caches between tests so provider switches take effect."""
    from src.llm import factory

    factory._build_tokenizer.cache_clear()
    factory._build_token_counter.cache_clear()
    factory._build_docling_tokenizer.cache_clear()
    yield
    factory._build_tokenizer.cache_clear()
    factory._build_token_counter.cache_clear()
    factory._build_docling_tokenizer.cache_clear()


def test_azure_path_uses_cl100k_base(monkeypatch):
    from src.llm import factory

    monkeypatch.setattr(factory.settings, "embedding_provider", "azure_openai")

    tokenizer = factory.get_tokenizer()
    counter = factory.get_token_counter()

    tokens = tokenizer("hello world")
    assert isinstance(tokens, list)
    assert len(tokens) == 2
    assert counter("hello world") == 2


def test_vllm_path_uses_hf_tokenizer(monkeypatch):
    from src.llm import factory

    monkeypatch.setattr(factory.settings, "embedding_provider", "vllm")
    monkeypatch.setattr(factory.settings, "embedding_model_name", "fake/model")

    class FakeTokenizer:
        def encode(self, text, add_special_tokens=False):  # noqa: ARG002
            return list(range(len(text.split())))

    class FakeAutoTokenizer:
        @staticmethod
        def from_pretrained(name):
            assert name == "fake/model"
            return FakeTokenizer()

    fake_transformers = type("M", (), {"AutoTokenizer": FakeAutoTokenizer})
    monkeypatch.setitem(__import__("sys").modules, "transformers", fake_transformers)

    counter = factory.get_token_counter()
    assert counter("one two three four") == 4


def test_vllm_fallback_on_load_failure(monkeypatch, caplog):
    from src.llm import factory

    monkeypatch.setattr(factory.settings, "embedding_provider", "vllm")
    monkeypatch.setattr(factory.settings, "embedding_model_name", "does/not/exist")

    class BrokenAutoTokenizer:
        @staticmethod
        def from_pretrained(name):
            raise RuntimeError("boom")

    fake_transformers = type("M", (), {"AutoTokenizer": BrokenAutoTokenizer})
    monkeypatch.setitem(__import__("sys").modules, "transformers", fake_transformers)

    with caplog.at_level("WARNING"):
        counter = factory.get_token_counter()
        result = counter("12345678")  # 8 chars -> 2 via char/4

    assert result == 2
    assert any("HF tokenizer load failed" in rec.message for rec in caplog.records)


def test_fallback_semantics(monkeypatch):
    """Fallback returns ceil(len/4) for non-empty, 0 for empty."""
    from src.llm import factory

    monkeypatch.setattr(factory.settings, "embedding_provider", "vllm")
    monkeypatch.setattr(factory.settings, "embedding_model_name", "does/not/exist")

    class BrokenAutoTokenizer:
        @staticmethod
        def from_pretrained(name):
            raise RuntimeError("boom")

    fake_transformers = type("M", (), {"AutoTokenizer": BrokenAutoTokenizer})
    monkeypatch.setitem(__import__("sys").modules, "transformers", fake_transformers)

    counter = factory.get_token_counter()
    assert counter("a") == 1      # ceil(1/4) = 1
    assert counter("abcd") == 1   # ceil(4/4) = 1
    assert counter("abcde") == 2  # ceil(5/4) = 2
    assert counter("") == 0
