from __future__ import annotations

from typing import List

import tiktoken

from app.core.config import settings
from app.core.logger import logger


class Tokenizer:
    """Token counting and truncation for LLM context assembly."""

    def __init__(self, encoding_name: str = settings.TOKENIZER_MODEL):
        try:
            self.encoding = tiktoken.get_encoding(encoding_name)
            self.encoding_name = encoding_name
        except Exception as exc:
            logger.error("Failed to load tokenizer {name}: {error}", name=encoding_name, error=str(exc))
            self.encoding = tiktoken.get_encoding("cl100k_base")
            self.encoding_name = "cl100k_base"

    def count_tokens(self, text: str) -> int:
        try:
            return len(self.encoding.encode(text))
        except Exception as exc:
            logger.error("Error counting tokens: {error}", error=str(exc))
            return 0

    def count_tokens_batch(self, texts: List[str]) -> List[int]:
        return [self.count_tokens(text) for text in texts]

    def truncate_text(self, text: str, max_tokens: int) -> str:
        try:
            tokens = self.encoding.encode(text)
            if len(tokens) <= max_tokens:
                return text
            return self.encoding.decode(tokens[:max_tokens])
        except Exception as exc:
            logger.error("Error truncating text: {error}", error=str(exc))
            return text[: max(0, max_tokens * 4)]


_tokenizer: Tokenizer | None = None


def get_tokenizer() -> Tokenizer:
    global _tokenizer
    if _tokenizer is None:
        _tokenizer = Tokenizer()
    return _tokenizer
