From 3630cc7e1ae98c5594eaa2ccbc042a0f0af95727 Mon Sep 17 00:00:00 2001 From: Chaitany Patel Date: Thu, 10 Sep 2026 11:59:37 +0530 Subject: [PATCH] feat(context): Add TokenBudget and pluggable tokenizers Create the feast.context module with the token accounting that prompt assembly in an OnDemandFeatureView builds on. Tokenizers resolve the way online store types do in repo_config: a built-in name maps to a class path in TOKENIZER_CLASS_FOR_TYPE, and anything else is itself the path of a Tokenizer subclass with a no-argument constructor, loaded through import_class. Adding a tokenizer is therefore what adding a vector store is - write the class, pass its path, no registration call and no mutable global state. Feast ships cl100k_base, o200k_base and a dependency-free character estimate. TokenBudget is a frozen dataclass holding a resolved Tokenizer: consume() returns a new budget rather than mutating, and try_consume() returns None instead of raising, which is the primitive priority_select() will use to fill a budget greedily. TokenBudget.of() resolves a name, class path or instance and can refuse the fallback; is_approximate reports whether counts are exact. Tokenizers compare by value, so budgets built from the same name are interchangeable as dict keys and set members. When tiktoken cannot be loaded, get_tokenizer() degrades to the character estimate with a logged warning; fallback=False makes it fatal for paths that need exact counts. TiktokenTokenizer counts with disallowed_special=(), so a feature value containing a literal "<|endoftext|>" is treated as ordinary text instead of raising. tiktoken is not declared as a dependency yet, so the tests that assert exact encodings skip when it is absent. Part of RHOAIENG-80116 (PR 1/9). Signed-off-by: Chaitany Patel --- sdk/python/feast/context/__init__.py | 39 +++ sdk/python/feast/context/errors.py | 53 ++++ sdk/python/feast/context/token_budget.py | 163 ++++++++++++ sdk/python/feast/context/tokenizer.py | 237 ++++++++++++++++++ sdk/python/tests/unit/context/__init__.py | 0 sdk/python/tests/unit/context/conftest.py | 22 ++ .../tests/unit/context/test_token_budget.py | 200 +++++++++++++++ .../tests/unit/context/test_tokenizer.py | 141 +++++++++++ 8 files changed, 855 insertions(+) create mode 100644 sdk/python/feast/context/__init__.py create mode 100644 sdk/python/feast/context/errors.py create mode 100644 sdk/python/feast/context/token_budget.py create mode 100644 sdk/python/feast/context/tokenizer.py create mode 100644 sdk/python/tests/unit/context/__init__.py create mode 100644 sdk/python/tests/unit/context/conftest.py create mode 100644 sdk/python/tests/unit/context/test_token_budget.py create mode 100644 sdk/python/tests/unit/context/test_tokenizer.py diff --git a/sdk/python/feast/context/__init__.py b/sdk/python/feast/context/__init__.py new file mode 100644 index 00000000000..62de13ebead --- /dev/null +++ b/sdk/python/feast/context/__init__.py @@ -0,0 +1,39 @@ +"""Token-budget-aware helpers for assembling LLM context in feature views. + +Plain functions over immutable values, so an OnDemandFeatureView calling them +gets the same prompt offline and online. +""" + +from feast.context.errors import ( + ContextError, + TokenBudgetExceededError, + TokenizerNotFoundError, + TokenizerUnavailableError, +) +from feast.context.token_budget import TokenBudget +from feast.context.tokenizer import ( + ApproximateTokenizer, + Cl100kBaseTokenizer, + O200kBaseTokenizer, + TiktokenTokenizer, + Tokenizer, + TokenizerName, + TokenizerSpec, + get_tokenizer, +) + +__all__ = [ + "ApproximateTokenizer", + "Cl100kBaseTokenizer", + "ContextError", + "O200kBaseTokenizer", + "TiktokenTokenizer", + "TokenBudget", + "TokenBudgetExceededError", + "Tokenizer", + "TokenizerName", + "TokenizerNotFoundError", + "TokenizerSpec", + "TokenizerUnavailableError", + "get_tokenizer", +] diff --git a/sdk/python/feast/context/errors.py b/sdk/python/feast/context/errors.py new file mode 100644 index 00000000000..2504eabebd9 --- /dev/null +++ b/sdk/python/feast/context/errors.py @@ -0,0 +1,53 @@ +from typing import Iterable + +from fastapi import status as HttpStatusCode + +from feast.errors import FeastError + + +class ContextError(FeastError): + """Base class for all feast.context failures.""" + + +class TokenizerNotFoundError(ContextError): + def __init__(self, name: str, available: Iterable[str]): + super().__init__( + f"Unknown tokenizer '{name}'. Built-in tokenizers: " + f"{', '.join(sorted(available))}. For your own, pass the class " + f"path of a Tokenizer subclass, e.g. 'my_pkg.MyTokenizer'." + ) + + def http_status_code(self) -> int: + return HttpStatusCode.HTTP_400_BAD_REQUEST + + +class TokenizerUnavailableError(ContextError): + """A known tokenizer cannot be loaded. + + Covers a missing dependency, and one present but unable to load its + encoding — tiktoken fetching a BPE file on a host with no network, say. + """ + + def __init__(self, name: str, reason: str, install_hint: str): + super().__init__( + f"Tokenizer '{name}' is unavailable: {reason}. Install its " + f"dependencies with: {install_hint}" + ) + self.tokenizer_name = name + + def http_status_code(self) -> int: + return HttpStatusCode.HTTP_400_BAD_REQUEST + + +class TokenBudgetExceededError(ContextError): + def __init__(self, requested: int, remaining: int, max_tokens: int): + super().__init__( + f"Cannot consume {requested} tokens: only {remaining} of " + f"{max_tokens} tokens remain in the budget." + ) + self.requested = requested + self.remaining = remaining + self.max_tokens = max_tokens + + def http_status_code(self) -> int: + return HttpStatusCode.HTTP_400_BAD_REQUEST diff --git a/sdk/python/feast/context/token_budget.py b/sdk/python/feast/context/token_budget.py new file mode 100644 index 00000000000..2b2ccc8657d --- /dev/null +++ b/sdk/python/feast/context/token_budget.py @@ -0,0 +1,163 @@ +"""Token accounting for prompt assembly.""" + +from dataclasses import dataclass, field +from typing import Optional + +from feast.context.errors import TokenBudgetExceededError +from feast.context.tokenizer import ( + ApproximateTokenizer, + Tokenizer, + TokenizerName, + TokenizerSpec, + get_tokenizer, +) + + +@dataclass(frozen=True) +class TokenBudget: + """An immutable token allowance measured with a specific tokenizer. + + ``consume()`` returns a new budget instead of mutating this one, so a + budget is safe to share across threads and to reuse between assemblies:: + + budget = TokenBudget.of(4096, "cl100k_base") + budget.count_tokens("Hello world") # -> 2 + budget = budget.consume("Hello world") + budget.remaining # -> 4094 + + Attributes: + max_tokens: Size of the allowance. Zero admits only empty content. + tokenizer: The tokenizer doing the counting. A name is accepted too + and resolved at construction; :meth:`of` types that properly and + can forbid the fallback to the estimate. + consumed: Tokens already spent. + """ + + max_tokens: int + tokenizer: Tokenizer = field(default_factory=get_tokenizer) + consumed: int = 0 + + def __post_init__(self) -> None: + if self.max_tokens < 0: + raise ValueError(f"max_tokens must not be negative, got {self.max_tokens}.") + if self.consumed < 0: + raise ValueError(f"consumed must not be negative, got {self.consumed}.") + if self.consumed > self.max_tokens: + raise ValueError( + f"consumed ({self.consumed}) must not exceed max_tokens " + f"({self.max_tokens})." + ) + if not isinstance(self.tokenizer, Tokenizer): + # A name from config lands here; resolve it once, so counting + # never re-enters resolution and equality stays value-based. + object.__setattr__(self, "tokenizer", get_tokenizer(self.tokenizer)) + + @classmethod + def of( + cls, + max_tokens: int, + tokenizer: TokenizerSpec = TokenizerName.CL100K_BASE, + *, + consumed: int = 0, + fallback: bool = True, + ) -> "TokenBudget": + """Build a budget from a tokenizer name, instance, or class path. + + Args: + fallback: Set False to raise instead of degrading to the character + estimate when the tokenizer cannot be loaded. + """ + return cls( + max_tokens=max_tokens, + tokenizer=get_tokenizer(tokenizer, fallback=fallback), + consumed=consumed, + ) + + @property + def remaining(self) -> int: + """Tokens still available.""" + return self.max_tokens - self.consumed + + @property + def is_exhausted(self) -> bool: + """Whether the allowance is fully spent.""" + return self.remaining == 0 + + @property + def tokenizer_name(self) -> str: + """Name of the tokenizer actually counting.""" + return self.tokenizer.name + + @property + def is_approximate(self) -> bool: + """Whether counts are estimated rather than exact. + + True when the requested tokenizer could not be loaded and the budget + fell back to counting characters. + """ + return isinstance(self.tokenizer, ApproximateTokenizer) + + def count_tokens(self, text: str) -> int: + """Tokens ``text`` would occupy, regardless of what remains.""" + return self.tokenizer.count_tokens(text) + + def fits(self, text: str) -> bool: + """Whether ``text`` fits in the remaining allowance.""" + return self.count_tokens(text) <= self.remaining + + def fits_tokens(self, tokens: int) -> bool: + """Whether ``tokens`` more tokens fit in the remaining allowance.""" + _check_not_negative(tokens) + return tokens <= self.remaining + + def consume(self, text: str) -> "TokenBudget": + """Charge ``text`` against the budget, returning the new budget. + + Raises: + TokenBudgetExceededError: ``text`` does not fit; guard with + ``fits()`` or use ``try_consume()``. + """ + return self.consume_tokens(self.count_tokens(text)) + + def consume_tokens(self, tokens: int) -> "TokenBudget": + """Charge ``tokens`` against the budget, returning the new budget. + + Raises: + TokenBudgetExceededError: ``tokens`` exceeds what remains. + """ + if not self.fits_tokens(tokens): + raise TokenBudgetExceededError(tokens, self.remaining, self.max_tokens) + return self._replace_consumed(self.consumed + tokens) + + def try_consume(self, text: str) -> Optional["TokenBudget"]: + """Charge ``text`` if it fits, else return None. + + Lets a caller keep whichever candidate sections still fit, without + catching exceptions. + """ + tokens = self.count_tokens(text) + if not self.fits_tokens(tokens): + return None + return self._replace_consumed(self.consumed + tokens) + + def reset(self) -> "TokenBudget": + """A budget with the same limit and tokenizer, nothing consumed.""" + return self._replace_consumed(0) + + def _replace_consumed(self, consumed: int) -> "TokenBudget": + return TokenBudget( + max_tokens=self.max_tokens, + tokenizer=self.tokenizer, + consumed=consumed, + ) + + def __str__(self) -> str: + return ( + f"TokenBudget({self.consumed}/{self.max_tokens} tokens used, " + f"tokenizer={self.tokenizer_name})" + ) + + +def _check_not_negative(tokens: int) -> None: + if tokens < 0: + raise ValueError(f"tokens must not be negative, got {tokens}.") diff --git a/sdk/python/feast/context/tokenizer.py b/sdk/python/feast/context/tokenizer.py new file mode 100644 index 00000000000..09af77def68 --- /dev/null +++ b/sdk/python/feast/context/tokenizer.py @@ -0,0 +1,237 @@ +"""Tokenizers and how a name resolves to one. + +Resolution follows the same convention as online store types in +``repo_config``: a built-in name maps to a class path, and anything else is +itself the fully-qualified path of a :class:`Tokenizer` subclass with a +no-argument constructor:: + + TokenBudget(max_tokens=4096, tokenizer="cl100k_base") + TokenBudget(max_tokens=4096, tokenizer="my_pkg.tokenizers.LlamaTokenizer") +""" + +import logging +import math +from abc import ABC, abstractmethod +from enum import Enum +from functools import lru_cache +from typing import Any, Union + +from feast.context.errors import TokenizerNotFoundError, TokenizerUnavailableError +from feast.errors import FeastInvalidBaseClass +from feast.importer import import_class + +logger = logging.getLogger(__name__) + +#: Average characters per token for English prose. +DEFAULT_CHARS_PER_TOKEN = 4 + +#: Quoted in error messages when tiktoken is missing. +TIKTOKEN_INSTALL_HINT = "pip install tiktoken" + + +class TokenizerName(str, Enum): + """Built-in tokenizer names.""" + + #: GPT-4, GPT-3.5-turbo, text-embedding-3-*. + CL100K_BASE = "cl100k_base" + #: GPT-4o and the o-series models. + O200K_BASE = "o200k_base" + #: Character estimate; no third-party dependency. + APPROXIMATE = "approximate" + + +TOKENIZER_CLASS_FOR_TYPE = { + TokenizerName.CL100K_BASE.value: "feast.context.tokenizer.Cl100kBaseTokenizer", + TokenizerName.O200K_BASE.value: "feast.context.tokenizer.O200kBaseTokenizer", + TokenizerName.APPROXIMATE.value: "feast.context.tokenizer.ApproximateTokenizer", +} + + +class Tokenizer(ABC): + """Counts the tokens a model would consume for a piece of text. + + Implementations must be stateless and thread-safe: one instance per name is + shared across every :class:`TokenBudget`. A third-party subclass needs a + no-argument constructor and a class name ending in ``Tokenizer``, which is + what resolution by class path looks for. + + Two tokenizers that count identically compare equal, so budgets built from + the same name are interchangeable as dict keys and set members. + """ + + @property + @abstractmethod + def name(self) -> str: + """Identifier reported in budgets and logs.""" + + @abstractmethod + def count_tokens(self, text: str) -> int: + """Number of tokens in ``text``. Returns 0 for the empty string.""" + + def _identity(self) -> tuple[object, ...]: + """Whatever makes two instances count the same way.""" + return (self.name,) + + def __eq__(self, other: object) -> bool: + return isinstance(other, Tokenizer) and self._identity() == other._identity() + + def __hash__(self) -> int: + return hash(self._identity()) + + def __repr__(self) -> str: + return f"{type(self).__name__}(name={self.name!r})" + + +class ApproximateTokenizer(Tokenizer): + """Character-count estimate used when no real tokenizer is available. + + Counts ``ceil(len(text) / chars_per_token)``: close enough for budgeting, + not for exact context-window arithmetic. Code and non-Latin scripts drift + furthest from the ratio. + """ + + def __init__(self, chars_per_token: int = DEFAULT_CHARS_PER_TOKEN) -> None: + if chars_per_token <= 0: + raise ValueError( + f"chars_per_token must be positive, got {chars_per_token}." + ) + self._chars_per_token = chars_per_token + + @property + def name(self) -> str: + return TokenizerName.APPROXIMATE.value + + @property + def chars_per_token(self) -> int: + return self._chars_per_token + + def count_tokens(self, text: str) -> int: + return math.ceil(len(text) / self._chars_per_token) + + def _identity(self) -> tuple[object, ...]: + return (self.name, self._chars_per_token) + + +class TiktokenTokenizer(Tokenizer): + """Exact token counts from a tiktoken encoding. + + Raises: + TokenizerUnavailableError: tiktoken is missing or the encoding failed + to load. + TokenizerNotFoundError: tiktoken does not know ``encoding_name``. + """ + + def __init__(self, encoding_name: str) -> None: + self._encoding_name = encoding_name + self._encoding = _load_tiktoken_encoding(encoding_name) + + @property + def name(self) -> str: + return self._encoding_name + + def count_tokens(self, text: str) -> int: + if not text: + return 0 + # Feature values are arbitrary text: count a literal "<|endoftext|>" + # rather than raise, which is tiktoken's default for special tokens. + return len(self._encoding.encode(text, disallowed_special=())) + + +class Cl100kBaseTokenizer(TiktokenTokenizer): + """cl100k_base, in the no-argument form that resolution by name needs.""" + + def __init__(self) -> None: + super().__init__(TokenizerName.CL100K_BASE.value) + + +class O200kBaseTokenizer(TiktokenTokenizer): + """o200k_base, in the no-argument form that resolution by name needs.""" + + def __init__(self) -> None: + super().__init__(TokenizerName.O200K_BASE.value) + + +def _load_tiktoken_encoding(encoding_name: str) -> Any: + """Load a tiktoken encoding. tiktoken memoizes these internally.""" + try: + import tiktoken + except ImportError as e: + raise TokenizerUnavailableError( + encoding_name, f"tiktoken is not installed ({e})", TIKTOKEN_INSTALL_HINT + ) from e + + try: + return tiktoken.get_encoding(encoding_name) + except ValueError as e: + raise TokenizerNotFoundError( + encoding_name, tiktoken.list_encoding_names() + ) from e + except Exception as e: + # Encodings download on first use; an offline host fails here. + raise TokenizerUnavailableError( + encoding_name, + f"tiktoken failed to load the encoding ({e})", + TIKTOKEN_INSTALL_HINT, + ) from e + + +#: What callers may pass wherever a tokenizer is expected. +TokenizerSpec = Union[str, TokenizerName, Tokenizer] + + +def get_tokenizer( + spec: TokenizerSpec = TokenizerName.CL100K_BASE, + *, + fallback: bool = True, +) -> Tokenizer: + """Resolve ``spec`` to a tokenizer instance. + + A :class:`Tokenizer` passes through; a name is built once, then cached. + + Args: + spec: Tokenizer instance, built-in name, or the class path of a + ``Tokenizer`` subclass with a no-argument constructor. + fallback: On missing dependencies, return + :class:`ApproximateTokenizer` with a warning instead of raising. + Set False when exact counts are required. + + Raises: + TokenizerNotFoundError: ``spec`` is neither a built-in name nor a + class path ending in ``Tokenizer``. + TokenizerUnavailableError: the tokenizer cannot be built and + ``fallback`` is False. + """ + if isinstance(spec, Tokenizer): + return spec + name = spec.value if isinstance(spec, TokenizerName) else spec + if not isinstance(name, str) or not name.strip(): + raise TypeError(f"Tokenizer name must be a non-empty string, got {spec!r}.") + + try: + return _build_tokenizer(name.strip()) + except TokenizerUnavailableError as e: + if not fallback: + raise + logger.warning( + "%s Falling back to a character-based estimate; token counts will " + "be approximate.", + e, + ) + # Through the cache, so budgets that fell back still compare equal. + return _build_tokenizer(TokenizerName.APPROXIMATE.value) + + +@lru_cache(maxsize=None) +def _build_tokenizer(name: str) -> Tokenizer: + """Import and instantiate the tokenizer ``name`` refers to.""" + tokenizer_type = TOKENIZER_CLASS_FOR_TYPE.get(name, name) + if "." not in tokenizer_type or not tokenizer_type.endswith("Tokenizer"): + raise TokenizerNotFoundError(name, TOKENIZER_CLASS_FOR_TYPE.keys()) + + module_name, class_name = tokenizer_type.rsplit(".", 1) + tokenizer = import_class(module_name, class_name, "Tokenizer")() + if not isinstance(tokenizer, Tokenizer): + # import_class checks the base class by name, so an unrelated class + # called "Tokenizer" gets this far. + raise FeastInvalidBaseClass(tokenizer_type, "Tokenizer") + return tokenizer diff --git a/sdk/python/tests/unit/context/__init__.py b/sdk/python/tests/unit/context/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/sdk/python/tests/unit/context/conftest.py b/sdk/python/tests/unit/context/conftest.py new file mode 100644 index 00000000000..2bed5b15656 --- /dev/null +++ b/sdk/python/tests/unit/context/conftest.py @@ -0,0 +1,22 @@ +import builtins + +import pytest + +from feast.context import tokenizer as tokenizer_module + + +@pytest.fixture +def no_tiktoken(monkeypatch): + """Make ``import tiktoken`` fail, as it does when it is not installed.""" + real_import = builtins.__import__ + + def fake_import(name, *args, **kwargs): + if name == "tiktoken": + raise ImportError("No module named 'tiktoken'") + return real_import(name, *args, **kwargs) + + monkeypatch.setattr(builtins, "__import__", fake_import) + # Drop tokenizers an earlier test may have built and cached. + tokenizer_module._build_tokenizer.cache_clear() + yield + tokenizer_module._build_tokenizer.cache_clear() diff --git a/sdk/python/tests/unit/context/test_token_budget.py b/sdk/python/tests/unit/context/test_token_budget.py new file mode 100644 index 00000000000..9480c488b53 --- /dev/null +++ b/sdk/python/tests/unit/context/test_token_budget.py @@ -0,0 +1,200 @@ +import dataclasses +import inspect + +import pytest + +from feast.context import TokenBudget +from feast.context.errors import TokenBudgetExceededError, TokenizerNotFoundError +from feast.context.tokenizer import ( + ApproximateTokenizer, + Tokenizer, + TokenizerName, + get_tokenizer, +) + +# Four characters per token: "abcd" is one token under the estimate. +ESTIMATING = ApproximateTokenizer() + + +class HalvingTokenizer(Tokenizer): + """A third-party tokenizer, resolved by class path.""" + + @property + def name(self) -> str: + return "halving" + + def count_tokens(self, text: str) -> int: + return len(text) // 2 + + +def budget(max_tokens: int = 100, consumed: int = 0) -> TokenBudget: + """A budget on the deterministic estimate, so counts need no tiktoken.""" + return TokenBudget(max_tokens=max_tokens, tokenizer=ESTIMATING, consumed=consumed) + + +class TestConstruction: + def test_defaults_to_cl100k_base(self): + field = TokenBudget.__dataclass_fields__["tokenizer"] + assert field.default_factory is get_tokenizer + assert ( + inspect.signature(get_tokenizer).parameters["spec"].default + is TokenizerName.CL100K_BASE + ) + + def test_accepts_a_tokenizer_name(self): + assert TokenBudget.of(10, "approximate").tokenizer is get_tokenizer( + "approximate" + ) + + def test_accepts_a_tokenizer_instance(self): + assert TokenBudget(max_tokens=10, tokenizer=ESTIMATING).tokenizer is ESTIMATING + + def test_resolves_a_name_passed_to_the_constructor(self): + # Config supplies a string; the field still holds a Tokenizer. + resolved = TokenBudget(max_tokens=10, tokenizer="approximate").tokenizer # type: ignore[arg-type] + assert resolved is get_tokenizer("approximate") + + def test_resolves_a_custom_tokenizer_by_class_path(self): + b = TokenBudget.of(10, "tests.unit.context.test_token_budget.HalvingTokenizer") + assert b.tokenizer_name == "halving" + assert b.count_tokens("abcdefgh") == 4 + assert b.consume("abcdefgh").remaining == 6 + + def test_unknown_tokenizer_raises(self): + with pytest.raises(TokenizerNotFoundError): + TokenBudget.of(10, "gpt-42") + + def test_rejects_negative_max_tokens(self): + with pytest.raises(ValueError, match="max_tokens"): + budget(max_tokens=-1) + + def test_rejects_negative_consumed(self): + with pytest.raises(ValueError, match="consumed"): + budget(consumed=-1) + + def test_rejects_consumed_above_max_tokens(self): + with pytest.raises(ValueError, match="must not exceed"): + budget(max_tokens=10, consumed=11) + + def test_is_immutable(self): + with pytest.raises(dataclasses.FrozenInstanceError): + budget().max_tokens = 5 # type: ignore[misc] + + def test_budgets_with_the_same_state_are_equal_and_hash_alike(self): + # Default-constructed, so a tokenizer compared by identity would fail. + assert TokenBudget(max_tokens=10) == TokenBudget(max_tokens=10) + assert len({TokenBudget(max_tokens=10), TokenBudget(max_tokens=10)}) == 1 + + def test_replace_keeps_the_resolved_tokenizer(self): + replaced = dataclasses.replace(budget(max_tokens=10), consumed=4) + assert replaced.tokenizer is ESTIMATING + assert replaced.remaining == 6 + + +class TestCounting: + def test_counts_tokens(self): + assert budget().count_tokens("abcdefgh") == 2 + + def test_empty_string_costs_nothing(self): + assert budget().count_tokens("") == 0 + + def test_is_approximate_flags_the_estimate(self): + assert budget().is_approximate is True + + +class TestFits: + @pytest.mark.parametrize( + "max_tokens, text, expected", + [ + (10, "abcdefgh", True), + (2, "abcdefgh", True), # exactly filling the budget fits + (1, "a" * 40, False), + (0, "", True), # nothing always fits + ], + ) + def test_fits_compares_against_what_remains(self, max_tokens, text, expected): + assert budget(max_tokens=max_tokens).fits(text) is expected + + def test_fits_tokens_compares_against_what_remains(self): + b = budget(max_tokens=10, consumed=8) + assert b.fits_tokens(2) is True + assert b.fits_tokens(3) is False + + def test_fits_tokens_rejects_negative_counts(self): + with pytest.raises(ValueError, match="tokens"): + budget().fits_tokens(-1) + + +class TestConsume: + def test_returns_a_new_budget(self): + original = budget(max_tokens=10) + after = original.consume("abcdefgh") + assert after is not original + assert original.remaining == 10 + assert after.remaining == 8 + assert after.consumed == 2 + + def test_exhausts_the_budget_exactly(self): + after = budget(max_tokens=2).consume("abcdefgh") + assert after.remaining == 0 + assert after.is_exhausted is True + + def test_over_budget_raises_and_leaves_the_original_untouched(self): + b = budget(max_tokens=2) + with pytest.raises(TokenBudgetExceededError, match="only 2 of 2 tokens remain"): + b.consume("a" * 40) + assert b.remaining == 2 + + def test_consume_tokens_charges_a_count_directly(self): + assert budget(max_tokens=10).consume_tokens(4).remaining == 6 + + def test_consume_tokens_rejects_negative_counts(self): + with pytest.raises(ValueError, match="tokens"): + budget().consume_tokens(-1) + + +class TestTryConsume: + def test_returns_the_new_budget_when_content_fits(self): + after = budget(max_tokens=10).try_consume("abcdefgh") + assert after is not None + assert after.remaining == 8 + + def test_returns_none_when_content_does_not_fit(self): + assert budget(max_tokens=2).try_consume("a" * 40) is None + + def test_greedy_selection_keeps_what_fits(self): + b = budget(max_tokens=3) + kept = [] + for section in ["abcd", "a" * 40, "abcdefgh"]: + candidate = b.try_consume(section) + if candidate is not None: + b = candidate + kept.append(section) + assert kept == ["abcd", "abcdefgh"] + assert b.remaining == 0 + + +class TestReset: + def test_clears_consumption_and_keeps_limit_and_tokenizer(self): + after = budget(max_tokens=10).consume("abcdefgh").reset() + assert after.consumed == 0 + assert after == budget(max_tokens=10) + + +class TestTiktokenBudget: + @pytest.mark.parametrize( + "name", [TokenizerName.CL100K_BASE.value, TokenizerName.O200K_BASE.value] + ) + def test_counts_exact_tokens(self, name): + pytest.importorskip("tiktoken") + b = TokenBudget.of(4096, name) + assert b.tokenizer_name == name + assert b.is_approximate is False + assert b.count_tokens("Hello world") == 2 + assert b.consume("Hello world").remaining == 4094 + + def test_the_default_budget_counts_cl100k_base(self): + pytest.importorskip("tiktoken") + assert TokenBudget(max_tokens=4096).tokenizer_name == ( + TokenizerName.CL100K_BASE.value + ) diff --git a/sdk/python/tests/unit/context/test_tokenizer.py b/sdk/python/tests/unit/context/test_tokenizer.py new file mode 100644 index 00000000000..2b909cad65b --- /dev/null +++ b/sdk/python/tests/unit/context/test_tokenizer.py @@ -0,0 +1,141 @@ +import pytest + +from feast.context.errors import TokenizerNotFoundError, TokenizerUnavailableError +from feast.context.tokenizer import ( + ApproximateTokenizer, + Cl100kBaseTokenizer, + O200kBaseTokenizer, + TiktokenTokenizer, + TokenizerName, + get_tokenizer, +) +from feast.errors import FeastInvalidBaseClass + + +class FakeTokenizer: + """Named like a tokenizer, but not a Tokenizer subclass.""" + + +class TestApproximateTokenizer: + @pytest.mark.parametrize( + "text, expected", + [ + ("abcdefgh", 2), + ("abcde", 2), # partial tokens round up + ("", 0), + ("a" * 100_000, 25_000), + ], + ) + def test_counts_four_characters_per_token(self, text, expected): + assert ApproximateTokenizer().count_tokens(text) == expected + + def test_chars_per_token_is_configurable(self): + assert ApproximateTokenizer(chars_per_token=2).count_tokens("abcd") == 2 + + @pytest.mark.parametrize("chars_per_token", [0, -1]) + def test_rejects_non_positive_chars_per_token(self, chars_per_token): + with pytest.raises(ValueError): + ApproximateTokenizer(chars_per_token=chars_per_token) + + +class TestTiktokenTokenizer: + @pytest.mark.parametrize( + "encoding", [TokenizerName.CL100K_BASE.value, TokenizerName.O200K_BASE.value] + ) + def test_counts_tokens_for_supported_encodings(self, encoding): + pytest.importorskip("tiktoken") + tokenizer = TiktokenTokenizer(encoding) + assert tokenizer.name == encoding + assert tokenizer.count_tokens("Hello world") == 2 + assert tokenizer.count_tokens("") == 0 + # Unicode and long text: exact counts are the encoding's business, but + # neither may crash or come back empty. + assert tokenizer.count_tokens("北京の天気") > 0 + assert tokenizer.count_tokens("word " * 10_000) >= 10_000 + + def test_special_token_text_is_counted_not_rejected(self): + pytest.importorskip("tiktoken") + tokenizer = Cl100kBaseTokenizer() + assert tokenizer.count_tokens("<|endoftext|> appears in this review") > 0 + + def test_unknown_encoding_raises(self): + pytest.importorskip("tiktoken") + with pytest.raises(TokenizerNotFoundError): + TiktokenTokenizer("no_such_encoding") + + +class TestEquality: + def test_same_name_and_settings_are_equal(self): + assert ApproximateTokenizer() == ApproximateTokenizer() + assert len({ApproximateTokenizer(), ApproximateTokenizer()}) == 1 + + def test_different_settings_are_not_equal(self): + assert ApproximateTokenizer() != ApproximateTokenizer(chars_per_token=3) + + def test_same_encoding_from_different_classes_is_equal(self): + pytest.importorskip("tiktoken") + assert Cl100kBaseTokenizer() == TiktokenTokenizer( + TokenizerName.CL100K_BASE.value + ) + assert Cl100kBaseTokenizer() != O200kBaseTokenizer() + + +class TestGetTokenizer: + @pytest.mark.parametrize( + "name, expected_class", + [ + (TokenizerName.CL100K_BASE.value, Cl100kBaseTokenizer), + (TokenizerName.O200K_BASE.value, O200kBaseTokenizer), + (TokenizerName.APPROXIMATE.value, ApproximateTokenizer), + ], + ) + def test_resolves_builtin_names(self, name, expected_class): + if expected_class is not ApproximateTokenizer: + pytest.importorskip("tiktoken") + assert isinstance(get_tokenizer(name), expected_class) + + def test_resolves_a_class_path(self): + tokenizer = get_tokenizer("feast.context.tokenizer.ApproximateTokenizer") + assert isinstance(tokenizer, ApproximateTokenizer) + + def test_rejects_a_class_path_that_is_not_a_tokenizer(self): + with pytest.raises(FeastInvalidBaseClass): + get_tokenizer("tests.unit.context.test_tokenizer.FakeTokenizer") + + def test_rejects_a_path_that_is_not_a_tokenizer_class_name(self): + with pytest.raises(TokenizerNotFoundError): + get_tokenizer("feast.repo_config.RepoConfig") + + def test_unknown_name_raises_with_builtin_names(self): + with pytest.raises(TokenizerNotFoundError, match="cl100k_base"): + get_tokenizer("gpt-42") + + def test_builds_once_and_caches(self): + assert get_tokenizer("approximate") is get_tokenizer(" approximate ") + + def test_accepts_enum_member_and_equivalent_string(self): + assert get_tokenizer(TokenizerName.APPROXIMATE) is get_tokenizer("approximate") + + def test_tokenizer_instance_passes_through(self): + tokenizer = ApproximateTokenizer() + assert get_tokenizer(tokenizer) is tokenizer + + @pytest.mark.parametrize("name", [" ", None]) + def test_rejects_invalid_names(self, name): + with pytest.raises(TypeError): + get_tokenizer(name) + + +class TestFallback: + def test_falls_back_to_estimate_without_tiktoken(self, no_tiktoken, caplog): + tokenizer = get_tokenizer(TokenizerName.CL100K_BASE) + + assert isinstance(tokenizer, ApproximateTokenizer) + assert tokenizer.count_tokens("abcdefgh") == 2 + # The shared instance, so budgets that fell back stay comparable. + assert tokenizer is get_tokenizer(TokenizerName.APPROXIMATE) + assert any("tiktoken" in record.message for record in caplog.records) + + def test_fallback_disabled_raises(self, no_tiktoken): + with pytest.raises(TokenizerUnavailableError, match="pip install tiktoken"): + get_tokenizer(TokenizerName.CL100K_BASE, fallback=False)