diff --git a/litellm/caching/caching.py b/litellm/caching/caching.py index 12aaa051041..4070f41fd32 100644 --- a/litellm/caching/caching.py +++ b/litellm/caching/caching.py @@ -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: diff --git a/tests/integration/authorization/test_team_scoped_models.py b/tests/integration/authorization/test_team_scoped_models.py index 2b347a6f4f8..aea6b9e4ba4 100644 --- a/tests/integration/authorization/test_team_scoped_models.py +++ b/tests/integration/authorization/test_team_scoped_models.py @@ -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"]] diff --git a/tests/integration/cost_calculation/cost_tracking_cases.json b/tests/integration/cost_calculation/cost_tracking_cases.json index 8242c8683b5..f03893c4244 100644 --- a/tests/integration/cost_calculation/cost_tracking_cases.json +++ b/tests/integration/cost_calculation/cost_tracking_cases.json @@ -30790,7 +30790,8 @@ "role": "user", "content": "proxy behaviour probe" } - ] + ], + "max_tokens": 412 }, "response": { "content_type": "application/json", diff --git a/tests/integration/cost_calculation/test_cost_tracking.py b/tests/integration/cost_calculation/test_cost_tracking.py index cfe0a99ef4b..379b2c13f5b 100644 --- a/tests/integration/cost_calculation/test_cost_tracking.py +++ b/tests/integration/cost_calculation/test_cost_tracking.py @@ -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"}, diff --git a/tests/integration/management/test_vector_store_file_list_managed_ids.py b/tests/integration/management/test_vector_store_file_list_managed_ids.py index 7ad68e884e6..c6124519c9e 100644 --- a/tests/integration/management/test_vector_store_file_list_managed_ids.py +++ b/tests/integration/management/test_vector_store_file_list_managed_ids.py @@ -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") diff --git a/tests/integration/observability/test_otel_excluded_services.py b/tests/integration/observability/test_otel_excluded_services.py index 8768052b488..09f40630391 100644 --- a/tests/integration/observability/test_otel_excluded_services.py +++ b/tests/integration/observability/test_otel_excluded_services.py @@ -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}" + ) diff --git a/tests/integration/observability/test_otel_v1_request_trace.py b/tests/integration/observability/test_otel_v1_request_trace.py index c99ddc11ced..468aceb4860 100644 --- a/tests/integration/observability/test_otel_v1_request_trace.py +++ b/tests/integration/observability/test_otel_v1_request_trace.py @@ -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 diff --git a/tests/integration/pricing/test_per_second_pricing.py b/tests/integration/pricing/test_per_second_pricing.py index ad44a631054..a4958bd18c5 100644 --- a/tests/integration/pricing/test_per_second_pricing.py +++ b/tests/integration/pricing/test_per_second_pricing.py @@ -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( diff --git a/tests/integration/routing/test_stale_cost_map_boot.py b/tests/integration/routing/test_stale_cost_map_boot.py index d2eeb2bdc6c..13b3d777587 100644 --- a/tests/integration/routing/test_stale_cost_map_boot.py +++ b/tests/integration/routing/test_stale_cost_map_boot.py @@ -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"}]}, } ] diff --git a/tests/integration/spend/test_roi_branch_spend.py b/tests/integration/spend/test_roi_branch_spend.py index 9efb22c43be..cc9cfad8cd2 100644 --- a/tests/integration/spend/test_roi_branch_spend.py +++ b/tests/integration/spend/test_roi_branch_spend.py @@ -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),) diff --git a/tests/unit/caching/test_unit_test_caching.py b/tests/unit/caching/test_unit_test_caching.py index fd9f4bb9e89..d720838c507 100644 --- a/tests/unit/caching/test_unit_test_caching.py +++ b/tests/unit/caching/test_unit_test_caching.py @@ -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"] diff --git a/tests/unit/proxy/proxy_server/test_routes_model_cost_map.py b/tests/unit/proxy/proxy_server/test_routes_model_cost_map.py index 40bd66ea91c..31f69d686d4 100644 --- a/tests/unit/proxy/proxy_server/test_routes_model_cost_map.py +++ b/tests/unit/proxy/proxy_server/test_routes_model_cost_map.py @@ -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(