diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index f009390576a..bed2503d211 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -8,6 +8,7 @@ from typing import TYPE_CHECKING, Annotated, Any, Final, Literal, NamedTuple, Ty import httpx from pydantic import ( + AfterValidator, BaseModel, BeforeValidator, ConfigDict, @@ -2743,6 +2744,41 @@ class ScheduledJobStaggerSettings(LiteLLMPydanticObjectBase): ) +SPEND_LOGS_METADATA_ALWAYS_KEPT_FIELDS: Final = frozenset({"status", "cold_storage_object_key"}) + + +def _known_spend_logs_metadata_field(name: str) -> str: + if name not in SpendLogsMetadata.__annotations__: + raise ValueError(f"{name!r} is not a LiteLLM_SpendLogs.metadata field") + return name + + +SpendLogsMetadataFieldName: TypeAlias = Annotated[str, AfterValidator(_known_spend_logs_metadata_field)] + + +class SpendLogsMetadataFields(BaseModel): + model_config = ConfigDict(extra="forbid", frozen=True) + + include: tuple[SpendLogsMetadataFieldName, ...] | None = None + exclude: tuple[SpendLogsMetadataFieldName, ...] | None = None + + @model_validator(mode="after") + def _exactly_one_list(self) -> "SpendLogsMetadataFields": + if (self.include is None) == (self.exclude is None): + raise ValueError("set exactly one of 'include' or 'exclude'") + always_kept_excluded: Final = sorted(SPEND_LOGS_METADATA_ALWAYS_KEPT_FIELDS.intersection(self.exclude or ())) + if always_kept_excluded: + raise ValueError(f"{always_kept_excluded} are always kept and cannot be excluded") + return self + + def keeps(self, name: str) -> bool: + if name in SPEND_LOGS_METADATA_ALWAYS_KEPT_FIELDS: + return True + if self.include is not None: + return name in self.include + return name not in (self.exclude or ()) + + DEFAULT_RESPONSES_WEBSOCKET_SESSION_LIMIT_SECONDS: Final[float] = 3600.0 @@ -3091,6 +3127,10 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase): None, description="If True, stores request messages and responses in spend logs. Default is False.", ) + spend_logs_metadata_fields: SpendLogsMetadataFields | None = Field( + None, + description="Which keys of LiteLLM_SpendLogs.metadata are written to the database. Set exactly one of 'include' (write only these keys) or 'exclude' (drop these keys). 'status' and 'cold_storage_object_key' are always written. Daily spend tables, budgets and logging callbacks still see every key. Unset writes every key", + ) disable_auto_add_proxy_admin_to_teams: bool | None = Field( None, description="By default, the user calling /team/new is automatically added to the new team as a team admin. If True, proxy admins are no longer auto-added; members explicitly listed in members_with_roles are unaffected. Default is False.", diff --git a/litellm/proxy/db/baseline_accounting.py b/litellm/proxy/db/baseline_accounting.py index 486717f9b93..f7a2be654c8 100644 --- a/litellm/proxy/db/baseline_accounting.py +++ b/litellm/proxy/db/baseline_accounting.py @@ -243,18 +243,19 @@ WHERE scope = $1 AND revision > $2::bigint AND ($5::float8 IS NULL OR publication::jsonb->>'status' = 'estimated') ORDER BY started_at, request_id """ +_PUBLISHED_LOG_FIELDS: Final = ("autorouter_savings_estimate", "autorouter_savings") _UPDATE_LOGS: Final = """ WITH changes AS ( SELECT request_id, publication::jsonb AS publication FROM jsonb_to_recordset($1::jsonb) AS x(request_id text, publication jsonb) ) UPDATE "LiteLLM_SpendLogs" AS logs -SET metadata = (COALESCE(logs.metadata::jsonb, '{}'::jsonb) - 'autorouter_baseline_observation') || jsonb_build_object( +SET metadata = (COALESCE(logs.metadata::jsonb, '{}'::jsonb) - 'autorouter_baseline_observation') || (jsonb_build_object( 'autorouter_savings_estimate', changes.publication, 'autorouter_savings', CASE WHEN changes.publication->>'status' = 'estimated' THEN (changes.publication->>'baseline_spend')::float8 - (changes.publication->>'actual_spend')::float8 ELSE NULL END -) +) - ARRAY(SELECT jsonb_array_elements_text($2::jsonb))) FROM changes WHERE logs.request_id = changes.request_id """ _UPDATE_PUBLICATIONS: Final = """ @@ -397,7 +398,11 @@ async def _publish(db: SupportsRawQueries, changes: Sequence[_Change]) -> None: if not changes: return serialized: Final = json.dumps(tuple(change.model_dump(mode="json") for change in changes), separators=(",", ":")) - await db.execute_raw(_UPDATE_LOGS, serialized) + from litellm.proxy.spend_tracking.spend_tracking_utils import configured_spend_logs_metadata_fields + + fields: Final = configured_spend_logs_metadata_fields() + unstored: Final = tuple(name for name in _PUBLISHED_LOG_FIELDS if fields is not None and not fields.keeps(name)) + await db.execute_raw(_UPDATE_LOGS, serialized, json.dumps(unstored)) await db.execute_raw(_UPDATE_SESSIONS, serialized) if any(change.user_id for change in changes): await db.execute_raw(_UPDATE_USER_SESSIONS, serialized) diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index 9fb8cc20981..5bafe19cd0b 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -654,7 +654,14 @@ class DBSpendUpdateWriter: from litellm.repositories.table_repositories import SpendLogsRepository request_id: Final = payload["request_id"] - row: Final = _batch_cost_row_to_write(payload, disable_spend_logs) + from litellm.proxy.spend_tracking.spend_tracking_utils import ( + configured_spend_logs_metadata_fields, + spend_log_row_with_retained_metadata, + ) + + row: Final = spend_log_row_with_retained_metadata( + _batch_cost_row_to_write(payload, disable_spend_logs), configured_spend_logs_metadata_fields() + ) spend_logs: Final = SpendLogsRepository(prisma_client).table try: claimed: Final = await spend_logs.create_many( diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 90e760093b4..e035f061db9 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -6950,6 +6950,13 @@ class ProxyConfig: ).user_api_key_cache_max_size ) + if "spend_logs_metadata_fields" in general_settings: + _ = ConfigGeneralSettings.model_validate( + MappingProxyType( + {"spend_logs_metadata_fields": typed_general_settings["spend_logs_metadata_fields"]} + ) + ) + ### PKCE MULTI-INSTANCE PREREQUISITE CHECK ### # PKCE verifiers are stored in redis_usage_cache when available so they can # be read back by any instance (not just the one that started the auth flow). diff --git a/litellm/proxy/spend_tracking/spend_tracking_utils.py b/litellm/proxy/spend_tracking/spend_tracking_utils.py index 9d52f7d71bb..99bd2fbc480 100644 --- a/litellm/proxy/spend_tracking/spend_tracking_utils.py +++ b/litellm/proxy/spend_tracking/spend_tracking_utils.py @@ -9,7 +9,7 @@ from functools import reduce from types import MappingProxyType from typing import TYPE_CHECKING, Final, Literal, Protocol, cast, runtime_checkable -from pydantic import BaseModel, JsonValue +from pydantic import BaseModel, JsonValue, TypeAdapter, ValidationError import litellm from litellm._logging import verbose_proxy_logger @@ -47,7 +47,12 @@ from litellm.litellm_core_utils.ptu_pricing import azure_spillover from litellm.litellm_core_utils.safe_json_dumps import safe_dumps, strip_null_bytes from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker from litellm.llms.anthropic.common_utils import resolve_used_client_oauth_token -from litellm.proxy._types import SpendLogsMetadata, SpendLogsPayload, SpendLogsRouterMetadata +from litellm.proxy._types import ( + SpendLogsMetadata, + SpendLogsMetadataFields, + SpendLogsPayload, + SpendLogsRouterMetadata, +) from litellm.proxy.route_llm_request import ProxyModelNotFoundError from litellm.proxy.spend_tracking.spend_log_error_logger import spend_log_error from litellm.proxy.utils import PrismaClient, hash_token @@ -1726,6 +1731,35 @@ def should_store_prompts_and_responses_in_spend_logs() -> bool: return get_secret_bool("STORE_PROMPTS_IN_SPEND_LOGS") is True +_SPEND_LOGS_METADATA_FIELDS_ADAPTER: Final[TypeAdapter[SpendLogsMetadataFields | None]] = TypeAdapter( + SpendLogsMetadataFields | None +) +_SPEND_LOGS_METADATA_ADAPTER: Final = TypeAdapter(dict[str, JsonValue]) + + +def configured_spend_logs_metadata_fields() -> SpendLogsMetadataFields | None: + from litellm.proxy.proxy_server import general_settings_view + + try: + return _SPEND_LOGS_METADATA_FIELDS_ADAPTER.validate_python( + general_settings_view().get("spend_logs_metadata_fields") + ) + except ValidationError as e: + verbose_proxy_logger.error("Ignoring invalid general_settings.spend_logs_metadata_fields: %s", e) + return None + + +def spend_log_row_with_retained_metadata( + row: Mapping[str, object], fields: SpendLogsMetadataFields | None +) -> Mapping[str, object]: + metadata_json: Final = row.get("metadata") + if fields is None or not isinstance(metadata_json, str): + return row + metadata: Final = _SPEND_LOGS_METADATA_ADAPTER.validate_json(metadata_json) + retained: Final = {name: value for name, value in metadata.items() if fields.keeps(name)} + return {**row, "metadata": safe_dumps(retained)} + + def _get_status_for_spend_log( metadata: dict, ) -> Literal["success", "failure"]: diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 8f796f065a6..4ff23624c1b 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -7446,6 +7446,13 @@ class ProxyUpdateSpend: "Spend tracking - processing %d spend logs for DB write", len(logs_to_process), ) + from litellm.proxy.spend_tracking.spend_tracking_utils import ( + configured_spend_logs_metadata_fields, + spend_log_row_with_retained_metadata, + ) + + retention: Final = configured_spend_logs_metadata_fields() + rows_to_write: Final = [spend_log_row_with_retained_metadata(row, retention) for row in logs_to_process] start_time: Final = time.time() try: for i in range(n_retry_times + 1): @@ -7455,7 +7462,7 @@ class ProxyUpdateSpend: if not base_url.endswith("/"): base_url += "/" verbose_proxy_logger.debug("base_url: %s", base_url) - json_data = json.dumps(logs_to_process) + json_data = json.dumps(rows_to_write) response = await db_writer_client.post( url=base_url + "spend/update", data=json_data, @@ -7466,8 +7473,8 @@ class ProxyUpdateSpend: # Items already removed from queue at start of function pass else: - for j in range(0, len(logs_to_process), BATCH_SIZE): - batch = logs_to_process[j : j + BATCH_SIZE] + for j in range(0, len(rows_to_write), BATCH_SIZE): + batch = rows_to_write[j : j + BATCH_SIZE] batch_with_dates = [prisma_client.jsonify_object({**entry}) for entry in batch] isolation_budget = MAX_SPEND_LOG_ISOLATION_FAILURES_PER_BATCH for statement_rows in spend_log_write_batches( diff --git a/tests/integration/spend/test_batch_completion_accounting.py b/tests/integration/spend/test_batch_completion_accounting.py index 4ecab10f942..bbc0300b3f0 100644 --- a/tests/integration/spend/test_batch_completion_accounting.py +++ b/tests/integration/spend/test_batch_completion_accounting.py @@ -63,7 +63,7 @@ def _failed_line(index: int) -> str: ) -def _batch_routes(model: str) -> RoutedResponse: +def batch_routes(model: str) -> RoutedResponse: output_lines: Final = ( _succeeded_line(1, model, **FIRST_LINE), _succeeded_line(2, model, **SECOND_LINE), @@ -139,7 +139,7 @@ def _batches_total_tokens(endpoints: dict[str, object] | None) -> int | None: return int(total_tokens) if isinstance(total_tokens, (int, float, str)) else None -def _input_file(model: str) -> bytes: +def batch_input_file(model: str) -> bytes: return ( "\n".join( json.dumps( @@ -166,13 +166,13 @@ def test_completed_batch_spend_row_records_reasoning_tokens_and_error_file_failu with gateway.scenario() as scenario: key: Final = scenario.key() scenario_id: Final = f"batch-accounting-{uuid.uuid4().hex[:12]}" - handle: Final = register_scenario(scenario_id, _batch_routes("gpt-4o-mini")) + handle: Final = register_scenario(scenario_id, batch_routes("gpt-4o-mini")) scenario.cleanups.callback(delete_scenario, handle) model: Final = scenario.model(api_base=handle.api_base()) file_response: Final = gateway.request_multipart( "/v1/files", {"purpose": "batch", "model": model}, - {"file": ("in.jsonl", _input_file(model), "application/jsonl")}, + {"file": ("in.jsonl", batch_input_file(model), "application/jsonl")}, key=key, ) assert file_response.status_code == 200, file_response.text @@ -235,7 +235,7 @@ BATCH_SPEND: Final = (BATCH_PROMPT_TOKENS * INPUT_COST_PER_TOKEN + BATCH_COMPLET def test_completed_batch_spend_lands_under_batches_in_team_endpoint_activity(gateway: Gateway) -> None: with gateway.scenario() as scenario: scenario_id: Final = f"batch-endpoint-{uuid.uuid4().hex[:12]}" - handle: Final = register_scenario(scenario_id, _batch_routes("gpt-4o-mini")) + handle: Final = register_scenario(scenario_id, batch_routes("gpt-4o-mini")) scenario.cleanups.callback(delete_scenario, handle) model: Final = scenario.model( api_base=handle.api_base(), @@ -247,7 +247,7 @@ def test_completed_batch_spend_lands_under_batches_in_team_endpoint_activity(gat file_response: Final = gateway.request_multipart( "/v1/files", {"purpose": "batch", "model": model}, - {"file": ("in.jsonl", _input_file(model), "application/jsonl")}, + {"file": ("in.jsonl", batch_input_file(model), "application/jsonl")}, key=key, ) assert file_response.status_code == 200, file_response.text diff --git a/tests/integration/spend/test_spend_logs_metadata_fields.py b/tests/integration/spend/test_spend_logs_metadata_fields.py new file mode 100644 index 00000000000..d7f1dbfd998 --- /dev/null +++ b/tests/integration/spend/test_spend_logs_metadata_fields.py @@ -0,0 +1,509 @@ +import json +from collections.abc import Callable, Generator, Mapping +from contextlib import contextmanager +from dataclasses import dataclass +from hashlib import sha256 +from pathlib import Path +from typing import Final +from uuid import uuid4 + +import httpx +import pytest +from integration._support.client import JSON_OBJECT, Gateway, Scenario, eventually, object_value, string_value +from integration._support.database import read_rows, scratch_database +from integration._support.process import owned_proxy +from integration._support.upstream import delete_scenario, register_scenario +from integration._support.wire import Reply, Request, Wire, wire_server +from integration.spend.test_batch_completion_accounting import batch_input_file, batch_routes +from integration.streaming.test_stream_contracts import text_stream +from pydantic import JsonValue + +CACHED_PROMPT_TOKENS: Final = 4 + + +def _config( + tmp_path: Path, + spend_logs_metadata_fields: Mapping[str, JsonValue] | None, + *, + model_list: tuple[Mapping[str, JsonValue], ...] = (), + litellm_settings: Mapping[str, JsonValue] | None = None, + guardrails: tuple[Mapping[str, JsonValue], ...] = (), +) -> Path: + retention: Final = ( + {} if spend_logs_metadata_fields is None else {"spend_logs_metadata_fields": dict(spend_logs_metadata_fields)} + ) + config: Final = tmp_path / f"spend-logs-metadata-fields-{uuid4()}.json" + config.write_text( + json.dumps( + { + "model_list": [dict(model) for model in model_list], + "general_settings": { + "master_key": "os.environ/LITELLM_MASTER_KEY", + "database_url": "os.environ/DATABASE_URL", + "store_model_in_db": True, + "proxy_batch_write_at": 1, + "proxy_batch_polling_interval": 1, + **retention, + }, + "litellm_settings": dict(litellm_settings or {}), + "guardrails": [dict(guardrail) for guardrail in guardrails], + } + ) + ) + return config + + +def _respond(request: Request) -> Reply: + if request.method == "GET": + return Reply(body=b'{"object":"list","data":[]}') + return Reply( + body=json.dumps( + { + "id": f"chatcmpl-{uuid4()}", + "object": "chat.completion", + "created": 1, + "model": "gpt-4o-mini", + "choices": [{"index": 0, "message": {"role": "assistant", "content": "ok"}, "finish_reason": "stop"}], + "usage": { + "prompt_tokens": 10, + "completion_tokens": 2, + "total_tokens": 12, + "prompt_tokens_details": {"cached_tokens": CACHED_PROMPT_TOKENS}, + }, + } + ).encode() + ) + + +def _stored_row_after_one_chat( + gateway: Gateway, tmp_path: Path, spend_logs_metadata_fields: Mapping[str, JsonValue] +) -> tuple[dict[str, JsonValue], str]: + with ( + wire_server(_respond) as wire, + owned_proxy(gateway, tmp_path, {}, config=_config(tmp_path, spend_logs_metadata_fields)) as isolated, + isolated.scenario() as scenario, + ): + model: Final = scenario.model(model="openai/gpt-4o-mini", api_base=wire.url, api_key="synthetic-openai-key") + api_key: Final = scenario.key(key_alias=f"metadata-fields-{uuid4()}", models=[model]) + response_id: Final = string_value(isolated.chat(model, key=api_key)["id"]) + rows: Final = eventually( + lambda: read_rows( + 'SELECT metadata, proxy_server_request, response FROM "LiteLLM_SpendLogs" WHERE request_id=%s', + (response_id,), + ), + lambda values: len(values) == 1, + seconds=70, + ) + return rows[0], sha256(api_key.encode()).hexdigest() + + +def test_excluded_metadata_fields_are_not_stored_but_still_reach_daily_spend(gateway: Gateway, tmp_path: Path) -> None: + row, hashed_key = _stored_row_after_one_chat( + gateway, tmp_path, {"exclude": ["model_map_information", "usage_object"]} + ) + + metadata: Final = object_value(row["metadata"]) + assert "model_map_information" not in metadata, metadata + assert "usage_object" not in metadata, metadata + assert {"status", "cold_storage_object_key"} <= set(metadata), metadata + assert string_value(metadata["user_api_key_alias"]).startswith("metadata-fields-") + assert row["proxy_server_request"] == {} + assert row["response"] == {} + daily: Final = eventually( + lambda: read_rows( + 'SELECT prompt_tokens, cache_read_input_tokens FROM "LiteLLM_DailyUserSpend" WHERE api_key=%s', + (hashed_key,), + ), + lambda values: len(values) == 1, + seconds=70, + ) + assert daily[0]["cache_read_input_tokens"] == CACHED_PROMPT_TOKENS, daily + + +def test_included_metadata_fields_are_the_only_ones_stored_besides_always_kept( + gateway: Gateway, tmp_path: Path +) -> None: + row, _ = _stored_row_after_one_chat(gateway, tmp_path, {"include": ["user_api_key_alias"]}) + + assert set(object_value(row["metadata"])) == {"status", "cold_storage_object_key", "user_api_key_alias"} + + +def test_guardrail_usage_is_tracked_when_guardrail_information_is_not_stored(gateway: Gateway, tmp_path: Path) -> None: + guardrail_name: Final = f"metadata-fields-guardrail-{uuid4()}" + with ( + wire_server(lambda _: Reply(body=b'{"action":"NONE"}')) as policy, + wire_server(_respond) as wire, + owned_proxy( + gateway, + tmp_path, + {}, + config=_config( + tmp_path, + {"exclude": ["guardrail_information"]}, + guardrails=( + { + "guardrail_name": guardrail_name, + "litellm_params": { + "guardrail": "generic_guardrail_api", + "mode": "pre_call", + "default_on": True, + "api_base": policy.url, + "api_key": "synthetic-guardrail-key", + }, + }, + ), + ), + ) as isolated, + isolated.scenario() as scenario, + ): + model: Final = scenario.model(model="openai/gpt-4o-mini", api_base=wire.url, api_key="synthetic-openai-key") + response_id: Final = string_value(isolated.chat(model, key=scenario.key(models=[model]))["id"]) + assert len(policy.drain()) == 1 + rows: Final = eventually( + lambda: read_rows('SELECT metadata FROM "LiteLLM_SpendLogs" WHERE request_id=%s', (response_id,)), + lambda values: len(values) == 1, + seconds=70, + ) + assert "guardrail_information" not in object_value(rows[0]["metadata"]), rows[0] + indexed: Final = eventually( + lambda: read_rows( + 'SELECT guardrail_id FROM "LiteLLM_SpendLogGuardrailIndex" WHERE request_id=%s', (response_id,) + ), + lambda values: len(values) == 1, + seconds=70, + ) + metrics: Final = eventually( + lambda: read_rows( + 'SELECT requests_evaluated, passed_count FROM "LiteLLM_DailyGuardrailMetrics" WHERE guardrail_id=%s', + (string_value(indexed[0]["guardrail_id"]),), + ), + lambda values: values == [{"requests_evaluated": 1, "passed_count": 1}], + seconds=70, + ) + assert metrics == [{"requests_evaluated": 1, "passed_count": 1}], metrics + + +EXCLUDED: Final = ("model_map_information", "user_api_key_alias") + + +@dataclass(frozen=True, slots=True) +class _Isolated: + proxy: Gateway + scenario: Scenario + key: str + database_url: str | None = None + + @property + def hashed_key(self) -> str: + return sha256(self.key.encode()).hexdigest() + + def rows(self, count: int, where: str = "TRUE", seconds: int = 70) -> tuple[dict[str, JsonValue], ...]: + return tuple( + eventually( + lambda: read_rows( + "SELECT request_id, call_type, status, spend, cache_hit, metadata " + f'FROM "LiteLLM_SpendLogs" WHERE api_key=%s AND {where} ORDER BY "startTime"', + (self.hashed_key,), + database_url=self.database_url, + ), + lambda values: len(values) == count, + seconds=seconds, + ) + ) + + +@contextmanager +def _isolated(gateway: Gateway, config: Path, tmp_path: Path, database_url: str | None = None) -> Generator[_Isolated]: + database: Final = {} if database_url is None else {"DATABASE_URL": database_url} + replica: Final = () if database_url is None else ("DATABASE_URL_READ_REPLICA",) + with ( + owned_proxy(gateway, tmp_path, database, config=config, remove_environment=replica) as proxy, + proxy.scenario() as scenario, + ): + key: Final = scenario.key(key_alias=f"metadata-fields-{uuid4()}") + yield _Isolated(proxy, scenario, key, database_url) + + +def _assert_filtered(row: Mapping[str, JsonValue], *kept: str) -> dict[str, JsonValue]: + metadata: Final = object_value(row["metadata"]) + assert not set(EXCLUDED) & set(metadata), metadata + assert {"status", "cold_storage_object_key", *kept} <= set(metadata), metadata + return metadata + + +def _anthropic_event(event: Mapping[str, JsonValue]) -> bytes: + return f"event: {event['type']}\ndata: {json.dumps(event)}\n\n".encode() + + +def _anthropic_stream(model: JsonValue) -> tuple[bytes, ...]: + message: Final = {**_anthropic_message(model), "content": [], "stop_reason": None} + events: Final[tuple[Mapping[str, JsonValue], ...]] = ( + {"type": "message_start", "message": message}, + {"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}}, + {"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "ok"}}, + {"type": "content_block_stop", "index": 0}, + {"type": "message_delta", "delta": {"stop_reason": "end_turn"}, "usage": {"output_tokens": 2}}, + {"type": "message_stop"}, + ) + return tuple(_anthropic_event(event) for event in events) + + +def _respond_streaming_or_not(request: Request) -> Reply: + if request.method == "GET": + return Reply(body=b'{"object":"list","data":[]}') + if request.target.startswith("/v1/messages"): + return _anthropic(request) + if JSON_OBJECT.validate_json(request.body).get("stream") is True: + return Reply(content_type="text/event-stream", chunks=text_stream(f"chatcmpl-{uuid4()}")) + return _respond(request) + + +def _stream(proxy: Gateway, key: str, path: str, body: Mapping[str, JsonValue]) -> None: + with proxy.client.stream("POST", path, json=dict(body), headers={"Authorization": f"Bearer {key}"}) as response: + assert response.status_code == 200, response.read().decode() + lines: Final = tuple(response.iter_lines()) + assert any(line.startswith("data:") for line in lines), lines + + +def _call(proxy: Gateway, key: str, path: str, body: Mapping[str, JsonValue]) -> None: + response: Final = proxy.request("POST", path, body, key=key) + assert response.status_code == 200, response.text + + +def test_every_inference_surface_and_stream_stores_filtered_metadata(gateway: Gateway, tmp_path: Path) -> None: + with ( + wire_server(_respond_streaming_or_not) as wire, + _isolated(gateway, _config(tmp_path, {"exclude": list(EXCLUDED)}), tmp_path) as isolated, + ): + model: Final = isolated.scenario.model( + model="deepseek/gpt-4o-mini", api_base=wire.url + "/v1", api_key="synthetic-key" + ) + anthropic: Final = isolated.scenario.model( + model="anthropic/claude-sonnet-4-6", api_base=wire.url, api_key="synthetic-key" + ) + prompt: Final[list[JsonValue]] = [{"role": "user", "content": f"surfaces {uuid4()}"}] + calls: Final[tuple[Callable[[Gateway, str, str, Mapping[str, JsonValue]], None], ...]] = ( + _stream, + _call, + _stream, + _call, + _stream, + ) + requests: Final[tuple[tuple[str, Mapping[str, JsonValue]], ...]] = ( + ("/v1/chat/completions", {"model": model, "messages": prompt, "stream": True}), + ("/v1/messages", {"model": anthropic, "max_tokens": 16, "messages": prompt}), + ("/v1/messages", {"model": anthropic, "max_tokens": 16, "messages": prompt, "stream": True}), + ("/v1/responses", {"model": model, "input": f"surfaces {uuid4()}"}), + ("/v1/responses", {"model": model, "input": f"surfaces {uuid4()}", "stream": True}), + ) + for send, (path, body) in zip(calls, requests, strict=True): + send(isolated.proxy, isolated.key, path, body) + + rows: Final = isolated.rows(len(requests)) + for row in rows: + assert row["status"] == "success", row + _assert_filtered(row, "usage_object") + + +def test_failed_request_keeps_its_error_and_status_but_drops_excluded_metadata( + gateway: Gateway, tmp_path: Path +) -> None: + def fail(request: Request) -> Reply: + if request.method == "GET": + return Reply(body=b'{"object":"list","data":[]}') + return Reply(status=500, body=b'{"error":{"message":"scripted upstream failure","type":"server_error"}}') + + with ( + wire_server(fail) as wire, + _isolated(gateway, _config(tmp_path, {"exclude": list(EXCLUDED)}), tmp_path) as isolated, + ): + model: Final = isolated.scenario.model(model="openai/gpt-4o-mini", api_base=wire.url, api_key="synthetic") + response: Final = isolated.proxy.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": "fail"}]}, + key=isolated.key, + ) + assert response.status_code >= 500, response.text + + row: Final = isolated.rows(1)[0] + assert row["status"] == "failure", row + metadata: Final = _assert_filtered(row, "error_information") + assert metadata["status"] == "failure", metadata + assert "scripted upstream failure" in json.dumps(metadata["error_information"]), metadata + + +def test_response_cache_hit_row_is_filtered_and_charged_nothing(gateway: Gateway, tmp_path: Path) -> None: + config: Final = _config( + tmp_path, {"exclude": list(EXCLUDED)}, litellm_settings={"cache": True, "cache_params": {"type": "local"}} + ) + with wire_server(_respond) as wire, _isolated(gateway, config, tmp_path) as isolated: + model: Final = isolated.scenario.model(model="openai/gpt-4o-mini", api_base=wire.url, api_key="synthetic") + text: Final = f"cache {uuid4()}" + first: Final = string_value(isolated.proxy.chat(model, key=isolated.key, text=text)["id"]) + isolated.proxy.chat(model, key=isolated.key, text=text) + + rows: Final = isolated.rows(2) + assert len([request for request in wire.drain() if request.method == "POST"]) == 1 + paid, hit = sorted(rows, key=lambda row: row["cache_hit"] == "True") + assert paid["request_id"] == first and float(str(paid["spend"])) > 0, rows + assert hit["cache_hit"] == "True" and string_value(hit["request_id"]).startswith(first + "_cache_hit"), rows + assert float(str(hit["spend"])) == 0, rows + for row in rows: + _assert_filtered(row, "usage_object") + + +def test_batch_cost_row_is_filtered_and_charged_once(gateway: Gateway, tmp_path: Path) -> None: + with _isolated(gateway, _config(tmp_path, {"exclude": list(EXCLUDED)}), tmp_path) as isolated: + handle: Final = register_scenario(f"metadata-fields-batch-{uuid4().hex[:12]}", batch_routes("gpt-4o-mini")) + isolated.scenario.cleanups.callback(delete_scenario, handle) + model: Final = isolated.scenario.model(api_base=handle.api_base()) + uploaded: Final = isolated.proxy.request_multipart( + "/v1/files", + {"purpose": "batch", "model": model}, + {"file": ("in.jsonl", batch_input_file(model), "application/jsonl")}, + key=isolated.key, + ) + assert uploaded.status_code == 200, uploaded.text + batch: Final = isolated.proxy.post( + "/v1/batches", + { + "input_file_id": string_value(JSON_OBJECT.validate_json(uploaded.content)["id"]), + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + "model": model, + }, + key=isolated.key, + ) + batch_path: Final = f"/v1/batches/{string_value(batch['id'])}" + retrievals: Final = tuple(isolated.proxy.request("GET", batch_path, key=isolated.key) for _ in range(2)) + assert all(r.status_code == 200 and r.json()["status"] == "completed" for r in retrievals), retrievals + + row: Final = isolated.rows(1, "call_type='aretrieve_batch'")[0] + assert float(str(row["spend"])) > 0, row + metadata: Final = _assert_filtered(row, "usage_object") + assert (metadata["batch_successful_requests"], metadata["batch_failed_requests"]) == (2, 3), metadata + + +def _update_retention(proxy: Gateway, value: JsonValue) -> httpx.Response: + return proxy.request( + "POST", + "/config/field/update", + {"field_name": "spend_logs_metadata_fields", "field_value": value, "config_type": "general_settings"}, + ) + + +def test_runtime_retention_update_rejects_invalid_values_and_applies_valid_ones_without_restart( + gateway: Gateway, tmp_path: Path +) -> None: + with ( + scratch_database() as database_url, + wire_server(_respond) as wire, + _isolated(gateway, _config(tmp_path, None), tmp_path, database_url) as isolated, + ): + model: Final = isolated.scenario.model(model="openai/gpt-4o-mini", api_base=wire.url, api_key="synthetic") + invalid: Final[tuple[JsonValue, ...]] = ( + {"include": ["usage_object"], "exclude": ["model_map_information"]}, + {}, + {"exclude": ["not_a_metadata_field"]}, + {"exclude": ["status"]}, + ) + + isolated.proxy.chat(model, key=isolated.key) + assert set(EXCLUDED) <= set(object_value(isolated.rows(1)[0]["metadata"])) + for value in invalid: + assert _update_retention(isolated.proxy, value).status_code == 400, value + isolated.proxy.chat(model, key=isolated.key) + assert set(EXCLUDED) <= set(object_value(isolated.rows(2)[-1]["metadata"])) + + assert _update_retention(isolated.proxy, {"exclude": list(EXCLUDED)}).status_code == 200 + isolated.proxy.chat(model, key=isolated.key) + _assert_filtered(isolated.rows(3)[-1]) + + for value in invalid: + assert _update_retention(isolated.proxy, value).status_code == 400, value + isolated.proxy.chat(model, key=isolated.key) + _assert_filtered(isolated.rows(4)[-1]) + + +SAVINGS_FIELDS: Final = ("autorouter_savings_estimate", "autorouter_savings") + + +def _anthropic_message(model: JsonValue) -> dict[str, JsonValue]: + return { + "id": f"msg_{uuid4().hex}", + "type": "message", + "role": "assistant", + "model": model, + "content": [{"type": "text", "text": "ok"}], + "stop_reason": "end_turn", + "stop_sequence": None, + "usage": { + "input_tokens": 12, + "output_tokens": 2, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 0, + "cache_creation": {"ephemeral_5m_input_tokens": 0, "ephemeral_1h_input_tokens": 0}, + }, + } + + +def _anthropic(request: Request) -> Reply: + if request.target.startswith("/v1/messages/count_tokens"): + return Reply(body=b'{"input_tokens":12}') + body: Final = JSON_OBJECT.validate_json(request.body) + if body.get("stream") is True: + return Reply(content_type="text/event-stream", chunks=_anthropic_stream(body["model"])) + return Reply(body=json.dumps(_anthropic_message(body["model"])).encode()) + + +def _router_models(wire: Wire, router: str) -> tuple[Mapping[str, JsonValue], ...]: + tiers: Final = {"cheap": "anthropic/claude-sonnet-4-6", "frontier": "anthropic/claude-opus-4-8"} + return ( + *( + {"model_name": name, "litellm_params": {"model": model, "api_base": wire.url, "api_key": "synthetic"}} + for name, model in tiers.items() + ), + { + "model_name": router, + "litellm_params": { + "model": "auto_router/complexity_router", + "complexity_router_config": { + "tiers": {"SIMPLE": "cheap", "MEDIUM": "cheap", "COMPLEX": "frontier", "REASONING": "frontier"}, + }, + }, + }, + ) + + +@pytest.mark.parametrize( + ("retention", "stored"), + [(None, SAVINGS_FIELDS), ({"exclude": list(SAVINGS_FIELDS)}, ())], + ids=["unset-keeps-savings", "excluded-savings-stay-out"], +) +@pytest.mark.timeout(240) +def test_delayed_autorouter_savings_publication_respects_retention( + gateway: Gateway, tmp_path: Path, retention: Mapping[str, JsonValue] | None, stored: tuple[str, ...] +) -> None: + router: Final = f"router-{uuid4().hex[:8]}" + with wire_server(_anthropic) as wire: + config: Final = _config(tmp_path, retention, model_list=_router_models(wire, router)) + with _isolated(gateway, config, tmp_path) as isolated: + response: Final = isolated.proxy.request( + "POST", + "/v1/messages", + {"model": router, "max_tokens": 16, "messages": [{"role": "user", "content": "hi"}]}, + key=isolated.key, + headers={"x-litellm-session-id": f"session-{uuid4()}"}, + ) + assert response.status_code == 200, response.text + + published: Final = isolated.rows( + 1, + "metadata::jsonb ? 'routing_decision' AND NOT metadata::jsonb ? 'autorouter_baseline_observation' " + 'AND request_id IN (SELECT request_id FROM "LiteLLM_AutoRouterBaselineObservation" ' + "WHERE publication IS NOT NULL)", + seconds=150, + )[0] + metadata: Final = object_value(published["metadata"]) + assert {name for name in SAVINGS_FIELDS if name in metadata} == set(stored), metadata diff --git a/tests/unit/proxy/spend_tracking/test_spend_tracking_utils.py b/tests/unit/proxy/spend_tracking/test_spend_tracking_utils.py index c5bb1ed5b61..9aa9110f7a4 100644 --- a/tests/unit/proxy/spend_tracking/test_spend_tracking_utils.py +++ b/tests/unit/proxy/spend_tracking/test_spend_tracking_utils.py @@ -18,7 +18,7 @@ from litellm.constants import LITELLM_TRUNCATED_PAYLOAD_FIELD, LITELLM_TRUNCATIO from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup from litellm.llms.base_llm.ocr.transformation import OCRResponse, OCRUsageInfo from litellm.litellm_core_utils.safe_json_dumps import safe_dumps -from litellm.proxy._types import SpendLogsPayload, UserAPIKeyAuth +from litellm.proxy._types import SpendLogsMetadataFields, SpendLogsPayload, UserAPIKeyAuth from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup from litellm.proxy.route_llm_request import ProxyModelNotFoundError from litellm.proxy.spend_tracking.spend_tracking_utils import ( @@ -37,9 +37,11 @@ from litellm.proxy.spend_tracking.spend_tracking_utils import ( _sanitize_guardrail_information_for_spend_logs, _sanitize_request_body_for_spend_logs_payload, _scrub_raw_model_from_error_information, + configured_spend_logs_metadata_fields, get_logging_payload, get_spend_logs_id, should_store_prompts_and_responses_in_spend_logs, + spend_log_row_with_retained_metadata, ) from litellm.proxy.utils import hash_token from litellm.types.router import GenericLiteLLMParams @@ -79,10 +81,16 @@ def test_classifier_audit_spend_storage_obeys_privacy_and_truncation(monkeypatch "classifier_input": {"system": "rubric" * 1000, "messages": [{"role": "user", "content": "ask"}]}, "originating_request_masked": {"input": "source-only", "api_key": "REDACTED"}, } - stored: Final = json.loads(_get_proxy_server_request_for_spend_logs_payload( - metadata={}, litellm_params={"proxy_server_request": {"body": {"model": "classifier"}}}, - kwargs={"standard_logging_object": audit, "standard_callback_dynamic_params": {"turn_off_message_logging": redact}}, - )) + stored: Final = json.loads( + _get_proxy_server_request_for_spend_logs_payload( + metadata={}, + litellm_params={"proxy_server_request": {"body": {"model": "classifier"}}}, + kwargs={ + "standard_logging_object": audit, + "standard_callback_dynamic_params": {"turn_off_message_logging": redact}, + }, + ) + ) if not store_prompts or redact: assert "classifier_input" not in stored assert "originating_request_masked" not in stored @@ -208,9 +216,7 @@ def test_batch_lifecycle_rows_derive_the_same_session_from_the_batch_id(): from litellm.proxy.spend_tracking.spend_tracking_utils import _get_batch_trace_session_id create_session: Final = _get_batch_trace_session_id(call_type="acreate_batch", request_id="batch-uid-1") - cost_session: Final = _get_batch_trace_session_id( - call_type="aretrieve_batch", request_id="batch-uid-1_batch_cost" - ) + cost_session: Final = _get_batch_trace_session_id(call_type="aretrieve_batch", request_id="batch-uid-1_batch_cost") assert create_session == cost_session == "batch-uid-1" @@ -2911,7 +2917,11 @@ def test_proxy_server_request_payload_redacts_provider_credentials(mock_should_s "extra_headers": {"Authorization": "Bearer canary-extra-header"}, "tools": [ {"type": "function", "function": {"name": "f", "parameters": tool_parameters}}, - {"type": "mcp", "server_url": "https://mcp.example.com", "headers": {"Authorization": "canary-mcp"}}, + { + "type": "mcp", + "server_url": "https://mcp.example.com", + "headers": {"Authorization": "canary-mcp"}, + }, ], "fallbacks": [{"model": "azure-b", **credentials}], "metadata": metadata, @@ -5465,7 +5475,7 @@ ANTHROPIC_MESSAGES_SSE_CHUNKS: Final = ( 'event: content_block_stop\ndata: {"type":"content_block_stop","index":0}\n\n', 'event: message_delta\ndata: {"type":"message_delta","delta":{"stop_reason":"end_turn"},' '"usage":{"output_tokens":4}}\n\n', - "event: message_stop\ndata: {\"type\":\"message_stop\"}\n\n", + 'event: message_stop\ndata: {"type":"message_stop"}\n\n', ) @@ -5503,9 +5513,7 @@ def test_spend_log_request_id_is_the_message_id_a_non_streaming_messages_caller_ """ logging_obj = _anthropic_messages_logging_obj(stream=False) - logged_response = logging_obj._handle_anthropic_messages_response_logging( - result=ANTHROPIC_MESSAGES_RESPONSE - ) + logged_response = logging_obj._handle_anthropic_messages_response_logging(result=ANTHROPIC_MESSAGES_RESPONSE) assert logged_response.id == "msg_01Lit6806NonStreaming" assert ( @@ -5581,9 +5589,7 @@ def test_spend_log_request_id_still_falls_back_to_litellm_call_id_without_a_prov end_time=datetime.datetime.now(timezone.utc), logging_obj=logging_obj, ) - assert logging_obj.model_call_details["complete_streaming_response"].id == ( - "6806cafe-0000-4000-8000-000000000001" - ) + assert logging_obj.model_call_details["complete_streaming_response"].id == ("6806cafe-0000-4000-8000-000000000001") def test_spend_log_request_id_for_chat_completions_is_untouched(): @@ -5665,6 +5671,7 @@ def test_failed_agent_request_keeps_registered_display_name(): assert payload["status"] == "failure" assert payload["model_id"] == "registered-agent" + _CLI_SESSION_ALIAS: Final = "cli-session-alice" _CLI_SESSION_TOKEN: Final = "cli-session-Qm7xJ2kP9sLw4vT1nR8yAa" @@ -5796,20 +5803,109 @@ def test_baseline_estimate_metadata_comes_from_the_logging_stamp() -> None: def test_untrusted_agent_label_cannot_replace_verified_billing_identity(billing_agent: str | None) -> None: kwargs = { "model": "gpt-4", - "litellm_params": {"metadata": { - "user_api_key": "test-key", - "agent_id": "header-selected-agent", - "billing_agent_id": billing_agent, - }}, + "litellm_params": { + "metadata": { + "user_api_key": "test-key", + "agent_id": "header-selected-agent", + "billing_agent_id": billing_agent, + } + }, } payload = get_logging_payload( - kwargs=kwargs, response_obj={"id": "request"}, - start_time=datetime.datetime.now(timezone.utc), end_time=datetime.datetime.now(timezone.utc), + kwargs=kwargs, + response_obj={"id": "request"}, + start_time=datetime.datetime.now(timezone.utc), + end_time=datetime.datetime.now(timezone.utc), ) assert payload["agent_id"] == "header-selected-agent" assert payload["billing_agent_id"] == billing_agent +def _spend_log_row(metadata: Mapping[str, object]) -> Mapping[str, object]: + return MappingProxyType({"request_id": "req-1", "spend": 0.5, "metadata": json.dumps(dict(metadata))}) + + +_STORED_METADATA: Final = MappingProxyType( + { + "status": "success", + "cold_storage_object_key": "logs/req-1.json", + "model_map_information": {"model_map_key": "gpt-4o", "model_map_value": {"max_tokens": 10}}, + "usage_object": {"prompt_tokens": 3}, + "user_api_key_alias": "alias", + } +) + + +def test_spend_log_row_keeps_every_metadata_field_when_unconfigured() -> None: + row: Final = _spend_log_row(_STORED_METADATA) + + assert spend_log_row_with_retained_metadata(row, None) is row + + +def test_spend_log_row_drops_excluded_metadata_fields_only() -> None: + row: Final = _spend_log_row(_STORED_METADATA) + + stored: Final = spend_log_row_with_retained_metadata( + row, SpendLogsMetadataFields(exclude=("model_map_information", "user_api_key_alias")) + ) + + assert json.loads(cast(str, stored["metadata"])) == { + "status": "success", + "cold_storage_object_key": "logs/req-1.json", + "usage_object": {"prompt_tokens": 3}, + } + assert {name: value for name, value in stored.items() if name != "metadata"} == { + "request_id": "req-1", + "spend": 0.5, + } + + +def test_spend_log_row_include_keeps_listed_and_always_kept_fields() -> None: + stored: Final = spend_log_row_with_retained_metadata( + _spend_log_row(_STORED_METADATA), SpendLogsMetadataFields(include=("usage_object",)) + ) + + assert json.loads(cast(str, stored["metadata"])) == { + "status": "success", + "cold_storage_object_key": "logs/req-1.json", + "usage_object": {"prompt_tokens": 3}, + } + + +@pytest.mark.parametrize( + "configured", + [ + {"include": ["usage_object"], "exclude": ["model_map_information"]}, + {}, + {"exclude": ["model_map_informaton"]}, + {"include": ["usage_object", "not_a_field"]}, + {"exclude": ["status"]}, + {"exclude": ["cold_storage_object_key"]}, + {"exclude": ["model_map_information"], "drop": ["usage_object"]}, + ], +) +def test_spend_logs_metadata_fields_rejects_ambiguous_or_lossy_config(configured: dict[str, list[str]]) -> None: + from pydantic import ValidationError + + from litellm.proxy._types import ConfigGeneralSettings + + with pytest.raises(ValidationError): + ConfigGeneralSettings.model_validate({"spend_logs_metadata_fields": configured}) + + +def test_configured_spend_logs_metadata_fields_ignores_invalid_runtime_value() -> None: + with patch( + "litellm.proxy.proxy_server.general_settings", + {"spend_logs_metadata_fields": {"include": ["usage_object"], "exclude": ["status"]}}, + ): + assert configured_spend_logs_metadata_fields() is None + with patch( + "litellm.proxy.proxy_server.general_settings", + {"spend_logs_metadata_fields": {"exclude": ["model_map_information"]}}, + ): + assert configured_spend_logs_metadata_fields() == SpendLogsMetadataFields(exclude=("model_map_information",)) + + class _OcrUsageInfoDict(TypedDict, total=False): pages_processed: ReadOnly[int] doc_size_bytes: ReadOnly[int] diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index cd85908979e..fdc63268fe5 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -30429,6 +30429,8 @@ export interface components { search_tool_deny_by_default: boolean; /** @description Daily check of the spend LiteLLM captured against the provider's own bill (OpenAI via OPENAI_ADMIN_KEY). Publishes litellm_spend_capture_rate per provider and alerts when the ratio over the lookback window falls under the threshold (default 0.9). Off unless set. */ spend_capture_rate_check?: components["schemas"]["SpendCaptureRateCheckSettings"] | null; + /** @description Which keys of LiteLLM_SpendLogs.metadata are written to the database. Set exactly one of 'include' (write only these keys) or 'exclude' (drop these keys). 'status' and 'cold_storage_object_key' are always written. Daily spend tables, budgets and logging callbacks still see every key. Unset writes every key */ + spend_logs_metadata_fields?: components["schemas"]["SpendLogsMetadataFields"] | null; /** * Store Model In Db * @description If True, models and config are stored in and loaded from the database. Default is False. @@ -46277,6 +46279,13 @@ export interface components { */ threshold: number; }; + /** SpendLogsMetadataFields */ + SpendLogsMetadataFields: { + /** Exclude */ + exclude?: string[] | null; + /** Include */ + include?: string[] | null; + }; /** SpendMetrics */ SpendMetrics: { /**