mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
parent
eeb192d4a7
commit
3286782dea
12 changed files with 108 additions and 24 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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"]]
|
||||
|
||||
|
|
|
|||
|
|
@ -30790,7 +30790,8 @@
|
|||
"role": "user",
|
||||
"content": "proxy behaviour probe"
|
||||
}
|
||||
]
|
||||
],
|
||||
"max_tokens": 412
|
||||
},
|
||||
"response": {
|
||||
"content_type": "application/json",
|
||||
|
|
|
|||
|
|
@ -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"},
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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"}]},
|
||||
}
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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),)
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue