mirror of
https://github.com/headroomlabs-ai/headroom.git
synced 2026-08-27 14:17:10 -04:00
fix(providers/openai): bound tiktoken vocab loads with the guarded loader (#2554)
## Description
`headroom/providers/openai.py::_get_encoding` calls
`tiktoken.get_encoding` directly. tiktoken downloads missing
vocabularies via `requests.get` with **no timeout**, so on a network
that blackholes the vocab CDN (corporate firewall, SSL-intercepting
proxy), whichever thread first counts tokens for an OpenAI model — proxy
startup included — blocks indefinitely.
This is the provider-path hole left by #956: the tokenizer registry
already routes through a bounded loader
(`headroom/tokenizers/tiktoken_counter.py`, worker-thread load +
`HEADROOM_TIKTOKEN_LOAD_TIMEOUT_SECONDS`, default 10s) and falls back to
estimation, but the OpenAI provider path never got the same treatment.
Observed in production (Headroom Desktop fleet, Sentry): a proxy that
never finished booting, with a faulthandler dump wedged in
`tiktoken/registry.py` `get_encoding` on the main thread, reached from
the `headroom` CLI entrypoint via click. The desktop app now also
pre-seeds a persistent `TIKTOKEN_CACHE_DIR`, but the unbounded load
affects every deployment of the proxy, so it should be fixed here too.
Follow-up to #956.
## Type of Change
- [x] Bug fix (non-breaking change that fixes an issue)
- [ ] New feature (non-breaking change that adds functionality)
- [ ] Breaking change (fix or feature that would cause existing
functionality to change)
- [ ] Documentation update
- [ ] Performance improvement
- [ ] Code refactoring (no functional changes)
## Changes Made
- `_get_encoding` now routes through the bounded `load_encoding` from
`headroom.tokenizers.tiktoken_counter` instead of calling
`tiktoken.get_encoding` directly, so a stalled vocab download raises
`TiktokenLoadError` after the timeout instead of hanging the calling
thread.
- `OpenAIProvider.get_token_counter` catches `TiktokenLoadError` and
falls back to `EstimatingTokenCounter`, cached per model so later
requests never re-block on the same failed download — mirroring
`TokenizerRegistry._create_tiktoken`.
- `TIKTOKEN_AVAILABLE` uses `importlib.util.find_spec` (the module-level
`import tiktoken` became unused; same pattern as `LITELLM_AVAILABLE`).
- Two regression tests (`TestGuardedEncodingLoad`) covering the
bounded-raise path and the cached estimation fallback.
## Testing
- [x] Unit tests pass (`pytest`)
- [x] Linting passes (`ruff check .`)
- [x] Type checking passes (`mypy headroom`)
- [x] New tests added for new functionality
- [ ] Manual testing performed
### Test Output
```text
$ uv run --frozen --extra dev pytest tests/test_providers/ tests/test_tokenizers/
======================= 117 passed, 4 warnings in 27.95s =======================
$ uvx ruff check headroom/providers/openai.py tests/test_providers/test_openai.py
All checks passed!
$ uvx ruff format --check headroom/providers/openai.py tests/test_providers/test_openai.py
2 files already formatted
$ uv run --frozen --extra dev mypy headroom/providers/openai.py
Success: no issues found in 1 source file
```
## Real Behavior Proof
- Environment: macOS (arm64), Python 3.12, uv-managed venv, branch off
`upstream/main` (a6d4921e).
- Exact command / steps: the stalled download is simulated in
`TestGuardedEncodingLoad` by monkeypatching
`tiktoken_counter.load_encoding` to raise `TiktokenLoadError`;
`OpenAITokenCounter("gpt-4o")` then raises the bounded error, and
`OpenAIProvider.get_token_counter("gpt-4o")` returns a working
`EstimatingTokenCounter` and reuses the same instance on the second
call.
- Observed result: tests pass (see output above); with an empty tiktoken
on-disk cache, the encoding load path is identical to the one in the
production faulthandler dump.
- Not tested: an end-to-end run against a genuinely blackholed vocab CDN
(needs a firewalled network); the timeout mechanism itself is #956's
code, already covered by its own tests.
## Review Readiness
- [x] I have performed a self-review
- [x] This PR is ready for human review
## Checklist
- [x] My code follows the project's style guidelines
- [x] I have performed a self-review of my code
- [x] I commented my code, particularly in hard-to-understand areas
- [x] I made corresponding changes to the documentation
- [x] My changes generate no new warnings
- [x] I added tests that prove my fix is effective or that my feature
works
- [x] New and existing unit tests pass locally with my changes
- [ ] I updated CHANGELOG.md if applicable (N/A — release-please
generates it from the PR title; changelog-guard forbids manual edits)
## Screenshots (if applicable)
N/A — no UI change.
## Additional Notes
The estimation fallback is sticky for the process lifetime (per-model
cache + the loader's fail-fast set from #956): a network that recovers
mid-session keeps estimation until restart. That matches the registry's
existing behavior, so no new divergence is introduced.
🤖 Generated with [Claude Code](https://claude.com/claude-code)
---------
Co-authored-by: Claude Fable 5 <noreply@anthropic.com>
Co-authored-by: JD Davis <jd@jds-macbook-air.tail2a279.ts.net>
This commit is contained in:
parent
7f24d695ee
commit
0805e8e410
2 changed files with 79 additions and 12 deletions
|
|
@ -32,12 +32,7 @@ _UNKNOWN_MODEL_WARNINGS: set[str] = set()
|
|||
# Models whose price came from the built-in table rather than LiteLLM.
|
||||
_PRICING_FALLBACK_WARNINGS: set[str] = set()
|
||||
|
||||
try:
|
||||
import tiktoken
|
||||
|
||||
TIKTOKEN_AVAILABLE = True
|
||||
except ImportError:
|
||||
TIKTOKEN_AVAILABLE = False
|
||||
TIKTOKEN_AVAILABLE = importlib.util.find_spec("tiktoken") is not None
|
||||
|
||||
LITELLM_AVAILABLE = importlib.util.find_spec("litellm") is not None
|
||||
|
||||
|
|
@ -285,12 +280,21 @@ def _check_pricing_staleness() -> str | None:
|
|||
|
||||
@lru_cache(maxsize=8)
|
||||
def _get_encoding(encoding_name: str) -> Any:
|
||||
"""Get tiktoken encoding, cached."""
|
||||
"""Get tiktoken encoding, cached.
|
||||
|
||||
Routes through the bounded loader so a stalled vocab download raises
|
||||
:class:`~headroom.tokenizers.tiktoken_counter.TiktokenLoadError` after a
|
||||
timeout instead of hanging the caller indefinitely — ``tiktoken`` fetches
|
||||
vocabularies with no network timeout, and this runs on whatever thread
|
||||
first counts tokens for a model, including proxy startup (GH #956).
|
||||
"""
|
||||
if not TIKTOKEN_AVAILABLE:
|
||||
raise RuntimeError(
|
||||
"tiktoken is required for OpenAI provider. Install with: pip install tiktoken"
|
||||
)
|
||||
return tiktoken.get_encoding(encoding_name)
|
||||
from ..tokenizers.tiktoken_counter import load_encoding
|
||||
|
||||
return load_encoding(encoding_name)
|
||||
|
||||
|
||||
def _lookup_encoding_name(model: str, custom_encodings: dict[str, str] | None = None) -> str | None:
|
||||
|
|
@ -343,6 +347,8 @@ class OpenAITokenCounter:
|
|||
|
||||
Raises:
|
||||
RuntimeError: If tiktoken is not installed.
|
||||
TiktokenLoadError: If the encoding's vocabulary can't be loaded
|
||||
within the bounded timeout (e.g. stalled download).
|
||||
"""
|
||||
self.model = model
|
||||
encoding_name = _get_encoding_name_for_model(model, custom_encodings)
|
||||
|
|
@ -507,7 +513,9 @@ class OpenAIProvider(Provider):
|
|||
the proxy pipeline resolves its tokenizer through this provider while
|
||||
handlers resolve through the tokenizer registry, the two disagree about
|
||||
the same request — savings become a difference of two rulers. Defer to
|
||||
the registry so each model has exactly one tokenizer.
|
||||
the registry so each model has exactly one tokenizer. For OpenAI models,
|
||||
fall back to estimation when a vocabulary cannot load within the bounded
|
||||
timeout; the cached fallback prevents subsequent requests from blocking.
|
||||
"""
|
||||
if model not in self._token_counters:
|
||||
if _lookup_encoding_name(model, self._encodings) is None:
|
||||
|
|
@ -515,9 +523,19 @@ class OpenAIProvider(Provider):
|
|||
|
||||
self._token_counters[model] = cast(Any, get_tokenizer(model))
|
||||
else:
|
||||
self._token_counters[model] = OpenAITokenCounter(
|
||||
model=model, custom_encodings=self._encodings
|
||||
)
|
||||
from ..tokenizers.tiktoken_counter import TiktokenLoadError
|
||||
|
||||
try:
|
||||
self._token_counters[model] = OpenAITokenCounter(
|
||||
model=model, custom_encodings=self._encodings
|
||||
)
|
||||
except TiktokenLoadError as exc:
|
||||
logger.warning(
|
||||
"tiktoken unavailable for %s (%s); using estimation.", model, exc
|
||||
)
|
||||
from ..tokenizers.estimator import EstimatingTokenCounter
|
||||
|
||||
self._token_counters[model] = EstimatingTokenCounter()
|
||||
return self._token_counters[model]
|
||||
|
||||
def get_context_limit(self, model: str) -> int:
|
||||
|
|
|
|||
|
|
@ -130,3 +130,52 @@ class TestEncodingSelection:
|
|||
# Unknown models now get a fallback encoding instead of raising
|
||||
encoding = _get_encoding_name_for_model("completely-unknown")
|
||||
assert encoding == "o200k_base" # Default fallback
|
||||
|
||||
|
||||
class TestGuardedEncodingLoad:
|
||||
"""The provider must never hang on tiktoken's unbounded vocab download.
|
||||
|
||||
Regression for the OpenAI-provider hole in GH #956: `_get_encoding` called
|
||||
`tiktoken.get_encoding` directly, so a stalled vocab download blocked the
|
||||
calling thread (proxy startup included) forever instead of timing out.
|
||||
"""
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _clear_encoding_cache(self):
|
||||
from headroom.providers import openai as openai_module
|
||||
|
||||
openai_module._get_encoding.cache_clear()
|
||||
yield
|
||||
openai_module._get_encoding.cache_clear()
|
||||
|
||||
def test_get_encoding_routes_through_bounded_loader(self, monkeypatch):
|
||||
from headroom.providers.openai import OpenAITokenCounter
|
||||
from headroom.tokenizers import tiktoken_counter
|
||||
|
||||
seen: list[str] = []
|
||||
|
||||
def fake_load_encoding(name: str):
|
||||
seen.append(name)
|
||||
raise tiktoken_counter.TiktokenLoadError(f"{name} load timed out")
|
||||
|
||||
monkeypatch.setattr(tiktoken_counter, "load_encoding", fake_load_encoding)
|
||||
with pytest.raises(tiktoken_counter.TiktokenLoadError):
|
||||
OpenAITokenCounter(model="gpt-4o")
|
||||
assert seen == ["o200k_base"]
|
||||
|
||||
def test_get_token_counter_falls_back_to_estimation(self, monkeypatch):
|
||||
from headroom.providers.openai import OpenAIProvider
|
||||
from headroom.tokenizers import tiktoken_counter
|
||||
from headroom.tokenizers.estimator import EstimatingTokenCounter
|
||||
|
||||
def fake_load_encoding(name: str):
|
||||
raise tiktoken_counter.TiktokenLoadError(f"{name} load timed out")
|
||||
|
||||
monkeypatch.setattr(tiktoken_counter, "load_encoding", fake_load_encoding)
|
||||
provider = OpenAIProvider()
|
||||
counter = provider.get_token_counter("gpt-4o")
|
||||
assert isinstance(counter, EstimatingTokenCounter)
|
||||
assert counter.count_text("hello world") > 0
|
||||
# Cached per model: later requests reuse the fallback instead of
|
||||
# re-blocking on the failed download.
|
||||
assert provider.get_token_counter("gpt-4o") is counter
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue