diff --git a/headroom/proxy/savings_tracker.py b/headroom/proxy/savings_tracker.py index 4b7d82c44..ccffb81c1 100644 --- a/headroom/proxy/savings_tracker.py +++ b/headroom/proxy/savings_tracker.py @@ -16,6 +16,7 @@ import tempfile import threading from csv import DictWriter from datetime import datetime, timedelta, timezone +from functools import lru_cache from io import StringIO from pathlib import Path from typing import Any @@ -164,6 +165,23 @@ def _normalize_model(value: Any) -> str: return cleaned or MODEL_UNKNOWN +# `_resolve_litellm_model` is called on every savings-tracking update (i.e. +# every request), and `model` is client-controlled — it comes straight off +# the request body. For a model LiteLLM can't price (a custom / local / +# gateway name), the uncached fallback below calls `litellm.cost_per_token` +# purely to probe resolvability, which prints LiteLLM's noisy "Provider +# List: https://docs.litellm.ai/docs/providers" banner on every failed probe +# (#2851). Cache the resolution per model name so that probe runs at most +# once per distinct model — bounded, not a plain dict: a request-facing +# proxy must not let a caller grow an unbounded cache for free by sending a +# fresh model string on every request. `maxsize` caps memory; LRU eviction +# means a model that stops being sent eventually falls out and simply +# re-probes if it's ever sent again — never a correctness issue, only +# whether the probe (and its noisy failure banner) reruns. +_MODEL_RESOLUTION_CACHE_MAXSIZE = 256 + + +@lru_cache(maxsize=_MODEL_RESOLUTION_CACHE_MAXSIZE) def _resolve_litellm_model(model: str) -> str: """Resolve model name to one LiteLLM recognizes. @@ -173,6 +191,12 @@ def _resolve_litellm_model(model: str) -> str: "claude-opus" identically to the live /stats path. Uses the shared result only when it maps to a priced model_cost key; otherwise falls through to the bare-prefix logic below. Fail-soft: pricing never breaks bookkeeping. + + Bounded LRU cache, keyed by model name — see + ``_MODEL_RESOLUTION_CACHE_MAXSIZE`` above. Tests that mock the LiteLLM + module across calls with the same model name must call + ``_resolve_litellm_model.cache_clear()`` between cases, or results from + an earlier case leak in. """ litellm = _get_litellm_module() if litellm is None: diff --git a/tests/conftest.py b/tests/conftest.py index 2aecdbc4a..2a75bd260 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -91,6 +91,29 @@ def _reset_copilot_routing_flag(): reset_request_routed_to_copilot() +# `savings_tracker._resolve_litellm_model` is an `lru_cache`d, module-global, +# process-lifetime cache keyed by model name (bounded — see #2860). Many test +# files monkeypatch `savings_tracker.litellm` to a fake with different +# `model_cost`/`cost_per_token` behavior per test, but reuse common model +# names like "gpt-4o" across them. Without a reset, whichever test resolves +# "gpt-4o" first "wins" the cache entry for the rest of the run, and later +# tests silently stop exercising their own fake — a real-not-hypothetical +# order-dependence bug once the cache is process-lifetime instead of per-call. +# Clear before AND after so a test's own within-test resolutions never leak +# in from, or leak out to, a neighboring test either. +@pytest.fixture(autouse=True) +def _reset_litellm_model_resolution_cache(): + try: + from headroom.proxy.savings_tracker import _resolve_litellm_model + except ModuleNotFoundError: + yield + return + + _resolve_litellm_model.cache_clear() + yield + _resolve_litellm_model.cache_clear() + + # ============================================================================= # Global test hooks # ============================================================================= diff --git a/tests/test_savings_tracker_litellm_resolution_cache.py b/tests/test_savings_tracker_litellm_resolution_cache.py new file mode 100644 index 000000000..04300bf12 --- /dev/null +++ b/tests/test_savings_tracker_litellm_resolution_cache.py @@ -0,0 +1,88 @@ +"""Regression: `_resolve_litellm_model`'s cache must be bounded (PR #2860 review). + +A plain unbounded dict cache keyed by a client-controlled model string is a +memory-retention path on a request-facing proxy: a caller can grow it without +limit by sending a new model name on every request. The fix uses a bounded +`functools.lru_cache`. These tests pin the three properties that actually +matter, independent of the litellm pricing behavior covered elsewhere: + +- repeated resolution of the same unresolvable model only probes litellm once +- the cache never grows past its bound, no matter how many distinct model + names get resolved +- an evicted name is transparently re-probed (never silently wrong or stuck) + rather than growing the cache further +""" + +from __future__ import annotations + +import types + +from headroom.proxy import savings_tracker as st + + +def _fake_litellm_always_unresolvable(probe_calls: dict[str, int]) -> types.SimpleNamespace: + """A fake litellm where every model is unpriced and unresolvable. + + `cost_per_token` always raises — exactly what a real custom/local model + litellm has never heard of does — which is the call this cache exists to + memoize (see the comment above `_resolve_litellm_model` in + savings_tracker.py: that raise is also where real litellm prints its + noisy "Provider List" banner, #2851). + """ + + def cost_per_token(*, model, prompt_tokens, completion_tokens): + probe_calls[model] = probe_calls.get(model, 0) + 1 + raise RuntimeError("unknown model") + + return types.SimpleNamespace(model_cost={}, cost_per_token=cost_per_token) + + +def test_resolve_litellm_model_probes_unknown_model_once(monkeypatch): + probe_calls: dict[str, int] = {} + monkeypatch.setattr( + st, "_get_litellm_module", lambda: _fake_litellm_always_unresolvable(probe_calls) + ) + + for _ in range(5): + resolved = st._resolve_litellm_model("widget-local-model") + assert resolved == "widget-local-model" + + assert probe_calls == {"widget-local-model": 1} + + +def test_resolve_litellm_model_cache_is_bounded(monkeypatch): + probe_calls: dict[str, int] = {} + monkeypatch.setattr( + st, "_get_litellm_module", lambda: _fake_litellm_always_unresolvable(probe_calls) + ) + + extra_beyond_bound = 50 + for i in range(st._MODEL_RESOLUTION_CACHE_MAXSIZE + extra_beyond_bound): + st._resolve_litellm_model(f"widget-local-model-{i}") + + info = st._resolve_litellm_model.cache_info() + assert info.maxsize == st._MODEL_RESOLUTION_CACHE_MAXSIZE + # However many distinct names were resolved, the cache itself never + # grows past its bound -- this is the actual memory-retention fix. + assert info.currsize == st._MODEL_RESOLUTION_CACHE_MAXSIZE + + +def test_resolve_litellm_model_evicted_name_reprobes(monkeypatch): + probe_calls: dict[str, int] = {} + monkeypatch.setattr( + st, "_get_litellm_module", lambda: _fake_litellm_always_unresolvable(probe_calls) + ) + + st._resolve_litellm_model("seed-model") + assert probe_calls["seed-model"] == 1 + + # Push exactly `maxsize` new distinct names through without ever touching + # "seed-model" again -- LRU eviction must push it out to make room. + for i in range(st._MODEL_RESOLUTION_CACHE_MAXSIZE): + st._resolve_litellm_model(f"filler-model-{i}") + + # A resolvable name being evicted is not a correctness bug (it just + # re-probes) -- the assertion that matters is that it *does* re-probe + # rather than silently reusing a slot it no longer legitimately owns. + st._resolve_litellm_model("seed-model") + assert probe_calls["seed-model"] == 2