Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/spend_tracking/input_tokens.py: 20%
78 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
1"""Input-token counting for the budget reservation path.
3Tokenizing is the reservation path's dominant CPU cost and is O(prompt), so
4counting a large prompt inline stalls every other request on the worker.
5Models whose tokenizer the Rust bridge ports (Anthropic, tiktoken cl100k_base
6and o200k_base) are counted from the raw body by the bridge, once per distinct
7tokenizer, which parses and tokenizes with the GIL released. Everything it
8declines, and every model with no Rust tokenizer, is counted in Python, large
9prompts in a worker thread.
10"""
12from __future__ import annotations
14import asyncio
15import json
16from collections.abc import Mapping, Sequence
17from types import MappingProxyType
18from typing import Final
20import litellm
21from litellm._logging import verbose_proxy_logger
22from litellm.rust_bridge import runtime
23from litellm.rust_bridge.catalog import Route, RouteContext
24from litellm.rust_bridge.token_counter import (
25 TOKEN_COUNTER,
26 RustTokenCounterFactory,
27 RustTokenizer,
28 native_count,
29 rust_tokenizer,
30)
32TOKENIZE_OFF_EVENT_LOOP_MIN_CHARS: Final = 30_000
34_INPUT_SIZE_FIELDS: Final = ("messages", "prompt", "input", "query", "documents", "tools", "tool_choice")
37def _approximate_input_size(request_body: Mapping[str, object]) -> int:
38 """Length of the request's input text, a cheap stand-in for tokenizing cost.
40 Every field count_input_tokens_for_model hands the tokenizer is sized here,
41 and rendering rather than walking keeps mapping keys in the total, which a
42 tool schema's property names are."""
43 return sum(len(str(request_body.get(field, ""))) for field in _INPUT_SIZE_FIELDS)
46async def count_input_tokens(
47 request_body: dict,
48 raw_body: bytes | None,
49 models: Sequence[str],
50) -> Mapping[str, int]:
51 """Input-token count per model, sharing one native count across models that
52 select the same tokenizer."""
53 tokenizers: Final[tuple[tuple[str, RustTokenizer | None], ...]] = tuple(
54 (model, rust_tokenizer(model)) for model in models
55 )
56 groups: Final[tuple[RustTokenizer | None, ...]] = tuple(dict.fromkeys(tokenizer for _, tokenizer in tokenizers))
57 group_counts: Final = [
58 await _count_group(
59 request_body=request_body,
60 raw_body=raw_body,
61 tokenizer=tokenizer,
62 models=tuple(model for model, selected in tokenizers if selected == tokenizer),
63 )
64 for tokenizer in groups
65 ]
66 counts: Final = MappingProxyType({model: tokens for group in group_counts for model, tokens in group.items()})
67 verbose_proxy_logger.debug("input token counts: %s", dict(counts))
68 return counts
71async def _count_group(
72 request_body: dict,
73 raw_body: bytes | None,
74 tokenizer: RustTokenizer | None,
75 models: tuple[str, ...],
76) -> Mapping[str, int]:
77 async def python() -> Mapping[str, int]:
78 if _approximate_input_size(request_body) < TOKENIZE_OFF_EVENT_LOOP_MIN_CHARS:
79 return _count_input_tokens_for_models(request_body=request_body, models=models)
80 return await asyncio.to_thread(
81 _count_input_tokens_for_models,
82 request_body=request_body,
83 models=models,
84 )
86 if tokenizer is None or raw_body is None:
87 return await python()
88 try:
89 return await runtime.arun(
90 RouteContext(Route.TOKEN_COUNTER, provider=tokenizer),
91 binding=TOKEN_COUNTER,
92 native=lambda factory: _native_counts(factory, tokenizer, raw_body, models),
93 python=python,
94 )
95 except (RuntimeError, ValueError) as error:
96 from litellm.rust_bridge.fork_guard import ForkedAfterNativeRuntimeStarted, ProcessReservedForForking
98 if isinstance(error, (ForkedAfterNativeRuntimeStarted, ProcessReservedForForking)):
99 raise
100 verbose_proxy_logger.debug("Rust token counter (%s) failed, counting in Python: %s", tokenizer, error)
101 return await python()
104async def _native_counts(
105 factory: RustTokenCounterFactory,
106 tokenizer: RustTokenizer,
107 raw_body: bytes,
108 models: tuple[str, ...],
109) -> Mapping[str, int]:
110 count: Final = await native_count(factory, tokenizer, raw_body)
111 verbose_proxy_logger.debug("Rust token counter (%s) counted %d input tokens", tokenizer, count.input_tokens)
112 return MappingProxyType({model: count.input_tokens for model in models})
115def _count_input_tokens_for_models(
116 request_body: dict,
117 models: Sequence[str],
118) -> Mapping[str, int]:
119 return MappingProxyType(
120 {
121 model: tokens
122 for model in models
123 if (tokens := count_input_tokens_for_model(request_body=request_body, model=model)) is not None
124 }
125 )
128def count_input_tokens_for_model(request_body: dict, model: str) -> int | None:
129 try:
130 if "messages" in request_body:
131 try:
132 return litellm.token_counter(
133 model=model,
134 messages=request_body.get("messages") or (),
135 tools=request_body.get("tools"),
136 tool_choice=request_body.get("tool_choice"),
137 )
138 except ValueError:
139 return _count_text_tokens(model=model, text=request_body.get("messages"))
140 if "prompt" in request_body:
141 return _count_text_tokens(model=model, text=request_body.get("prompt"))
142 if "input" in request_body:
143 return _count_text_tokens(model=model, text=request_body.get("input"))
144 if "query" in request_body or "documents" in request_body:
145 query_tokens: Final = _count_text_tokens(model=model, text=request_body.get("query"))
146 document_tokens: Final = _count_text_tokens(
147 model=model,
148 text=request_body.get("documents"),
149 )
150 return query_tokens + document_tokens
151 except Exception:
152 verbose_proxy_logger.debug("Unable to count input tokens for budget reservation", exc_info=True)
153 return None
156def _count_text_tokens(model: str, text: object) -> int:
157 if text is None:
158 return 0
160 token_count = 0
161 stack: Final = [text]
162 while stack:
163 item = stack.pop()
164 if item is None:
165 continue
166 if isinstance(item, list):
167 stack.extend(item)
168 continue
169 if isinstance(item, dict):
170 token_count += litellm.token_counter(model=model, text=json.dumps(item))
171 continue
172 token_count += litellm.token_counter(model=model, text=str(item))
173 return token_count