fix(caching): key response cache by router model group in litellm_metadata (#44542)

* test: align integration fixtures with current behavior

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(integration): wait for the spend flush before asserting its trace placement

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(caching): key response-cache entries by model group from litellm_metadata

/v1/responses routes through the router with model_group in litellm_metadata, which the cache key ignored, so identical requests to different model groups sharing one underlying model hit each other's cached responses

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

---------

Co-authored-by: mateo <mateo@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-10-05 17:06:21 +00:00 • committed by GitHub
parent eeb192d4a7
commit 3286782dea
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
12 changed files with 108 additions and 24 deletions

View file

@ -446,11 +446,19 @@ class Cache:
2. Else if a model_group is set, then return the model_group as the model. This is used for all requests sent through the litellm.Router()
3. Else use the `model` passed in kwargs
"""
metadata: Final[dict] = kwargs.get("metadata", {}) or {}
litellm_params: Final[dict] = kwargs.get("litellm_params", {}) or {}
metadata_in_litellm_params: Final[dict] = litellm_params.get("metadata", {}) or {}
model_group: Final[str | None] = metadata.get("model_group") or metadata_in_litellm_params.get("model_group")
caching_group: Final = self._get_caching_group(metadata, model_group)
metadata_sources: Final[tuple[dict, ...]] = (
kwargs.get("metadata") or {},
kwargs.get("litellm_metadata") or {},
litellm_params.get("metadata") or {},
litellm_params.get("litellm_metadata") or {},
)
model_group: Final[str | None] = next(
(source["model_group"] for source in metadata_sources if source.get("model_group")), None
)
caching_group: Final = next(
(group for source in metadata_sources if (group := self._get_caching_group(source, model_group))), None
)
return caching_group or model_group or kwargs["model"]
def _get_caching_group(self, metadata: dict, model_group: str | None) -> str | None:

View file

@ -66,7 +66,7 @@ def test_a_team_model_is_listed_and_served_only_for_keys_of_its_team(gateway: Ga
def _v2_team_public_names(gateway: Gateway, key: str, model: str) -> list[JsonValue]:
response: Final = gateway.request("GET", "/v2/model/info", key=key, params={"model_name": model})
response: Final = gateway.request("GET", "/v2/model/info", key=key, params={"model": model})
assert response.status_code == 200, response.text
return [entry["model_info"].get("team_public_model_name") for entry in response.json()["data"]]

View file

@ -30790,7 +30790,8 @@
"role": "user",
"content": "proxy behaviour probe"
}
]
],
"max_tokens": 412
},
"response": {
"content_type": "application/json",

View file

@ -281,6 +281,8 @@ def test_case_bills_expected_cost(gateway: Gateway, case: CostTrackingTestCase)
{
"path": f"/{scenario_id}/v1/decisions",
"authorization": "Bearer sk-scripted-provider",
"method": "POST",
"api_key": "",
"body": {
"model": "pplx-decider-v1-27b",
"state": {"source": "cost-tracking"},

View file

@ -602,14 +602,23 @@ def test_two_different_after_values_forward_the_last_one(gateway: Gateway) -> No
assert _ids(page) == (managed_b,), page
@pytest.mark.parametrize("status", (401, 404, 500))
def test_provider_errors_reach_the_caller_and_other_models_keep_mapping(gateway: Gateway, status: int) -> None:
@pytest.mark.parametrize(
("status", "expected_error_type"),
((401, "authentication_error"), (404, "invalid_request_error"), (500, "internal_server_error")),
)
def test_provider_errors_reach_the_caller_and_other_models_keep_mapping(
gateway: Gateway, status: int, expected_error_type: str
) -> None:
message: Final = f"provider refused listing {uuid.uuid4().hex[:8]}"
with _rig(gateway, listing=_error_listing(status, message)) as failing, _rig(gateway, "a.txt") as healthy:
member: Final = _member(failing.scenario, failing.model, healthy.model)
managed_a: Final = healthy.upload(member.key, "a.txt")
failed: Final = failing.list(member.key, {"model": failing.model})
assert _json(failed) == _provider_error(status, message), failed.text
assert failed.status_code == status, failed.text
error: Final = object_value(_json(failed)["error"])
assert message in string_value(error["message"]), failed.text
assert error["code"] == str(status), failed.text
assert error["type"] == expected_error_type, failed.text
assert len(failing.list_requests()) == 1
assert _ids(healthy.listed(member.key)) == (managed_a,)
liveliness: Final = gateway.request("GET", "/health/liveliness")

View file

@ -170,7 +170,9 @@ def _assert_tenant_keeps_redis_without_postgres(
tenant_start, _ = recorded_spans(audit_sinks.tenant)
operator_start, _ = recorded_spans(audit_sinks.operator)
traffic: Final = _drive(candidate, langfuse_vars)
_await_db_span(audit_sinks.operator, None, "batch_write_to_db", seconds=60, since=operator_start)
_await_db_span(
audit_sinks.operator, None, "postgres.update LiteLLM_VerificationToken", seconds=60, since=operator_start
)
tenant_trace: Final = _trace_id(audit_sinks.tenant, traffic)
_await_db_span(audit_sinks.tenant, tenant_trace, "redis", seconds=60)
systems: Final = _db_systems(_trace_spans(audit_sinks.tenant, tenant_trace, seconds=15))
@ -232,7 +234,9 @@ def test_excluded_services_drops_db_spans_at_tenant_only(
_, all_tenant = recorded_spans(audit_sinks.tenant, ten_start)
names: Final = sorted(str(span["name"]) for span in all_tenant)
assert _db_systems(all_tenant) == set(), f"aux db spans reached tenant: {names}"
assert not any("batch_write_to_db" in name for name in names), f"spend writer reached tenant: {names}"
assert not any("postgres.update LiteLLM_VerificationToken" in name for name in names), (
f"spend writer reached tenant: {names}"
)
@pytest.mark.timeout(180)
@ -247,8 +251,8 @@ def test_without_excluded_services_the_tenant_still_gets_redis_and_postgres_span
with owned_proxy(gateway, tmp_path, {"LITELLM_OTEL_V2": "1"}, config=config, workers=2) as candidate:
tenant_start, _ = recorded_spans(audit_sinks.tenant)
traffic: Final = _drive(candidate, langfuse_vars)
_await_db_span(audit_sinks.tenant, None, "batch_write_to_db", seconds=60, since=tenant_start)
tenant_trace: Final = _trace_id(audit_sinks.tenant, traffic)
_await_db_span(audit_sinks.tenant, tenant_trace, "postgresql", seconds=60, since=tenant_start)
_await_db_span(audit_sinks.tenant, tenant_trace, "redis", seconds=60)
_assert_core_spans_present(_trace_spans(audit_sinks.tenant, tenant_trace, seconds=15))
_, all_tenant = recorded_spans(audit_sinks.tenant, tenant_start)
@ -574,7 +578,7 @@ def test_bogus_excluded_services_env_logs_and_drops_without_otel_callback(
_assert_tenant_keeps_redis_without_postgres(owned.gateway, audit_sinks, langfuse_vars)
def test_postgres_exclusion_covers_batch_write_to_db(
def test_postgres_exclusion_covers_spend_flush(
gateway: Gateway,
audit_sinks: SpanSinks,
otel_audit_config: AuditConfigWriter,
@ -586,11 +590,19 @@ def test_postgres_exclusion_covers_batch_write_to_db(
op_start, _ = recorded_spans(audit_sinks.operator)
ten_start, _ = recorded_spans(audit_sinks.tenant)
traffic: Final = _drive(candidate, langfuse_vars)
_await_db_span(audit_sinks.operator, None, "batch_write_to_db", seconds=60, since=op_start)
_await_db_span(audit_sinks.operator, None, "postgres.update LiteLLM_VerificationToken", seconds=60, since=op_start)
operator_trace: Final = _trace_id(audit_sinks.operator, traffic)
_, operator_spans = recorded_spans(audit_sinks.operator, op_start)
assert any(
span["name"] == "postgres.update LiteLLM_VerificationToken" and span["trace_id"] != operator_trace
for span in operator_spans
), tuple((span["name"], span["trace_id"]) for span in operator_spans)
tenant_trace: Final = _trace_id(audit_sinks.tenant, traffic)
_await_db_span(audit_sinks.tenant, tenant_trace, "redis", seconds=60)
tenant_spans: Final = _trace_spans(audit_sinks.tenant, tenant_trace, seconds=15)
_, all_tenant = recorded_spans(audit_sinks.tenant, ten_start)
names: Final = sorted(str(span["name"]) for span in all_tenant)
assert "redis" in _db_systems(tenant_spans), f"redis spans missing at tenant: {names}"
assert not any("batch_write_to_db" in name for name in names), f"spend writer reached tenant: {names}"
assert not any("postgres.update LiteLLM_VerificationToken" in name for name in names), (
f"spend writer reached tenant: {names}"
)

View file

@ -4,7 +4,7 @@ from pathlib import Path
from typing import Final
import pytest
from integration._support.client import Gateway, eventually, gateway_from_environment
from integration._support.client import Gateway, eventually, gateway_from_environment, object_value
from integration._support.otlp_sink import Span, SpanSinks, recorded_spans
from integration._support.process import owned_proxy
from pydantic import JsonValue
@ -25,7 +25,7 @@ def _traces(spans: tuple[Span, ...]) -> dict[str, frozenset[str]]:
return {trace: frozenset(span["name"] for span in spans if span["trace_id"] == trace) for trace in trace_ids}
def test_default_otel_logger_puts_datastore_model_and_spend_writer_spans_in_the_request_trace(
def test_default_otel_logger_keeps_spend_flush_outside_the_request_trace(
gateway: Gateway, audit_sinks: SpanSinks, otel_audit_config: AuditConfigWriter, tmp_path: Path
) -> None:
config: Final = otel_audit_config(tmp_path, {})
@ -41,7 +41,7 @@ def test_default_otel_logger_puts_datastore_model_and_spend_writer_spans_in_the_
key=key,
)
assert response.status_code == 200, response.text
expected: Final = frozenset({"postgres", "redis", "raw_gen_ai_request", "batch_write_to_db"})
expected: Final = frozenset({"postgres", "redis", "raw_gen_ai_request"})
traces: Final = eventually(
lambda: _traces(recorded_spans(audit_sinks.operator, start)[1]),
lambda grouped: any(expected <= names for names in grouped.values()),
@ -51,3 +51,22 @@ def test_default_otel_logger_puts_datastore_model_and_spend_writer_spans_in_the_
assert any(expected <= names for names in traces.values()), {
trace: sorted(names) for trace, names in traces.items()
}
request_trace: Final = next(trace for trace, names in traces.items() if expected <= names)
key_info: Final = eventually(
lambda: candidate.request("GET", "/key/info", key=key, params={"key": key}),
lambda response: response.status_code == 200
and float(str(object_value(response.json()["info"])["spend"])) > 0,
seconds=60,
)
assert key_info.status_code == 200, key_info.text
assert float(str(object_value(key_info.json()["info"])["spend"])) > 0
request_spans: Final = tuple(
span for span in recorded_spans(audit_sinks.operator, start)[1] if span["trace_id"] == request_trace
)
assert not any(span["name"] == "batch_write_to_db" for span in request_spans), request_spans
assert not any(
span["name"] == "postgres"
and span["attributes"].get("call_type") == "commit_spend_updates"
and span["attributes"].get("table_name") == "LiteLLM_VerificationToken"
for span in request_spans
), request_spans

View file

@ -70,8 +70,11 @@ def _clear_observations(upstream: httpx.Client) -> None:
def _observed_request_body(upstream: httpx.Client) -> dict[str, JsonValue]:
observations: Final = JSON_OBJECT.validate_json(upstream.get("/__observations").content)["requests"]
assert isinstance(observations, list)
assert len(observations) == 1
return object_value(object_value(observations[0])["body"])
post_observations: Final = tuple(
observation for observation in observations if isinstance(observation, dict) and observation.get("method") == "POST"
)
assert len(post_observations) == 1
return object_value(object_value(post_observations[0])["body"])
@pytest.mark.parametrize(

View file

@ -81,6 +81,8 @@ def test_config_deployment_dropped_by_stale_boot_cost_map_is_restored_after_relo
{
"path": "/v1/chat/completions",
"authorization": "Bearer sk-upstream",
"method": "POST",
"api_key": "",
"body": {"model": model, "messages": [{"role": "user", "content": "stale cost map control"}]},
}
]

View file

@ -123,18 +123,19 @@ def test_documented_header_and_body_tags_reach_recorded_branch_and_pr_cost(gatew
({"tags": tags}, {}),
({}, {"x-litellm-tags": ", ".join(tags + tags)}),
)
for payload, headers in examples:
for index, (payload, headers) in enumerate(examples):
response: Final = gateway.request(
"POST",
"/v1/chat/completions",
{
"model": model,
"messages": [{"role": "user", "content": "tag attribution"}],
"messages": [{"role": "user", "content": f"tag attribution {index}"}],
**payload,
},
headers=headers,
)
assert response.status_code == 200, response.text
assert len(upstream.drain()) == 3
rows: Final = eventually(
lambda: read_rows(
'SELECT spend FROM "LiteLLM_SpendLogs" WHERE request_tags @> %s::jsonb', (json.dumps(tags),)

View file

@ -251,6 +251,33 @@ def test_get_model_param_value():
assert cache._get_model_param_value(kwargs) == "not-in-caching-group-gpt-3.5-turbo"
def test_get_model_param_value_reads_model_group_from_litellm_metadata():
cache = Cache()
request = {
"model": "openai/gpt-5.6",
"input": "search this text",
"tools": [{"type": "web_search_preview", "search_context_size": "medium"}],
}
assert cache._get_model_param_value({**request, "litellm_metadata": {"model_group": "group-a"}}) == "group-a"
assert (
cache._get_model_param_value({**request, "litellm_params": {"litellm_metadata": {"model_group": "group-a"}}})
== "group-a"
)
assert cache._get_model_param_value(
{
**request,
"litellm_metadata": {
"model_group": "group-a",
"caching_groups": [("group-a", "group-b")],
},
}
) == "('group-a', 'group-b')"
assert cache.get_cache_key(**request, litellm_metadata={"model_group": "group-a"}) != cache.get_cache_key(
**request, litellm_metadata={"model_group": "group-b"}
)
def test_preset_cache_key():
"""
Test that the preset cache key is used if it is set in kwargs["litellm_params"]

View file

@ -103,7 +103,7 @@ def test_reload_model_cost_map_surfaces_the_blob_id_of_the_bytes_served_on_every
import httpx
import litellm
from litellm.litellm_core_utils.get_model_cost_map import git_blob_id
from litellm.litellm_core_utils.get_model_cost_map import _finalize_model_cost_map, git_blob_id
from litellm.proxy import proxy_server as ps
from litellm.proxy._types import LitellmUserRoles
@ -142,7 +142,7 @@ def test_reload_model_cost_map_surfaces_the_blob_id_of_the_bytes_served_on_every
assert {key: status_response.json()[key] for key in expected} == expected
assert public_response.status_code == 200
assert "gpt-4o" in public_response.json()
assert reload_body["models_count"] == len(litellm.model_cost)
assert reload_body["models_count"] == len(_finalize_model_cost_map(json.loads(body)))
def test_reload_model_cost_map_fetch_failure_502_keeps_map(