fix(cache_warming): give the gates the same body the original request had

Two findings with one shape: the ceilings warming routes through read their inputs
off the request body, so anything the original request carried that a replay does
not is not refused by those ceilings, it is invisible to them.

Caller tags were dropped. Tag budgets, tag budget reservation and the limiter's tag
descriptors all resolve tags through one owner, get_tags_from_request_body, reading
them off the body. Warming already enters through all three entry points, so the
fix is not a tag check of its own; it is capturing the caller's tags and presenting
them again, which is what those owners were always going to read. The warming
marker stays on spend_logs_metadata rather than joining them, because tags feed
deployment selection and operators can forbid requests from carrying them.

The session id is caller-controlled and was embedded verbatim in the record key,
the index member, the touched key and every warmth key, none of which
max_payload_bytes bounds, so a caller could hold far more Redis memory than the
payload bound implies. It is now hashed inside record_key, which is the one place
capture and the warm-aware pick both derive keys from, so they cannot disagree; the
record body keeps the real value because a replay carries it for deployment
affinity to pin on. A session id past a sane ceiling is not captured at all rather
than truncated, since a truncated one would pin against traffic that does not exist.
This commit is contained in:
Tin Chi Lo 2026-07-30 19:58:08 -07:00
parent 5bfd6432fd
commit ab8b034e0f
7 changed files with 71 additions and 3 deletions

View file

@ -28,6 +28,8 @@ if TYPE_CHECKING:
from litellm.router_strategy.complexity_router.complexity_router import ComplexityRouter
_MAX_UNCOMPRESSED_RATIO = 8
# A session id rides in the record body and in every derived Redis key; the payload bound does not cover it
MAX_SESSION_ID_CHARS = 256
# The partition a record is stored under and the identity it is attributed to must cover the same
# dimensions, or two principals differing only in an omitted one share a record. Both are derived from
# proxy_identity_fields() for that reason; user_api_key is the key-state lookup handle rather than a
@ -35,6 +37,16 @@ _MAX_UNCOMPRESSED_RATIO = 8
_ATTRIBUTION_KEYS = ("user_api_key", *proxy_identity_fields())
@lru_cache(maxsize=64)
def _warn_session_id_too_long(auto_router_model_name: str) -> None:
verbose_router_logger.warning(
"cache_warming: auto-router %s saw a session_id longer than %s characters; it is caller-controlled "
"and is retained on the record, so such sessions are not captured",
auto_router_model_name,
MAX_SESSION_ID_CHARS,
)
@lru_cache(maxsize=64)
def _warn_privacy_gate_blocked(auto_router_model_name: str) -> None:
verbose_router_logger.warning(
@ -89,6 +101,18 @@ def _capture_allowed(kwargs: Mapping[str, object]) -> bool:
return _should_store_prompts_and_responses_in_spend_logs()
def _caller_tags(metadata_dicts: Sequence[Mapping[str, object]]) -> tuple[str, ...]:
"""The caller's own request tags, kept so a replay presents the same body the gates read them from.
Tag budgets, tag budget reservation and the limiter's tag descriptors all resolve tags through
get_tags_from_request_body, so a replay that drops them is not refused by those ceilings, it is simply
invisible to them."""
for metadata in metadata_dicts:
tags = metadata.get("tags")
if isinstance(tags, list):
return tuple(tag for tag in tags if isinstance(tag, str))
return ()
def _is_replay(metadata_dicts: Sequence[Mapping[str, object]]) -> bool:
for metadata in metadata_dicts:
if metadata.get(CACHE_WARMING_REPLAY_MARKER_KEY):
@ -204,6 +228,9 @@ async def capture_session(
session_id = get_request_metadata_field(request_kwargs, "session_id")
if session_id is None:
return
if len(session_id) > MAX_SESSION_ID_CHARS:
_warn_session_id_too_long(strategy.model_name)
return
payload = _build_payload(request_kwargs, messages, _call_surface(request_kwargs), routed_model)
if payload is None:
return
@ -227,6 +254,7 @@ async def capture_session(
await store.upsert_session(
caller_scope=caller_scope,
session_id=session_id,
tags=_caller_tags(metadata_dicts),
payload_compressed=blob,
payload_sha256=sha,
token_estimate=token_estimate,

View file

@ -152,6 +152,7 @@ def _replay_body(
data[get_metadata_variable_name_from_kwargs(data)] = { # mutable-ok: request metadata, never retained
CACHE_WARMING_REPLAY_MARKER_KEY: True,
**({"session_id": record.session_id} if record.session_id is not None else {}),
**({"tags": list(record.tags)} if record.tags else {}),
"spend_logs_metadata": {CACHE_WARMING_REPLAY_TAG: "true"}, # mutable-ok: request metadata, never retained
}
return data

View file

@ -1,3 +1,4 @@
import hashlib
import time
from collections.abc import Awaitable, Mapping
from functools import lru_cache
@ -123,8 +124,16 @@ class CacheWarmingStore:
def record_key(auto_router_model_name: str, caller_scope: str, session_id: str) -> str:
"""Scoped by auto-router as well as caller and session. The record hash is already per-router through
its hash-tagged container, but warmth keys are derived from this identity and live at the top level,
so without the router in it two warming auto-routers sharing one Redis read each other's warmth."""
return f"{auto_router_model_name}:{caller_scope}:{session_id}"
so without the router in it two warming auto-routers sharing one Redis read each other's warmth.
The session id is caller-controlled and would otherwise be embedded verbatim in this key, the index
member, the touched key and every warmth key, none of which the payload bound covers, so a caller
could hold far more Redis memory than max_payload_bytes implies. Hashing it here rather than at the
call sites is what keeps capture and the warm-aware pick agreeing on the same key; the record body
keeps the real value, which is what a replay carries so deployment affinity pins it alongside the
caller's own traffic."""
session_digest = hashlib.sha256(session_id.encode("utf-8")).hexdigest()[:32]
return f"{auto_router_model_name}:{caller_scope}:{session_digest}"
@staticmethod
def touched_key(record_key: str) -> str:
@ -180,6 +189,7 @@ class CacheWarmingStore:
payload_sha256: str,
token_estimate: int,
served_model: str,
tags: tuple[str, ...],
attribution: CacheWarmingAttribution,
ttl_seconds: int,
max_sessions: int,
@ -197,6 +207,7 @@ class CacheWarmingStore:
last_activity=now,
served_model=served_model,
session_id=session_id,
tags=tags,
attribution=attribution,
auto_router_model_name=self.auto_router_model_name,
)

View file

@ -80,6 +80,7 @@ class CacheWarmingRecord(BaseModel):
last_activity: float
served_model: str
session_id: str | None = None
tags: tuple[str, ...] = ()
attribution: CacheWarmingAttribution
auto_router_model_name: str

View file

@ -518,3 +518,21 @@ async def test_a_session_is_dropped_once_its_key_leaves_the_tenant_it_was_captur
keys = FakeKeyDirectory({"k": key_state(token="k", team_id=current_team)})
await tick(llm_router, active=refresher(keys=keys))
assert llm_router.completion_calls == []
@pytest.mark.asyncio
async def test_a_replay_presents_the_callers_own_tags_to_the_gates_that_read_them():
"""Tag budgets, tag budget reservation and the limiter's tag descriptors all resolve tags through one
owner, get_tags_from_request_body, which reads them off the request body. Warming already enters through
all three, so a replay that drops the caller's tags is not refused by those ceilings, it is invisible to
them, and its spend lands outside the tag it belongs to. The marker stays out of this channel because
tags feed deployment selection."""
from litellm.proxy.common_utils.http_parsing_utils import get_tags_from_request_body
llm_router, redis = warming_rig(redis=FakeRedisCache())
seed_session(redis, tags=("cost-center-7",), touched=_VISITED_BOTH_TIERS)
await tick(llm_router)
assert llm_router.completion_calls, "expected a replay"
for call in llm_router.completion_calls:
assert get_tags_from_request_body(call) == ["cost-center-7"]
assert CACHE_WARMING_REPLAY_TAG not in get_tags_from_request_body(call)

View file

@ -109,6 +109,7 @@ def _record_json(**overrides: object) -> str:
token_estimate=2048,
last_activity=1000.0,
served_model="sonnet",
tags=(),
attribution=CacheWarmingAttribution(user_api_key="hashed"),
auto_router_model_name="smart-router",
).model_dump()
@ -127,6 +128,7 @@ async def _upsert(store: CacheWarmingStore, session_id: str = "s1", max_sessions
payload_sha256="sha2",
token_estimate=4096,
served_model="sonnet",
tags=(),
attribution=CacheWarmingAttribution(),
ttl_seconds=1800,
max_sessions=max_sessions,
@ -138,7 +140,6 @@ def test_key_shapes_are_scoped_and_hash_tagged():
so two warming auto-routers on one Redis cannot read each other's warmth and a session's record, index
entry and warmth stamps stay on one node."""
record = CacheWarmingStore.record_key("smart-router", "keyhash", "session-1")
assert record == "smart-router:keyhash:session-1"
other_router = CacheWarmingStore.record_key("other-router", "keyhash", "session-1")
assert CacheWarmingStore.warmth_key(record, "opus") != CacheWarmingStore.warmth_key(other_router, "opus")
store = _store(None)
@ -147,6 +148,12 @@ def test_key_shapes_are_scoped_and_hash_tagged():
assert store.index_key() == f"{slot}:index"
assert CacheWarmingStore.warmth_key(record, "opus").startswith(f"{slot}:")
# the session id is caller-controlled and reaches this key, the index member, the touched key and every
# warmth key, none of which max_payload_bytes bounds, so it is hashed rather than embedded
huge = CacheWarmingStore.record_key("smart-router", "keyhash", "s" * 10_000)
assert len(huge) == len(record) and "s" * 100 not in huge
assert record != CacheWarmingStore.record_key("smart-router", "keyhash", "session-2")
@pytest.mark.asyncio
async def test_cap_enforced_atomically_with_the_record_write():

View file

@ -221,6 +221,7 @@ def seed_session(
org_id: str | None = None,
project_id: str | None = None,
touched: tuple[str, ...] | None = None,
tags: tuple[str, ...] = (),
) -> str:
payload = CacheWarmingPayload(
model=served_model,
@ -239,6 +240,7 @@ def seed_session(
last_activity=last_activity if last_activity is not None else time.time(),
served_model=served_model,
session_id=session_id,
tags=tags,
attribution=CacheWarmingAttribution(
user_api_key=user_api_key,
user_api_key_team_id=team_id,