diff --git a/.github/workflows/image-scan.yml b/.github/workflows/image-scan.yml index 8faf3ef6229..d798df4c3a4 100644 --- a/.github/workflows/image-scan.yml +++ b/.github/workflows/image-scan.yml @@ -17,6 +17,8 @@ on: - backend/Dockerfile - backend/main.py - docker/component_entrypoint.sh + - docker/entrypoint.sh + - litellm/proxy/prisma_migration.py - litellm-proxy-extras/** - tests/proxy_migration_tests/** - uv.lock diff --git a/Dockerfile b/Dockerfile index 66ce3af4a65..700b0d6525e 100644 --- a/Dockerfile +++ b/Dockerfile @@ -1,10 +1,10 @@ # syntax=docker/dockerfile:1.7 # Base image for building -ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:42df77a9974d6ec8b17a5ee8bc23b532600a44d705acef2409e0933c1251b45f +ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:a31344ab2cb8618db84f535eec56f76f6178b142cb92cb2e48676cc2dcebea72 # Runtime image -ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:42df77a9974d6ec8b17a5ee8bc23b532600a44d705acef2409e0933c1251b45f +ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:a31344ab2cb8618db84f535eec56f76f6178b142cb92cb2e48676cc2dcebea72 ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:240fb85ab0f263ef12f492d8476aa3a2e4e1e333f7d67fbdd923d00a506a516a # Pinned by digest like the other base images; bump explicitly on Node upgrades. ARG UI_BUILD_IMAGE=node:24.19-alpine3.24@sha256:d32cdf619f63fe0471182d08996dd516c6275bb5fd31ae06e55a570bd9e1ad43 diff --git a/backend/Dockerfile b/backend/Dockerfile index 853c74b05ca..4ca40944606 100644 --- a/backend/Dockerfile +++ b/backend/Dockerfile @@ -1,5 +1,5 @@ -ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:42df77a9974d6ec8b17a5ee8bc23b532600a44d705acef2409e0933c1251b45f -ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:42df77a9974d6ec8b17a5ee8bc23b532600a44d705acef2409e0933c1251b45f +ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:a31344ab2cb8618db84f535eec56f76f6178b142cb92cb2e48676cc2dcebea72 +ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:a31344ab2cb8618db84f535eec56f76f6178b142cb92cb2e48676cc2dcebea72 ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:240fb85ab0f263ef12f492d8476aa3a2e4e1e333f7d67fbdd923d00a506a516a FROM $UV_IMAGE AS uvbin diff --git a/docker/Dockerfile.database b/docker/Dockerfile.database index 4bf3ae2b417..f0d6d02fccf 100644 --- a/docker/Dockerfile.database +++ b/docker/Dockerfile.database @@ -1,10 +1,10 @@ # syntax=docker/dockerfile:1.7 # Base image for building -ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:42df77a9974d6ec8b17a5ee8bc23b532600a44d705acef2409e0933c1251b45f +ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:a31344ab2cb8618db84f535eec56f76f6178b142cb92cb2e48676cc2dcebea72 # Runtime image -ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:42df77a9974d6ec8b17a5ee8bc23b532600a44d705acef2409e0933c1251b45f +ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:a31344ab2cb8618db84f535eec56f76f6178b142cb92cb2e48676cc2dcebea72 ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:240fb85ab0f263ef12f492d8476aa3a2e4e1e333f7d67fbdd923d00a506a516a # Pinned by digest like the other base images; bump explicitly on Node upgrades. ARG UI_BUILD_IMAGE=node:24.19-alpine3.24@sha256:d32cdf619f63fe0471182d08996dd516c6275bb5fd31ae06e55a570bd9e1ad43 diff --git a/docker/Dockerfile.non_root b/docker/Dockerfile.non_root index 7392cc09a0d..4a5df6ecd69 100644 --- a/docker/Dockerfile.non_root +++ b/docker/Dockerfile.non_root @@ -1,8 +1,8 @@ # syntax=docker/dockerfile:1.7 # Base images -ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:42df77a9974d6ec8b17a5ee8bc23b532600a44d705acef2409e0933c1251b45f -ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:42df77a9974d6ec8b17a5ee8bc23b532600a44d705acef2409e0933c1251b45f +ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:a31344ab2cb8618db84f535eec56f76f6178b142cb92cb2e48676cc2dcebea72 +ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:a31344ab2cb8618db84f535eec56f76f6178b142cb92cb2e48676cc2dcebea72 ARG PROXY_EXTRAS_SOURCE=published ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:240fb85ab0f263ef12f492d8476aa3a2e4e1e333f7d67fbdd923d00a506a516a # Pinned by digest like the other base images; bump explicitly on Node upgrades. diff --git a/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py b/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py index 4bb00408fc3..76e92538aaa 100644 --- a/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py +++ b/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py @@ -966,6 +966,16 @@ class CheckBatchCost: ) elif response.status in PROVIDER_TERMINAL_BATCH_STATUSES: + from litellm.proxy.openai_files_endpoints.common_utils import ( + _completed_batch_safe_to_retire, + ) + + if response.status in ("completed", "complete") and not _completed_batch_safe_to_retire(response): + verbose_proxy_logger.info( + f"CheckBatchCost: batch {batch_id} is completed but its output file id " + f"has not appeared yet; leaving job {job.id} for the next poll cycle" + ) + continue await self._finalize_unbilled_terminal_job(job, response) # Record polling run metrics (always, even if nothing was processed) diff --git a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py index c986e835e4f..39f8de0b0cc 100644 --- a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py +++ b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py @@ -45,6 +45,8 @@ from litellm.proxy._types import ( UserAPIKeyAuth, ) from litellm.proxy.openai_files_endpoints.common_utils import ( + FILE_LIST_CONTINUATION_CHUNK_SIZE, + MAX_FILE_LIST_LIMIT, _is_base64_encoded_unified_file_id, apply_unified_file_ids, ensure_batch_response_managed_file_ids, @@ -54,6 +56,8 @@ from litellm.proxy.openai_files_endpoints.common_utils import ( map_raw_file_ids_to_unified, normalize_mime_type_for_provider, resolve_managed_output_file_model_name, + validate_file_list_limit, + validate_file_list_purpose, ) from litellm.proxy.pass_through_endpoints.llm_provider_handlers.batch_attribution import ( request_tags_from_metadata, @@ -63,9 +67,9 @@ from litellm.types.llms.openai import ( # pyright: ignore[reportAttributeAccess AsyncCursorPage, ChatCompletionFileObject, CreateFileRequest, + FileListPage, FileObject, OpenAIFileObject, - OpenAIFilesPurpose, ResponsesAPIResponse, ) from litellm.types.utils import ( @@ -144,7 +148,14 @@ class _ManagedFileRow(Protocol): class _ManagedFileTableActions(Protocol): async def find_first(self, where: Mapping[str, object]) -> Optional[_ManagedFileRow]: ... - async def find_many(self, where: Mapping[str, object]) -> Sequence[_ManagedFileRow]: ... + async def find_many( + self, + where: Mapping[str, object], + take: int = ..., + order: Union[Mapping[str, str], Sequence[Mapping[str, str]]] = ..., + cursor: Mapping[str, str] = ..., + skip: int = ..., + ) -> Sequence[_ManagedFileRow]: ... async def upsert(self, where: Mapping[str, str], data: Mapping[str, Mapping[str, object]]) -> _ManagedFileRow: ... @@ -1365,12 +1376,76 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): async def afile_list( self, - purpose: Optional[OpenAIFilesPurpose], + purpose: Optional[str], litellm_parent_otel_span: Optional[Span], + user_api_key_dict: UserAPIKeyAuth, + limit: Optional[int] = None, + after: Optional[str] = None, **data: Dict, - ) -> List[OpenAIFileObject]: - """Handled in files_endpoints.py""" - return [] + ) -> FileListPage: + """List the managed files the caller owns, newest first. + + Pagination is keyset based on ``unified_file_id`` so a key that owns + every file on the proxy still reads one bounded page at a time. + ``purpose`` is applied after parsing, because the managed file table + keeps it inside the ``file_object`` blob instead of a column, and rows + whose blob will not parse drop out there too, so a chunk of rows can + yield fewer matches than the page holds. Successive chunks are read + until the page is full or the caller's rows run out, which keeps + ``data`` non-empty while matches remain and its last id usable as the + next cursor. A first chunk that fills the page costs one query; once a + scan has to continue past it, the chunk widens to + ``FILE_LIST_CONTINUATION_CHUNK_SIZE``, so the walk costs one query per + that many rows instead of one per page. That bound is per query, not + per request: the work is still linear in the rows the caller owns, and + a filter matching nothing reads every one of them, with no index + covering either the owner filter or the sort. + """ + validate_file_list_limit(limit) + validate_file_list_purpose(purpose) + + owner_filter: Final = build_owner_filter(user_api_key_dict) + if owner_filter is None: + return FileListPage(**build_list_page([])) + + if after: + cursor_row = await _managed_file_table(self.prisma_client).find_first( + where={**owner_filter, "unified_file_id": after} + ) + if cursor_row is None: + raise ProxyException( + message=f"Invalid 'after' cursor: no file found with id '{after}'.", + type="invalid_request_error", + param="after", + code=400, + openai_code="invalid_value", + ) + + page_size: Final = min(limit or MAX_FILE_LIST_LIMIT, MAX_FILE_LIST_LIMIT) + matches: Final[List[OpenAIFileObject]] = [] + cursor_id = after + chunk_size = page_size + 1 + + while len(matches) <= page_size: + cursor_args: _CursorPageArgs = {"cursor": {"unified_file_id": cursor_id}, "skip": 1} if cursor_id else {} + chunk = await _managed_file_table(self.prisma_client).find_many( + where=owner_filter, + take=chunk_size, + order=[{"created_at": "desc"}, {"unified_file_id": "desc"}], + **cursor_args, + ) + matches.extend( + parsed_file_object.model_copy(update={"id": row.unified_file_id}) + for row in chunk + if (parsed_file_object := _parse_managed_file_object(row.file_object, row.unified_file_id)) is not None + and (purpose is None or parsed_file_object.purpose == purpose) + ) + if len(chunk) < chunk_size: + break + cursor_id = chunk[-1].unified_file_id + chunk_size = max(chunk_size, FILE_LIST_CONTINUATION_CHUNK_SIZE) + + return FileListPage(**build_list_page(matches[:page_size], has_more=len(matches) > page_size)) def _is_batch_polling_enabled(self) -> bool: """ diff --git a/gateway/Dockerfile b/gateway/Dockerfile index 223df524d7c..4a2e32e186e 100644 --- a/gateway/Dockerfile +++ b/gateway/Dockerfile @@ -1,5 +1,5 @@ -ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:42df77a9974d6ec8b17a5ee8bc23b532600a44d705acef2409e0933c1251b45f -ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:42df77a9974d6ec8b17a5ee8bc23b532600a44d705acef2409e0933c1251b45f +ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:a31344ab2cb8618db84f535eec56f76f6178b142cb92cb2e48676cc2dcebea72 +ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:a31344ab2cb8618db84f535eec56f76f6178b142cb92cb2e48676cc2dcebea72 ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:240fb85ab0f263ef12f492d8476aa3a2e4e1e333f7d67fbdd923d00a506a516a FROM $UV_IMAGE AS uvbin diff --git a/litellm/integrations/otel/__init__.py b/litellm/integrations/otel/__init__.py index 9c1205bb277..d7627d4d63d 100644 --- a/litellm/integrations/otel/__init__.py +++ b/litellm/integrations/otel/__init__.py @@ -49,6 +49,7 @@ from litellm.integrations.otel.model.semconv import ( Error, GenAI, GenAIOperation, + GenAIOutputType, GenAIProvider, JsonRpc, LiteLLM, @@ -60,6 +61,7 @@ from litellm.integrations.otel.model.semconv import ( RpcSystem, Server, resolve_operation, + resolve_output_type, resolve_provider, ) from litellm.integrations.otel.model.spans import ( @@ -84,6 +86,7 @@ __all__ = [ "Error", "GenAI", "GenAIOperation", + "GenAIOutputType", "GenAIProvider", "GuardrailSpanData", "JsonRpc", @@ -116,6 +119,7 @@ __all__ = [ "is_otel_v2_enabled", "promoted_baggage", "resolve_operation", + "resolve_output_type", "resolve_provider", "span_role_for_service", "validate_registry", diff --git a/litellm/integrations/otel/mappers/genai.py b/litellm/integrations/otel/mappers/genai.py index 79487e69ac4..5e3401cd62c 100644 --- a/litellm/integrations/otel/mappers/genai.py +++ b/litellm/integrations/otel/mappers/genai.py @@ -42,6 +42,7 @@ class GenAIMapper: _LLM_CALL_ATTRS: dict[str, Callable[[LLMCallSpanData], AttrValue | None]] = { GenAI.OPERATION_NAME: lambda d: d.operation.value, GenAI.PROVIDER_NAME: lambda d: d.provider or None, + GenAI.OUTPUT_TYPE: lambda d: d.output_type.value if d.output_type else None, GenAI.REQUEST_MODEL: lambda d: d.request_model or None, GenAI.REQUEST_TEMPERATURE: lambda d: d.request_params.temperature, GenAI.REQUEST_TOP_P: lambda d: d.request_params.top_p, @@ -65,6 +66,7 @@ class GenAIMapper: Server.ADDRESS: lambda d: d.server.address if d.server else None, Server.PORT: lambda d: d.server.port if d.server else None, LiteLLM.CALL_ID: lambda d: d.identity.call_id or None, + LiteLLM.CALL_TYPE: lambda d: d.call_type, # The provider/underlying model is only known once routing has picked a # deployment, so it can't ride identity Baggage (seeded at auth, before # routing) onto the boundary-born LLM span — stamp it directly here. diff --git a/litellm/integrations/otel/model/payloads.py b/litellm/integrations/otel/model/payloads.py index aba9cc80240..4e4ed4b7513 100644 --- a/litellm/integrations/otel/model/payloads.py +++ b/litellm/integrations/otel/model/payloads.py @@ -15,8 +15,10 @@ from litellm.integrations.otel.model.metadata import ( ) from litellm.integrations.otel.model.semconv import ( GenAIOperation, + GenAIOutputType, MCPMethod, resolve_operation, + resolve_output_type, resolve_provider, ) from litellm.integrations.otel.model.utils import ( @@ -310,6 +312,11 @@ class LLMCallSpanData: choices_out: tuple[Mapping[str, object], ...] = () system_fingerprint: str | None = None time_to_first_chunk_seconds: float | None = None + # The requested output modality, set only on the routes that pin one (image + # generation, speech, transcription, OCR), and the litellm route itself, which + # keeps routes the convention folds into one operation distinguishable. + output_type: GenAIOutputType | None = None + call_type: str | None = None @classmethod def from_standard_logging_payload( @@ -334,8 +341,9 @@ class LLMCallSpanData: # otherwise the content-bearing mappers receive empty sequences and emit # no prompt/response text. finish_reasons: Final = _finish_reasons(choices_out) + call_type: Final = as_str(payload.get("call_type")) return cls( - operation=resolve_operation(as_str(payload.get("call_type"))), + operation=resolve_operation(call_type), provider=resolve_provider(as_str(payload.get("custom_llm_provider"))), request_model=context.request_model, response_model=context.response_model, @@ -358,6 +366,8 @@ class LLMCallSpanData: choices_out=choices_out if capture_content else (), system_fingerprint=as_str(response.get("system_fingerprint")), time_to_first_chunk_seconds=time_to_first_chunk_seconds, + output_type=resolve_output_type(call_type), + call_type=call_type or None, ) diff --git a/litellm/integrations/otel/model/semconv.py b/litellm/integrations/otel/model/semconv.py index ada2822ba66..1647e0a5bd1 100644 --- a/litellm/integrations/otel/model/semconv.py +++ b/litellm/integrations/otel/model/semconv.py @@ -3,7 +3,9 @@ Keys follow the OpenTelemetry GenAI semantic conventions (experimental). Anythin without a semconv equivalent lives under the ``litellm.*`` vendor namespace. """ +from collections.abc import Mapping from enum import Enum +from types import MappingProxyType from typing import Final from litellm._logging import verbose_logger @@ -30,6 +32,21 @@ class GenAIOperation(str, Enum): EXECUTE_TOOL = "execute_tool" # MCP tool-call spans LITELLM_VECTOR_STORE_MANAGEMENT = "litellm.vector_store_management" LITELLM_VECTOR_STORE_FILE_MANAGEMENT = "litellm.vector_store_file_management" + LITELLM_MODERATION = "litellm.moderation" + + +class GenAIOutputType(str, Enum): + """Values for ``gen_ai.output.type``, the modality the client asked for. + + It is what separates the inference routes that share ``generate_content``: + image generation requests ``image``, speech requests ``speech``, and + transcription and OCR both request ``text``. + """ + + TEXT = "text" + JSON = "json" + IMAGE = "image" + SPEECH = "speech" class GenAIProvider(str, Enum): @@ -258,6 +275,11 @@ class LiteLLM: """Vendor-extension keys (no semconv equivalent). Always ``litellm.*``.""" CALL_ID: Final = "litellm.call_id" + # The litellm route that produced the call. Needed because the convention maps + # several routes onto one operation: transcription and OCR are both + # ``generate_content`` with a ``text`` output type, so this is the only thing + # that tells them apart. + CALL_TYPE: Final = "litellm.call_type" COST_PREFIX: Final = "litellm.cost." METADATA_PREFIX: Final = "litellm.metadata." TEAM_ID: Final = "litellm.team.id" @@ -352,6 +374,16 @@ _OPERATION_BY_CALL_TYPE: Final[dict[str, GenAIOperation]] = { "aembedding": GenAIOperation.EMBEDDINGS, "responses": GenAIOperation.CHAT, "aresponses": GenAIOperation.CHAT, + "image_generation": GenAIOperation.GENERATE_CONTENT, + "aimage_generation": GenAIOperation.GENERATE_CONTENT, + "moderation": GenAIOperation.LITELLM_MODERATION, + "amoderation": GenAIOperation.LITELLM_MODERATION, + "ocr": GenAIOperation.GENERATE_CONTENT, + "aocr": GenAIOperation.GENERATE_CONTENT, + "speech": GenAIOperation.GENERATE_CONTENT, + "aspeech": GenAIOperation.GENERATE_CONTENT, + "transcription": GenAIOperation.GENERATE_CONTENT, + "atranscription": GenAIOperation.GENERATE_CONTENT, "call_mcp_tool": GenAIOperation.EXECUTE_TOOL, "vector_store_search": GenAIOperation.RETRIEVAL, "avector_store_search": GenAIOperation.RETRIEVAL, @@ -385,6 +417,23 @@ _OPERATION_BY_CALL_TYPE: Final[dict[str, GenAIOperation]] = { } +# litellm ``call_type`` -> ``gen_ai.output.type``. Only the call types whose route +# fixes the requested modality are listed; the attribute is conditionally required +# on a request that asks for an output format, so anything else is left unstamped. +_OUTPUT_TYPE_BY_CALL_TYPE: Final[Mapping[str, GenAIOutputType]] = MappingProxyType( + { + "image_generation": GenAIOutputType.IMAGE, + "aimage_generation": GenAIOutputType.IMAGE, + "speech": GenAIOutputType.SPEECH, + "aspeech": GenAIOutputType.SPEECH, + "transcription": GenAIOutputType.TEXT, + "atranscription": GenAIOutputType.TEXT, + "ocr": GenAIOutputType.TEXT, + "aocr": GenAIOutputType.TEXT, + } +) + + def resolve_provider(custom_llm_provider: str | None) -> str: """Map a litellm provider string to a ``gen_ai.provider.name`` value. @@ -416,3 +465,11 @@ def resolve_operation(call_type: str | None) -> GenAIOperation: GenAIOperation.CHAT.value, ) return GenAIOperation.CHAT + + +def resolve_output_type(call_type: str | None) -> GenAIOutputType | None: + """Map a litellm ``call_type`` to a ``gen_ai.output.type`` value, or ``None`` + for a route that doesn't pin the output modality.""" + if not call_type: + return None + return _OUTPUT_TYPE_BY_CALL_TYPE.get(call_type.lower()) diff --git a/litellm/integrations/prometheus.py b/litellm/integrations/prometheus.py index 76066f4a305..f9195db1d67 100644 --- a/litellm/integrations/prometheus.py +++ b/litellm/integrations/prometheus.py @@ -215,7 +215,9 @@ class PrometheusLogger(CustomLogger): # request latency metrics self.litellm_request_total_latency_metric = self._histogram_factory( "litellm_request_total_latency_metric", - "Total latency (seconds) for a request to LiteLLM", + "End-to-end latency (seconds) for a request to LiteLLM Proxy Server, from the moment " + "the request reached the proxy through the end of processing -- includes " + "authentication, pre-call hooks, the LLM API call, and post-call processing", labelnames=self.get_labels_for_metric("litellm_request_total_latency_metric"), buckets=self.latency_buckets, ) @@ -458,7 +460,8 @@ class PrometheusLogger(CustomLogger): # Request queue time metric self.litellm_request_queue_time_metric = self._histogram_factory( "litellm_request_queue_time_seconds", - "Time spent in request queue before processing starts (seconds)", + "Time (seconds) from request arrival at the proxy to the start of pre-call " + "processing -- includes authentication and any ASGI-level queueing", labelnames=self.get_labels_for_metric("litellm_request_queue_time_seconds"), buckets=self.latency_buckets, ) @@ -2078,27 +2081,37 @@ class PrometheusLogger(CustomLogger): _labels, ) - # total request latency + # request queue time (time from arrival to processing start) -- read first so + # it can be folded into the total-latency metric below. start_time/end_time + # only span from after auth completes, so without this the "total" latency + # metric silently excludes auth and pre-call hook time. + _litellm_params: Final = kwargs.get("litellm_params", {}) or {} + queue_time_seconds: Final = (_litellm_params.get("metadata") or {}).get("queue_time_seconds") + + # total request latency: true end-to-end, from request arrival (queue_time_seconds, + # when available) through the end of processing. total_time_seconds: Final = self._safe_duration_seconds( start_time=start_time, end_time=end_time, ) if total_time_seconds is not None: + _observed_total_time_seconds: Final = ( + total_time_seconds + queue_time_seconds + if queue_time_seconds is not None and queue_time_seconds >= 0 + else total_time_seconds + ) _labels = prometheus_label_factory( supported_enum_labels=self.get_labels_for_metric(metric_name="litellm_request_total_latency_metric"), enum_values=enum_values, label_context=label_context, ) - self.litellm_request_total_latency_metric.labels(**_labels).observe(total_time_seconds) + self.litellm_request_total_latency_metric.labels(**_labels).observe(_observed_total_time_seconds) self._track_end_user_metric_series( self.litellm_request_total_latency_metric, "litellm_request_total_latency_metric", _labels, ) - # request queue time (time from arrival to processing start) - _litellm_params: Final = kwargs.get("litellm_params", {}) or {} - queue_time_seconds: Final = (_litellm_params.get("metadata") or {}).get("queue_time_seconds") if queue_time_seconds is not None and queue_time_seconds >= 0: _labels = prometheus_label_factory( supported_enum_labels=self.get_labels_for_metric(metric_name="litellm_request_queue_time_seconds"), diff --git a/litellm/integrations/websearch_interception/ARCHITECTURE.md b/litellm/integrations/websearch_interception/ARCHITECTURE.md index ce7f01c5a2a..4ea7a7ae527 100644 --- a/litellm/integrations/websearch_interception/ARCHITECTURE.md +++ b/litellm/integrations/websearch_interception/ARCHITECTURE.md @@ -207,6 +207,59 @@ response = await litellm.messages.acreate( --- +## Loop Ceiling + +One intercepted request can chain several follow-up model calls, since the model often searches again after +reading the first set of results. `max_agentic_loops` caps how many of those follow-ups run, and it defaults +to 3. LiteLLM also breaks the loop early when the model asks for the exact same tool call twice in a row. + +Set the ceiling on the feature, which the interceptor applies to `/v1/messages` requests: + +```yaml +litellm_settings: + websearch_interception_params: + enabled_providers: ["bedrock"] + max_agentic_loops: 5 +``` + +Or per deployment, which wins over the feature-level setting: + +```yaml +model_list: + - model_name: claude-sonnet-4-5 + litellm_params: + model: bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0 + max_agentic_loops: 5 +``` + +Clients cannot set it. `max_agentic_loops` is on the proxy's untrusted-field list, so a request body that +carries it is ignored and one request can never drive an unbounded number of upstream model calls. + +Both places are validated at config load, and a value that is not an integer of at least 1 stops the proxy +from starting rather than surfacing later. The per-deployment one is checked while the model list is read, +not on `LiteLLM_Params`, because the proxy builds its router with `ignore_invalid_deployments=True` and a +validator down there would drop the deployment silently instead of refusing to start. + +When the ceiling is reached on a non-streaming `/v1/messages` request, the turn ends there and the client gets +the last response back with the internal `litellm_web_search` tool call removed and `stop_reason: end_turn`. +The client never declared that tool, so leaving the block in would hand it a tool call it has no way to answer. +The answer can be less complete than it would have been with more loops, which is the tradeoff the ceiling +buys. Where the refused call was the only block left, the turn comes back with no text in it at all. + +Non-streaming is not a limitation on the client here, because a client that asked for a stream gets the same +treatment. Interception converts an intercepted `stream=True` request to non-streaming before the loop runs and +rebuilds the SSE stream from the finalized turn afterwards, so the ceiling is always reached on a response the +client has not seen yet. `AgenticStreamingIterator` is the one caller that reaches the loop with its events +already on the wire, and it keeps raising, because a finalized turn would arrive there as a second message +rather than as a replacement. + +Two other surfaces do not get that treatment yet. `/v1/responses` returns its own shape that the finalizer does +not rewrite, so it still hands back the internal call. And `/v1/chat/completions` runs its own copy of these +rails in `litellm_core_utils/chat_completion_agentic_loop.py`, which still raises rather than ending the turn. +Both are tracked separately + +--- + ## Streaming Support WebSearch interception works transparently with both streaming and non-streaming requests. diff --git a/litellm/integrations/websearch_interception/handler.py b/litellm/integrations/websearch_interception/handler.py index e59ef0449d0..13a16947fb4 100644 --- a/litellm/integrations/websearch_interception/handler.py +++ b/litellm/integrations/websearch_interception/handler.py @@ -31,6 +31,9 @@ from litellm.integrations.websearch_interception.tools import ( from litellm.integrations.websearch_interception.transformation import ( WebSearchTransformation, ) +from litellm.litellm_core_utils.agentic_loop_settings import ( + validated_max_agentic_loops, +) from litellm.llms.base_llm.search.transformation import SearchResponse from litellm.types.integrations.custom_logger import ( CHAT_COMPLETION_AGENTIC_SURFACE, @@ -122,6 +125,7 @@ class WebSearchInterceptionLogger(CustomLogger): self, enabled_providers: list[LlmProviders | str] | None = None, search_tool_name: str | None = None, + max_agentic_loops: int | None = None, ): """ Args: @@ -131,6 +135,9 @@ class WebSearchInterceptionLogger(CustomLogger): Default: None (all providers enabled) search_tool_name: Name of search tool configured in router's search_tools. If None, will attempt to use first available search tool. + max_agentic_loops: How many follow-up model calls one intercepted request + may chain before the loop is refused and the turn ends. + If None, LiteLLM's default of 3 applies. """ super().__init__() # Convert enum values to strings for comparison @@ -139,8 +146,16 @@ class WebSearchInterceptionLogger(CustomLogger): else: self.enabled_providers = [p.value if isinstance(p, LlmProviders) else p for p in enabled_providers] self.search_tool_name = search_tool_name + self.max_agentic_loops = self._validated_max_agentic_loops(max_agentic_loops) self._request_has_websearch = False # Track if current request has web search + @staticmethod + def _validated_max_agentic_loops(max_agentic_loops: object) -> int | None: + """ + Reject loop ceilings the agentic loop cannot honor, at config load time. + """ + return validated_max_agentic_loops(max_agentic_loops, field="websearch_interception_params.max_agentic_loops") + async def try_short_circuit_search( self, model: str, @@ -398,6 +413,7 @@ class WebSearchInterceptionLogger(CustomLogger): websearch_interception_params: enabled_providers: ["bedrock"] search_tool_name: "my-perplexity-search" + max_agentic_loops: 5 Usage: config = litellm_settings.get("websearch_interception_params", {}) @@ -406,6 +422,7 @@ class WebSearchInterceptionLogger(CustomLogger): # Extract parameters from config enabled_providers_str: Final = config.get("enabled_providers", None) search_tool_name: Final = config.get("search_tool_name", None) + max_agentic_loops: Final = config.get("max_agentic_loops", None) # Convert string provider names to LlmProviders enum values enabled_providers: list[LlmProviders | str] | None = None @@ -423,6 +440,7 @@ class WebSearchInterceptionLogger(CustomLogger): return cls( enabled_providers=enabled_providers, search_tool_name=search_tool_name, + max_agentic_loops=max_agentic_loops, ) @staticmethod @@ -493,6 +511,10 @@ class WebSearchInterceptionLogger(CustomLogger): verbose_logger.debug("WebSearchInterception: Pre-request hook triggered for provider=%s", custom_llm_provider) + deployment_max_agentic_loops: Final = kwargs.get("max_agentic_loops") + if self.max_agentic_loops is not None and deployment_max_agentic_loops is None: + kwargs["max_agentic_loops"] = self.max_agentic_loops # rebind-ok: this hook returns the kwargs it edits + # If the client sent an Anthropic-native web_search_* tool, mark the # request so the agentic loop emits native web_search_tool_result # blocks in the final response (for citations panels, etc.). The flag diff --git a/litellm/litellm_core_utils/agentic_loop_settings.py b/litellm/litellm_core_utils/agentic_loop_settings.py new file mode 100644 index 00000000000..3dd8d437aef --- /dev/null +++ b/litellm/litellm_core_utils/agentic_loop_settings.py @@ -0,0 +1,59 @@ +""" +Shared validation for the agentic loop ceiling. + +``max_agentic_loops`` can be set in two places, and the two disagreed about +what a bad value means. The feature-level +``litellm_settings.websearch_interception_params.max_agentic_loops`` was +checked at config load, while a per-deployment +``model_list[].litellm_params.max_agentic_loops`` was passed straight through +to ``int(... or 3)``. That let a per-deployment ``0`` read as the default 3, +turning the tightest ceiling into the loosest one, and let a per-deployment +``"three"`` boot the proxy and then fail every request to that model. + +Both settings now go through :func:`validated_max_agentic_loops`, which names +the field it rejected so the error says which line of the config to fix. + +Anything that spells a whole number is still accepted, because the old +``int(... or 3)`` accepted those and a ceiling is routinely parameterized as +``max_agentic_loops: os.environ/MAX_AGENTIC_LOOPS``, which resolves to a +string. Rejecting ``"5"`` would stop such a proxy from booting on upgrade. +""" + +from typing import Final + +DEFAULT_MAX_AGENTIC_LOOPS: Final = 3 + + +def _as_whole_number(value: object) -> int | None: + """ + Return ``value`` as an int when it spells a whole number, else ``None``. + + ``bool`` is excluded explicitly because it is an ``int`` subclass, so + ``max_agentic_loops: true`` would otherwise be read as a ceiling of 1. + """ + if isinstance(value, bool): + return None + if isinstance(value, int): + return value + if isinstance(value, float): + return int(value) if value.is_integer() else None + if isinstance(value, str): + try: + return int(value.strip()) + except ValueError: + return None + return None + + +def validated_max_agentic_loops(max_agentic_loops: object, field: str) -> int | None: + """ + Return ``max_agentic_loops`` as an int, or raise naming ``field``. + """ + if max_agentic_loops is None: + return None + ceiling: Final = _as_whole_number(max_agentic_loops) + if ceiling is None: + raise TypeError(f"{field} must be an integer, got {max_agentic_loops!r}") + if ceiling < 1: + raise ValueError(f"{field} must be at least 1, got {ceiling}") + return ceiling diff --git a/litellm/litellm_core_utils/chat_completion_agentic_loop.py b/litellm/litellm_core_utils/chat_completion_agentic_loop.py index b91c1785a54..07bed1f88ad 100644 --- a/litellm/litellm_core_utils/chat_completion_agentic_loop.py +++ b/litellm/litellm_core_utils/chat_completion_agentic_loop.py @@ -5,6 +5,10 @@ from typing import Final, cast from litellm._logging import verbose_logger from litellm.integrations.custom_logger import CustomLogger +from litellm.litellm_core_utils.agentic_loop_settings import ( + DEFAULT_MAX_AGENTIC_LOOPS, + validated_max_agentic_loops, +) from litellm.types.integrations.custom_logger import ( CHAT_COMPLETION_AGENTIC_SURFACE, NON_CODE_INTERPRETER_INTERCEPTION_INTERNAL_PREFIXES, @@ -52,7 +56,10 @@ def _coerce_int(value: object, default: int) -> int: def _agentic_loop_settings(kwargs: dict[str, object]) -> tuple[int, int, list[str]]: depth: Final = _coerce_int(kwargs.get("_agentic_loop_depth"), 0) - max_loops: Final = max(_coerce_int(kwargs.get("max_agentic_loops"), 3), 1) + configured: Final = validated_max_agentic_loops( + kwargs.get("max_agentic_loops"), field="litellm_params.max_agentic_loops" + ) + max_loops: Final = DEFAULT_MAX_AGENTIC_LOOPS if configured is None else configured raw_fingerprints: Final = kwargs.get("_agentic_loop_fingerprints") fingerprints: Final = [str(fp) for fp in raw_fingerprints] if isinstance(raw_fingerprints, list) else [] return depth, max_loops, fingerprints diff --git a/litellm/litellm_core_utils/logging_utils.py b/litellm/litellm_core_utils/logging_utils.py index a17415f3ab8..91c8ba36b26 100644 --- a/litellm/litellm_core_utils/logging_utils.py +++ b/litellm/litellm_core_utils/logging_utils.py @@ -3,6 +3,7 @@ import functools import inspect import re import time +from collections.abc import Mapping from datetime import datetime from typing import TYPE_CHECKING, Any, Final @@ -268,6 +269,16 @@ def _set_duration_in_model_call_details( verbose_logger.warning("Error setting `llm_api_duration_ms`: %s", e) +def speech_request_body(model: str, voice: str, optional_params: Mapping[str, object]) -> Mapping[str, object]: + """Speech request body for telemetry, without the caller headers the provider SDKs + take as request kwargs rather than body fields.""" + return { # mutable-ok: loggers isinstance-check the request body as a dict + "model": model, + "voice": voice, + **{key: value for key, value in optional_params.items() if key != "extra_headers"}, + } + + def track_llm_api_timing(): """ Decorator to track LLM API call timing for both sync and async functions. diff --git a/litellm/llms/anthropic/experimental_pass_through/messages/fake_stream_iterator.py b/litellm/llms/anthropic/experimental_pass_through/messages/fake_stream_iterator.py index 215d4a5b42b..14f1b7697cf 100644 --- a/litellm/llms/anthropic/experimental_pass_through/messages/fake_stream_iterator.py +++ b/litellm/llms/anthropic/experimental_pass_through/messages/fake_stream_iterator.py @@ -113,6 +113,14 @@ class FakeAnthropicMessagesStreamIterator: } chunks.append(f"event: content_block_delta\ndata: {json.dumps(content_block_delta)}\n\n".encode()) + else: + passthrough_start: Final = { + "type": "content_block_start", + "index": index, + "content_block": block_dict, + } + chunks.append(f"event: content_block_start\ndata: {json.dumps(passthrough_start)}\n\n".encode()) + content_block_stop: Final = {"type": "content_block_stop", "index": index} chunks.append(f"event: content_block_stop\ndata: {json.dumps(content_block_stop)}\n\n".encode()) return chunks diff --git a/litellm/llms/azure/azure.py b/litellm/llms/azure/azure.py index c8f94b575ad..980b27cda55 100644 --- a/litellm/llms/azure/azure.py +++ b/litellm/llms/azure/azure.py @@ -17,7 +17,7 @@ from openai import ( import litellm from litellm.constants import AZURE_OPERATION_POLLING_TIMEOUT, DEFAULT_MAX_RETRIES from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj -from litellm.litellm_core_utils.logging_utils import track_llm_api_timing +from litellm.litellm_core_utils.logging_utils import speech_request_body, track_llm_api_timing from litellm.litellm_core_utils.url_utils import SSRFError, assert_same_origin from litellm.llms.custom_httpx.http_handler import ( AsyncHTTPHandler, @@ -1352,6 +1352,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): organization: str | None, max_retries: int, timeout: float | httpx.Timeout, + logging_obj: LiteLLMLoggingObj, azure_ad_token: str | None = None, azure_ad_token_provider: Callable | None = None, aspeech: bool | None = None, @@ -1373,6 +1374,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): azure_ad_token_provider=azure_ad_token_provider, max_retries=max_retries, timeout=timeout, + logging_obj=logging_obj, client=client, litellm_params=litellm_params, ) @@ -1387,6 +1389,15 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): litellm_params=litellm_params, ) + logging_obj.pre_call( + input=input, + api_key=api_key, + additional_args={ # mutable-ok: loggers isinstance-check this payload as a dict + "complete_input_dict": speech_request_body(model, voice, optional_params), + "api_base": str(azure_client.base_url), + }, + ) + response: Final = azure_client.audio.speech.create( model=model, voice=voice, @@ -1408,6 +1419,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): azure_ad_token_provider: Callable | None, max_retries: int, timeout: float | httpx.Timeout, + logging_obj: LiteLLMLoggingObj, client=None, litellm_params: dict | None = None, ) -> HttpxBinaryResponseContent: @@ -1421,6 +1433,15 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): litellm_params=litellm_params, ) + logging_obj.pre_call( + input=input, + api_key=api_key, + additional_args={ # mutable-ok: loggers isinstance-check this payload as a dict + "complete_input_dict": speech_request_body(model, voice, optional_params), + "api_base": str(azure_client.base_url), + }, + ) + azure_response: Final = await azure_client.audio.speech.create( model=model, voice=voice, diff --git a/litellm/llms/base_llm/files/transformation.py b/litellm/llms/base_llm/files/transformation.py index 174be93448b..b20fe0f1560 100644 --- a/litellm/llms/base_llm/files/transformation.py +++ b/litellm/llms/base_llm/files/transformation.py @@ -11,9 +11,9 @@ from litellm.types.llms.openai import ( AllMessageValues, CreateFileRequest, FileContentRequest, + FileListPage, OpenAICreateFileRequestOptionalParams, OpenAIFileObject, - OpenAIFilesPurpose, ) from litellm.types.utils import LlmProviders, ModelResponse @@ -240,10 +240,13 @@ class BaseFileEndpoints(ABC): @abstractmethod async def afile_list( self, - purpose: OpenAIFilesPurpose | None, + purpose: str | None, litellm_parent_otel_span: Span | None, + user_api_key_dict: UserAPIKeyAuth, + limit: int | None = None, + after: str | None = None, **data: dict, - ) -> list[OpenAIFileObject]: + ) -> FileListPage: pass @abstractmethod diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 369e150f6bd..ed079197513 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -19,6 +19,10 @@ import litellm.types.utils from litellm._logging import _redact_string, verbose_logger from litellm.anthropic_beta_headers_manager import update_headers_with_filtered_beta from litellm.constants import REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES +from litellm.litellm_core_utils.agentic_loop_settings import ( + DEFAULT_MAX_AGENTIC_LOOPS, + validated_max_agentic_loops, +) from litellm.litellm_core_utils.asyncify import run_async_function from litellm.litellm_core_utils.realtime_errors import realtime_error_event, websocket_close_reason from litellm.litellm_core_utils.realtime_streaming import RealTimeStreaming @@ -89,6 +93,7 @@ from litellm.types.files import StreamingMediaUploadConfig, TwoStepFileUploadCon from litellm.types.integrations.custom_logger import ( AgenticLoopPlan, AgenticLoopRequestPatch, + AgenticLoopSafetyError, ) from litellm.types.llms.anthropic_messages.anthropic_response import ( AnthropicMessagesResponse, @@ -5077,9 +5082,12 @@ class BaseLLMHTTPHandler: @staticmethod def _get_agentic_loop_settings(kwargs: dict) -> tuple[int, int, list[str]]: depth: Final = int(kwargs.get("_agentic_loop_depth", 0) or 0) - max_loops: Final = int(kwargs.get("max_agentic_loops", 3) or 3) + configured: Final = validated_max_agentic_loops( + kwargs.get("max_agentic_loops"), field="litellm_params.max_agentic_loops" + ) + max_loops: Final = DEFAULT_MAX_AGENTIC_LOOPS if configured is None else configured fingerprints: Final = list(kwargs.get("_agentic_loop_fingerprints", []) or []) - return depth, max(max_loops, 1), fingerprints + return depth, max_loops, fingerprints @staticmethod def _has_agentic_completion_hook(logging_obj: LiteLLMLoggingObj) -> bool: @@ -5122,7 +5130,8 @@ class BaseLLMHTTPHandler: """ Evaluate agentic-loop safety guards (fingerprint cycle / max depth). - Raises ValueError on abort. Returns the current fingerprint on success. + Raises AgenticLoopSafetyError on abort. Returns the current fingerprint + on success. These checks must not be swallowed by the per-callback ``except Exception`` block that wraps callback dispatch — they are bounded-loop / cycle-break @@ -5130,9 +5139,9 @@ class BaseLLMHTTPHandler: """ fingerprint: Final = BaseLLMHTTPHandler._fingerprint_agentic_tools(tool_calls) if fingerprint in fingerprints: - raise ValueError("Agentic loop detected repeated tool-call fingerprint; aborting rerun") + raise AgenticLoopSafetyError("Agentic loop detected repeated tool-call fingerprint; aborting rerun") if depth >= max_loops: - raise ValueError(f"Exceeded max_agentic_loops={max_loops} for model={model}") + raise AgenticLoopSafetyError(f"Exceeded max_agentic_loops={max_loops} for model={model}") return fingerprint @staticmethod @@ -5142,6 +5151,97 @@ class BaseLLMHTTPHandler: except Exception: return str(tools) + @staticmethod + def _refused_agentic_tool_identifiers(tool_calls: object) -> tuple[frozenset[str], frozenset[str]]: + """ + Collect the ids and names of the tool calls a safety rail just refused. + + Callbacks hand back either a bare list of tool calls or a dict wrapping + that list under ``tool_calls``, and both the anthropic and responses + shapes carry an ``id`` (or ``call_id``) plus a ``name``. + """ + calls: Final = tool_calls.get("tool_calls") if isinstance(tool_calls, dict) else tool_calls + if not isinstance(calls, list): + return frozenset(), frozenset() + dict_calls: Final = (call for call in calls if isinstance(call, dict)) + fields: Final = tuple((call.get("id"), call.get("call_id"), call.get("name")) for call in dict_calls) + ids: Final = frozenset( + value for call_id, caller_id, _ in fields for value in (call_id, caller_id) if isinstance(value, str) + ) + names: Final = frozenset(name for _, _, name in fields if isinstance(name, str)) + return ids, names + + @staticmethod + def _is_refused_tool_use_block(block: object, refused_ids: frozenset[str], refused_names: frozenset[str]) -> bool: + """ + Whether this response block belongs to a tool call the rail refused. + + An id settles it on its own, so a block carrying one is matched on the id + alone and a client's own tool call survives even where it happens to + share a name with a refused one. The name is only consulted for tool call + shapes that arrive without an id. + """ + if not isinstance(block, dict) or block.get("type") != "tool_use": + return False + block_id: Final = block.get("id") + if isinstance(block_id, str) and refused_ids: + return block_id in refused_ids + return block.get("name") in refused_names + + @staticmethod + def _can_replace_turn_with_terminal_response(stream: bool, api_surface: str) -> bool: + """ + Whether a refused rerun can still be answered with a finalized turn. + + Only the anthropic messages surface can. The responses surface carries a + pydantic model the finalizer does not rewrite, so it keeps raising, which + is what every surface did before this path learned to end the turn. + + The messages and responses call sites pass ``stream=False``, because + interception converts an intercepted stream to non-streaming before the + loop runs and rebuilds the SSE stream from the finalized turn + afterwards. ``AgenticStreamingIterator`` passes ``stream=True``, and + that path keeps raising: its events are already on the wire, so a + finalized turn would reach the client as a second message rather than + as a replacement. + """ + return not stream and api_surface == "anthropic_messages" + + @staticmethod + def _finalize_refused_agentic_response(response: object, tool_calls: object) -> object: + """ + Turn the response into a terminal turn after a safety rail refused the rerun. + + The refused tool calls target tools LiteLLM injected on the client's + behalf, so a client that never declared them cannot send back a matching + ``tool_result``. Their blocks are dropped and a ``tool_use`` stop reason + is closed out as ``end_turn``, which is what a provider-native web search + turn returns once it stops calling tools. + + A ``tool_use`` block the client itself declared is left alone, and while + one is still in the response the stop reason stays ``tool_use`` so the + client knows to answer it. + """ + if not isinstance(response, dict): + return response + + refused_ids, refused_names = BaseLLMHTTPHandler._refused_agentic_tool_identifiers(tool_calls) + finalized: Final = dict(response) + content: Final = finalized.get("content") + if isinstance(content, list): + kept_blocks: Final = [ + block + for block in content + if not BaseLLMHTTPHandler._is_refused_tool_use_block(block, refused_ids, refused_names) + ] + finalized["content"] = kept_blocks + client_tool_use_remains: Final = any( + isinstance(block, dict) and block.get("type") == "tool_use" for block in kept_blocks + ) + if not client_tool_use_remains and finalized.get("stop_reason") == "tool_use": + finalized["stop_reason"] = "end_turn" + return finalized + async def _execute_anthropic_agentic_plan( self, plan: AgenticLoopPlan, @@ -5507,14 +5607,30 @@ class BaseLLMHTTPHandler: continue # Safety guards must run OUTSIDE the callback try/except — they are - # bounded-loop / cycle-break rails that must propagate to the caller. - fingerprint = self._check_agentic_loop_safety( - tool_calls=tool_calls, - fingerprints=fingerprints, - depth=depth, - max_loops=max_loops, - model=model, - ) + # bounded-loop / cycle-break rails, not callback bugs. + try: + fingerprint = self._check_agentic_loop_safety( + tool_calls=tool_calls, + fingerprints=fingerprints, + depth=depth, + max_loops=max_loops, + model=model, + ) + except AgenticLoopSafetyError as e: + if not self._can_replace_turn_with_terminal_response(stream, api_surface): + raise + _call_id = getattr(logging_obj, "litellm_call_id", "unknown") + verbose_logger.warning( + "LiteLLM.AgenticLoopRefused: ending turn [call_id=%s model=%s]: %s", + _call_id, + model, + str(e), + ) + return self._maybe_wrap_in_fake_stream( + self._finalize_refused_agentic_response(response=response, tool_calls=tool_calls), + logging_obj, + api_surface, + ) try: kwargs_with_provider = hook_kwargs.copy() diff --git a/litellm/llms/openai/openai.py b/litellm/llms/openai/openai.py index 4fc6655ca54..ee0efb88a38 100644 --- a/litellm/llms/openai/openai.py +++ b/litellm/llms/openai/openai.py @@ -22,7 +22,7 @@ from litellm._logging import verbose_logger from litellm.constants import DEFAULT_MAX_RETRIES from litellm.files.types import FileContentStreamingResult from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj -from litellm.litellm_core_utils.logging_utils import track_llm_api_timing +from litellm.litellm_core_utils.logging_utils import speech_request_body, track_llm_api_timing from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException from litellm.llms.bedrock.chat.invoke_handler import MockResponseIterator @@ -1365,9 +1365,21 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): client=client, ) - if headers: - data["extra_headers"] = headers - response = await openai_aclient.images.generate(**data, timeout=timeout) + logging_obj.pre_call( + input=prompt, + api_key=openai_aclient.api_key, + additional_args={ # mutable-ok: loggers isinstance-check this payload as a dict + "headers": {"Authorization": f"Bearer {openai_aclient.api_key}"}, # mutable-ok: logged header map + "api_base": str(openai_aclient.base_url), + "acompletion": True, + "complete_input_dict": data, + }, + ) + + request_data: Final = ( # mutable-ok: the OpenAI SDK takes the request body as a dict + {**data, "extra_headers": headers} if headers else data + ) + response = await openai_aclient.images.generate(**request_data, timeout=timeout) stringified_response: Final = response.model_dump() ## LOGGING logging_obj.post_call( @@ -1450,9 +1462,10 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): ) ## COMPLETION CALL - if headers: - data["extra_headers"] = headers - _response: Final = openai_client.images.generate(**data, timeout=timeout) + request_data: Final = ( # mutable-ok: the OpenAI SDK takes the request body as a dict + {**data, "extra_headers": headers} if headers else data + ) + _response: Final = openai_client.images.generate(**request_data, timeout=timeout) response: Final = _response.model_dump() ## LOGGING @@ -1501,6 +1514,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): project: str | None, max_retries: int, timeout: float | httpx.Timeout, + logging_obj: LiteLLMLoggingObj, aspeech: bool | None = None, client=None, shared_session: Optional["ClientSession"] = None, @@ -1517,6 +1531,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): project=project, max_retries=max_retries, timeout=timeout, + logging_obj=logging_obj, client=client, shared_session=shared_session, ) @@ -1531,7 +1546,17 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): shared_session=shared_session, ) - response: Final = cast(OpenAI, openai_client).audio.speech.create( + sync_client: Final = cast(OpenAI, openai_client) + logging_obj.pre_call( + input=input, + api_key=api_key, + additional_args={ # mutable-ok: loggers isinstance-check this payload as a dict + "complete_input_dict": speech_request_body(model, voice, optional_params), + "api_base": str(sync_client.base_url), + }, + ) + + response: Final = sync_client.audio.speech.create( model=model, voice=voice, input=input, @@ -1551,6 +1576,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): project: str | None, max_retries: int, timeout: float | httpx.Timeout, + logging_obj: LiteLLMLoggingObj, client=None, shared_session: Optional["ClientSession"] = None, ) -> HttpxBinaryResponseContent: @@ -1567,6 +1593,15 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): ), ) + logging_obj.pre_call( + input=input, + api_key=api_key, + additional_args={ # mutable-ok: loggers isinstance-check this payload as a dict + "complete_input_dict": speech_request_body(model, voice, optional_params), + "api_base": str(openai_client.base_url), + }, + ) + response: Final = await openai_client.audio.speech.create( model=model, voice=voice, diff --git a/litellm/main.py b/litellm/main.py index 52785e7a393..2cf53833c5a 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -7537,6 +7537,15 @@ async def amoderation( }, custom_llm_provider=custom_llm_provider, ) + moderation_request: Final = {"input": input, "model": model} # mutable-ok: logged as the raw request body + litellm_logging_obj.pre_call( + input=input, + api_key=api_key, + additional_args={ # mutable-ok: loggers isinstance-check this payload as a dict + "complete_input_dict": moderation_request, + "api_base": str(_openai_client.base_url), + }, + ) if model is not None: response = await _openai_client.moderations.create(input=input, model=model) @@ -8042,6 +8051,7 @@ def speech( project=project, max_retries=max_retries, timeout=timeout, + logging_obj=logging_obj, client=client, # pass AsyncOpenAI, OpenAI client aspeech=aspeech, shared_session=shared_session, @@ -8120,6 +8130,7 @@ def speech( organization=organization, max_retries=max_retries, timeout=timeout, + logging_obj=logging_obj, client=client, # pass AsyncOpenAI, OpenAI client aspeech=aspeech, litellm_params=litellm_params_dict, diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 12d6b44a648..e7b98b3cc7f 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -19,6 +19,7 @@ from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Protocol, cast from fastapi import HTTPException, Request, status from pydantic import BaseModel +from typing_extensions import ReadOnly, TypedDict import litellm from litellm._logging import verbose_proxy_logger @@ -2512,6 +2513,27 @@ async def delete_cache_key_objects( await publish_auth_cache_invalidation(cache_key=hashed_token) +class _TeamNotFoundDetail(TypedDict): + error: ReadOnly[str] + + +class TeamNotFoundError(HTTPException): + """The team row is provably absent, as opposed to merely unreadable. + + ``get_team_object`` reports every failure as a 404, so a deleted team and a + database that would not answer are indistinguishable to its callers. Callers + that must not treat a degraded read as a definitive answer, such as the + authorization fallback in ``user_api_key_auth``, key on this subclass. It + stays a 404 carrying the same detail, so every other caller is unaffected. + """ + + def __init__(self, team_id: str) -> None: + detail: Final[_TeamNotFoundDetail] = { + "error": f"Team doesn't exist in db. Team={team_id}. Create team via `/team/new` call." + } + super().__init__(status_code=404, detail=detail) + + @log_db_metrics async def _get_team_db_check( team_id: str, prisma_client: PrismaClient, team_id_upsert: bool | None = None @@ -2557,6 +2579,10 @@ async def _get_team_object_from_user_api_key_cache( ) if should_check_db: response = await _get_team_db_check(team_id=team_id, prisma_client=prisma_client, team_id_upsert=team_id_upsert) + # The database answered and the row is not there. Distinct from every + # other failure here, which leaves the team's grant unknown. + if response is None: + raise TeamNotFoundError(team_id=team_id) else: response = None @@ -2678,6 +2704,8 @@ async def get_team_object( key=key, team_id_upsert=team_id_upsert, ) + except TeamNotFoundError: + raise except Exception: raise HTTPException( status_code=404, diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index fe4f1ee4ae5..658d176f6a7 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -35,6 +35,7 @@ from litellm.litellm_core_utils.dot_notation_indexing import get_nested_value from litellm.proxy._types import * from litellm.proxy.auth.auth_checks import ( ExperimentalUIJWTToken, + TeamNotFoundError, _cache_key_object, _can_object_call_model, _check_end_user_budget, @@ -49,6 +50,7 @@ from litellm.proxy.auth.auth_checks import ( common_checks, get_end_user_object, get_jwt_key_mapping_object, + get_object_permission, get_project_object, get_team_object, get_user_object, @@ -86,6 +88,7 @@ from litellm.proxy.common_utils.http_parsing_utils import ( ) from litellm.proxy.common_utils.realtime_utils import _realtime_request_body from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache +from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup from litellm.proxy.utils import ( PrismaClient, @@ -1066,6 +1069,26 @@ async def _resolve_jwt_to_virtual_key( return None +def _ensure_litellm_received_at_on_request_state(request: Request) -> datetime: + """Idempotently stamp ``request.state.litellm_received_at`` with the moment + litellm's own code started handling this request -- the first line of + ``user_api_key_auth``, before any auth/pre-call work runs. This is the + basis for the request-latency Prometheus metrics (see + ``litellm/integrations/prometheus.py``), and unlike the OTEL SERVER span + below, it is set unconditionally so those metrics don't depend on OTEL + being configured. + """ + existing_received_at: Final[datetime | None] = getattr(request.state, "litellm_received_at", None) + if existing_received_at is not None: + return existing_received_at + received_at: Final = datetime.now(timezone.utc) + try: + request.state.litellm_received_at = received_at + except Exception: + pass + return received_at + + def _ensure_parent_otel_span_on_request_state(request: Request) -> None: """Idempotently create the OTEL SERVER span and stash it on ``request.state.parent_otel_span``. Safe to call multiple times. @@ -1076,15 +1099,12 @@ def _ensure_parent_otel_span_on_request_state(request: Request) -> None: """ from litellm.proxy.proxy_server import open_telemetry_logger + start_time: Final = _ensure_litellm_received_at_on_request_state(request) + if open_telemetry_logger is None: return if getattr(request.state, "parent_otel_span", None) is not None: return - start_time: Final = datetime.now(timezone.utc) - try: - request.state.litellm_received_at = start_time - except Exception: - pass parent_otel_span: Final = open_telemetry_logger.create_litellm_proxy_request_started_span( start_time=start_time, headers=_safe_get_request_headers(request), @@ -1142,6 +1162,26 @@ async def _record_unparsable_body_failure( verbose_proxy_logger.exception("Failed to log the request rejected for an unparsable body: %s", e) +async def _resolve_object_permission_for_unresolvable_team( + object_permission_id: str | None, + prisma_client: PrismaClient | None, + user_api_key_cache: UserApiKeyCache, + parent_otel_span: Span | None, + proxy_logging_obj: ProxyLogging, +) -> LiteLLM_ObjectPermissionTable | None: + """Re-resolve a team's object permission by id when the team row itself is unreadable, so the + token-derived fallback doesn't silently drop it.""" + if object_permission_id is None or prisma_client is None: + return None + return await get_object_permission( + object_permission_id=object_permission_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + parent_otel_span=parent_otel_span, + proxy_logging_obj=proxy_logging_obj, + ) + + async def _user_api_key_auth_builder( request: Request, api_key: str, @@ -2114,6 +2154,13 @@ async def _user_api_key_auth_builder( models=valid_token.team_models, metadata=valid_token.team_metadata, object_permission_id=valid_token.team_object_permission_id, + object_permission=await _resolve_object_permission_for_unresolvable_team( + object_permission_id=valid_token.team_object_permission_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + parent_otel_span=parent_otel_span, + proxy_logging_obj=proxy_logging_obj, + ), ) else: _team_obj = None @@ -2262,6 +2309,36 @@ def _team_obj_from_token(valid_token: UserAPIKeyAuth) -> LiteLLM_TeamTableCached ) +def _token_can_vouch_for_team(valid_token: UserAPIKeyAuth, lookup_error: BaseException) -> bool: + """Whether the token's own team fields may stand in for a team that failed to + resolve, without widening access. + + The UI dashboard mints every session key against the ``UI_TEAM_ID`` sentinel, + which by design never has a team row, so a failed lookup for it is not a + degraded read to be treated with suspicion; it always vouches, exactly as it + always safely has (these keys are restricted elsewhere to UI-only routes). + + For every other team, a team that is provably gone is a definitive answer, + not a degraded read, so nothing may stand in for it and no setting may + override that. + + Otherwise the team's grant is merely unknown. A token carrying one may vouch, + since replaying a recorded grant cannot widen it and denying every team key + while the row is briefly unreadable would trade the widening for an outage. A + token carrying none may not: ``team_models=[]`` reads as every model and + ``team_blocked=False`` as unblocked. ``allow_requests_on_db_unavailable`` opts + back out, and is only consulted here because the failure is known by this + point to be a degraded read. + """ + if valid_token.team_id == UI_TEAM_ID: + return True + if isinstance(lookup_error, TeamNotFoundError): + return False + if valid_token.team_models: + return True + return PrismaDBExceptionHandler.should_allow_request_on_db_unavailable() + + @tracer.wrap() async def _run_centralized_common_checks( user_api_key_auth_obj: UserAPIKeyAuth, @@ -2466,7 +2543,12 @@ async def _run_centralized_common_checks( if isinstance(team_result, BaseException): # Token-derived fallback only valid when a team_id is set; # _team_obj_from_token asserts that precondition. - team_object = _team_obj_from_token(user_api_key_auth_obj) if user_api_key_auth_obj.team_id is not None else None + if user_api_key_auth_obj.team_id is None: + team_object = None + elif _token_can_vouch_for_team(user_api_key_auth_obj, team_result): + team_object = _team_obj_from_token(user_api_key_auth_obj) + else: + raise team_result else: team_object = team_result diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 3fb09cde931..dbbf9cb673e 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -3,7 +3,6 @@ import contextlib import json import logging import math -import time import traceback from collections.abc import AsyncGenerator, Awaitable, Callable, Mapping from datetime import datetime @@ -178,6 +177,7 @@ else: ProxyConfig = Any from litellm.proxy.litellm_pre_call_utils import ( add_litellm_data_to_request, + refresh_proxy_server_request_body_snapshot, reject_url_valued_destination, ) from litellm.types.utils import ( @@ -1725,13 +1725,17 @@ class ProxyBaseLLMRequestProcessing: ) # Calculate request queue time after add_litellm_data_to_request - # which sets arrival_time in proxy_server_request + # which sets arrival_time in proxy_server_request. Ends at start_time + # (not a freshly captured time.time() here) so this window is exactly + # [arrival_time, start_time], with zero overlap with the + # litellm_request_total_latency_metric window of [start_time, end_time] -- + # otherwise the few lines of add_litellm_data_to_request's own work would + # be double-counted across both metrics. proxy_server_request: Final = self.data.get("proxy_server_request", {}) arrival_time: Final = proxy_server_request.get("arrival_time") queue_time_seconds = None if arrival_time is not None: - processing_start_time: Final = time.time() - queue_time_seconds = processing_start_time - arrival_time + queue_time_seconds = start_time.timestamp() - arrival_time # Store queue time in metadata after add_litellm_data_to_request to ensure it's preserved if queue_time_seconds is not None: @@ -1862,6 +1866,12 @@ class ProxyBaseLLMRequestProcessing: call_type=route_type, ) + # Refresh AFTER pre_call_hook: guardrails (e.g. Presidio PII masking) may + # have mutated `self.data` in place, and the audit-trail snapshot taken in + # add_litellm_data_to_request predates that mutation. + refresh_proxy_server_request_body_snapshot(self.data) + verbose_proxy_logger.debug("receiving data: %s", self.data) + if "messages" in self.data and self.data["messages"]: logging_obj.update_messages(self.data["messages"]) diff --git a/litellm/proxy/guardrails/_content_utils.py b/litellm/proxy/guardrails/_content_utils.py index 6ed6f0013df..ae92adcb1ee 100644 --- a/litellm/proxy/guardrails/_content_utils.py +++ b/litellm/proxy/guardrails/_content_utils.py @@ -8,7 +8,7 @@ skip the other shapes — these helpers normalise that so every hook sees every text fragment. """ -from collections.abc import Callable, Iterator +from collections.abc import Callable, Iterator, Mapping from typing import Any, Final # Call types whose body carries free-form chat / prompt text that @@ -33,7 +33,9 @@ def is_text_content_call_type(call_type: str) -> bool: return call_type in TEXT_CONTENT_CALL_TYPES -TEXT_PART_TYPES: Final[frozenset[str]] = frozenset({"text", "input_text", "output_text"}) +TEXT_PART_TYPES: Final[frozenset[str]] = frozenset( + {"text", "input_text", "output_text", "summary_text", "reasoning_text"} +) # Responses-API item types whose ``output`` field carries user/tool text # that guardrails should inspect. ``function_call_output`` is the @@ -42,6 +44,16 @@ TEXT_PART_TYPES: Final[frozenset[str]] = frozenset({"text", "input_text", "outpu _OUTPUT_ITEM_TYPES: Final[frozenset[str]] = frozenset({"function_call_output", "custom_tool_call_output"}) +def _part_text(part: Mapping[str, object]) -> str | None: + """Return non-empty plaintext from any content part that carries ``text``.""" + if not isinstance(part, dict): + return None + text = part.get("text") + if isinstance(text, str) and text: + return text + return None + + def _iter_text_parts_in_content(content: Any) -> Iterator[str]: """Yield text fragments from a ``message.content`` value (string or multimodal list). Non-text parts (images, audio, …) are skipped.""" @@ -58,10 +70,9 @@ def _iter_text_parts_in_content(content: Any) -> Iterator[str]: continue if not isinstance(part, dict): continue - if part.get("type") in TEXT_PART_TYPES: - text = part.get("text") - if isinstance(text, str) and text: - yield text + text = _part_text(part) + if text is not None: + yield text def _coerce_input_to_messages(input_value: Any) -> list[dict[str, Any]]: @@ -75,8 +86,23 @@ def _coerce_input_to_messages(input_value: Any) -> list[dict[str, Any]]: if isinstance(item, str): messages.append({"role": "user", "content": item}) elif isinstance(item, dict): - if item.get("type") in TEXT_PART_TYPES: + if _part_text(item) is not None: messages.append({"role": item.get("role") or "user", "content": [item]}) + elif item.get("type") == "reasoning": + if "content" in item: + messages.append( + { # mutable-ok: append reasoning content + "role": item.get("role") or "assistant", + "content": item["content"], + } + ) + if isinstance(item.get("summary"), list): + messages.append( + { # mutable-ok: append reasoning summary + "role": item.get("role") or "assistant", + "content": item["summary"], + } + ) elif "content" in item: messages.append({"role": item.get("role") or "user", "content": item["content"]}) elif item.get("type") in _OUTPUT_ITEM_TYPES and "output" in item: @@ -126,12 +152,7 @@ def walk_user_text(data: dict[str, Any], visit: Callable[[str], str]) -> int: if isinstance(part, str) and part: visited += 1 new_parts.append(visit(part)) - elif ( - isinstance(part, dict) - and part.get("type") in TEXT_PART_TYPES - and isinstance(part.get("text"), str) - and part["text"] - ): + elif isinstance(part, dict) and _part_text(part) is not None: visited += 1 new_parts.append({**part, "text": visit(part["text"])}) else: @@ -158,10 +179,14 @@ def walk_user_text(data: dict[str, Any], visit: Callable[[str], str]) -> int: visited += 1 input_value[idx] = visit(item) elif isinstance(item, dict): - if item.get("type") in TEXT_PART_TYPES: - if isinstance(item.get("text"), str) and item["text"]: - visited += 1 - input_value[idx] = {**item, "text": visit(item["text"])} + if _part_text(item) is not None: + visited += 1 + input_value[idx] = {**item, "text": visit(item["text"])} # mutable-ok: rewrite text part in place + elif item.get("type") == "reasoning": + if "content" in item: + item["content"] = _rewrite_content(item["content"]) + if isinstance(item.get("summary"), list): + item["summary"] = _rewrite_content(item["summary"]) elif "content" in item: item["content"] = _rewrite_content(item["content"]) elif item.get("type") in _OUTPUT_ITEM_TYPES and "output" in item: diff --git a/litellm/proxy/guardrails/guardrail_hooks/presidio.py b/litellm/proxy/guardrails/guardrail_hooks/presidio.py index c3b7498d9ec..bcee45355e3 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/presidio.py +++ b/litellm/proxy/guardrails/guardrail_hooks/presidio.py @@ -788,7 +788,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): async def async_logging_hook(self, kwargs: dict, result: Any, call_type: str) -> tuple[dict, Any]: """ - Masks the input before logging to langfuse, datadog, etc. + Masks the input and output before logging to langfuse, datadog, etc. """ if call_type == "completion" or call_type == "acompletion": # /chat/completions requests messages: Final[list | None] = kwargs.get("messages", None) @@ -847,6 +847,19 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): verbose_proxy_logger.debug("Presidio PII Masking: Redacted pii message: %s", messages) kwargs["messages"] = messages + if ( + isinstance(result, ModelResponse) + and result.choices + and not isinstance(result.choices[0], StreamingChoices) + ): + await self._process_response_for_pii(response=result, request_data=kwargs, mode="mask") + elif self._is_anthropic_message_response(result): + await self._process_anthropic_response_for_pii( + response=cast(dict, result), # cast-ok: _is_anthropic_message_response narrows via isinstance + request_data=kwargs, + mode="mask", + ) + return kwargs, result async def async_post_call_success_hook( diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index c1099081867..4794da05a3e 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -1622,6 +1622,32 @@ class LiteLLMProxyRequestSetup: ) +def refresh_proxy_server_request_body_snapshot( + data: dict, # mutable-ok: mutates proxy_server_request.body in place on the shared request dict +) -> None: + """ + Re-snapshot ``data["proxy_server_request"]["body"]`` from the current state of ``data``. + + ``add_litellm_data_to_request`` takes the initial snapshot before guardrails + (pre_call_hook) run. A guardrail that masks PII/PCI in place (e.g. Presidio) + mutates ``data`` afterward, so callers that persist ``proxy_server_request.body`` + for audit/spend-tracking purposes must call this again post-guardrail, or the + persisted body silently bypasses whatever masking the guardrail applied. + + By the time a caller refreshes post-guardrail, ``litellm.utils.function_setup`` + has already stamped ``data["litellm_logging_obj"]`` with a live (non-serializable) + ``Logging`` instance, so it must be excluded here the same way ``secret_fields`` + and ``proxy_server_request`` are. + """ + proxy_server_request = data.get("proxy_server_request") + if not isinstance(proxy_server_request, dict): + return + _body_snapshot_exclude = ( + frozenset({"secret_fields", "proxy_server_request", "litellm_logging_obj"}) | _TRANSPORT_ONLY_CREDENTIAL_KEYS + ) + proxy_server_request["body"] = {k: v for k, v in data.items() if k not in _body_snapshot_exclude} + + async def add_litellm_data_to_request( data: dict, request: Request, @@ -1718,11 +1744,17 @@ async def add_litellm_data_to_request( # Init - Proxy Server Request # we do this as soon as entering so we track the original request ########################################################## - # Track arrival time for queue time metric. The body snapshot is filled - # in after the admin-injection strip below so the audit / spend-tracking - # consumers of proxy_server_request["body"] see the cleaned metadata - # rather than attacker-forged user_api_key_* fields. - arrival_time: Final = time.time() + # Track arrival time for queue time metric. Prefer the timestamp stamped at + # the top of user_api_key_auth (request.state.litellm_received_at): by the + # time this function runs, auth has already completed, so time.time() here + # would silently exclude the entire auth phase from the queue-time window. + # Falls back to time.time() for callers that never went through + # user_api_key_auth. The body snapshot is filled in after the + # admin-injection strip below so the audit / spend-tracking consumers of + # proxy_server_request["body"] see the cleaned metadata rather than + # attacker-forged user_api_key_* fields. + _litellm_received_at: Final = getattr(request.state, "litellm_received_at", None) + arrival_time: Final = _litellm_received_at.timestamp() if _litellm_received_at is not None else time.time() data["proxy_server_request"] = { "url": str(request.url), "method": request.method, @@ -1802,8 +1834,6 @@ async def add_litellm_data_to_request( cache_dict: Final = parse_cache_control(cache_control_header) data["ttl"] = cache_dict.get("s-maxage") - verbose_proxy_logger.debug("receiving data: %s", data) - # requester_metadata is snapshotted AFTER the strip below so # downstream consumers (e.g. PANW guardrail reading user_ip / # profile_id) don't see attacker-injected admin slots preserved in @@ -1863,9 +1893,7 @@ async def add_litellm_data_to_request( # self-reference — body.proxy_server_request.body would be the same # dict as body, producing an infinite traversal loop for any consumer # that walks the structure. - _body_snapshot_exclude = frozenset({"secret_fields", "proxy_server_request"}) | _TRANSPORT_ONLY_CREDENTIAL_KEYS - _body_snapshot: Final = {k: v for k, v in data.items() if k not in _body_snapshot_exclude} - data["proxy_server_request"]["body"] = _body_snapshot + refresh_proxy_server_request_body_snapshot(data) # Snapshot the requester-supplied metadata for downstream consumers. # Taking the deepcopy after the user_api_key_* / _pipeline_managed_guardrails diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index bf42aeeec05..34a91dc59da 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -20,7 +20,7 @@ import secrets import traceback from collections.abc import Awaitable, Callable, Mapping, Sequence from datetime import datetime, timedelta, timezone -from typing import Any, Final, Literal, Optional, Protocol, TypeVar, cast +from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Protocol, TypeVar, cast import fastapi import yaml @@ -148,6 +148,9 @@ from litellm.types.utils import ( TeamUIKeyGenerationConfig, ) +if TYPE_CHECKING: + from prisma import Prisma + _PrismaRowT = TypeVar("_PrismaRowT") _RepositoryModelT = TypeVar("_RepositoryModelT", bound=BaseModel) @@ -4337,10 +4340,19 @@ def _transform_verification_tokens_to_deleted_records( async def _save_deleted_verification_token_records( records: Sequence[Mapping[str, object]], prisma_client: PrismaClient, + tx: "Prisma | None" = None, ) -> None: - """Save deleted verification token records to the database.""" + """Save deleted verification token records to the database. + + ``tx`` runs the write on that transaction's connection instead of a fresh + one, so a caller batching this with other writes gets one all-or-nothing + commit. + """ if not records: return + if tx is not None: + await tx.litellm_deletedverificationtoken.create_many(data=records) + return await _deleted_verification_token_table(prisma_client).create_many(data=records) @@ -4349,6 +4361,7 @@ async def _persist_deleted_verification_tokens( prisma_client: PrismaClient, user_api_key_dict: UserAPIKeyAuth, litellm_changed_by: str | None = None, + tx: "Prisma | None" = None, ) -> None: """Persist deleted verification token records by transforming and saving them.""" records: Final = _transform_verification_tokens_to_deleted_records( @@ -4359,6 +4372,7 @@ async def _persist_deleted_verification_tokens( await _save_deleted_verification_token_records( records=records, prisma_client=prisma_client, + tx=tx, ) diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index a8e545a8551..01254d5c064 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -3286,15 +3286,6 @@ async def team_member_delete( _db_new_team_members: Final[list[dict]] = [m.model_dump() for m in new_team_members] - _ = await _team_db(prisma_client).update( - where={ - "team_id": data.team_id, - }, - data={"members_with_roles": json.dumps(_db_new_team_members)}, - ) - - _emit_team_members_metric(existing_team_row) - ## DELETE TEAM ID from USER ROW, IF EXISTS ## # get user row removed_user_ids: Final = frozenset(m.user_id for m in removed_team_members if m.user_id is not None) @@ -3303,52 +3294,62 @@ async def team_member_delete( ) existing_user_rows: Final[Sequence[LiteLLM_UserTable]] = await _user_db(prisma_client).find_many(where=key_val) - for existing_user in existing_user_rows: - if data.team_id in existing_user.teams: - await _user_db(prisma_client).update( - where={ - "user_id": existing_user.user_id, - }, - data={"teams": {"set": [team for team in existing_user.teams if team != data.team_id]}}, - ) - # Also clean up any existing team membership rows for this user and team user_ids_to_delete: Final = removed_user_ids.union( (data.user_id,) if data.user_id is not None else (), (user.user_id for user in existing_user_rows if user.user_id), ) - for _uid in sorted(user_ids_to_delete): - await _team_membership_db(prisma_client).delete_many(where={"team_id": data.team_id, "user_id": _uid}) - ## DELETE KEYS CREATED BY USER FOR THIS TEAM - if user_ids_to_delete: - from litellm.proxy.management_endpoints.key_management_endpoints import ( - _persist_deleted_verification_tokens, + # Fetch keys before deletion so their audit records can be persisted alongside the delete. + # An empty user_ids_to_delete still resolves cleanly: prisma's "in": [] matches no rows. + keys_to_delete: Final[list[LiteLLM_VerificationToken]] = await _tokens_db(prisma_client).find_many( + where={ + "user_id": {"in": sorted(user_ids_to_delete)}, + "team_id": data.team_id, + } + ) + + # All four cleanups run on one connection so a failure between them leaves + # no partial removal: either every write below lands, or none of them do. + async with prisma_client.tx() as tx: + await tx.litellm_teamtable.update( + where={"team_id": data.team_id}, + data={"members_with_roles": json.dumps(_db_new_team_members)}, ) - # Fetch keys before deletion to persist them - keys_to_delete: Final[list[LiteLLM_VerificationToken]] = await _tokens_db(prisma_client).find_many( - where={ - "user_id": {"in": sorted(user_ids_to_delete)}, - "team_id": data.team_id, - } - ) + for existing_user in existing_user_rows: + if data.team_id in existing_user.teams: + await tx.litellm_usertable.update( + where={"user_id": existing_user.user_id}, + data={"teams": {"set": [team for team in existing_user.teams if team != data.team_id]}}, + ) - if keys_to_delete: - await _persist_deleted_verification_tokens( - keys=keys_to_delete, - prisma_client=prisma_client, - user_api_key_dict=user_api_key_dict, - litellm_changed_by=None, + for _uid in sorted(user_ids_to_delete): + await tx.litellm_teammembership.delete_many(where={"team_id": data.team_id, "user_id": _uid}) + + if user_ids_to_delete: + if keys_to_delete: + from litellm.proxy.management_endpoints.key_management_endpoints import ( + _persist_deleted_verification_tokens, + ) + + await _persist_deleted_verification_tokens( + keys=keys_to_delete, + prisma_client=prisma_client, + user_api_key_dict=user_api_key_dict, + litellm_changed_by=None, + tx=tx, + ) + + await tx.litellm_verificationtoken.delete_many( + where={ + "user_id": {"in": sorted(user_ids_to_delete)}, + "team_id": data.team_id, + } ) - await _tokens_db(prisma_client).delete_many( - where={ - "user_id": {"in": sorted(user_ids_to_delete)}, - "team_id": data.team_id, - } - ) + _emit_team_members_metric(existing_team_row) return existing_team_row diff --git a/litellm/proxy/openai_files_endpoints/common_utils.py b/litellm/proxy/openai_files_endpoints/common_utils.py index b9af01e9aea..142aced4a38 100644 --- a/litellm/proxy/openai_files_endpoints/common_utils.py +++ b/litellm/proxy/openai_files_endpoints/common_utils.py @@ -4,12 +4,14 @@ import re from collections.abc import Mapping from dataclasses import dataclass, field from types import MappingProxyType -from typing import TYPE_CHECKING, Final, Literal, Optional, Protocol, runtime_checkable +from typing import TYPE_CHECKING, Final, Literal, Optional, Protocol, get_args, runtime_checkable +from litellm.proxy._types import ProxyException from litellm.repositories.table_repositories import ( ManagedFileRepository, ManagedObjectRepository, ) +from litellm.types.llms.openai import OpenAIFilesPurpose from litellm.types.utils import SpecialEnums if TYPE_CHECKING: @@ -22,6 +24,50 @@ if TYPE_CHECKING: from litellm.types.utils import LiteLLMBatch +MAX_FILE_LIST_LIMIT: Final = 10000 + +FILE_LIST_CONTINUATION_CHUNK_SIZE: Final = 500 + + +def validate_file_list_limit(limit: int | None) -> None: + """Reject a ``limit`` outside the range OpenAI documents for GET /v1/files.""" + if limit is None or 1 <= limit <= MAX_FILE_LIST_LIMIT: + return + bound, expected, openai_code = ( + ("below minimum", ">= 1", "integer_below_min_value") + if limit < 1 + else ("above maximum", f"<= {MAX_FILE_LIST_LIMIT}", "integer_above_max_value") + ) + raise ProxyException( + message=f"Invalid 'limit': integer {bound} value. Expected a value {expected}, but got {limit} instead.", + type="invalid_request_error", + param="limit", + code=400, + openai_code=openai_code, + ) + + +def validate_file_list_purpose(purpose: str | None) -> None: + """Reject a ``purpose`` filter no upload to this proxy could have stored. + + An unknown purpose matches no file, so filtering on it would report an + empty page for what is really a bad request. Rejecting it keeps a managed + listing consistent with the upload route, which refuses the same values + against this same set. The provider-backed listings do not: they pass + ``purpose`` upstream, so a purpose OpenAI accepts before it is added here + is rejected on the managed path while still working on those. + """ + valid_purposes: Final = get_args(OpenAIFilesPurpose) + if purpose is None or purpose in valid_purposes: + return + raise ProxyException( + message=f"Invalid purpose: {purpose}. Must be one of: {valid_purposes}", + type="invalid_request_error", + param="purpose", + code=400, + ) + + @runtime_checkable class ManagedResourceAccessChecker(Protocol): async def can_user_call_unified_file_id( @@ -1288,6 +1334,25 @@ def batch_cost_poller_is_active() -> bool: return False +def _completed_batch_safe_to_retire(response: "LiteLLMBatch") -> bool: + """Whether a "completed" batch may be retired from cost recovery. + + ``batch_processed=True`` is the sole re-pickup gate for CheckBatchCost's + cost-recovery poller, so setting it retires the batch permanently. A batch can + reach ``status="completed"`` while ``output_file_id`` is still ``None`` (the + provider response briefly lags before the output id populates). Retiring in that + window loses the spend record forever. Retire only once we can prove there is + nothing left to recover: the output file has actually arrived, or the provider + reports no successful request lines. When counts are unknown, stay eligible so + the next poller pass revisits it. (#37713) + """ + if getattr(response, "output_file_id", None) is not None: + return True + request_counts = getattr(response, "request_counts", None) + completed = getattr(request_counts, "completed", None) + return completed == 0 + + async def update_batch_in_database( batch_id: str, unified_batch_id: str | Literal[False], @@ -1369,7 +1434,7 @@ async def update_batch_in_database( } poller_owns: Final = batch_cost_poller_is_active() if poller_owns_accounting is None else poller_owns_accounting - if db_status == "complete" and not poller_owns: + if db_status == "complete" and not poller_owns and _completed_batch_safe_to_retire(response): update_data["batch_processed"] = True try: diff --git a/litellm/proxy/openai_files_endpoints/files_endpoints.py b/litellm/proxy/openai_files_endpoints/files_endpoints.py index 813ce9630a5..92bbd58ed90 100644 --- a/litellm/proxy/openai_files_endpoints/files_endpoints.py +++ b/litellm/proxy/openai_files_endpoints/files_endpoints.py @@ -65,6 +65,7 @@ from litellm.proxy.openai_files_endpoints.common_utils import ( get_credentials_for_model, handle_model_based_routing, prepare_data_with_credentials, + validate_file_list_limit, validate_managed_files_requirement, validate_managed_id_requirement, ) @@ -1410,6 +1411,8 @@ async def list_files( provider: str | None = None, target_model_names: str | None = None, purpose: str | None = None, + limit: int | None = None, + after: str | None = None, ): """ Returns information about a specific file. that can be used across - Assistants API, Batch API @@ -1434,6 +1437,8 @@ async def list_files( data: dict = {} try: + validate_file_list_limit(limit) + # Include original request and headers in the data base_llm_response_processor: Final = ProxyBaseLLMRequestProcessing(data=data) ( @@ -1500,24 +1505,30 @@ async def list_files( or get_custom_llm_provider_from_request_headers(request=request) or get_custom_llm_provider_from_request_query(request=request) or await get_custom_llm_provider_from_request_body(request=request) - or "openai" ) + managed_files_obj: Final = proxy_logging_obj.get_proxy_hook("managed_files") + if custom_llm_provider is None and isinstance(managed_files_obj, BaseFileEndpoints): + response = await managed_files_obj.afile_list( + purpose=purpose, + litellm_parent_otel_span=user_api_key_dict.parent_otel_span, + user_api_key_dict=user_api_key_dict, + limit=limit, + after=after, + ) + else: + resolved_custom_llm_provider: Final = custom_llm_provider or "openai" + apply_team_provider_credentials( + data=data, + llm_router=llm_router, + user_api_key_dict=user_api_key_dict, + custom_llm_provider=resolved_custom_llm_provider, + ) - # No model/target_model_names pinned: resolve upstream credentials from - # the team's deployment for this provider so the call is authenticated - # against the team's own account (e.g. the team's openai deployment). - apply_team_provider_credentials( - data=data, - llm_router=llm_router, - user_api_key_dict=user_api_key_dict, - custom_llm_provider=custom_llm_provider, - ) - - response = await litellm.afile_list( - custom_llm_provider=custom_llm_provider, - purpose=purpose, - **data, - ) + response = await litellm.afile_list( + custom_llm_provider=resolved_custom_llm_provider, + purpose=purpose, + **data, + ) if response is None: raise HTTPException( @@ -1561,6 +1572,8 @@ async def list_files( ) verbose_proxy_logger.error("litellm.proxy.proxy_server.list_files(): Exception occured - %s", e) verbose_proxy_logger.debug(traceback.format_exc()) + if isinstance(e, ProxyException): + raise if isinstance(e, HTTPException): raise ProxyException( message=getattr(e, "message", str(e.detail)), diff --git a/litellm/proxy/pass_through_endpoints/managed_id_rewriter.py b/litellm/proxy/pass_through_endpoints/managed_id_rewriter.py index e8b5fab626f..e08d277788f 100644 --- a/litellm/proxy/pass_through_endpoints/managed_id_rewriter.py +++ b/litellm/proxy/pass_through_endpoints/managed_id_rewriter.py @@ -50,7 +50,7 @@ from litellm.repositories.table_repositories import ( ManagedFileRepository, ManagedObjectRepository, ) -from litellm.types.llms.openai import OpenAIFileObject +from litellm.types.llms.openai import BATCH_GUARDRAIL_RESPONSE_FIELD, OpenAIFileObject from litellm.types.passthrough_endpoints.managed_id_rewriter import ( ManagedFileIdReader, ManagedFileIdWriter, @@ -980,6 +980,7 @@ def _serialize_file_list_item(row: ManagedFileRow) -> dict[str, JsonValue]: file_object: Final = _parse_file_object(row.file_object) if isinstance(file_object, dict): item.update(file_object) + item.pop(BATCH_GUARDRAIL_RESPONSE_FIELD, None) item["id"] = row.unified_file_id # managed ID always wins over stored raw id return item diff --git a/litellm/proxy/prisma_migration.py b/litellm/proxy/prisma_migration.py index 373c3811949..1b95d24c011 100644 --- a/litellm/proxy/prisma_migration.py +++ b/litellm/proxy/prisma_migration.py @@ -1,8 +1,9 @@ """Standalone entrypoint for applying database migrations and generating the Prisma client. -The entrypoint enforces migration failures by default. Set -ENFORCE_PRISMA_MIGRATION_CHECK=false to preserve log-only behavior for migration and -Prisma generate failures. +Migration failures fail the entrypoint by default; set ENFORCE_PRISMA_MIGRATION_CHECK=false +for log-only behavior. A failed 'prisma generate' is always log-only: every shipped image +bakes the client at build time, and refreshing it writes into site-packages, which an +arbitrary non-root uid or a read-only root filesystem cannot do. """ import os @@ -30,13 +31,13 @@ def main() -> int: verbose_proxy_logger.info("Running 'prisma generate'...") result: Final = subprocess.run(("prisma", "generate"), capture_output=True, text=True) verbose_proxy_logger.info("'prisma generate' stdout: %s", result.stdout) - exit_code: Final = result.returncode - if exit_code != 0: - verbose_proxy_logger.info("'prisma generate' failed with exit code %s.", exit_code) - verbose_proxy_logger.error("'prisma generate' stderr: %s", result.stderr) - if enforce_prisma_migration_check: - return exit_code + if result.returncode != 0: + verbose_proxy_logger.warning( + "'prisma generate' exited %s; continuing with the client baked at image build time. stderr: %s", + result.returncode, + result.stderr, + ) return 0 diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 9f16584340d..7dced4e26b6 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -254,6 +254,9 @@ from litellm.exceptions import RejectedRequestError from litellm.integrations.custom_guardrail import ModifyResponseException from litellm.integrations.custom_logger import CustomLogger from litellm.integrations.SlackAlerting.slack_alerting import SlackAlerting +from litellm.litellm_core_utils.agentic_loop_settings import ( + validated_max_agentic_loops, +) from litellm.litellm_core_utils.asyncify import asyncify from litellm.litellm_core_utils.core_helpers import ( _get_parent_otel_span_from_kwargs, @@ -4081,6 +4084,27 @@ def resolve_complexity_router_plugins( complexity_router_config["classifier_plugin"] = resolved_classifier # rebind-ok: out-param, resolved in place +def validate_deployment_max_agentic_loops(model: Mapping[str, Any]) -> None: + """ + Reject a per-deployment `max_agentic_loops` the agentic loop cannot honor. + + Checked here rather than on `LiteLLM_Params` because the proxy builds its + router with `ignore_invalid_deployments=True`, so a validator down there + turns a bad value into a silently missing model instead of a refusal to + start. Left unchecked entirely, a `0` used to read as the default ceiling + of 3 and a non-integer failed every request to that model instead. + """ + litellm_params: Final = model.get("litellm_params") or {} + if "max_agentic_loops" not in litellm_params: + return + + model_name: Final = model.get("model_name", "") + validated_max_agentic_loops( + litellm_params["max_agentic_loops"], + field=f"litellm_params.max_agentic_loops on model {model_name!r}", + ) + + def pin_complexity_router_model_id(model: dict) -> None: # mutable-ok: out-param, model_info is stamped in place """ Stamps `model_info.id` from the raw litellm_params before plugin resolution swaps @@ -5416,6 +5440,7 @@ class ProxyConfig: for k, v in model["litellm_params"].items(): if isinstance(v, str) and v.startswith("os.environ/"): model["litellm_params"][k] = get_secret(v) + validate_deployment_max_agentic_loops(model) pin_complexity_router_model_id(model) complexity_router_config = model["litellm_params"].get("complexity_router_config") if isinstance(complexity_router_config, dict): diff --git a/litellm/responses/litellm_completion_transformation/session_handler.py b/litellm/responses/litellm_completion_transformation/session_handler.py index 59ff492a79f..2008006e6cf 100644 --- a/litellm/responses/litellm_completion_transformation/session_handler.py +++ b/litellm/responses/litellm_completion_transformation/session_handler.py @@ -114,6 +114,7 @@ class ResponsesSessionHandler: chat_completion_messages = LiteLLMCompletionResponsesConfig.transform_responses_api_input_to_messages( input=response_input_param, responses_api_request=proxy_server_request_dict or {}, + replay_reasoning=True, ) chat_completion_message_history.extend(chat_completion_messages) @@ -126,6 +127,7 @@ class ResponsesSessionHandler: chat_completion_messages = LiteLLMCompletionResponsesConfig.transform_responses_api_input_to_messages( input=_messages, responses_api_request=proxy_server_request_dict or {}, + replay_reasoning=True, ) chat_completion_message_history.extend(chat_completion_messages) diff --git a/litellm/responses/litellm_completion_transformation/streaming_iterator.py b/litellm/responses/litellm_completion_transformation/streaming_iterator.py index a8092edc625..92bbca9ee5b 100644 --- a/litellm/responses/litellm_completion_transformation/streaming_iterator.py +++ b/litellm/responses/litellm_completion_transformation/streaming_iterator.py @@ -48,6 +48,22 @@ from litellm.types.utils import ( ) +def _output_items_with_id(items: tuple[Any, ...], item_type: str, item_id: str | None) -> tuple[Any, ...]: + if item_id is None: + return items + + target_index: Final = next( + (index for index, item in enumerate(items) if getattr(item, "type", None) == item_type), + None, + ) + if target_index is None: + return items + + return tuple( + item.model_copy(update={"id": item_id}) if index == target_index else item for index, item in enumerate(items) + ) + + class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): """ Async iterator for processing streaming responses from the Responses API. @@ -912,9 +928,10 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): reasoning_content = "".join(self._accumulated_reasoning_content_parts) # Ensure we have a valid reasoning_item_id - reasoning_item_id = ( + self._cached_reasoning_item_id = ( self._reasoning_item_id or self._cached_reasoning_item_id or f"rs_{uuid.uuid4()}" ) + reasoning_item_id = self._cached_reasoning_item_id # Create text.done event first with its own sequence number self._sequence_number += 1 @@ -1034,9 +1051,9 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): and the ReasoningSummaryTextDeltaEvent, which is used by the responses API to emit reasoning content. It also handles emitting annotation.added events when annotations are detected in the chunk. """ - if self._cached_item_id is None and chunk.id: - self._cached_item_id = chunk.id - item_id: Final = self._cached_item_id or chunk.id + if self._cached_item_id is None: + self._cached_item_id = f"msg_{uuid.uuid4()}" + item_id: Final = self._cached_item_id # Check if this chunk has annotations first (before processing text/reasoning) # This ensures we detect and queue annotation events from the annotation chunk @@ -1071,9 +1088,12 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): ): reasoning_content: Final = chunk.choices[0].delta.reasoning_content + if self._cached_reasoning_item_id is None: + self._cached_reasoning_item_id = f"rs_{uuid.uuid4()}" + return ReasoningSummaryTextDeltaEvent( type=ResponsesAPIStreamEvents.REASONING_SUMMARY_TEXT_DELTA, - item_id=f"rs_{hash(str(reasoning_content))}", + item_id=self._cached_reasoning_item_id, output_index=0, delta=reasoning_content, ) @@ -1124,6 +1144,19 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): chat_completion_delta: Final[ChatCompletionDelta] = choice.delta return chat_completion_delta.content or "" + def _output_with_streamed_item_ids(self, responses_api_response: ResponsesAPIResponse) -> tuple[Any, ...]: + """ + Reuse the item IDs already emitted by the incremental streaming events in the + ``response.completed`` snapshot, so a streaming client that replays the snapshot + sends back the same IDs it observed mid-stream. + """ + message_aligned: Final = _output_items_with_id( + tuple(responses_api_response.output or ()), + "message", + self._cached_item_id, + ) + return _output_items_with_id(message_aligned, "reasoning", self._cached_reasoning_item_id) + def _emit_response_completed_event(self, litellm_model_response: ModelResponse) -> ResponseCompletedEvent | None: if litellm_model_response: # Add cost to usage object if include_cost_in_streaming_usage is True @@ -1149,6 +1182,8 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): if self._cached_response_id: responses_api_response.id = self._cached_response_id + responses_api_response.output = list(self._output_with_streamed_item_ids(responses_api_response)) + # Encode the response ID to match non-streaming behavior encoded_response: Final = self._with_encoded_response_id(responses_api_response) diff --git a/litellm/responses/litellm_completion_transformation/transformation.py b/litellm/responses/litellm_completion_transformation/transformation.py index 64084bfb063..b8d7b726a28 100644 --- a/litellm/responses/litellm_completion_transformation/transformation.py +++ b/litellm/responses/litellm_completion_transformation/transformation.py @@ -4,6 +4,7 @@ Handles transforming from Responses API -> LiteLLM completion (Chat Completion import json import re +import uuid from collections.abc import Iterator, Mapping, Sequence from types import MappingProxyType from typing import ( @@ -42,8 +43,10 @@ from litellm.types.llms.openai import ( AllMessageValues, ChatCompletionImageObject, ChatCompletionImageUrlObject, + ChatCompletionRedactedThinkingBlock, ChatCompletionResponseMessage, ChatCompletionSystemMessage, + ChatCompletionThinkingBlock, ChatCompletionToolCallChunk, ChatCompletionToolCallFunctionChunk, ChatCompletionToolMessage, @@ -294,6 +297,7 @@ class LiteLLMCompletionResponsesConfig: "messages": LiteLLMCompletionResponsesConfig.transform_responses_api_input_to_messages( input=input, responses_api_request=responses_api_request, + replay_reasoning=True, ), "model": model, "tool_choice": LiteLLMCompletionResponsesConfig._transform_tool_choice( @@ -338,6 +342,7 @@ class LiteLLMCompletionResponsesConfig: def transform_responses_api_input_to_messages( input: str | ResponseInputParam, responses_api_request: ResponsesAPIOptionalRequestParams | dict, + replay_reasoning: bool = False, ) -> list[ AllMessageValues | GenericChatCompletionMessage @@ -347,6 +352,16 @@ class LiteLLMCompletionResponsesConfig: ]: """ Transform a Responses API input into a list of messages + + ``replay_reasoning`` belongs to callers whose messages are about to be + sent to a model: prior-turn ``reasoning`` items are then rebuilt as + assistant ``reasoning_content`` and signed ``thinking_blocks`` so the + provider gets its own chain-of-thought back instead of reading it as + visible text. + + Callers that only inspect the messages (token counting, rate limiting, + guardrail scanning) leave it off, because they need every piece of text + in the request to stay readable as message ``content``. """ messages: list[ AllMessageValues @@ -365,6 +380,7 @@ class LiteLLMCompletionResponsesConfig: messages.extend( LiteLLMCompletionResponsesConfig._transform_response_input_param_to_chat_completion_message( input=input, + replay_reasoning=replay_reasoning, ) ) @@ -441,11 +457,15 @@ class LiteLLMCompletionResponsesConfig: @staticmethod def _transform_response_input_param_to_chat_completion_message( input: str | ResponseInputParam, + replay_reasoning: bool = False, ) -> list[ AllMessageValues | GenericChatCompletionMessage | ChatCompletionMessageToolCall | ChatCompletionResponseMessage ]: """ Transform a ResponseInputParam into a Chat Completion message + + See ``transform_responses_api_input_to_messages`` for what + ``replay_reasoning`` means. """ messages: list[ AllMessageValues @@ -461,7 +481,8 @@ class LiteLLMCompletionResponsesConfig: for _input in input: chat_completion_messages = ( LiteLLMCompletionResponsesConfig._transform_responses_api_input_item_to_chat_completion_message( - input_item=_input + input_item=_input, + replay_reasoning=replay_reasoning, ) ) @@ -557,7 +578,160 @@ class LiteLLMCompletionResponsesConfig: continue messages.extend(chat_completion_messages) - return messages + if not replay_reasoning: + return messages + return LiteLLMCompletionResponsesConfig._merge_reasoning_only_assistant_messages(messages) + + @staticmethod + def _reasoning_only_assistant_message( + reasoning_text: str | None, + thinking_blocks: Sequence[ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock] | None, + ) -> ChatCompletionResponseMessage: + """ + Build the assistant message that carries a prior turn's reasoning and + nothing else, so a reasoning item never reaches the provider as visible + assistant ``content``. + """ + message: Final = ChatCompletionResponseMessage(role="assistant", content=None) + if reasoning_text: + message["reasoning_content"] = reasoning_text + if thinking_blocks: + message["thinking_blocks"] = list( # mutable-ok: thinking_blocks is a list on the message contract + thinking_blocks + ) + return message + + @staticmethod + def _merge_reasoning_only_assistant_messages( + messages: list[ # mutable-ok: input sequence + AllMessageValues + | GenericChatCompletionMessage + | ChatCompletionMessageToolCall + | ChatCompletionResponseMessage + ], + ) -> list[ # mutable-ok: fresh merged list + AllMessageValues | GenericChatCompletionMessage | ChatCompletionMessageToolCall | ChatCompletionResponseMessage + ]: + """ + Responses API emits prior-turn reasoning as its own ``reasoning`` input + item, which becomes a standalone assistant message with + ``content=None`` + ``reasoning_content``. Chat-completions providers + (e.g. DeepSeek V4, Kimi K2.6) expect the chain-of-thought on the + assistant message that carries the answer or tool calls. This pass + merges standalone reasoning-only assistant messages into the + immediately following assistant message. + + Signed ``thinking_blocks`` decoded from ``encrypted_content`` travel the + same way and are placed ahead of any thinking blocks the target message + already carries, because Anthropic and Bedrock verify signatures against + the original block order. + + If the reasoning item is not followed by an assistant message (e.g. a + stateless chain replays ``reasoning`` + ``user``), the standalone + reasoning message is preserved so the reasoning is still passed back. + """ + + def _role(msg: object) -> str: + if isinstance(msg, dict): + return str(msg.get("role") or "") + return str(getattr(msg, "role", "") or "") + + def _reasoning_text(msg: object) -> str | None: + if isinstance(msg, dict): + value = msg.get("reasoning_content") # rebind-ok: branch lookup + else: + value = getattr(msg, "reasoning_content", None) # rebind-ok: branch lookup + return value if isinstance(value, str) and value else None + + def _thinking_blocks( + msg: object, + ) -> tuple[ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock, ...] | None: + if isinstance(msg, dict): + value = msg.get("thinking_blocks") # rebind-ok: branch lookup + else: + value = getattr(msg, "thinking_blocks", None) # rebind-ok: branch lookup + return tuple(value) if isinstance(value, list) and value else None + + def _content(msg: object) -> object | None: + if isinstance(msg, dict): + return msg.get("content") + return getattr(msg, "content", None) + + def _tool_calls(msg: object) -> object | None: + if isinstance(msg, dict): + return msg.get("tool_calls") + return getattr(msg, "tool_calls", None) + + def _apply_pending( + msg: object, + pending_items: Sequence[ + tuple[ + str | None, + tuple[ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock, ...] | None, + ] + ], + ) -> None: + pending_texts: Final = tuple(text for text, _ in pending_items if text) + pending_blocks: Final = tuple(block for _, blocks in pending_items for block in blocks or ()) + if pending_texts: + existing_text: Final = _reasoning_text(msg) + combined: Final = "\n".join(pending_texts + ((existing_text,) if existing_text else ())) + if isinstance(msg, dict): + cast(dict[str, Any], msg)["reasoning_content"] = combined # cast-ok: mutable reasoning carrier + else: + setattr(msg, "reasoning_content", combined) # noqa: B010 # attribute name is fixed, not dynamic + if pending_blocks: + replayed: Final = list( # mutable-ok: thinking_blocks is a list on the message contract + pending_blocks + (_thinking_blocks(msg) or ()) + ) + if isinstance(msg, dict): + cast(dict[str, Any], msg)["thinking_blocks"] = replayed # cast-ok: mutable reasoning carrier + else: + setattr(msg, "thinking_blocks", replayed) # noqa: B010 # attribute name is fixed, not dynamic + + _standalone: Final = LiteLLMCompletionResponsesConfig._reasoning_only_assistant_message + + merged: list[ # mutable-ok: accumulator # rebind-ok: accumulator + AllMessageValues + | GenericChatCompletionMessage + | ChatCompletionMessageToolCall + | ChatCompletionResponseMessage + ] = [] # mutable-ok: accumulator + pending: list[ # mutable-ok: accumulator # rebind-ok: accumulator + tuple[ + str | None, + tuple[ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock, ...] | None, + ] + ] = [] # mutable-ok: accumulator + + for msg in messages: + if ( + _role(msg) == "assistant" + and _content(msg) is None + and not _tool_calls(msg) + and (_reasoning_text(msg) is not None or _thinking_blocks(msg) is not None) + ): + pending.append((_reasoning_text(msg), _thinking_blocks(msg))) + continue + + if pending and _role(msg) == "assistant": + _apply_pending(msg, pending) + pending = [] # mutable-ok: reset accumulator + elif pending: + # Not followed by an assistant message — keep the reasoning + # standalone instead of dropping it. + merged.extend( # mutable-ok: append reasoning messages + [_standalone(text, blocks) for text, blocks in pending] # mutable-ok: append reasoning messages + ) + pending = [] # mutable-ok: reset accumulator + + merged.append(msg) + + merged.extend( # mutable-ok: append trailing reasoning + [_standalone(text, blocks) for text, blocks in pending] # mutable-ok: append trailing reasoning + ) + + return merged @staticmethod def _merged_trailing_assistant_message( @@ -998,6 +1172,7 @@ class LiteLLMCompletionResponsesConfig: @staticmethod def _transform_responses_api_input_item_to_chat_completion_message( input_item: Any, + replay_reasoning: bool = False, ) -> list[AllMessageValues | GenericChatCompletionMessage | ChatCompletionResponseMessage]: """ Transform a Responses API input item into a Chat Completion message @@ -1026,6 +1201,49 @@ class LiteLLMCompletionResponsesConfig: return LiteLLMCompletionResponsesConfig._transform_responses_api_function_call_to_chat_completion_message( function_call=input_item ) + elif input_item.get("type") == "reasoning": + # A ResponseReasoningItemParam carries the prior-turn chain-of-thought. + # Chat-completions providers (DeepSeek V4, Kimi K2.6, ...) expect this + # to be replayed as `reasoning_content` on an assistant message, not as + # visible `content` (prompt pollution) and not dropped (DeepSeek V4 + # rejects multi-turn requests with a missing `reasoning_content`). + # Callers that only inspect the request keep reading the text as + # message `content`, summary-only items included: whatever the + # provider-bound branch below replays must stay scannable. + if not replay_reasoning: + # `content` wins only when it is what the provider-bound branch + # would replay; an empty or block-only `content` falls back to + # the summary text, which is what that branch replays instead. + inspectable: Final[object] = ( + input_item.get("content") + if LiteLLMCompletionResponsesConfig._reasoning_text_from_content(input_item) is not None + else LiteLLMCompletionResponsesConfig._reasoning_text_from_summary(input_item) + or input_item.get("content") + ) + if inspectable is None: + return [] # mutable-ok: empty drop result + return [ # mutable-ok: single message result + GenericChatCompletionMessage( + role=input_item.get("role") or "user", + content=LiteLLMCompletionResponsesConfig._transform_responses_api_content_to_chat_completion_content( + inspectable + ), + ) + ] + reasoning_text = LiteLLMCompletionResponsesConfig._extract_reasoning_text_from_input_item( # rebind-ok: extraction result + input_item + ) + thinking_blocks = LiteLLMCompletionResponsesConfig._decode_thinking_blocks_from_input_item( # rebind-ok: extraction result + input_item + ) + if not reasoning_text and not thinking_blocks: + return [] # mutable-ok: empty drop result + return [ # mutable-ok: single message result + LiteLLMCompletionResponsesConfig._reasoning_only_assistant_message( + reasoning_text=reasoning_text, + thinking_blocks=thinking_blocks, + ) + ] else: content: Final[object] = input_item.get("content") # Handle None content: Responses API allows None content, but GenericChatCompletionMessage requires content @@ -1041,6 +1259,116 @@ class LiteLLMCompletionResponsesConfig: ) ] + @staticmethod + def _reasoning_text_from_content(input_item: Mapping[str, object]) -> str | None: + """ + Plaintext a ResponseReasoningItemParam carries in ``content``. + + Handles content as a string and content as a list of blocks + (output_text / summary_text / text). Returns None when the item has + no content, or only opaque blocks (e.g. encrypted_content). + """ + content: Final[object] = input_item.get("content") + if isinstance(content, str) and content.strip(): + return content + if isinstance(content, list): + text_parts: Final[list[str]] = [] # mutable-ok: text accumulator + for block in content: + if not isinstance(block, Mapping): + continue + block_type = block.get("type") + if block_type in ("encrypted_content", "redacted_thinking"): + continue + text = block.get("text") + if isinstance(text, str) and text.strip(): + text_parts.append(text.strip()) + if text_parts: + return "\n".join(text_parts) + return None + + @staticmethod + def _reasoning_text_from_summary(input_item: Mapping[str, object]) -> str | None: + """ + Plaintext a ResponseReasoningItemParam carries in ``summary``. + + Guardrail traversal in litellm/proxy/guardrails/_content_utils.py + inspects and rewrites these summary blocks before they are forwarded. + """ + summary: Final[object] = input_item.get("summary") + if not isinstance(summary, list): + return None + text_parts: Final[list[str]] = [] # mutable-ok: text accumulator + for block in summary: + if not isinstance(block, Mapping): + continue + text = block.get("text") + if isinstance(text, str) and text.strip(): + text_parts.append(text.strip()) + return "\n".join(text_parts) if text_parts else None + + @staticmethod + def _extract_reasoning_text_from_input_item(input_item: Mapping[str, object]) -> str | None: + """ + Extract plaintext reasoning from a ResponseReasoningItemParam. + + ``content`` wins, ``summary`` is the fallback. Returns None when only + opaque forms (e.g. encrypted_content) are present. + """ + return LiteLLMCompletionResponsesConfig._reasoning_text_from_content( + input_item + ) or LiteLLMCompletionResponsesConfig._reasoning_text_from_summary(input_item) + + @staticmethod + def _decode_thinking_blocks_from_input_item( + input_item: Mapping[str, object], + ) -> tuple[ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock, ...] | None: + """ + Decode ``encrypted_content`` written by ``_encode_thinking_blocks`` back + into the signed thinking blocks it serialized. + + LiteLLM writes this field itself for providers whose reasoning is signed + (Anthropic, Bedrock converse): it is a JSON array of the provider's own + ``thinking`` / ``redacted_thinking`` blocks, not an opaque OpenAI blob. + Replaying the blocks on the assistant message is what lets the provider + verify the signature and keep the prior chain-of-thought. + + Returns None for anything this deployment did not write, so a genuinely + opaque blob is still skipped rather than forwarded as garbage. + """ + encrypted_content: Final[object] = input_item.get("encrypted_content") + if not isinstance(encrypted_content, str) or not encrypted_content.strip(): + return None + try: + decoded: Final[object] = json.loads(encrypted_content) + except ValueError: + return None + if not isinstance(decoded, list): + return None + + blocks: Final = tuple( + cast( # cast-ok: shape validated by _is_replayable_thinking_block + ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock, + block, + ) + for block in decoded + if isinstance(block, Mapping) and LiteLLMCompletionResponsesConfig._is_replayable_thinking_block(block) + ) + return blocks or None + + @staticmethod + def _is_replayable_thinking_block(block: Mapping[str, object]) -> bool: + """ + A thinking block is only worth replaying when the provider can verify + it: a ``thinking`` block needs its signature, a ``redacted_thinking`` + block needs its opaque data. + """ + block_type: Final[object] = block.get("type") + if block_type == "thinking": + return bool(block.get("signature")) + if block_type == "redacted_thinking": + return bool(block.get("data")) + return False + @staticmethod def _is_input_item_tool_call_output(input_item: Mapping[str, object]) -> bool: """ @@ -2017,7 +2345,7 @@ class LiteLLMCompletionResponsesConfig: return [ GenericResponseOutputItem( type="reasoning", - id=f"rs_{hash(reasoning_content or encrypted_content)}", + id=f"rs_{uuid.uuid4()}", status=LiteLLMCompletionResponsesConfig._map_chat_completion_finish_reason_to_responses_status( choice.finish_reason ), @@ -2038,7 +2366,6 @@ class LiteLLMCompletionResponsesConfig: @staticmethod def _extract_image_generation_output_items( - chat_completion_response: ModelResponse, choice: Choices, ) -> list[OutputImageGenerationCall]: """ @@ -2054,7 +2381,7 @@ class LiteLLMCompletionResponsesConfig: To Responses API format: { 'type': 'image_generation_call', - 'id': 'img_...', + 'id': 'ig_...', 'status': 'completed', 'result': 'iVBORw0...' # Pure base64 without data: prefix } @@ -2065,7 +2392,7 @@ class LiteLLMCompletionResponsesConfig: if not images: return image_generation_items - for idx, image_item in enumerate(_DICT_ITEMS_LIST_ADAPTER.validate_python(images)): + for image_item in _DICT_ITEMS_LIST_ADAPTER.validate_python(images): # Extract base64 from data URL image_url = _TEXT_ADAPTER.validate_python( _ANY_KEY_DICT_ADAPTER.validate_python(image_item.get("image_url", {})).get("url", "") @@ -2076,7 +2403,7 @@ class LiteLLMCompletionResponsesConfig: image_generation_items.append( OutputImageGenerationCall( type="image_generation_call", - id=f"{chat_completion_response.id}_img_{idx}", + id=f"ig_{uuid.uuid4()}", status=LiteLLMCompletionResponsesConfig._map_finish_reason_to_image_generation_status( choice.finish_reason ), @@ -2141,7 +2468,6 @@ class LiteLLMCompletionResponsesConfig: if hasattr(choice.message, "images") and choice.message.images: # Extract image generation output image_generation_items = LiteLLMCompletionResponsesConfig._extract_image_generation_output_items( - chat_completion_response=chat_completion_response, choice=choice, ) message_output_items.extend(image_generation_items) @@ -2150,7 +2476,7 @@ class LiteLLMCompletionResponsesConfig: message_output_items.append( GenericResponseOutputItem( type="message", - id=chat_completion_response.id, + id=f"msg_{uuid.uuid4()}", status=LiteLLMCompletionResponsesConfig._map_chat_completion_finish_reason_to_responses_status( choice.finish_reason ), diff --git a/litellm/types/integrations/custom_logger.py b/litellm/types/integrations/custom_logger.py index 89b85bc5114..2cca16351af 100644 --- a/litellm/types/integrations/custom_logger.py +++ b/litellm/types/integrations/custom_logger.py @@ -23,6 +23,21 @@ def is_interception_internal_key( return any(key.startswith(prefix) for prefix in prefixes) +class AgenticLoopSafetyError(ValueError): + """ + Raised when an agentic-loop safety rail refuses a rerun. + + Covers both rails: the bounded-loop cap (``max_agentic_loops``) and the + repeated tool-call fingerprint cycle break. Subclasses ``ValueError`` so + callers that already catch the broader type keep working. + + Only the anthropic messages loop raises this today. The chat completions + loop in ``litellm_core_utils/chat_completion_agentic_loop.py`` still raises + a plain ``ValueError`` from its own copy of the same rails, so catching + this type alone will not cover that surface until it is moved over. + """ + + class StandardCustomLoggerInitParams(BaseModel): """ Params for initializing a CustomLogger. diff --git a/litellm/types/integrations/websearch_interception.py b/litellm/types/integrations/websearch_interception.py index 90713b270be..7926b9eee0a 100644 --- a/litellm/types/integrations/websearch_interception.py +++ b/litellm/types/integrations/websearch_interception.py @@ -5,6 +5,7 @@ Type definitions for WebSearch Interception integration. from typing import Literal, TypedDict from pydantic import BaseModel +from typing_extensions import ReadOnly class AnthropicSearchQuery(BaseModel): @@ -35,6 +36,7 @@ class WebSearchInterceptionConfig(TypedDict, total=False): websearch_interception_params: enabled_providers: ["bedrock"] search_tool_name: "my-perplexity-search" + max_agentic_loops: 5 """ enabled_providers: list[str] @@ -42,3 +44,6 @@ class WebSearchInterceptionConfig(TypedDict, total=False): search_tool_name: str | None """Name of search tool configured in router's search_tools. If None, uses first available.""" + + max_agentic_loops: ReadOnly[int | None] + """How many follow-up model calls one intercepted request may chain. If None, LiteLLM's default of 3 applies.""" diff --git a/litellm/types/llms/openai.py b/litellm/types/llms/openai.py index 1588c650177..e7a3f825455 100644 --- a/litellm/types/llms/openai.py +++ b/litellm/types/llms/openai.py @@ -65,9 +65,12 @@ from pydantic import ( BaseModel, ConfigDict, Discriminator, + Field, PrivateAttr, + SerializerFunctionWrapHandler, field_serializer, field_validator, + model_serializer, ) from typing_extensions import ( NotRequired, @@ -275,6 +278,7 @@ OpenAIFilesPurpose = Literal[ "fine-tune-results", "vision", "user_data", + "evals", "messages", ] @@ -313,6 +317,9 @@ class BatchGuardrailReport(BaseModel): """Every record that was redacted or dropped, in file order.""" +BATCH_GUARDRAIL_RESPONSE_FIELD: Final = "litellm_batch_guardrail" + + class OpenAIFileObject(BaseModel): id: str """The file identifier, which can be referenced in the API endpoints.""" @@ -361,6 +368,17 @@ class OpenAIFileObject(BaseModel): _hidden_params: dict = {"response_cost": 0.0} # no cost for writing a file + @model_serializer(mode="wrap") + def _omit_absent_batch_guardrail( # noqa: ANN202 # annotating it replaces the model's serialization schema + self, handler: SerializerFunctionWrapHandler + ): + serialized: Final[Mapping[str, object]] = handler(self) + if self.litellm_batch_guardrail is not None: + return serialized + return { # mutable-ok: pydantic's json serializer rejects a mapping that is not a dict + key: value for key, value in serialized.items() if key != BATCH_GUARDRAIL_RESPONSE_FIELD + } + def __contains__(self, key) -> bool: # Define custom behavior for the 'in' operator return hasattr(self, key) @@ -381,6 +399,21 @@ class OpenAIFileObject(BaseModel): return self.dict() +class FileListPage(BaseModel): + """A page of files, as `GET /v1/files` returns it. + + Post-call hooks and logging callbacks are handed the listing response, and + the provider SDKs hand them a page object rather than a mapping, so this + exposes the same ``.data`` attribute while serializing to an identical body. + """ + + object: Literal["list"] = "list" + data: list[OpenAIFileObject] = Field(default_factory=list) + first_id: str | None = None + last_id: str | None = None + has_more: bool = False + + CREATE_FILE_REQUESTS_PURPOSE = Literal["assistants", "batch", "fine-tune", "messages"] diff --git a/litellm/types/utils.py b/litellm/types/utils.py index ac2ab1c8363..94526de0757 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -2846,7 +2846,7 @@ class StandardLoggingRoutingDecision(TypedDict, total=False): classifier_cost: float escalated: bool tier_boundaries: StandardLoggingRoutingDecisionTierBoundaries - reasoning_override_min_score: ReadOnly[float] + reasoning_override_min_score: float # writable-ok: Pydantic warns on ReadOnly TypedDict fields conversation_continuing: bool savings_baseline_model: str savings_baseline_deployment_id: str diff --git a/migrations/Dockerfile b/migrations/Dockerfile index 52795d426ec..6335e6f6bd8 100644 --- a/migrations/Dockerfile +++ b/migrations/Dockerfile @@ -1,5 +1,5 @@ -ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:42df77a9974d6ec8b17a5ee8bc23b532600a44d705acef2409e0933c1251b45f -ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:42df77a9974d6ec8b17a5ee8bc23b532600a44d705acef2409e0933c1251b45f +ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:a31344ab2cb8618db84f535eec56f76f6178b142cb92cb2e48676cc2dcebea72 +ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:a31344ab2cb8618db84f535eec56f76f6178b142cb92cb2e48676cc2dcebea72 ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:240fb85ab0f263ef12f492d8476aa3a2e4e1e333f7d67fbdd923d00a506a516a FROM $UV_IMAGE AS uvbin diff --git a/tests/e2e/e2e_config.py b/tests/e2e/e2e_config.py index 0266c75e1a7..21a7a8c478a 100644 --- a/tests/e2e/e2e_config.py +++ b/tests/e2e/e2e_config.py @@ -78,6 +78,7 @@ DD_SEARCH_INTERVAL = float(os.environ.get("E2E_DD_SEARCH_INTERVAL", "10")) POLL_TIMEOUT = float(os.environ.get("E2E_POLL_TIMEOUT", "120")) POLL_INTERVAL = float(os.environ.get("E2E_POLL_INTERVAL", "5")) REQUEST_TIMEOUT = float(os.environ.get("E2E_REQUEST_TIMEOUT", "60")) +SLOW_PROVIDER_TIMEOUT_SECONDS = float(os.environ.get("E2E_SLOW_PROVIDER_TIMEOUT", "180")) # How long a control-plane write (/model/new, /guardrails, /v1/agents) may take to # reach EVERY replica. Distinct from POLL_TIMEOUT, which is sized for spend-row diff --git a/tests/e2e/llm_translation/endpoints_client.py b/tests/e2e/llm_translation/endpoints_client.py index fa33737467e..4d2c73e7078 100644 --- a/tests/e2e/llm_translation/endpoints_client.py +++ b/tests/e2e/llm_translation/endpoints_client.py @@ -12,6 +12,7 @@ from __future__ import annotations from dataclasses import dataclass from typing import Literal +from e2e_config import SLOW_PROVIDER_TIMEOUT_SECONDS from e2e_http import BinaryStream, Result, StreamingResponse from models import CacheControl, ChatMessage, LiteLLMParamsBody, RichMessage, TextBlock from proxy_client import ProxyClient @@ -74,12 +75,14 @@ class ResponsesRequest(BaseModel): stream: bool = False tools: list[ResponsesFunctionTool] | None = None guardrails: list[str] | None = None + cache: dict[str, bool] | None = {"no-cache": True} class MessagesRequest(BaseModel): model: str max_tokens: int messages: list[ChatMessage] + cache: dict[str, bool] | None = {"no-cache": True} class RichMessagesRequest(BaseModel): @@ -87,18 +90,20 @@ class RichMessagesRequest(BaseModel): max_tokens: int = 64 system: list[TextBlock] messages: list[RichMessage] - cache: dict[str, bool] = {"no-cache": True} + cache: dict[str, bool] | None = {"no-cache": True} class CompletionsRequest(BaseModel): model: str prompt: str max_tokens: int = 32 + cache: dict[str, bool] | None = {"no-cache": True} class EmbeddingsRequest(BaseModel): model: str input: str + cache: dict[str, bool] | None = {"no-cache": True} class RerankRequest(BaseModel): @@ -106,6 +111,7 @@ class RerankRequest(BaseModel): query: str documents: list[str] top_n: int + cache: dict[str, bool] | None = {"no-cache": True} class SpeechRequest(BaseModel): @@ -446,6 +452,7 @@ class EndpointsClient: file_content_type="image/png", file_field="image", response_type=ImagesResult, + timeout=SLOW_PROVIDER_TIMEOUT_SECONDS, ) def generate_content( diff --git a/tests/e2e/models.py b/tests/e2e/models.py index 7711ca92b48..5e2cb90958e 100644 --- a/tests/e2e/models.py +++ b/tests/e2e/models.py @@ -233,6 +233,7 @@ class ChatBody(BaseModel): tool_choice: str | None = None guardrails: list[str] | None = None response_format: dict[str, object] | None = None + cache: dict[str, bool] | None = {"no-cache": True} class RouterSettingsOverride(BaseModel): @@ -431,6 +432,7 @@ class AnthropicMessagesBody(BaseModel): stream: bool | None = None tools: list[AnthropicTool] | None = None guardrails: list[str] | None = None + cache: dict[str, bool] | None = {"no-cache": True} class CountTokensBody(BaseModel): @@ -496,6 +498,7 @@ class McpServerInfo(BaseModel): class EmbedBody(BaseModel): model: str input: str + cache: dict[str, bool] | None = {"no-cache": True} class EmbedResponse(BaseModel): diff --git a/tests/e2e/otel_client.py b/tests/e2e/otel_client.py index 41555590dec..b11fddebc9c 100644 --- a/tests/e2e/otel_client.py +++ b/tests/e2e/otel_client.py @@ -24,7 +24,7 @@ import pytest from pydantic import BaseModel, ConfigDict, Field from e2e_config import OTEL_QUERY_URL, POLL_INTERVAL, POLL_TIMEOUT -from e2e_http import URL, NoBody, Success, get +from e2e_http import URL, NetworkError, NoBody, Result, Success, get #: OTEL resource service.name the proxy exports under (OTEL_SERVICE_NAME default). JAEGER_SERVICE = "litellm" @@ -100,18 +100,20 @@ def _settled(trace: JaegerTrace, names: set[str], prefixes: set[str]) -> bool: class OtelReader: query_url: str - def traces_for_call(self, call_id: str) -> list[JaegerTrace]: - """Every trace holding a span tagged with this call id. Jaeger matches - spans server-side and returns their full traces; more than one hit for - one call IS the split-trace bug, so this never collapses to one.""" - result = get( + def _query_traces(self, call_id: str) -> Result[JaegerTracesPage]: + return get( URL(f"{self.query_url}/api/traces"), headers=NoBody(), params=_TracesQuery(service=JAEGER_SERVICE, tags=json.dumps({CALL_ID_TAG: call_id})), response_type=JaegerTracesPage, timeout=30.0, ) - match result: + + def traces_for_call(self, call_id: str) -> list[JaegerTrace]: + """Every trace holding a span tagged with this call id. Jaeger matches + spans server-side and returns their full traces; more than one hit for + one call IS the split-trace bug, so this never collapses to one.""" + match self._query_traces(call_id): case Success(data=page): return page.data case failure: @@ -128,11 +130,24 @@ class OtelReader: on a split trace this never settles and the orphan comes back.""" deadline = time.monotonic() + POLL_TIMEOUT hits: list[JaegerTrace] = [] + unreachable: NetworkError | None = None while time.monotonic() < deadline: - hits = self.traces_for_call(call_id) - if len(hits) == 1 and _settled(hits[0], settled_names, settled_prefixes): - return hits + match self._query_traces(call_id): + case Success(data=page): + unreachable = None + hits = page.data + if len(hits) == 1 and _settled(hits[0], settled_names, settled_prefixes): + return hits + case NetworkError() as failure: + unreachable = failure + case failure: + pytest.fail(f"Jaeger query API at {self.query_url} failed: {failure}") time.sleep(POLL_INTERVAL) + if unreachable is not None: + pytest.fail( + f"Jaeger query API at {self.query_url} stayed unreachable until the " + f"{POLL_TIMEOUT}s poll deadline: {unreachable}" + ) return hits diff --git a/tests/e2e/proxy_client.py b/tests/e2e/proxy_client.py index 6cdd3354bf7..d12364e1794 100644 --- a/tests/e2e/proxy_client.py +++ b/tests/e2e/proxy_client.py @@ -70,6 +70,7 @@ from e2e_config import ( POLL_TIMEOUT, PROXY_BASE_URL, REQUEST_TIMEOUT, + SLOW_PROVIDER_TIMEOUT_SECONDS, settle_propagation, ) from transport import HttpTransport, SplitTransport, Transport @@ -425,6 +426,7 @@ class ProxyClient: headers=self.transport.bearer(key), json=body, response_type=OcrResponse, + timeout=SLOW_PROVIDER_TIMEOUT_SECONDS, ) def count_tokens(self, key: str, body: CountTokensBody) -> Result[CountTokensResponse]: diff --git a/tests/e2e/pytest.ini b/tests/e2e/pytest.ini index 2998a4b83c6..8feb4505ce3 100644 --- a/tests/e2e/pytest.ini +++ b/tests/e2e/pytest.ini @@ -2,7 +2,7 @@ # Config when any e2e suite under tests/e2e/ is run directly, e.g. # uv run pytest tests/e2e/quota_management/spend_tracking/ -v # The e2e marker is also registered in conftest.py for runs rooted elsewhere. -addopts = --strict-markers --strict-config +addopts = --strict-markers --strict-config --reruns 1 --only-rerun "kind='network'" --only-rerun "status_code=5[0-9][0-9]" markers = e2e: live test that requires a running proxy and real provider keys load: heavy throughput/load test; collected last so it never perturbs latency-sensitive suites diff --git a/tests/e2e/quota_management/budgets/test_spend_counter_reseed_e2e.py b/tests/e2e/quota_management/budgets/test_spend_counter_reseed_e2e.py index 7cfd3e33fd6..4a69135cdd1 100644 --- a/tests/e2e/quota_management/budgets/test_spend_counter_reseed_e2e.py +++ b/tests/e2e/quota_management/budgets/test_spend_counter_reseed_e2e.py @@ -39,6 +39,7 @@ pytestmark = pytest.mark.e2e MODEL = "claude-haiku-4-5" ACCUMULATE_CALLS = 24 BURST = 6 +BURST_TOLERATED_FAILURES = 1 # proxy_batch_write_at (60s) flushes the spend to the DB and default_redis_ttl (20s) # expires the counter; this waits out both. COLD_WAIT_SECONDS = 80 @@ -174,9 +175,11 @@ def test_cold_counter_reseed_keeps_counter_equal_to_db_spend( with ThreadPoolExecutor(max_workers=BURST) as pool: burst_results = list(pool.map(one, range(BURST))) - assert all(r.ok for r in burst_results), ( - "some burst calls failed; cannot exercise concurrent reseed. " - f"statuses={[r.status_code for r in burst_results]}" + failed = [r for r in burst_results if not r.ok] + assert len(failed) <= BURST_TOLERATED_FAILURES, ( + "too many burst calls failed; cannot exercise concurrent reseed. " + f"statuses={[r.status_code for r in burst_results]} " + f"bodies={[r.body[:300] for r in failed]}" ) counter: float | None = None diff --git a/tests/e2e/quota_management/spend_tracking/spend_e2e_client.py b/tests/e2e/quota_management/spend_tracking/spend_e2e_client.py index b4f64ba2ac5..056799b8499 100644 --- a/tests/e2e/quota_management/spend_tracking/spend_e2e_client.py +++ b/tests/e2e/quota_management/spend_tracking/spend_e2e_client.py @@ -65,6 +65,7 @@ def _chat_body( tags: list[str] | None = None, user: str | None = None, stream: bool = False, + cache: dict[str, bool] | None = {"no-cache": True}, ) -> ChatBody: return ChatBody( model=model, @@ -73,6 +74,7 @@ def _chat_body( stream=stream, user=user, metadata=ChatMetadata(tags=tags) if tags else None, + cache=cache, ) @@ -89,9 +91,11 @@ class SpendClient: max_tokens: int | None = None, tags: list[str] | None = None, user: str | None = None, + cache: dict[str, bool] | None = {"no-cache": True}, ) -> Result[ChatResponse]: return self.proxy.chat( - key, _chat_body(model, content, max_tokens=max_tokens, tags=tags, user=user) + key, + _chat_body(model, content, max_tokens=max_tokens, tags=tags, user=user, cache=cache), ) def chat_stream( diff --git a/tests/e2e/quota_management/spend_tracking/test_spend_tracking_e2e.py b/tests/e2e/quota_management/spend_tracking/test_spend_tracking_e2e.py index 6a0032981fd..c5d76d44580 100644 --- a/tests/e2e/quota_management/spend_tracking/test_spend_tracking_e2e.py +++ b/tests/e2e/quota_management/spend_tracking/test_spend_tracking_e2e.py @@ -226,8 +226,8 @@ def test_cache_hit_is_zero_cost_and_suffixed( # populated. The marker keeps each run isolated - a fixed prompt would persist # in the shared response cache across runs and make both calls hit (flaky). prompt = f"What is the capital of France? Answer in one word. {unique_marker()}" - _ = unwrap(client.chat(scoped_key, "gemini-2.5-flash", prompt, max_tokens=16)) - _ = unwrap(client.chat(scoped_key, "gemini-2.5-flash", prompt, max_tokens=16)) + _ = unwrap(client.chat(scoped_key, "gemini-2.5-flash", prompt, max_tokens=16, cache=None)) + _ = unwrap(client.chat(scoped_key, "gemini-2.5-flash", prompt, max_tokens=16, cache=None)) rows = client.poll_logs_for_key( scoped_key, diff --git a/tests/e2e/router/reliability_support.py b/tests/e2e/router/reliability_support.py index cd70ac45da6..cc1c91c635b 100644 --- a/tests/e2e/router/reliability_support.py +++ b/tests/e2e/router/reliability_support.py @@ -47,6 +47,7 @@ def chat_override( content: str, override: RouterSettingsOverride | None = None, stream: bool = False, + cache: dict[str, bool] | None = {"no-cache": True}, ) -> StreamingResponse: """POST /chat/completions with an optional per-request router_settings_override, returning the raw outcome so tests read status, body, and reliability headers.""" @@ -59,6 +60,7 @@ def chat_override( max_tokens=64, stream=stream, router_settings_override=override, + cache=cache, ), stream=stream, ) diff --git a/tests/e2e/router/test_reliability_cache_e2e.py b/tests/e2e/router/test_reliability_cache_e2e.py index 78d8fcdc08f..4ea05a1ecca 100644 --- a/tests/e2e/router/test_reliability_cache_e2e.py +++ b/tests/e2e/router/test_reliability_cache_e2e.py @@ -23,13 +23,13 @@ class TestReliabilityCache: def test_exact_cache_returns_cached(self, client: ComplexityRouterClient, scoped_key: str) -> None: prompt = f"cache probe {unique_marker()}" - first = chat_override(client.proxy, scoped_key, "gpt-5.5", prompt) + first = chat_override(client.proxy, scoped_key, "gpt-5.5", prompt, cache=None) assert first.status_code == 200, f"first call should succeed, got {first.status_code}: {first.body[:300]}" assert "x-litellm-cache-key" not in first.headers, ( "first (uncached) call must not report a cache-key header" ) - second = chat_override(client.proxy, scoped_key, "gpt-5.5", prompt) + second = chat_override(client.proxy, scoped_key, "gpt-5.5", prompt, cache=None) assert second.status_code == 200, f"second call should succeed, got {second.status_code}: {second.body[:300]}" assert "x-litellm-cache-key" in second.headers, ( "second identical call should hit the response cache and report a cache-key header " diff --git a/tests/e2e/transport.py b/tests/e2e/transport.py index 27b11befc8e..44fdbaa3e41 100644 --- a/tests/e2e/transport.py +++ b/tests/e2e/transport.py @@ -25,7 +25,13 @@ from e2e_http import ( class Transport(Protocol): def post[R: BaseModel]( - self, path: str, *, headers: BaseModel, json: BaseModel, response_type: type[R] + self, + path: str, + *, + headers: BaseModel, + json: BaseModel, + response_type: type[R], + timeout: float | None = None, ) -> Result[R]: ... def stream( @@ -93,6 +99,7 @@ class Transport(Protocol): file_field: str = "file", params: BaseModel | None = None, response_type: type[R], + timeout: float | None = None, ) -> Result[R]: ... def download(self, path: str, *, headers: BaseModel) -> StreamingResponse: ... @@ -120,14 +127,22 @@ class HttpTransport: return self.bearer(self.master_key) def post[R: BaseModel]( - self, path: str, *, headers: BaseModel, json: BaseModel, response_type: type[R] + self, + path: str, + *, + headers: BaseModel, + json: BaseModel, + response_type: type[R], + timeout: float | None = None, ) -> Result[R]: + """`timeout` overrides the transport-wide request_timeout for this call, for + provider operations that legitimately outlive it (image edits, OCR).""" return e2e_http.post( self._url(path), headers=headers, json=json, response_type=response_type, - timeout=self.request_timeout, + timeout=self.request_timeout if timeout is None else timeout, ) def get[R: BaseModel]( @@ -250,6 +265,7 @@ class HttpTransport: file_field: str = "file", params: BaseModel | None = None, response_type: type[R], + timeout: float | None = None, ) -> Result[R]: return e2e_http.upload( self._url(path), @@ -261,7 +277,7 @@ class HttpTransport: file_field=file_field, params=params, response_type=response_type, - timeout=self.request_timeout, + timeout=self.request_timeout if timeout is None else timeout, ) def download(self, path: str, *, headers: BaseModel) -> StreamingResponse: @@ -327,10 +343,16 @@ class SplitTransport: return self.data.master def post[R: BaseModel]( - self, path: str, *, headers: BaseModel, json: BaseModel, response_type: type[R] + self, + path: str, + *, + headers: BaseModel, + json: BaseModel, + response_type: type[R], + timeout: float | None = None, ) -> Result[R]: return self._route(path).post( - path, headers=headers, json=json, response_type=response_type + path, headers=headers, json=json, response_type=response_type, timeout=timeout ) def get[R: BaseModel]( @@ -426,6 +448,7 @@ class SplitTransport: file_field: str = "file", params: BaseModel | None = None, response_type: type[R], + timeout: float | None = None, ) -> Result[R]: return self._route(path).upload( path, @@ -437,6 +460,7 @@ class SplitTransport: file_field=file_field, params=params, response_type=response_type, + timeout=timeout, ) def download(self, path: str, *, headers: BaseModel) -> StreamingResponse: diff --git a/tests/proxy_unit_tests/test_check_batch_cost.py b/tests/proxy_unit_tests/test_check_batch_cost.py index 0065dbebc59..a1864c5e480 100644 --- a/tests/proxy_unit_tests/test_check_batch_cost.py +++ b/tests/proxy_unit_tests/test_check_batch_cost.py @@ -1044,7 +1044,9 @@ class TestCheckBatchCost: Pre-fix it matched neither the completed-with-output branch nor the failed/expired/cancelled branch, so batch_processed stayed False and the row was re-selected on every poll cycle forever. It must now be marked terminal - exactly once, without being billed (no output means nothing to bill). + exactly once, without being billed: request_counts.completed == 0 proves the + missing output file means nothing to bill rather than a lagging output id + (#37713 keeps the lagging case eligible for the next cycle). """ import base64 from unittest.mock import patch @@ -1073,6 +1075,7 @@ class TestCheckBatchCost: mock_response.status = completed_status mock_response.output_file_id = None mock_response.error_file_id = "file-error-123" + mock_response.request_counts = MagicMock(completed=0, failed=3, total=3) mock_response.model_dump_json.return_value = ( f'{{"id":"batch-1","status":"{completed_status}"}}' ) @@ -1107,6 +1110,72 @@ class TestCheckBatchCost: mock_llm_router.get_deployment_credentials_with_provider.call_count == 0 ), "a batch with no output file must not enter the cost-tracking path" + @pytest.mark.asyncio + @pytest.mark.parametrize( + "request_counts", + [MagicMock(completed=7, failed=0, total=7), None], + ids=["lagging_output_id", "unknown_counts"], + ) + async def test_completed_with_lagging_output_file_left_for_next_cycle( + self, + check_batch_cost_instance, + mock_prisma_client, + mock_llm_router, + request_counts, + ): + """#37713 regression: a batch can report completed while its output_file_id is + still lagging behind at the provider. Retiring it in that window (or when the + request counts cannot prove there is nothing to bill) permanently loses the + spend record, so the poller must leave the row untouched and revisit it on the + next cycle once the output id has appeared. + """ + import base64 + from unittest.mock import patch + + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( + return_value=1 + ) + mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock() + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( + return_value=None + ) + + mock_job = MagicMock() + mock_job.id = "job-completed-lagging-output-1" + mock_job.unified_object_id = base64.urlsafe_b64encode( + b"litellm_proxy;model_id:model-123;llm_batch_id:batch-456" + ).decode() + mock_job.created_by = "user-1" + + assert check_batch_cost_instance._has_batch_processed_column is True + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( + return_value=[mock_job] + ) + + mock_response = MagicMock() + mock_response.status = "completed" + mock_response.output_file_id = None + mock_response.error_file_id = None + mock_response.request_counts = request_counts + + mock_llm_router.aretrieve_batch = AsyncMock(return_value=mock_response) + mock_llm_router.get_deployment_credentials_with_provider = MagicMock( + return_value={"api_key": "sk-test"} + ) + + with patch( + "litellm.files.main.afile_content", + new_callable=AsyncMock, + ) as mock_afile_content: + await check_batch_cost_instance.check_batch_cost() + + assert ( + mock_prisma_client.db.litellm_managedobjecttable.update.call_count == 0 + ), "a completed batch whose output id is still lagging must stay eligible for the next poll" + assert ( + mock_afile_content.await_count == 0 + ), "a batch with no output file must not be billed" + @pytest.mark.asyncio async def test_non_terminal_status_left_unprocessed( self, check_batch_cost_instance, mock_prisma_client, mock_llm_router diff --git a/tests/proxy_unit_tests/test_user_api_key_auth.py b/tests/proxy_unit_tests/test_user_api_key_auth.py index cc7de71aa56..0cdf3500d50 100644 --- a/tests/proxy_unit_tests/test_user_api_key_auth.py +++ b/tests/proxy_unit_tests/test_user_api_key_auth.py @@ -154,6 +154,7 @@ async def test_team_object_has_object_permission_id(): token=hashed_key, last_refreshed_at=time.time(), team_object_permission_id=permission_id, + team_models=["gpt-4o"], ) user_api_key_cache.set_cache(key=hashed_key, value=valid_token) @@ -242,6 +243,7 @@ async def test_aaauser_personal_budgets(key_ownership): user_id=_user_id, team_id="my-special-team", team_max_budget=100, + team_models=["gpt-4o"], spend=20, ) diff --git a/tests/test_litellm/enterprise/proxy/test_managed_files_hook.py b/tests/test_litellm/enterprise/proxy/test_managed_files_hook.py index fcd03e77aa2..eddfc4fbd34 100644 --- a/tests/test_litellm/enterprise/proxy/test_managed_files_hook.py +++ b/tests/test_litellm/enterprise/proxy/test_managed_files_hook.py @@ -14,8 +14,8 @@ import pytest from typing import Optional from unittest.mock import AsyncMock, MagicMock, patch -from litellm.proxy._types import UserAPIKeyAuth -from litellm.types.llms.openai import OpenAIFileObject +from litellm.proxy._types import LitellmUserRoles, ProxyException, UserAPIKeyAuth +from litellm.types.llms.openai import FileListPage, OpenAIFileObject from litellm.types.utils import LiteLLMBatch @@ -66,6 +66,114 @@ def _make_user_api_key_dict() -> UserAPIKeyAuth: ) +def _make_team_member_api_key_dict() -> UserAPIKeyAuth: + """The shape most real virtual keys carry: a user_id and a team_id.""" + return UserAPIKeyAuth( + api_key="sk-test", + user_id="test-user", + team_id="test-team", + parent_otel_span=None, + ) + + +def _make_service_account_api_key_dict() -> UserAPIKeyAuth: + return UserAPIKeyAuth( + api_key="sk-service", + team_id="test-team", + parent_otel_span=None, + ) + + +def _make_admin_api_key_dict() -> UserAPIKeyAuth: + return UserAPIKeyAuth( + api_key="sk-admin", + user_id="admin-user", + user_role=LitellmUserRoles.PROXY_ADMIN, + parent_otel_span=None, + ) + + +def _make_managed_file_row( + unified_file_id: str, + purpose: str = "batch_output", + created_by: str = "test-user", + team_id: Optional[str] = None, +) -> MagicMock: + file_object = _make_file_object(f"file-provider-{unified_file_id}").model_copy( + update={"purpose": purpose} + ) + return MagicMock( + unified_file_id=unified_file_id, + file_object=file_object.model_dump(), + created_by=created_by, + team_id=team_id, + ) + + +def _make_unparseable_managed_file_row( + unified_file_id: str, + created_by: str = "test-user", + team_id: Optional[str] = None, +) -> MagicMock: + """A row whose stored blob cannot be parsed back into a file object.""" + return MagicMock( + unified_file_id=unified_file_id, + file_object=None, + created_by=created_by, + team_id=team_id, + ) + + +def _row_matches_where(row, where) -> bool: + """Apply the Prisma ``where`` shapes build_owner_filter actually emits: + ``{}``, a single equality, and the ``OR`` of equalities a key carrying + both a user_id and a team_id produces.""" + for field, expected in where.items(): + if field == "OR": + if not any(_row_matches_where(row, clause) for clause in expected): + return False + elif getattr(row, field) != expected: + return False + return True + + +class _FakeManagedFileTable: + """In-memory stand-in for the managed file table, newest row first.""" + + def __init__(self, rows): + self.rows = list(rows) + self.find_many_calls = [] + self.find_first_calls = [] + + def _owned_rows(self, where): + return [row for row in self.rows if _row_matches_where(row, where)] + + async def find_first(self, where): + self.find_first_calls.append(where) + return next(iter(self._owned_rows(where)), None) + + async def find_many(self, where, take=None, order=None, cursor=None, skip=0): + self.find_many_calls.append( + {"where": where, "take": take, "order": order, "cursor": cursor, "skip": skip} + ) + rows = self._owned_rows(where) + if cursor is not None: + start = next( + index + for index, row in enumerate(rows) + if row.unified_file_id == cursor["unified_file_id"] + ) + rows = rows[start + skip :] + return rows if take is None else rows[:take] + + +def _make_managed_files_over_rows(rows): + managed_files = _make_managed_files_instance() + table = _FakeManagedFileTable(rows) + managed_files.prisma_client.db.litellm_managedfiletable = table + return managed_files, table + + def _make_managed_files_instance(): """Create a _PROXY_LiteLLMManagedFiles with storage methods mocked out.""" from litellm_enterprise.proxy.hooks.managed_files import ( @@ -190,6 +298,578 @@ async def test_get_user_created_file_ids_remaps_stored_raw_provider_id_to_unifie assert files[0].purpose == raw_provider_object.purpose +@pytest.mark.asyncio +async def test_afile_list_returns_owner_scoped_managed_files(): + managed_files = _make_managed_files_instance() + managed_files.prisma_client.db.litellm_managedfiletable.find_many = AsyncMock( + return_value=[ + MagicMock( + file_object=_make_file_object("file-provider-id").model_dump(), + unified_file_id="unified-file-id", + ), + MagicMock( + file_object=_make_file_object("file-other-purpose").model_copy( + update={"purpose": "batch"} + ).model_dump(), + unified_file_id="unified-other-purpose", + ), + ] + ) + + response = await managed_files.afile_list( + purpose="batch_output", + litellm_parent_otel_span=None, + user_api_key_dict=_make_user_api_key_dict(), + ) + + managed_files.prisma_client.db.litellm_managedfiletable.find_many.assert_awaited_once_with( + where={"created_by": "test-user"}, + take=10001, + order=[{"created_at": "desc"}, {"unified_file_id": "desc"}], + ) + assert [file.id for file in response.data] == ["unified-file-id"] + assert response.first_id == "unified-file-id" + assert response.last_id == "unified-file-id" + assert response.has_more is False + + +@pytest.mark.asyncio +async def test_afile_list_returns_a_page_object_callbacks_can_read(): + """Post-call hooks receive the listing and read ``.data`` off it, the way the + provider SDK's page lets them. The body on the wire stays a plain list page.""" + from fastapi.encoders import jsonable_encoder + + managed_files, _ = _make_managed_files_over_rows([_make_managed_file_row("unified-file-id")]) + + page = await managed_files.afile_list( + purpose=None, + litellm_parent_otel_span=None, + user_api_key_dict=_make_user_api_key_dict(), + ) + + assert isinstance(page, FileListPage) + assert [file.id for file in page.data] == ["unified-file-id"] + + body = jsonable_encoder(page) + assert list(body) == ["object", "data", "first_id", "last_id", "has_more"] + assert body["object"] == "list" + assert [file["id"] for file in body["data"]] == ["unified-file-id"] + assert body["first_id"] == "unified-file-id" + assert body["last_id"] == "unified-file-id" + assert body["has_more"] is False + + +@pytest.mark.asyncio +@pytest.mark.parametrize("purpose", ["nonexistent_purpose", "EVALS", "batch "]) +async def test_afile_list_rejects_a_purpose_the_files_api_never_accepts(purpose): + """No stored file can carry an undocumented purpose, so filtering on one is a + bad request rather than a legitimately empty page.""" + managed_files, table = _make_managed_files_over_rows([_make_managed_file_row("unified-file-id")]) + + with pytest.raises(ProxyException) as exc_info: + await managed_files.afile_list( + purpose=purpose, + litellm_parent_otel_span=None, + user_api_key_dict=_make_user_api_key_dict(), + ) + + assert exc_info.value.code == "400" + assert exc_info.value.type == "invalid_request_error" + assert exc_info.value.param == "purpose" + assert table.find_many_calls == [] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("purpose", ["batch", "assistants", "fine-tune", "evals", None]) +async def test_afile_list_accepts_every_documented_purpose(purpose): + managed_files, _ = _make_managed_files_over_rows([_make_managed_file_row("unified-file-id")]) + + page = await managed_files.afile_list( + purpose=purpose, + litellm_parent_otel_span=None, + user_api_key_dict=_make_user_api_key_dict(), + ) + + assert isinstance(page, FileListPage) + + +@pytest.mark.asyncio +async def test_afile_list_does_not_leak_another_callers_files(): + managed_files, table = _make_managed_files_over_rows( + [ + _make_managed_file_row("unified-mine-2"), + _make_managed_file_row("unified-theirs", created_by="other-user"), + _make_managed_file_row("unified-mine-1"), + ] + ) + + response = await managed_files.afile_list( + purpose=None, + litellm_parent_otel_span=None, + user_api_key_dict=_make_user_api_key_dict(), + ) + + assert [file.id for file in response.data] == ["unified-mine-2", "unified-mine-1"] + assert table.find_many_calls[0]["where"] == {"created_by": "test-user"} + + +@pytest.mark.asyncio +async def test_afile_list_returns_own_and_team_files_for_a_key_carrying_both_ids(): + managed_files, table = _make_managed_files_over_rows( + [ + _make_managed_file_row("unified-mine"), + _make_managed_file_row("unified-teammates", created_by="other-user", team_id="test-team"), + _make_managed_file_row("unified-outsiders", created_by="outsider", team_id="other-team"), + ] + ) + + response = await managed_files.afile_list( + purpose=None, + litellm_parent_otel_span=None, + user_api_key_dict=_make_team_member_api_key_dict(), + ) + + assert [file.id for file in response.data] == ["unified-mine", "unified-teammates"] + assert table.find_many_calls[0]["where"] == { + "OR": [{"created_by": "test-user"}, {"team_id": "test-team"}] + } + + +@pytest.mark.asyncio +async def test_afile_list_scopes_a_service_account_key_to_its_team(): + managed_files, table = _make_managed_files_over_rows( + [ + _make_managed_file_row("unified-teams", created_by="other-user", team_id="test-team"), + _make_managed_file_row("unified-outsiders", created_by="outsider", team_id="other-team"), + ] + ) + + response = await managed_files.afile_list( + purpose=None, + litellm_parent_otel_span=None, + user_api_key_dict=_make_service_account_api_key_dict(), + ) + + assert [file.id for file in response.data] == ["unified-teams"] + assert table.find_many_calls[0]["where"] == {"team_id": "test-team"} + + +@pytest.mark.asyncio +async def test_afile_list_returns_every_callers_files_for_a_proxy_admin(): + managed_files, table = _make_managed_files_over_rows( + [ + _make_managed_file_row("unified-mine"), + _make_managed_file_row("unified-theirs", created_by="other-user", team_id="other-team"), + ] + ) + + response = await managed_files.afile_list( + purpose=None, + litellm_parent_otel_span=None, + user_api_key_dict=_make_admin_api_key_dict(), + ) + + assert [file.id for file in response.data] == ["unified-mine", "unified-theirs"] + assert table.find_many_calls[0]["where"] == {} + + +@pytest.mark.asyncio +async def test_afile_list_pages_a_team_key_across_both_halves_of_its_filter(): + """Keyset pagination has to walk an OR filter as one ordered set, without + repeating a row across pages or dropping one between them.""" + managed_files, table = _make_managed_files_over_rows( + [ + _make_managed_file_row("unified-0"), + _make_managed_file_row("unified-1", created_by="other-user", team_id="test-team"), + _make_managed_file_row("unified-2"), + _make_managed_file_row("unified-3", created_by="outsider", team_id="other-team"), + _make_managed_file_row("unified-4", created_by="other-user", team_id="test-team"), + ] + ) + user_api_key_dict = _make_team_member_api_key_dict() + + seen = [] + cursor = None + for _ in range(4): + response = await managed_files.afile_list( + purpose=None, + litellm_parent_otel_span=None, + user_api_key_dict=user_api_key_dict, + limit=2, + after=cursor, + ) + seen.extend(file.id for file in response.data) + if not response.has_more: + break + cursor = response.last_id + + assert seen == ["unified-0", "unified-1", "unified-2", "unified-4"] + assert all( + call["where"] == {"OR": [{"created_by": "test-user"}, {"team_id": "test-team"}]} + for call in table.find_many_calls + ) + + +@pytest.mark.asyncio +async def test_afile_list_orders_newest_first_and_breaks_ties_on_the_cursor_column(): + managed_files, table = _make_managed_files_over_rows([_make_managed_file_row("unified-mine")]) + + await managed_files.afile_list( + purpose=None, + litellm_parent_otel_span=None, + user_api_key_dict=_make_user_api_key_dict(), + ) + + assert table.find_many_calls[0]["order"] == [ + {"created_at": "desc"}, + {"unified_file_id": "desc"}, + ] + + +@pytest.mark.asyncio +async def test_afile_list_denies_a_caller_without_a_user_or_team(): + managed_files, table = _make_managed_files_over_rows([_make_managed_file_row("unified-mine")]) + + response = await managed_files.afile_list( + purpose=None, + litellm_parent_otel_span=None, + user_api_key_dict=UserAPIKeyAuth(api_key="sk-test", parent_otel_span=None), + ) + + assert response.data == [] + assert response.has_more is False + assert table.find_many_calls == [] + + +@pytest.mark.asyncio +async def test_afile_list_filters_by_purpose(): + managed_files, _ = _make_managed_files_over_rows( + [ + _make_managed_file_row("unified-batch-output"), + _make_managed_file_row("unified-batch", purpose="batch"), + ] + ) + + response = await managed_files.afile_list( + purpose="batch", + litellm_parent_otel_span=None, + user_api_key_dict=_make_user_api_key_dict(), + ) + + assert [file.id for file in response.data] == ["unified-batch"] + + +async def _walk_afile_list(managed_files, user_api_key_dict, purpose, limit): + """Page through the listing the way the official SDK does, off ``data[-1].id``.""" + seen = [] + after = None + while True: + page = await managed_files.afile_list( + purpose=purpose, + litellm_parent_otel_span=None, + user_api_key_dict=user_api_key_dict, + limit=limit, + after=after, + ) + page_ids = [file.id for file in page.data] + assert not set(page_ids) & set(seen) + seen.extend(page_ids) + if not page.has_more: + return seen + assert page_ids, "an SDK stops paging on an empty page, so has_more must never ride one" + after = page_ids[-1] + + +@pytest.mark.asyncio +async def test_afile_list_fills_a_page_past_rows_the_purpose_filter_drops(): + """The newest rows do not match, so the page must reach past them rather than come back empty.""" + managed_files, _ = _make_managed_files_over_rows( + [ + _make_managed_file_row("unified-0"), + _make_managed_file_row("unified-1"), + _make_managed_file_row("unified-2", purpose="batch"), + _make_managed_file_row("unified-3"), + _make_managed_file_row("unified-4", purpose="batch"), + ] + ) + user_api_key_dict = _make_user_api_key_dict() + + first_page = await managed_files.afile_list( + purpose="batch", + litellm_parent_otel_span=None, + user_api_key_dict=user_api_key_dict, + limit=1, + ) + + assert [file.id for file in first_page.data] == ["unified-2"] + assert first_page.has_more is True + assert first_page.last_id == "unified-2" + + second_page = await managed_files.afile_list( + purpose="batch", + litellm_parent_otel_span=None, + user_api_key_dict=user_api_key_dict, + limit=1, + after=first_page.last_id, + ) + + assert [file.id for file in second_page.data] == ["unified-4"] + assert second_page.has_more is False + + +@pytest.mark.parametrize("limit", [1, 2, 3]) +@pytest.mark.asyncio +async def test_afile_list_walks_every_purpose_match_at_any_limit(limit): + managed_files, _ = _make_managed_files_over_rows( + [ + _make_managed_file_row("unified-0"), + _make_managed_file_row("unified-1"), + _make_managed_file_row("unified-2", purpose="batch"), + _make_managed_file_row("unified-3"), + _make_managed_file_row("unified-4", purpose="batch"), + _make_managed_file_row("unified-5", purpose="batch"), + _make_managed_file_row("unified-6"), + ] + ) + + seen = await _walk_afile_list(managed_files, _make_user_api_key_dict(), "batch", limit) + + assert seen == ["unified-2", "unified-4", "unified-5"] + + +@pytest.mark.asyncio +async def test_afile_list_fills_a_page_past_rows_that_do_not_parse(): + managed_files, _ = _make_managed_files_over_rows( + [ + _make_unparseable_managed_file_row("unified-0"), + _make_unparseable_managed_file_row("unified-1"), + _make_managed_file_row("unified-2"), + ] + ) + + page = await managed_files.afile_list( + purpose=None, + litellm_parent_otel_span=None, + user_api_key_dict=_make_user_api_key_dict(), + limit=1, + ) + + assert [file.id for file in page.data] == ["unified-2"] + assert page.has_more is False + + +_DEEP_SCAN_ROW_COUNT = 2000 +_DEEP_SCAN_QUERY_BUDGET = 10 + + +@pytest.mark.asyncio +async def test_afile_list_bounds_the_queries_a_deep_purpose_match_costs(): + """A tiny limit over rows the filter drops must not turn one request into thousands of queries.""" + managed_files, table = _make_managed_files_over_rows( + [_make_managed_file_row(f"unified-{index:05d}") for index in range(_DEEP_SCAN_ROW_COUNT)] + + [_make_managed_file_row("unified-match", purpose="batch")] + ) + + page = await managed_files.afile_list( + purpose="batch", + litellm_parent_otel_span=None, + user_api_key_dict=_make_user_api_key_dict(), + limit=1, + ) + + assert [file.id for file in page.data] == ["unified-match"] + assert page.has_more is False + assert len(table.find_many_calls) <= _DEEP_SCAN_QUERY_BUDGET + + +@pytest.mark.asyncio +async def test_afile_list_bounds_the_queries_a_deep_unparseable_run_costs(): + """Rows that will not parse drop out like a filter does, so they get the same bound.""" + managed_files, table = _make_managed_files_over_rows( + [_make_unparseable_managed_file_row(f"unified-{index:05d}") for index in range(_DEEP_SCAN_ROW_COUNT)] + + [_make_managed_file_row("unified-parses")] + ) + + page = await managed_files.afile_list( + purpose=None, + litellm_parent_otel_span=None, + user_api_key_dict=_make_user_api_key_dict(), + limit=1, + ) + + assert [file.id for file in page.data] == ["unified-parses"] + assert page.has_more is False + assert len(table.find_many_calls) <= _DEEP_SCAN_QUERY_BUDGET + + +@pytest.mark.asyncio +async def test_afile_list_reads_one_chunk_when_the_first_one_fills_the_page(): + """The widened chunk must stay off the common path, where the newest rows already fill the page.""" + managed_files, table = _make_managed_files_over_rows( + [_make_managed_file_row(f"unified-{index:05d}") for index in range(_DEEP_SCAN_ROW_COUNT)] + ) + + page = await managed_files.afile_list( + purpose=None, + litellm_parent_otel_span=None, + user_api_key_dict=_make_user_api_key_dict(), + limit=2, + ) + + assert [file.id for file in page.data] == ["unified-00000", "unified-00001"] + assert page.has_more is True + assert [call["take"] for call in table.find_many_calls] == [3] + + +@pytest.mark.asyncio +async def test_afile_list_reports_no_more_pages_when_nothing_matches(): + managed_files, _ = _make_managed_files_over_rows( + [_make_managed_file_row(f"unified-{index}") for index in range(5)] + ) + + page = await managed_files.afile_list( + purpose="batch", + litellm_parent_otel_span=None, + user_api_key_dict=_make_user_api_key_dict(), + limit=2, + ) + + assert page.data == [] + assert page.has_more is False + assert page.first_id is None + assert page.last_id is None + + +@pytest.mark.asyncio +async def test_afile_list_honors_limit_and_reports_more_pages(): + managed_files, table = _make_managed_files_over_rows( + [_make_managed_file_row(f"unified-{index}") for index in range(5)] + ) + + response = await managed_files.afile_list( + purpose=None, + litellm_parent_otel_span=None, + user_api_key_dict=_make_user_api_key_dict(), + limit=2, + ) + + assert [file.id for file in response.data] == ["unified-0", "unified-1"] + assert response.has_more is True + assert table.find_many_calls[0]["take"] == 3 + + +@pytest.mark.asyncio +async def test_afile_list_pages_through_every_file_without_overlap(): + managed_files, table = _make_managed_files_over_rows( + [_make_managed_file_row(f"unified-{index}") for index in range(5)] + ) + user_api_key_dict = _make_user_api_key_dict() + + seen = [] + after = None + while True: + page = await managed_files.afile_list( + purpose=None, + litellm_parent_otel_span=None, + user_api_key_dict=user_api_key_dict, + limit=2, + after=after, + ) + page_ids = [file.id for file in page.data] + assert not set(page_ids) & set(seen) + seen.extend(page_ids) + if not page.has_more: + break + after = page.last_id + + assert seen == [f"unified-{index}" for index in range(5)] + assert table.find_many_calls[1]["cursor"] == {"unified_file_id": "unified-1"} + assert table.find_many_calls[1]["skip"] == 1 + + +@pytest.mark.parametrize( + "unknown_cursor", + ["unified-theirs", "unified-nowhere"], + ids=["another-users-file", "no-such-file"], +) +@pytest.mark.asyncio +async def test_afile_list_rejects_an_after_cursor_outside_the_callers_files(unknown_cursor): + from litellm.proxy._types import ProxyException + + managed_files, table = _make_managed_files_over_rows( + [ + _make_managed_file_row("unified-mine"), + _make_managed_file_row("unified-theirs", created_by="other-user"), + ] + ) + + with pytest.raises(ProxyException) as exc_info: + await managed_files.afile_list( + purpose=None, + litellm_parent_otel_span=None, + user_api_key_dict=_make_user_api_key_dict(), + after=unknown_cursor, + ) + + assert exc_info.value.code == "400" + assert exc_info.value.type == "invalid_request_error" + assert exc_info.value.param == "after" + assert exc_info.value.message == f"Invalid 'after' cursor: no file found with id '{unknown_cursor}'." + assert table.find_first_calls[0] == { + "created_by": "test-user", + "unified_file_id": unknown_cursor, + } + assert table.find_many_calls == [] + + +@pytest.mark.parametrize( + "limit, bound, expected_range", + [ + (0, "below minimum", ">= 1"), + (-1, "below minimum", ">= 1"), + (10001, "above maximum", "<= 10000"), + ], +) +@pytest.mark.asyncio +async def test_afile_list_rejects_a_limit_outside_the_openai_range(limit, bound, expected_range): + from litellm.proxy._types import ProxyException + + managed_files, table = _make_managed_files_over_rows([_make_managed_file_row("unified-mine")]) + + with pytest.raises(ProxyException) as exc_info: + await managed_files.afile_list( + purpose=None, + litellm_parent_otel_span=None, + user_api_key_dict=_make_user_api_key_dict(), + limit=limit, + ) + + assert exc_info.value.code == "400" + assert exc_info.value.type == "invalid_request_error" + assert exc_info.value.param == "limit" + assert exc_info.value.message == ( + f"Invalid 'limit': integer {bound} value. Expected a value {expected_range}, but got {limit} instead." + ) + assert table.find_many_calls == [] + + +@pytest.mark.parametrize("limit", [1, 10000]) +@pytest.mark.asyncio +async def test_afile_list_accepts_the_ends_of_the_openai_limit_range(limit): + managed_files, table = _make_managed_files_over_rows([_make_managed_file_row("unified-mine")]) + + response = await managed_files.afile_list( + purpose=None, + litellm_parent_otel_span=None, + user_api_key_dict=_make_user_api_key_dict(), + limit=limit, + ) + + assert [file.id for file in response.data] == ["unified-mine"] + assert response.has_more is False + assert table.find_many_calls[0]["take"] == limit + 1 + + @pytest.mark.asyncio async def test_parse_managed_file_object_warning_omits_rejected_values(caplog): from litellm_enterprise.proxy.hooks.managed_files import ( diff --git a/tests/test_litellm/integrations/otel/test_otel_v2_sources_of_truth.py b/tests/test_litellm/integrations/otel/test_otel_v2_sources_of_truth.py index 19d0cfc0b18..2a66d5ee139 100644 --- a/tests/test_litellm/integrations/otel/test_otel_v2_sources_of_truth.py +++ b/tests/test_litellm/integrations/otel/test_otel_v2_sources_of_truth.py @@ -4,6 +4,7 @@ and the typed StandardLoggingPayload adapter. These need no OTel SDK.""" import logging import re from pathlib import Path +from typing import Final import pytest @@ -14,6 +15,7 @@ from litellm.integrations.otel import ( Error, GenAI, GenAIOperation, + GenAIOutputType, HTTP, LiteLLM, OpenTelemetryV2Config, @@ -21,8 +23,10 @@ from litellm.integrations.otel import ( is_otel_v2_enabled, promoted_baggage, resolve_operation, + resolve_output_type, resolve_provider, ) +from litellm.integrations.otel.mappers.genai import GenAIMapper from litellm.integrations.otel.model import spans as spans_mod from litellm.integrations.otel.model.payloads import ( LLMCallSpanData, @@ -264,6 +268,74 @@ def test_vector_store_file_management_is_not_chat(call_type): assert resolve_operation(call_type).value == "litellm.vector_store_file_management" +_NON_CHAT_ROUTES: Final = ( + ("image_generation", GenAIOperation.GENERATE_CONTENT, GenAIOutputType.IMAGE), + ("speech", GenAIOperation.GENERATE_CONTENT, GenAIOutputType.SPEECH), + ("transcription", GenAIOperation.GENERATE_CONTENT, GenAIOutputType.TEXT), + ("ocr", GenAIOperation.GENERATE_CONTENT, GenAIOutputType.TEXT), + ("moderation", GenAIOperation.LITELLM_MODERATION, None), +) + + +@pytest.mark.parametrize( + ("call_type", "operation", "output_type"), + [ + (f"{prefix}{call_type}", operation, output_type) + for call_type, operation, output_type in _NON_CHAT_ROUTES + for prefix in ("", "a") + ], +) +def test_non_chat_inference_routes_follow_genai_semconv(call_type, operation, output_type): + """Image generation, speech, transcription and OCR all produce content, so the + convention names them ``generate_content`` and separates them by the requested + output modality rather than by an invented operation. Moderation classifies + instead of generating and the convention names nothing for it, so it keeps a + vendor value. Either way the spans must not land in the chat series a dashboard + reads.""" + assert resolve_operation(call_type) is operation + assert resolve_output_type(call_type) is output_type + + +@pytest.mark.parametrize( + ("call_type", "operation", "output_type"), + [(f"a{call_type}", operation, output_type) for call_type, operation, output_type in _NON_CHAT_ROUTES], +) +def test_non_chat_route_spans_carry_semconv_name_and_modality(call_type, operation, output_type): + """The emitted span, not just the mapping table: name is + ``{gen_ai.operation.name} {gen_ai.request.model}``, the modality rides + ``gen_ai.output.type``, and the route stays recoverable from + ``litellm.call_type`` now that several routes share one operation.""" + data = LLMCallSpanData.from_standard_logging_payload( + _sample_payload(call_type=call_type, model="some-model", custom_llm_provider="openai") + ) + attrs = GenAIMapper().map(data) + + assert spans_mod.llm_call_span_name(data) == f"{operation.value} some-model" + assert attrs[GenAI.OPERATION_NAME] == operation.value + assert attrs[GenAI.PROVIDER_NAME] == "openai" + assert attrs[GenAI.REQUEST_MODEL] == "some-model" + assert attrs[LiteLLM.CALL_TYPE] == call_type + assert attrs.get(GenAI.OUTPUT_TYPE) == (output_type.value if output_type else None) + + +def test_non_chat_route_error_span_keeps_error_attributes(): + """Modality mapping must not cost the failure signal: a failed non-chat call + still carries the error type alongside the standardized operation.""" + data = LLMCallSpanData.from_standard_logging_payload( + _sample_payload( + call_type="aspeech", + model="tts-1", + status="failure", + error_information={"error_class": "BadRequestError"}, + ) + ) + attrs = GenAIMapper().map(data) + + assert attrs[GenAI.OPERATION_NAME] == GenAIOperation.GENERATE_CONTENT.value + assert attrs[GenAI.OUTPUT_TYPE] == GenAIOutputType.SPEECH.value + assert attrs[Error.TYPE] == "BadRequestError" + + def test_vendor_operation_values_are_namespaced(): """A vendor value must stay under the ``litellm.`` prefix: an unprefixed invented name could collide with a value the convention adds later, silently changing what diff --git a/tests/test_litellm/integrations/test_prometheus_queue_guardrail_metrics.py b/tests/test_litellm/integrations/test_prometheus_queue_guardrail_metrics.py index 85be9e32121..f04e8d0d2c7 100644 --- a/tests/test_litellm/integrations/test_prometheus_queue_guardrail_metrics.py +++ b/tests/test_litellm/integrations/test_prometheus_queue_guardrail_metrics.py @@ -229,6 +229,137 @@ class TestPrometheusQueueTimeMetric: ), "Queue time metric should not be recorded for negative values" +class TestPrometheusTotalLatencyMetric: + """litellm_request_total_latency_metric must be true end-to-end latency: start_time + (set after auth already completed, see LIT-6012) plus queue_time_seconds (the + auth + pre-call setup window queue_time_seconds itself covers), not start_time alone.""" + + @staticmethod + def _enum_values() -> UserAPIKeyLabelValues: + return UserAPIKeyLabelValues( + end_user=None, + hashed_api_key="test-key", + api_key_alias="test-alias", + requested_model="gpt-3.5-turbo", + model_group="gpt-3.5-turbo", + team=None, + team_alias=None, + user=None, + user_email=None, + status_code="200", + model="gpt-3.5-turbo", + litellm_model_name="gpt-3.5-turbo", + tags=[], + model_id="gpt-3.5-turbo", + api_base="https://api.openai.com", + api_provider="openai", + exception_status=None, + exception_class=None, + custom_metadata_labels={}, + route=None, + ) + + def test_total_latency_includes_queue_time_when_present(self): + """The observed total-latency value must be (end_time - start_time) + queue_time_seconds, + so auth/pre-call time (queue_time_seconds) is not silently excluded from "total" latency.""" + prometheus_logger = PrometheusLogger() + + mock_metric = MagicMock() + mock_labeled_metric = MagicMock() + mock_metric.labels.return_value = mock_labeled_metric + prometheus_logger.litellm_request_total_latency_metric = mock_metric + + start_time = datetime(2024, 1, 1, 0, 0, 0) + end_time = datetime(2024, 1, 1, 0, 0, 2) # 2.0s of LLM-call/post-call time + queue_time_seconds = 0.5 # auth + pre-call setup time + + kwargs = { + "litellm_params": {"metadata": {"queue_time_seconds": queue_time_seconds}}, + "model": "gpt-3.5-turbo", + "start_time": start_time, + "end_time": end_time, + } + + prometheus_logger._set_latency_metrics( + kwargs=kwargs, + model="gpt-3.5-turbo", + user_api_key="test-key", + user_api_key_alias="test-alias", + user_api_team=None, + user_api_team_alias=None, + enum_values=self._enum_values(), + ) + + observed_value = mock_labeled_metric.observe.call_args_list[0][0][0] + assert observed_value == pytest.approx(2.5) + + def test_total_latency_falls_back_to_start_end_delta_without_queue_time(self): + """Without queue_time_seconds (e.g. a non-proxy caller), the metric must still + observe the plain end_time - start_time delta rather than erroring or dropping it.""" + prometheus_logger = PrometheusLogger() + + mock_metric = MagicMock() + mock_labeled_metric = MagicMock() + mock_metric.labels.return_value = mock_labeled_metric + prometheus_logger.litellm_request_total_latency_metric = mock_metric + + start_time = datetime(2024, 1, 1, 0, 0, 0) + end_time = datetime(2024, 1, 1, 0, 0, 2) + + kwargs = { + "litellm_params": {"metadata": {}}, + "model": "gpt-3.5-turbo", + "start_time": start_time, + "end_time": end_time, + } + + prometheus_logger._set_latency_metrics( + kwargs=kwargs, + model="gpt-3.5-turbo", + user_api_key="test-key", + user_api_key_alias="test-alias", + user_api_team=None, + user_api_team_alias=None, + enum_values=self._enum_values(), + ) + + observed_value = mock_labeled_metric.observe.call_args_list[0][0][0] + assert observed_value == pytest.approx(2.0) + + def test_total_latency_ignores_negative_queue_time(self): + """A negative queue_time_seconds (clock skew / bad data) must not be added in -- + matches the existing >= 0 guard on the standalone queue-time metric.""" + prometheus_logger = PrometheusLogger() + + mock_metric = MagicMock() + mock_labeled_metric = MagicMock() + mock_metric.labels.return_value = mock_labeled_metric + prometheus_logger.litellm_request_total_latency_metric = mock_metric + + start_time = datetime(2024, 1, 1, 0, 0, 0) + end_time = datetime(2024, 1, 1, 0, 0, 2) + + kwargs = { + "litellm_params": {"metadata": {"queue_time_seconds": -0.1}}, + "model": "gpt-3.5-turbo", + "start_time": start_time, + "end_time": end_time, + } + + prometheus_logger._set_latency_metrics( + kwargs=kwargs, + model="gpt-3.5-turbo", + user_api_key="test-key", + user_api_key_alias="test-alias", + user_api_team=None, + user_api_team_alias=None, + enum_values=self._enum_values(), + ) + + observed_value = mock_labeled_metric.observe.call_args_list[0][0][0] + assert observed_value == pytest.approx(2.0) + + class TestPrometheusGuardrailMetrics: """Test guardrail metrics recording""" diff --git a/tests/test_litellm/integrations/websearch_interception/test_websearch_agentic_loop_cap.py b/tests/test_litellm/integrations/websearch_interception/test_websearch_agentic_loop_cap.py new file mode 100644 index 00000000000..40fd8c4e9e6 --- /dev/null +++ b/tests/test_litellm/integrations/websearch_interception/test_websearch_agentic_loop_cap.py @@ -0,0 +1,754 @@ +""" +Unit tests for what an intercepted request returns once a safety rail refuses +another agentic loop. + +The web search interception loop injects an internal tool (litellm_web_search) +that the client never declared. When the loop cap or the repeated-fingerprint +guard trips, the turn has to end with a terminal response: leaking that internal +tool_use block leaves the client holding a tool call it cannot answer. + +Also covers the max_agentic_loops knob on websearch_interception_params, from +config.yaml through to the settings the loop actually reads. +""" + +import json +from unittest.mock import MagicMock + +import pytest + +import litellm +from litellm.integrations.custom_logger import CustomLogger +from litellm.integrations.websearch_interception.handler import ( + WebSearchInterceptionLogger, +) +from litellm.llms.anthropic.experimental_pass_through.messages.fake_stream_iterator import ( + FakeAnthropicMessagesStreamIterator, +) +from litellm.litellm_core_utils.agentic_loop_settings import DEFAULT_MAX_AGENTIC_LOOPS +from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler +from litellm.secret_managers.main import get_secret +from litellm.types.integrations.custom_logger import ( + AgenticLoopPlan, + AgenticLoopRequestPatch, + AgenticLoopSafetyError, +) + +INTERNAL_TOOL_NAME = "litellm_web_search" + + +@pytest.fixture(autouse=True) +def only_the_callbacks_these_tests_register(monkeypatch): + """ + These tests drive the hooks with a callback of their own on the logging + object, so a logger another test left on litellm.callbacks would join the + run and change what the hooks do. + """ + monkeypatch.setattr(litellm, "callbacks", []) + + +def _internal_tool_use_block(block_id: str = "toolu_internal_1") -> dict: + return { + "id": block_id, + "type": "tool_use", + "name": INTERNAL_TOOL_NAME, + "input": {"query": "who won the world cup"}, + } + + +def _native_search_blocks(index: int = 1) -> list[dict]: + return [ + { + "type": "server_tool_use", + "id": f"srvtoolu_{index}", + "name": "web_search", + "input": {"query": "who won the world cup"}, + }, + { + "type": "web_search_tool_result", + "tool_use_id": f"srvtoolu_{index}", + "content": [{"type": "web_search_result", "url": "https://example.com", "title": "Result"}], + }, + ] + + +def _response_asking_for_another_search(block_id: str = "toolu_internal_1") -> dict: + return { + "id": "msg_123", + "type": "message", + "role": "assistant", + "model": "claude-sonnet-4-5", + "content": [ + *_native_search_blocks(index=1), + {"type": "text", "text": "Let me check one more source."}, + _internal_tool_use_block(block_id), + ], + "stop_reason": "tool_use", + "usage": {"input_tokens": 10, "output_tokens": 5}, + } + + +def _block_types(response: dict) -> list[str]: + return [block["type"] for block in response["content"]] + + +def _tool_use_names(response: dict) -> list[str]: + return [block.get("name") for block in response["content"] if block.get("type") == "tool_use"] + + +class _InterceptingCallback(CustomLogger): + """ + Stands in for the websearch interceptor: asks for another loop whenever the + response carries an internal web search tool_use block, and injects the + native block pair on the way back out. + """ + + def __init__(self): + self.plan_calls = 0 + self.post_hook_calls = 0 + + async def async_should_run_agentic_loop( + self, response, model, messages, tools, stream, custom_llm_provider, kwargs + ): + if not isinstance(response, dict): + return True, {"tool_calls": [_internal_tool_use_block()]} + tool_calls = [ + block + for block in response.get("content", []) + if block.get("type") == "tool_use" and block.get("name") == INTERNAL_TOOL_NAME + ] + if not tool_calls: + return False, {} + return True, {"tool_calls": tool_calls, "tool_type": "websearch"} + + async def async_build_agentic_loop_plan( + self, + tools, + model, + messages, + response, + anthropic_messages_provider_config, + anthropic_messages_optional_request_params, + logging_obj, + stream, + kwargs, + ): + self.plan_calls += 1 + return AgenticLoopPlan( + run_agentic_loop=True, + request_patch=AgenticLoopRequestPatch( + messages=[{"role": "user", "content": "here are the search results"}], + max_tokens=1024, + ), + ) + + async def async_post_agentic_loop_response_hook(self, response, plan, kwargs): + self.post_hook_calls += 1 + if isinstance(response, dict): + response["content"] = [*_native_search_blocks(index=2), *response.get("content", [])] + return response + + +def _logging_obj(callback: CustomLogger, converted_stream: bool = False) -> MagicMock: + logging_obj = MagicMock() + logging_obj.model_call_details = {"websearch_interception_converted_stream": converted_stream} + logging_obj.dynamic_success_callbacks = [callback] + logging_obj.litellm_call_id = "call-abc" + return logging_obj + + +async def _run_hooks( + handler: BaseLLMHTTPHandler, + callback: CustomLogger, + kwargs: dict, + response: object = None, + stream: bool = False, + converted_stream: bool = False, + api_surface: str = "anthropic_messages", +): + return await handler._call_agentic_completion_hooks( + response=_response_asking_for_another_search() if response is None else response, + model="claude-sonnet-4-5", + messages=[{"role": "user", "content": "who won the world cup"}], + anthropic_messages_provider_config=MagicMock(), + anthropic_messages_optional_request_params={}, + logging_obj=_logging_obj(callback, converted_stream=converted_stream), + stream=stream, + custom_llm_provider="anthropic", + kwargs=kwargs, + api_surface=api_surface, + ) + + +class TestCappedLoopReturnsTerminalResponse: + def setup_method(self): + self.handler = BaseLLMHTTPHandler() + self.callback = _InterceptingCallback() + + @pytest.mark.asyncio + async def test_internal_tool_use_block_is_dropped(self): + result = await _run_hooks( + self.handler, + self.callback, + kwargs={"_agentic_loop_depth": 3, "max_agentic_loops": 3}, + ) + + assert isinstance(result, dict) + assert INTERNAL_TOOL_NAME not in _tool_use_names(result) + + @pytest.mark.asyncio + async def test_stop_reason_is_closed_out(self): + result = await _run_hooks( + self.handler, + self.callback, + kwargs={"_agentic_loop_depth": 3, "max_agentic_loops": 3}, + ) + + assert result["stop_reason"] == "end_turn" + + @pytest.mark.asyncio + async def test_native_blocks_and_text_survive(self): + result = await _run_hooks( + self.handler, + self.callback, + kwargs={"_agentic_loop_depth": 3, "max_agentic_loops": 3}, + ) + + assert _block_types(result) == ["server_tool_use", "web_search_tool_result", "text"] + + @pytest.mark.asyncio + async def test_turn_carrying_only_the_refused_call_still_ends_cleanly(self): + """ + The refused call can be every block the model produced, which leaves the + turn with no content once it is dropped. That still has to come back as a + finished turn rather than as the leaked call, so the client stops instead + of waiting on a tool it cannot run, and the rest of the message survives + so the request is still billed and traceable. + + An empty turn renders as nothing, which is the ceiling being set too low + for the question rather than a malformed response. + """ + nothing_but_the_refused_call = { + "id": "msg_123", + "type": "message", + "role": "assistant", + "model": "claude-sonnet-4-5", + "content": [_internal_tool_use_block()], + "stop_reason": "tool_use", + "usage": {"input_tokens": 10, "output_tokens": 5}, + } + + result = await _run_hooks( + self.handler, + self.callback, + kwargs={"_agentic_loop_depth": 3, "max_agentic_loops": 3}, + response=nothing_but_the_refused_call, + ) + + assert result["content"] == [] + assert result["stop_reason"] == "end_turn" + assert result["usage"] == {"input_tokens": 10, "output_tokens": 5} + assert result["id"] == "msg_123" + + @pytest.mark.asyncio + async def test_no_follow_up_model_call_is_planned(self): + """ + The rail has to end the turn without planning another model call, and it + has to end it by returning rather than by raising, which is the half that + the caller's response depends on. + """ + result = await _run_hooks( + self.handler, + self.callback, + kwargs={"_agentic_loop_depth": 3, "max_agentic_loops": 3}, + ) + + assert self.callback.plan_calls == 0 + assert result["stop_reason"] == "end_turn" + + @pytest.mark.asyncio + async def test_original_response_is_not_mutated(self): + response = _response_asking_for_another_search() + + await _run_hooks( + self.handler, + self.callback, + kwargs={"_agentic_loop_depth": 3, "max_agentic_loops": 3}, + response=response, + ) + + assert response["stop_reason"] == "tool_use" + assert INTERNAL_TOOL_NAME in _tool_use_names(response) + + @pytest.mark.asyncio + async def test_repeated_fingerprint_guard_is_terminal_too(self): + tool_calls = {"tool_calls": [_internal_tool_use_block()], "tool_type": "websearch"} + seen = json.dumps(tool_calls, sort_keys=True, default=str) + + result = await _run_hooks( + self.handler, + self.callback, + kwargs={"_agentic_loop_depth": 0, "max_agentic_loops": 3, "_agentic_loop_fingerprints": [seen]}, + ) + + assert self.callback.plan_calls == 0 + assert INTERNAL_TOOL_NAME not in _tool_use_names(result) + assert result["stop_reason"] == "end_turn" + + @pytest.mark.asyncio + async def test_client_declared_tool_use_is_left_alone(self): + response = _response_asking_for_another_search() + client_tool_use = {"id": "toolu_client_1", "type": "tool_use", "name": "get_weather", "input": {}} + response["content"].append(client_tool_use) + + result = await _run_hooks( + self.handler, + self.callback, + kwargs={"_agentic_loop_depth": 3, "max_agentic_loops": 3}, + response=response, + ) + + assert _tool_use_names(result) == ["get_weather"] + assert result["stop_reason"] == "tool_use" + + def test_only_the_refused_tool_calls_are_dropped(self): + """ + A block is matched on the id the rail refused, not on the tool name, so a + second block sharing that name survives when the rail never listed it. A + callback that picks its tool calls out by name hands both over and both + go, which is its own call to make; this is about not widening it here. + """ + response = _response_asking_for_another_search() + response["content"].append( + {"id": "toolu_client_1", "type": "tool_use", "name": INTERNAL_TOOL_NAME, "input": {}} + ) + + result = BaseLLMHTTPHandler._finalize_refused_agentic_response( + response=response, + tool_calls={"tool_calls": [_internal_tool_use_block()]}, + ) + + assert [block["id"] for block in result["content"] if block.get("type") == "tool_use"] == ["toolu_client_1"] + assert result["stop_reason"] == "tool_use" + + def test_tool_calls_without_ids_still_match_by_name(self): + """ + Not every callback shape carries ids on its tool calls, so the name is + still what decides when the rail refused a call that has no id. + """ + result = BaseLLMHTTPHandler._finalize_refused_agentic_response( + response=_response_asking_for_another_search(), + tool_calls={"tool_calls": [{"name": INTERNAL_TOOL_NAME, "input": {}}]}, + ) + + assert _tool_use_names(result) == [] + assert result["stop_reason"] == "end_turn" + + @pytest.mark.asyncio + async def test_streaming_caller_is_left_to_its_existing_behavior(self): + """ + A streaming caller has already sent the original message to the client, so + a finalized turn would land as a second message rather than replace the + first. The rail keeps raising there and the caller handles it as before. + """ + with pytest.raises(AgenticLoopSafetyError): + await _run_hooks( + self.handler, + self.callback, + kwargs={"_agentic_loop_depth": 3, "max_agentic_loops": 3}, + stream=True, + ) + + assert self.callback.plan_calls == 0 + + @pytest.mark.asyncio + async def test_responses_surface_is_left_to_its_existing_behavior(self): + """ + The responses surface carries a pydantic model rather than the anthropic + dict this finalizer rewrites, so it keeps raising instead of being handed + a response that was never actually finalized. + """ + with pytest.raises(AgenticLoopSafetyError): + await _run_hooks( + self.handler, + self.callback, + kwargs={"_agentic_loop_depth": 3, "max_agentic_loops": 3}, + api_surface="responses", + ) + + @pytest.mark.asyncio + async def test_non_dict_response_is_returned_untouched(self): + response = MagicMock() + + result = await _run_hooks( + self.handler, + self.callback, + kwargs={"_agentic_loop_depth": 3, "max_agentic_loops": 3}, + response=response, + ) + + assert result is response + + @pytest.mark.asyncio + async def test_converted_stream_gets_a_terminal_fake_stream(self): + """ + A converted stream is wrapped back into an Anthropic SSE stream here, the + same as every other return in this function, so a streaming client gets a + terminal stream rather than a bare dict. The interceptor turns the client's + stream into a non-streaming upstream call, so stream is False on this path + and the converted flag on the logging object is what marks it. + """ + result = await _run_hooks( + self.handler, + self.callback, + kwargs={"_agentic_loop_depth": 3, "max_agentic_loops": 3}, + converted_stream=True, + ) + + assert isinstance(result, FakeAnthropicMessagesStreamIterator) + assert result.response["stop_reason"] == "end_turn" + assert INTERNAL_TOOL_NAME not in _tool_use_names(result.response) + + def test_rails_cannot_trip_in_the_outermost_frame(self): + """ + Backs the invariant the test above relies on: at depth 0 the fingerprint set + is empty and the ceiling is at least 1, so neither rail can refuse. + """ + depth, max_loops, fingerprints = BaseLLMHTTPHandler._get_agentic_loop_settings(kwargs={}) + + assert depth == 0 + assert fingerprints == [] + assert max_loops >= 1 + + depth, max_loops, fingerprints = BaseLLMHTTPHandler._get_agentic_loop_settings( + kwargs={"max_agentic_loops": 1} + ) + + assert max_loops == 1 + assert BaseLLMHTTPHandler._check_agentic_loop_safety( + tool_calls={"tool_calls": [_internal_tool_use_block()]}, + fingerprints=fingerprints, + depth=depth, + max_loops=max_loops, + model="claude-sonnet-4-5", + ) + + def test_safety_error_is_still_a_value_error(self): + assert issubclass(AgenticLoopSafetyError, ValueError) + + def test_safety_error_type_names_the_rail(self): + with pytest.raises(AgenticLoopSafetyError, match="max_agentic_loops"): + BaseLLMHTTPHandler._check_agentic_loop_safety( + tool_calls={"tool_calls": [_internal_tool_use_block()]}, + fingerprints=[], + depth=3, + max_loops=3, + model="claude-sonnet-4-5", + ) + + +class TestOuterFramePostHookStillRuns: + """ + The cap used to raise through the parent frame's await, which skipped the + parent's post-loop hook. The parent now gets its terminal response back and + finishes normally, so the blocks it was going to inject still land. + """ + + @pytest.mark.asyncio + async def test_parent_frame_injects_its_blocks_after_the_cap_trips(self, monkeypatch): + handler = BaseLLMHTTPHandler() + callback = _InterceptingCallback() + + async def fake_acreate(**call_kwargs): + return await handler._call_agentic_completion_hooks( + response=_response_asking_for_another_search(block_id="toolu_internal_2"), + model=call_kwargs["model"], + messages=call_kwargs["messages"], + anthropic_messages_provider_config=MagicMock(), + anthropic_messages_optional_request_params={}, + logging_obj=_logging_obj(callback), + stream=False, + custom_llm_provider="anthropic", + kwargs={ + key: call_kwargs[key] + for key in ("_agentic_loop_depth", "max_agentic_loops", "_agentic_loop_fingerprints") + if key in call_kwargs + }, + ) + + monkeypatch.setattr("litellm.anthropic_interface.messages.acreate", fake_acreate) + + result = await _run_hooks( + handler, + callback, + kwargs={"_agentic_loop_depth": 0, "max_agentic_loops": 1}, + ) + + assert callback.plan_calls == 1 + assert callback.post_hook_calls == 1 + assert _block_types(result)[:2] == ["server_tool_use", "web_search_tool_result"] + assert INTERNAL_TOOL_NAME not in _tool_use_names(result) + assert result["stop_reason"] == "end_turn" + + +class TestMaxAgenticLoopsConfigKnob: + def test_from_config_yaml_reads_the_knob(self): + logger = WebSearchInterceptionLogger.from_config_yaml( + {"enabled_providers": ["bedrock"], "max_agentic_loops": 7} + ) + + assert logger.max_agentic_loops == 7 + + def test_from_config_yaml_leaves_it_unset_by_default(self): + logger = WebSearchInterceptionLogger.from_config_yaml({"enabled_providers": ["bedrock"]}) + + assert logger.max_agentic_loops is None + + @pytest.mark.parametrize("bad_value", [0, -1]) + def test_out_of_range_ceilings_are_rejected_at_config_load(self, bad_value): + with pytest.raises(ValueError, match="max_agentic_loops"): + WebSearchInterceptionLogger.from_config_yaml( + {"enabled_providers": ["bedrock"], "max_agentic_loops": bad_value} + ) + + @pytest.mark.parametrize("bad_value", ["three", True, 2.5]) + def test_non_integer_ceilings_are_rejected_at_config_load(self, bad_value): + with pytest.raises(TypeError, match="max_agentic_loops"): + WebSearchInterceptionLogger.from_config_yaml( + {"enabled_providers": ["bedrock"], "max_agentic_loops": bad_value} + ) + + def test_a_ceiling_spelled_as_a_string_is_read_at_config_load(self): + """ + `max_agentic_loops: os.environ/MAX_AGENTIC_LOOPS` resolves to a string + before it reaches the knob, so refusing "5" would break a config that + works today. + """ + logger = WebSearchInterceptionLogger.from_config_yaml( + {"enabled_providers": ["bedrock"], "max_agentic_loops": "5"} + ) + + assert logger.max_agentic_loops == 5 + + @pytest.mark.asyncio + async def test_knob_reaches_the_loop_settings(self): + logger = WebSearchInterceptionLogger.from_config_yaml( + {"enabled_providers": ["bedrock"], "max_agentic_loops": 7} + ) + kwargs = { + "tools": [{"type": "web_search_20250305", "name": "web_search", "max_uses": 5}], + "litellm_params": {"custom_llm_provider": "bedrock"}, + } + + updated = await logger.async_pre_request_hook(model="claude-sonnet-4-5", messages=[], kwargs=kwargs) + + _, max_loops, _ = BaseLLMHTTPHandler._get_agentic_loop_settings(kwargs=updated) + assert max_loops == 7 + + @pytest.mark.asyncio + async def test_deployment_setting_wins_over_the_feature_setting(self): + logger = WebSearchInterceptionLogger.from_config_yaml( + {"enabled_providers": ["bedrock"], "max_agentic_loops": 7} + ) + kwargs = { + "tools": [{"type": "web_search_20250305", "name": "web_search", "max_uses": 5}], + "litellm_params": {"custom_llm_provider": "bedrock"}, + "max_agentic_loops": 2, + } + + updated = await logger.async_pre_request_hook(model="claude-sonnet-4-5", messages=[], kwargs=kwargs) + + _, max_loops, _ = BaseLLMHTTPHandler._get_agentic_loop_settings(kwargs=updated) + assert max_loops == 2 + + @pytest.mark.asyncio + async def test_default_ceiling_applies_when_the_knob_is_unset(self): + logger = WebSearchInterceptionLogger.from_config_yaml({"enabled_providers": ["bedrock"]}) + kwargs = { + "tools": [{"type": "web_search_20250305", "name": "web_search", "max_uses": 5}], + "litellm_params": {"custom_llm_provider": "bedrock"}, + } + + updated = await logger.async_pre_request_hook(model="claude-sonnet-4-5", messages=[], kwargs=kwargs) + + assert "max_agentic_loops" not in updated + _, max_loops, _ = BaseLLMHTTPHandler._get_agentic_loop_settings(kwargs=updated) + assert max_loops == 3 + + +def _stream_events(response: dict) -> list[dict]: + events: list[dict] = [] + for chunk in FakeAnthropicMessagesStreamIterator(response=response): + for line in chunk.decode().splitlines(): + if line.startswith("data: "): + events.append(json.loads(line[len("data: ") :])) + return events + + +class TestBothCeilingKnobsAreValidated: + """ + ``max_agentic_loops`` is settable per deployment and feature-wide, and the + per-deployment one wins. Only the feature-wide one used to be checked, so a + per-deployment ``0`` was swallowed by an ``or 3`` and read as the default 3, + handing the loosest ceiling to whoever asked for the tightest. + """ + + def test_a_per_deployment_zero_is_rejected_not_read_as_the_default(self): + with pytest.raises(ValueError, match="must be at least 1, got 0"): + BaseLLMHTTPHandler._get_agentic_loop_settings(kwargs={"max_agentic_loops": 0}) + + def test_a_per_deployment_non_integer_names_the_field_it_came_from(self): + with pytest.raises(TypeError, match=r"litellm_params\.max_agentic_loops must be an integer"): + BaseLLMHTTPHandler._get_agentic_loop_settings(kwargs={"max_agentic_loops": "three"}) + + def test_a_per_deployment_true_is_not_read_as_a_ceiling_of_one(self): + with pytest.raises(TypeError, match="must be an integer"): + BaseLLMHTTPHandler._get_agentic_loop_settings(kwargs={"max_agentic_loops": True}) + + def test_an_absent_ceiling_falls_back_to_the_shared_default(self): + _, max_loops, _ = BaseLLMHTTPHandler._get_agentic_loop_settings(kwargs={}) + + assert max_loops == DEFAULT_MAX_AGENTIC_LOOPS + + def test_an_explicit_none_falls_back_to_the_shared_default(self): + _, max_loops, _ = BaseLLMHTTPHandler._get_agentic_loop_settings(kwargs={"max_agentic_loops": None}) + + assert max_loops == DEFAULT_MAX_AGENTIC_LOOPS + + def test_a_valid_per_deployment_ceiling_is_passed_through(self): + _, max_loops, _ = BaseLLMHTTPHandler._get_agentic_loop_settings(kwargs={"max_agentic_loops": 6}) + + assert max_loops == 6 + + @pytest.mark.parametrize("rejected", [0, -1, "three", True]) + def test_the_two_knobs_reject_the_same_values(self, rejected): + with pytest.raises((TypeError, ValueError)): + WebSearchInterceptionLogger(max_agentic_loops=rejected) + with pytest.raises((TypeError, ValueError)): + BaseLLMHTTPHandler._get_agentic_loop_settings(kwargs={"max_agentic_loops": rejected}) + + def test_each_knob_names_its_own_config_field(self): + with pytest.raises(ValueError, match=r"websearch_interception_params\.max_agentic_loops"): + WebSearchInterceptionLogger(max_agentic_loops=0) + with pytest.raises(ValueError, match=r"litellm_params\.max_agentic_loops"): + BaseLLMHTTPHandler._get_agentic_loop_settings(kwargs={"max_agentic_loops": 0}) + + +class TestACeilingThatSpellsAWholeNumberStillWorks: + """ + The ceiling used to go through ``int(... or 3)``, which accepted anything + ``int()`` accepted. A ceiling is routinely parameterized as + ``max_agentic_loops: os.environ/MAX_AGENTIC_LOOPS``, and ``get_secret`` + hands that back as the string ``"5"``, so tightening the check to + ``isinstance(int)`` would stop such a proxy from booting on upgrade. + """ + + @pytest.mark.parametrize("spelled", ["5", " 5 ", 5.0]) + def test_a_ceiling_that_spells_five_is_accepted_by_both_knobs(self, spelled): + _, max_loops, _ = BaseLLMHTTPHandler._get_agentic_loop_settings(kwargs={"max_agentic_loops": spelled}) + + assert max_loops == 5 + assert WebSearchInterceptionLogger(max_agentic_loops=spelled).max_agentic_loops == 5 + + def test_an_env_var_sourced_ceiling_survives_secret_resolution(self, monkeypatch): + monkeypatch.setenv("MAX_AGENTIC_LOOPS_UNDER_TEST", "7") + resolved = get_secret("os.environ/MAX_AGENTIC_LOOPS_UNDER_TEST") + + assert isinstance(resolved, str) + _, max_loops, _ = BaseLLMHTTPHandler._get_agentic_loop_settings(kwargs={"max_agentic_loops": resolved}) + assert max_loops == 7 + + def test_a_spelled_zero_is_still_refused_and_reports_the_number(self): + with pytest.raises(ValueError, match="must be at least 1, got 0"): + BaseLLMHTTPHandler._get_agentic_loop_settings(kwargs={"max_agentic_loops": "0"}) + + def test_a_word_is_still_refused(self): + with pytest.raises(TypeError, match=r"litellm_params\.max_agentic_loops must be an integer"): + BaseLLMHTTPHandler._get_agentic_loop_settings(kwargs={"max_agentic_loops": "three"}) + + def test_a_fractional_ceiling_is_refused_rather_than_truncated(self): + with pytest.raises(TypeError, match="must be an integer"): + BaseLLMHTTPHandler._get_agentic_loop_settings(kwargs={"max_agentic_loops": 5.5}) + + +class TestRebuiltStreamIsWellFormed: + """ + A capped turn is rebuilt into SSE by FakeAnthropicMessagesStreamIterator. + + Anthropic's SDK accumulator appends on content_block_start and then indexes + content[event.index] on content_block_delta, so a block that stops without + ever starting shifts every later index and the accumulator raises + IndexError. A web search turn carries server_tool_use and + web_search_tool_result blocks, which is exactly where that used to happen. + """ + + @staticmethod + def _capped_search_turn() -> dict: + return { + "id": "msg_01", + "type": "message", + "role": "assistant", + "model": "claude-sonnet-4-5", + "stop_reason": "end_turn", + "content": [ + { + "type": "server_tool_use", + "id": "srvtoolu_01", + "name": "web_search", + "input": {"query": "on-demand H100 hourly price"}, + }, + { + "type": "web_search_tool_result", + "tool_use_id": "srvtoolu_01", + "content": [ + { + "type": "web_search_result", + "url": "https://example.com/h100", + "title": "H100 pricing", + } + ], + }, + {"type": "text", "text": "AWS lists the H100 at $12.29 an hour."}, + ], + "usage": {"input_tokens": 100, "output_tokens": 20}, + } + + def test_every_content_block_stop_has_a_matching_start(self): + events = _stream_events(self._capped_search_turn()) + + started = [event["index"] for event in events if event["type"] == "content_block_start"] + stopped = [event["index"] for event in events if event["type"] == "content_block_stop"] + + assert started == [0, 1, 2] + assert stopped == [0, 1, 2] + + def test_no_delta_indexes_past_the_blocks_started_before_it(self): + events = _stream_events(self._capped_search_turn()) + + blocks_started = 0 + for event in events: + if event["type"] == "content_block_start": + blocks_started += 1 + elif event["type"] == "content_block_delta": + assert event["index"] < blocks_started + + def test_search_blocks_reach_the_client(self): + events = _stream_events(self._capped_search_turn()) + + started_types = [ + event["content_block"]["type"] for event in events if event["type"] == "content_block_start" + ] + + assert started_types == ["server_tool_use", "web_search_tool_result", "text"] + + def test_the_search_result_survives_the_rebuild_intact(self): + events = _stream_events(self._capped_search_turn()) + + result_block = next( + event["content_block"] + for event in events + if event["type"] == "content_block_start" + and event["content_block"]["type"] == "web_search_tool_result" + ) + + assert result_block["tool_use_id"] == "srvtoolu_01" + assert result_block["content"][0]["url"] == "https://example.com/h100" diff --git a/tests/test_litellm/llms/openai/image_generation/test_openai_image_generation_extra_headers.py b/tests/test_litellm/llms/openai/image_generation/test_openai_image_generation_extra_headers.py index 06871edb773..55ef74abd7b 100644 --- a/tests/test_litellm/llms/openai/image_generation/test_openai_image_generation_extra_headers.py +++ b/tests/test_litellm/llms/openai/image_generation/test_openai_image_generation_extra_headers.py @@ -148,6 +148,58 @@ class TestImageGenerationExtraHeaders: _, kwargs = mock_openai_client.images.generate.call_args assert "extra_headers" not in kwargs + @pytest.mark.parametrize("is_async", [False, True]) + @pytest.mark.asyncio + async def test_caller_headers_never_reach_the_logged_request_body( + self, openai_chat_completions, mock_logging_obj, is_async + ): + """The body handed to pre_call is also what telemetry reads at close time, so + merging caller headers into that same dict would publish a customer's auth + header as a span attribute. The upstream call still gets them.""" + mock_image_data = MagicMock() + mock_image_data.model_dump.return_value = { + "created": 1700000000, + "data": [{"url": "https://example.com/image.png"}], + } + + mock_openai_client = MagicMock() + mock_openai_client.api_key = "test-key" + mock_openai_client._base_url._uri_reference = "https://api.openai.com" + + test_headers = {"cf-aig-authorization": "Bearer custom-token"} + + if is_async: + mock_openai_client.images.generate = AsyncMock(return_value=mock_image_data) + await openai_chat_completions.aimage_generation( + prompt="A white cat", + data={"model": "dall-e-3", "prompt": "A white cat"}, + model_response=MagicMock(), + timeout=60.0, + logging_obj=mock_logging_obj, + api_key="test-key", + headers=test_headers, + client=mock_openai_client, + ) + else: + mock_openai_client.images.generate.return_value = mock_image_data + openai_chat_completions.image_generation( + model="dall-e-3", + prompt="A white cat", + timeout=60.0, + optional_params={}, + logging_obj=mock_logging_obj, + api_key="test-key", + headers=test_headers, + client=mock_openai_client, + ) + + logged_body = mock_logging_obj.pre_call.call_args[1]["additional_args"][ + "complete_input_dict" + ] + assert "extra_headers" not in logged_body + _, kwargs = mock_openai_client.images.generate.call_args + assert kwargs.get("extra_headers") == test_headers + def test_sync_image_generation_forwards_headers_to_async( self, openai_chat_completions, mock_logging_obj ): diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index a34df54adfa..04f38b5e2ed 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -2455,6 +2455,53 @@ async def test_get_team_object_raises_404_when_not_found(): assert "Team doesn't exist in db" in str(exc_info.value.detail) +def _mock_prisma_for_team_lookup(find_unique): + from unittest.mock import MagicMock + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_teamtable.find_unique = find_unique + return mock_prisma_client + + +@pytest.mark.asyncio +async def test_get_team_object_distinguishes_absent_team_from_unreadable_row(): + """A deleted team and a database that would not answer both surface as a 404, + which leaves callers unable to tell a definitive answer from a degraded read. + Only the row being positively absent raises the subclass; anything else keeps + the plain 404 so every existing caller is unaffected.""" + from unittest.mock import AsyncMock, MagicMock + + from fastapi import HTTPException + + from litellm.proxy.auth.auth_checks import TeamNotFoundError, get_team_object + + mock_cache = MagicMock() + mock_cache.async_get_cache = AsyncMock(return_value=None) + + # The database answered, and the row is not there. + with pytest.raises(TeamNotFoundError) as absent_info: + await get_team_object( + team_id="absent-team-lit5522", + prisma_client=_mock_prisma_for_team_lookup(AsyncMock(return_value=None)), + user_api_key_cache=mock_cache, + check_db_only=True, + ) + assert absent_info.value.status_code == 404 + assert "Team doesn't exist in db" in str(absent_info.value.detail) + + # The database did not answer. Same status and detail, but not the subclass, + # so a caller keying on it does not read this as proof the team is gone. + with pytest.raises(HTTPException) as unreadable_info: + await get_team_object( + team_id="unreadable-team-lit5522", + prisma_client=_mock_prisma_for_team_lookup(AsyncMock(side_effect=ConnectionError("db unreachable"))), + user_api_key_cache=mock_cache, + check_db_only=True, + ) + assert unreadable_info.value.status_code == 404 + assert not isinstance(unreadable_info.value, TeamNotFoundError) + + # Reject Client-Side Metadata Tags Tests diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py index c1e235b77f6..6a117985820 100644 --- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py +++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py @@ -29,6 +29,8 @@ from litellm.proxy.auth.auth_checks import get_key_object, _cache_key_object from litellm.proxy.auth.route_checks import RouteChecks from litellm.proxy.auth.user_api_key_auth import ( _check_key_model_budget_with_fallback, + _ensure_litellm_received_at_on_request_state, + _ensure_parent_otel_span_on_request_state, _PendingAutoRegister, _matches_routing_override, _reserve_budget_after_common_checks, @@ -3575,6 +3577,201 @@ async def test_auth_flow_never_persists_fallback_team_object_lit_4391(): setattr(_proxy_server_mod, k, v) +@pytest.mark.asyncio +async def test_auth_flow_fallback_team_resolves_object_permission_by_id(): + """The unresolvable-team fallback resolves team_object_permission by its own id instead of leaving it unset.""" + from starlette.datastructures import URL + from starlette.requests import Request + from fastapi import HTTPException + + from litellm.proxy._types import ( + LiteLLM_ObjectPermissionTable, + LitellmUserRoles, + UserAPIKeyAuth, + ) + from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder + + api_key = "sk-test-fallback-team-object-permission" + valid_token = UserAPIKeyAuth( + api_key=api_key, + token=api_key, + user_role=LitellmUserRoles.INTERNAL_USER, + team_id="team-fallback-object-permission", + team_object_permission_id="op-fallback-object-permission", + ) + + restricted_object_permission = LiteLLM_ObjectPermissionTable( + object_permission_id="op-fallback-object-permission", + vector_stores=["vs-allowed-only"], + mcp_servers=["mcp-allowed-only"], + ) + + mock_cache = AsyncMock() + mock_cache.async_get_cache = AsyncMock(return_value=valid_token) + mock_cache.async_set_cache = AsyncMock(return_value=None) + + mock_proxy_logging_obj = MagicMock() + mock_proxy_logging_obj.internal_usage_cache = MagicMock() + mock_proxy_logging_obj.internal_usage_cache.dual_cache = AsyncMock() + mock_proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None) + + import litellm.proxy.proxy_server as _proxy_server_mod + + _attrs = { + "prisma_client": MagicMock(), + "user_api_key_cache": mock_cache, + "proxy_logging_obj": mock_proxy_logging_obj, + "master_key": "sk-master-key", + "general_settings": {}, + "llm_model_list": [], + "llm_router": None, + "open_telemetry_logger": None, + "model_max_budget_limiter": MagicMock(), + "user_custom_auth": None, + "jwt_handler": None, + "litellm_proxy_admin_name": "admin", + } + _originals = {k: getattr(_proxy_server_mod, k, None) for k in _attrs} + + try: + for k, v in _attrs.items(): + setattr(_proxy_server_mod, k, v) + + request = Request(scope={"type": "http"}) + request._url = URL(url="/chat/completions") + + with ( + patch( + "litellm.proxy.auth.resolvers.store.IdentityStore._resolve_key", + new_callable=AsyncMock, + return_value=valid_token, + ), + patch( + "litellm.proxy.auth.user_api_key_auth.get_team_object", + new_callable=AsyncMock, + side_effect=HTTPException( + status_code=404, + detail={"error": "Team doesn't exist in db."}, + ), + ), + patch( + "litellm.proxy.auth.user_api_key_auth.get_object_permission", + new_callable=AsyncMock, + return_value=restricted_object_permission, + ) as mock_get_object_permission, + ): + result = await _user_api_key_auth_builder( + request=request, + api_key=f"Bearer {api_key}", + azure_api_key_header="", + anthropic_api_key_header=None, + google_ai_studio_api_key_header=None, + azure_apim_header=None, + request_data={}, + ) + + mock_get_object_permission.assert_awaited_once() + assert mock_get_object_permission.await_args.kwargs["object_permission_id"] == "op-fallback-object-permission" + assert result.team_object_permission == restricted_object_permission + assert result.team_object_permission.vector_stores == ["vs-allowed-only"] + assert result.team_object_permission.mcp_servers == ["mcp-allowed-only"] + + finally: + for k, v in _originals.items(): + setattr(_proxy_server_mod, k, v) + + +@pytest.mark.asyncio +async def test_auth_flow_fallback_team_object_permission_none_when_unreadable(): + """When the object_permission row is also unreadable, the fallback leaves team_object_permission as None + instead of raising or fabricating a grant.""" + from starlette.datastructures import URL + from starlette.requests import Request + from fastapi import HTTPException + + from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder + + api_key = "sk-test-fallback-team-object-permission-unreadable" + valid_token = UserAPIKeyAuth( + api_key=api_key, + token=api_key, + user_role=LitellmUserRoles.INTERNAL_USER, + team_id="team-fallback-object-permission-unreadable", + team_object_permission_id="op-fallback-object-permission-unreadable", + ) + + mock_cache = AsyncMock() + mock_cache.async_get_cache = AsyncMock(return_value=valid_token) + mock_cache.async_set_cache = AsyncMock(return_value=None) + + mock_proxy_logging_obj = MagicMock() + mock_proxy_logging_obj.internal_usage_cache = MagicMock() + mock_proxy_logging_obj.internal_usage_cache.dual_cache = AsyncMock() + mock_proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None) + + import litellm.proxy.proxy_server as _proxy_server_mod + + _attrs = { + "prisma_client": MagicMock(), + "user_api_key_cache": mock_cache, + "proxy_logging_obj": mock_proxy_logging_obj, + "master_key": "sk-master-key", + "general_settings": {}, + "llm_model_list": [], + "llm_router": None, + "open_telemetry_logger": None, + "model_max_budget_limiter": MagicMock(), + "user_custom_auth": None, + "jwt_handler": None, + "litellm_proxy_admin_name": "admin", + } + _originals = {k: getattr(_proxy_server_mod, k, None) for k in _attrs} + + try: + for k, v in _attrs.items(): + setattr(_proxy_server_mod, k, v) + + request = Request(scope={"type": "http"}) + request._url = URL(url="/chat/completions") + + with ( + patch( + "litellm.proxy.auth.resolvers.store.IdentityStore._resolve_key", + new_callable=AsyncMock, + return_value=valid_token, + ), + patch( + "litellm.proxy.auth.user_api_key_auth.get_team_object", + new_callable=AsyncMock, + side_effect=HTTPException( + status_code=404, + detail={"error": "Team doesn't exist in db."}, + ), + ), + patch( + "litellm.proxy.auth.user_api_key_auth.get_object_permission", + new_callable=AsyncMock, + return_value=None, + ), + ): + result = await _user_api_key_auth_builder( + request=request, + api_key=f"Bearer {api_key}", + azure_api_key_header="", + anthropic_api_key_header=None, + google_ai_studio_api_key_header=None, + azure_apim_header=None, + request_data={}, + ) + + assert result.team_object_permission is None + + finally: + for k, v in _originals.items(): + setattr(_proxy_server_mod, k, v) + + # --------------------------------------------------------------------------- # _run_centralized_common_checks — centralized authz gate @@ -4498,6 +4695,283 @@ async def test_centralized_common_checks_team_404_does_not_zero_other_contexts() setattr(_proxy_server_mod, k, v) +@pytest.mark.asyncio +async def test_centralized_common_checks_unresolvable_team_without_grant_is_refused(): + """The store restricts the team to gpt-4o-mini and the read of it fails, so the + only surviving team record is the token's own, which carries ``team_models=[]`` + and reads as every model. The request must be refused with the original lookup + error. Pre-fix it was served.""" + import litellm.proxy.proxy_server as _proxy_server_mod + from fastapi import HTTPException, Request + from starlette.datastructures import URL + + token = UserAPIKeyAuth( + api_key="sk-test", + team_id="restricted-team", + models=[], + team_models=[], + ) + request = Request(scope={"type": "http"}) + request._url = URL(url="/chat/completions") + request._body = json.dumps({"model": "gpt-4.1"}).encode() + + team_read_failure = HTTPException( + status_code=404, + detail={"error": "Team doesn't exist in db. Team=restricted-team."}, + ) + + attrs = _proxy_attrs_for_centralized_checks(user_custom_auth=None) + originals = {a: getattr(_proxy_server_mod, a, None) for a in attrs} + try: + for k, v in attrs.items(): + setattr(_proxy_server_mod, k, v) + with patch( + "litellm.proxy.auth.user_api_key_auth.get_team_object", + new_callable=AsyncMock, + side_effect=team_read_failure, + ): + with pytest.raises(HTTPException) as exc_info: + await _run_centralized_common_checks( + user_api_key_auth_obj=token, + request=request, + request_data={"model": "gpt-4.1"}, + route="/chat/completions", + ) + assert exc_info.value is team_read_failure + finally: + for k, v in originals.items(): + setattr(_proxy_server_mod, k, v) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("token_team_models", [[], ["gpt-4.1"]]) +async def test_centralized_common_checks_absent_team_refused_despite_db_unavailable_optout(token_team_models): + """A team that is provably gone is a definitive answer, not a degraded read. + ``allow_requests_on_db_unavailable`` is a static settings read, so without the + absent-versus-unreadable distinction it would hand a deleted team's key the + old permissive fallback while the database is perfectly healthy. Refused in + both token shapes, including the one whose grant would otherwise vouch. + + Imported from the module under test rather than from ``auth_checks``: other + tests in this suite ``importlib.reload`` that module, which rebinds the class + and would leave this raising a type the guard has never seen.""" + import litellm.proxy.proxy_server as _proxy_server_mod + from fastapi import HTTPException, Request + from starlette.datastructures import URL + + from litellm.proxy.auth.user_api_key_auth import TeamNotFoundError + + token = UserAPIKeyAuth( + api_key="sk-test", + team_id="deleted-team", + models=[], + team_models=token_team_models, + ) + request = Request(scope={"type": "http"}) + request._url = URL(url="/chat/completions") + request._body = json.dumps({"model": "gpt-4.1"}).encode() + + team_absent = TeamNotFoundError(team_id="deleted-team") + + attrs = _proxy_attrs_for_centralized_checks(user_custom_auth=None) + attrs["general_settings"] = {"allow_requests_on_db_unavailable": True} + originals = {a: getattr(_proxy_server_mod, a, None) for a in attrs} + try: + for k, v in attrs.items(): + setattr(_proxy_server_mod, k, v) + with patch( + "litellm.proxy.auth.user_api_key_auth.get_team_object", + new_callable=AsyncMock, + side_effect=team_absent, + ): + with pytest.raises(HTTPException) as exc_info: + await _run_centralized_common_checks( + user_api_key_auth_obj=token, + request=request, + request_data={"model": "gpt-4.1"}, + route="/chat/completions", + ) + assert exc_info.value is team_absent + finally: + for k, v in originals.items(): + setattr(_proxy_server_mod, k, v) + + +@pytest.mark.asyncio +async def test_centralized_common_checks_unreadable_team_keeps_db_unavailable_optout(): + """The counterpart: an unreadable team leaves the grant unknown rather than + answered, so an operator who has accepted degraded authorization during a + database fault still gets the fallback. Without this the fix would trade the + widening for a lockout with no way out.""" + import litellm.proxy.proxy_server as _proxy_server_mod + from fastapi import HTTPException as _HTTPException + from fastapi import Request + from starlette.datastructures import URL + + from litellm.proxy._types import LiteLLM_TeamTableCachedObj + + token = UserAPIKeyAuth(api_key="sk-test", team_id="unreadable-team", models=[], team_models=[]) + request = Request(scope={"type": "http"}) + request._url = URL(url="/chat/completions") + request._body = json.dumps({"model": "gpt-4.1"}).encode() + + received_team_objects: list[LiteLLM_TeamTableCachedObj | None] = [] + + async def _capturing_common_checks(*_args, **kwargs) -> bool: + received_team_objects.append(kwargs.get("team_object")) + return True + + attrs = _proxy_attrs_for_centralized_checks(user_custom_auth=None) + attrs["general_settings"] = {"allow_requests_on_db_unavailable": True} + originals = {a: getattr(_proxy_server_mod, a, None) for a in attrs} + try: + for k, v in attrs.items(): + setattr(_proxy_server_mod, k, v) + with ( + patch( + "litellm.proxy.auth.user_api_key_auth.get_team_object", + new_callable=AsyncMock, + side_effect=_HTTPException(status_code=404, detail={"error": "team unreadable"}), + ), + patch( + "litellm.proxy.auth.user_api_key_auth.common_checks", + _capturing_common_checks, + ), + ): + await _run_centralized_common_checks( + user_api_key_auth_obj=token, + request=request, + request_data={"model": "gpt-4.1"}, + route="/chat/completions", + ) + assert len(received_team_objects) == 1 + received_team_object = received_team_objects[0] + assert received_team_object is not None + assert received_team_object.team_id == "unreadable-team" + finally: + for k, v in originals.items(): + setattr(_proxy_server_mod, k, v) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "requested_model, is_granted", + [("gpt-4o-mini", True), ("gpt-4.1", False)], +) +async def test_centralized_common_checks_unresolvable_team_with_grant_enforces_it(requested_model, is_granted): + """Mirror of the refusal above: a token that does carry a team model grant keeps + the fallback, and the reconstructed team must still enforce that grant rather + than wave the request through.""" + import litellm.proxy.proxy_server as _proxy_server_mod + from fastapi import HTTPException, Request + from starlette.datastructures import URL + + from litellm.proxy._types import ProxyErrorTypes, ProxyException + + token = UserAPIKeyAuth( + api_key="sk-test", + team_id="restricted-team", + models=[], + team_models=["gpt-4o-mini"], + ) + request = Request(scope={"type": "http"}) + request._url = URL(url="/chat/completions") + request._body = json.dumps({"model": requested_model}).encode() + + attrs = _proxy_attrs_for_centralized_checks(user_custom_auth=None) + originals = {a: getattr(_proxy_server_mod, a, None) for a in attrs} + try: + for k, v in attrs.items(): + setattr(_proxy_server_mod, k, v) + with patch( + "litellm.proxy.auth.user_api_key_auth.get_team_object", + new_callable=AsyncMock, + side_effect=HTTPException(status_code=404, detail={"error": "team unreadable"}), + ): + if is_granted: + await _run_centralized_common_checks( + user_api_key_auth_obj=token, + request=request, + request_data={"model": requested_model}, + route="/chat/completions", + ) + else: + with pytest.raises(ProxyException) as exc_info: + await _run_centralized_common_checks( + user_api_key_auth_obj=token, + request=request, + request_data={"model": requested_model}, + route="/chat/completions", + ) + assert exc_info.value.type == ProxyErrorTypes.team_model_access_denied + finally: + for k, v in originals.items(): + setattr(_proxy_server_mod, k, v) + + +@pytest.mark.asyncio +async def test_centralized_common_checks_ui_sentinel_team_vouches_despite_absent_row(): + """The Admin UI mints every session key against the ``UI_TEAM_ID`` sentinel, + which by design never has a ``LiteLLM_TeamTable`` row, so ``get_team_object`` + always raises ``TeamNotFoundError`` for it. That must NOT be read as "team + provably gone, refuse" the way it is for a real team_id: PR #36837 made that + exact mistake and PR #36982 reverted it because every dashboard request + 404'd. The sentinel must keep vouching from the token unconditionally.""" + import litellm.proxy.proxy_server as _proxy_server_mod + from fastapi import Request + from starlette.datastructures import URL + + from litellm.proxy._types import UI_TEAM_ID, LiteLLM_TeamTableCachedObj + from litellm.proxy.auth.user_api_key_auth import TeamNotFoundError + + token = UserAPIKeyAuth( + api_key="sk-test", + user_id="ui-session-user", + team_id=UI_TEAM_ID, + models=[], + team_models=[], + ) + request = Request(scope={"type": "http"}) + request._url = URL(url="/user/info") + request._body = b"{}" + + received_team_objects: list[LiteLLM_TeamTableCachedObj | None] = [] + + async def _capturing_common_checks(*_args, **kwargs) -> bool: + received_team_objects.append(kwargs.get("team_object")) + return True + + attrs = _proxy_attrs_for_centralized_checks(user_custom_auth=None) + originals = {a: getattr(_proxy_server_mod, a, None) for a in attrs} + try: + for k, v in attrs.items(): + setattr(_proxy_server_mod, k, v) + with ( + patch( + "litellm.proxy.auth.user_api_key_auth.get_team_object", + new_callable=AsyncMock, + side_effect=TeamNotFoundError(team_id=UI_TEAM_ID), + ), + patch( + "litellm.proxy.auth.user_api_key_auth.common_checks", + _capturing_common_checks, + ), + ): + await _run_centralized_common_checks( + user_api_key_auth_obj=token, + request=request, + request_data={}, + route="/user/info", + ) + assert len(received_team_objects) == 1 + received_team_object = received_team_objects[0] + assert received_team_object is not None + assert received_team_object.team_id == UI_TEAM_ID + finally: + for k, v in originals.items(): + setattr(_proxy_server_mod, k, v) + + @pytest.mark.asyncio async def test_centralized_common_checks_user_http_exception_isolates_to_user_only(): """Per-fetch isolation, mirror of the team case: an HTTPException @@ -6222,3 +6696,42 @@ async def test_unlicensed_jwt_auth_is_forbidden_not_unauthorized(): assert error.code == "403" assert "enterprise" in error.message.lower() + + +class TestLitellmReceivedAtStamping: + """request.state.litellm_received_at must be stamped unconditionally at the + top of auth (LIT-6012), so request-latency Prometheus metrics don't depend + on OTEL being configured to see a true request-arrival timestamp.""" + + def test_stamped_even_when_otel_is_not_configured(self, monkeypatch): + monkeypatch.setattr( + "litellm.proxy.proxy_server.open_telemetry_logger", None + ) + request = MagicMock() + request.state = SimpleNamespace() + + _ensure_parent_otel_span_on_request_state(request) + + assert isinstance(request.state.litellm_received_at, datetime) + + def test_helper_is_idempotent(self): + request = MagicMock() + request.state = SimpleNamespace() + + first = _ensure_litellm_received_at_on_request_state(request) + second = _ensure_litellm_received_at_on_request_state(request) + + assert first == second + assert request.state.litellm_received_at == first + + def test_does_not_overwrite_an_earlier_stamp(self): + """Body-parse failures must not shorten the measured window: a value + already on request.state (stamped earlier) must win.""" + request = MagicMock() + earlier = datetime(2020, 1, 1) + request.state = SimpleNamespace(litellm_received_at=earlier) + + result = _ensure_litellm_received_at_on_request_state(request) + + assert result == earlier + assert request.state.litellm_received_at == earlier diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_presidio.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_presidio.py index 60be3be5e8b..acb43bc5b74 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_presidio.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_presidio.py @@ -582,6 +582,52 @@ async def test_logging_hook_multiple_content_items(presidio_guardrail): print("✓ Logging hook multiple content items test passed") +@pytest.mark.asyncio +async def test_logging_hook_masks_the_response_too(presidio_guardrail): + """ + Regression: async_logging_hook only masked kwargs["messages"] (the request) and + left `result` (the model's response) completely untouched, so in `logging_only` + mode any PII in the assistant's reply was logged to langfuse/datadog/etc. in the + clear. The hook's own docstring promises masking "before logging" for both input + and output. + """ + + async def mock_check_pii(text, output_parse_pii, presidio_config, request_data): + return text.replace("4111-1111-1111-1111", "[CREDIT_CARD]") + + presidio_guardrail.check_pii = mock_check_pii + + test_kwargs = { + "messages": [{"role": "user", "content": "hello"}], + "model": "gpt-4", + } + response = ModelResponse( + id="1", + object="chat.completion", + created=0, + model="gpt-test", + choices=[ + Choices( + message=Message( + role="assistant", + content="Sure, your card is 4111-1111-1111-1111", + ), + index=0, + finish_reason="stop", + ) + ], + ) + + _, result_response = await presidio_guardrail.async_logging_hook( + kwargs=test_kwargs, + result=response, + call_type="completion", + ) + + assert "[CREDIT_CARD]" in result_response.choices[0].message.content + assert "4111-1111-1111-1111" not in result_response.choices[0].message.content + + @pytest.mark.asyncio async def test_logging_only_does_not_mask_pre_call_request( mock_user_api_key, mock_cache diff --git a/tests/test_litellm/proxy/guardrails/test_content_utils.py b/tests/test_litellm/proxy/guardrails/test_content_utils.py index 3dfb98c12ea..d9e079c6d92 100644 --- a/tests/test_litellm/proxy/guardrails/test_content_utils.py +++ b/tests/test_litellm/proxy/guardrails/test_content_utils.py @@ -149,6 +149,22 @@ def test_iter_message_text_responses_api_tool_call_taxonomy(): assert list(iter_message_text(data)) == ["hello", "sunny"] +def test_iter_message_text_inspects_reasoning_content_and_summary(): + """VERIA: reasoning items forwarded as ``reasoning_content`` must be + inspected, including ``summary`` blocks the bridge reads as a fallback.""" + data = { + "input": [ + { + "type": "reasoning", + "id": "rs_1", + "content": [{"type": "summary_text", "text": "content secret"}], + "summary": [{"type": "summary_text", "text": "summary secret"}], + } + ] + } + assert list(iter_message_text(data)) == ["content secret", "summary secret"] + + # ── walk_user_text ──────────────────────────────────────────────────────────── @@ -308,6 +324,27 @@ def test_walk_user_text_redacts_mixed_list_input(): assert data["input"][2] == {"type": "image_url", "image_url": {"url": "..."}} +def test_walk_user_text_redacts_reasoning_content_and_summary(): + """VERIA: in-place redaction must cover both plaintext shapes the bridge + forwards from a reasoning item.""" + data = { + "input": [ + { + "type": "reasoning", + "id": "rs_1", + "content": [{"type": "summary_text", "text": "AKIAEXAMPLE content"}], + "summary": [{"type": "summary_text", "text": "AKIAEXAMPLE summary"}], + } + ] + } + visited = walk_user_text(data, lambda s: s.replace("AKIAEXAMPLE", "[REDACTED]")) + assert visited == 2 + item = data["input"][0] + assert item["content"][0]["text"] == "[REDACTED] content" + assert item["summary"][0]["text"] == "[REDACTED] summary" + assert item["id"] == "rs_1" + + # ── build_inspection_messages ───────────────────────────────────────────────── @@ -462,6 +499,23 @@ def test_build_inspection_messages_empty_data(): assert build_inspection_messages({"input": ""}) == [] +def test_build_inspection_messages_includes_reasoning_summary(): + """VERIA: remote guardrail APIs must see reasoning summaries even when + the reasoning item has no ``content`` field.""" + data = { + "input": [ + { + "type": "reasoning", + "id": "rs_1", + "summary": [{"type": "summary_text", "text": "secret summary"}], + } + ] + } + assert build_inspection_messages(data) == [ + {"role": "assistant", "content": "secret summary"} + ] + + # ── has_non_string_content ──────────────────────────────────────────────────── diff --git a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py index 34b12aecfed..f6d74a189bc 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py @@ -80,6 +80,23 @@ def _wire_team_create_tx(prisma_client): prisma_client.db.tx = lambda *_args, **_kwargs: _tx() +def _wire_member_delete_tx(prisma_client): + """/team/member_delete's four cleanups run inside one transaction, so a mocked + client has to hand back its own table mocks out of `tx()` for the existing + per-table assertions to keep seeing the calls.""" + tx = SimpleNamespace( + litellm_teamtable=prisma_client.db.litellm_teamtable, + litellm_usertable=prisma_client.db.litellm_usertable, + litellm_teammembership=prisma_client.db.litellm_teammembership, + litellm_verificationtoken=prisma_client.db.litellm_verificationtoken, + litellm_deletedverificationtoken=prisma_client.db.litellm_deletedverificationtoken, + ) + tx_cm = MagicMock() + tx_cm.__aenter__ = AsyncMock(return_value=tx) + tx_cm.__aexit__ = AsyncMock(return_value=None) + prisma_client.tx = MagicMock(return_value=tx_cm) + + # Mock prisma_client mock_prisma_client = MagicMock() # Set up async mock for db operations @@ -4147,6 +4164,8 @@ async def test_team_member_delete_cleans_membership(mock_db_client, mock_admin_a return_value=MagicMock() ) + _wire_member_delete_tx(mock_db_client) + # Execute await team_member_delete( data=TeamMemberDeleteRequest(team_id=test_team_id, user_id=test_user_id), @@ -4205,6 +4224,8 @@ async def test_team_member_delete_cleans_verification_tokens( return_value=MagicMock() ) + _wire_member_delete_tx(mock_db_client) + await team_member_delete( data=TeamMemberDeleteRequest(team_id=test_team_id, user_id=test_user_id), user_api_key_dict=mock_admin_auth, @@ -4300,6 +4321,8 @@ async def test_team_member_delete_by_email_the_user_row_does_not_carry( return_value=MagicMock() ) + _wire_member_delete_tx(mock_db_client) + await team_member_delete( data=TeamMemberDeleteRequest(team_id=test_team_id, user_email=roster_email), user_api_key_dict=mock_admin_auth, @@ -4318,6 +4341,83 @@ async def test_team_member_delete_by_email_the_user_row_does_not_carry( ) +class _InjectedMemberDeleteFailure(Exception): + pass + + +@pytest.mark.asyncio +async def test_team_member_delete_is_atomic_across_its_four_writes( + mock_db_client, mock_admin_auth +): + """ + /team/member_delete's four cleanups (team roster, user.teams, team + membership, verification tokens) run as one transaction, so a failure + partway through must not leave the removal half applied. + + Failing the second write (the user's ``teams`` update) pins two things a + non-transactional implementation gets wrong: the roster write that already + ran has to land on the SAME transaction client the failure raises on (so a + real database rolls it back too), and the writes still queued behind the + failure (membership delete, token delete) must never be attempted at all. + """ + from litellm.proxy._types import TeamMemberDeleteRequest + from litellm.proxy.management_endpoints.team_endpoints import team_member_delete + + test_team_id = "team-del-atomic-123" + test_user_id = "user-atomic@example.com" + + mock_team_row = MagicMock() + mock_team_row.model_dump.return_value = { + "team_id": test_team_id, + "members_with_roles": [ + {"user_id": test_user_id, "user_email": None, "role": "user"} + ], + "team_member_permissions": [], + "metadata": {}, + "models": [], + "spend": 0.0, + } + mock_db_client.db.litellm_teamtable.find_unique = AsyncMock( + return_value=mock_team_row + ) + mock_db_client.db.litellm_teamtable.update = AsyncMock(return_value=mock_team_row) + + mock_user_row = MagicMock() + mock_user_row.user_id = test_user_id + mock_user_row.teams = [test_team_id] + mock_db_client.db.litellm_usertable.find_many = AsyncMock( + return_value=[mock_user_row] + ) + mock_db_client.db.litellm_usertable.update = AsyncMock( + side_effect=_InjectedMemberDeleteFailure("boom between writes 1 and 2") + ) + + mock_db_client.db.litellm_teammembership = MagicMock() + mock_db_client.db.litellm_teammembership.delete_many = AsyncMock() + + mock_db_client.db.litellm_verificationtoken = MagicMock() + mock_db_client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) + mock_db_client.db.litellm_verificationtoken.delete_many = AsyncMock() + + _wire_member_delete_tx(mock_db_client) + + with pytest.raises(_InjectedMemberDeleteFailure): + await team_member_delete( + data=TeamMemberDeleteRequest(team_id=test_team_id, user_id=test_user_id), + user_api_key_dict=mock_admin_auth, + ) + + # The roster write ran, but on the transaction the injected failure also raised on. + mock_db_client.db.litellm_teamtable.update.assert_awaited_once() + mock_db_client.tx.assert_called_once() + aexit_args = mock_db_client.tx.return_value.__aexit__.await_args.args + assert aexit_args[0] is _InjectedMemberDeleteFailure + + # Writes queued behind the failure inside that same transaction never ran. + mock_db_client.db.litellm_teammembership.delete_many.assert_not_awaited() + mock_db_client.db.litellm_verificationtoken.delete_many.assert_not_awaited() + + @pytest.mark.asyncio async def test_new_team_max_budget_exceeds_user_max_budget(): """ @@ -7806,6 +7906,8 @@ async def test_team_member_delete_persists_deleted_keys(monkeypatch): mock_create_many_keys ) + _wire_member_delete_tx(mock_prisma_client) + monkeypatch.setattr( "litellm.proxy.proxy_server.prisma_client", mock_prisma_client, diff --git a/tests/test_litellm/proxy/openai_files_endpoint/test_files_common_utils.py b/tests/test_litellm/proxy/openai_files_endpoint/test_files_common_utils.py index 2161e345b40..7f84407f8b3 100644 --- a/tests/test_litellm/proxy/openai_files_endpoint/test_files_common_utils.py +++ b/tests/test_litellm/proxy/openai_files_endpoint/test_files_common_utils.py @@ -428,3 +428,47 @@ def test_add_internal_model_credentials_survives_a_failing_deployment_lookup(): add_internal_model_credentials(data=data, llm_router=router, model_id="deployment-gone") assert data == {"batch_id": "unified-batch-id"} + + +from litellm.proxy.openai_files_endpoints.common_utils import ( + _completed_batch_safe_to_retire, +) + + +def _completed_batch_for_retire( + output_file_id: str | None, completed: int | None = None +) -> LiteLLMBatch: + kwargs = dict( + id="batch-1", + completion_window="24h", + created_at=1234567890, + endpoint="/v1/chat/completions", + input_file_id="file-in", + object="batch", + status="completed", + output_file_id=output_file_id, + error_file_id=None, + ) + if completed is not None: + kwargs["request_counts"] = {"total": completed, "completed": completed, "failed": 0} + return LiteLLMBatch(**kwargs) + + +class TestCompletedBatchSafeToRetire: + """A completed batch is only safe to retire from cost recovery once its output + file has arrived or the provider proves no successful lines (#37713).""" + + def test_output_file_present_is_safe(self): + assert _completed_batch_safe_to_retire(_completed_batch_for_retire("file-out")) is True + + def test_no_output_and_no_successful_lines_is_safe(self): + # Every request line errored -> nothing left to recover. + assert _completed_batch_safe_to_retire(_completed_batch_for_retire(None, completed=0)) is True + + def test_no_output_but_successful_lines_is_not_safe(self): + # The bug: output_file_id is lagging; retiring here loses the spend record. + assert _completed_batch_safe_to_retire(_completed_batch_for_retire(None, completed=5)) is False + + def test_no_output_and_unknown_counts_is_not_safe(self): + # Counts unknown -> stay eligible so the next poller pass revisits it. + assert _completed_batch_safe_to_retire(_completed_batch_for_retire(None)) is False diff --git a/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py b/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py index 237a3092035..1b16e036ad7 100644 --- a/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py +++ b/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py @@ -1,5 +1,5 @@ import json -from typing import List +from typing import Final, List from unittest.mock import ANY, AsyncMock import pytest @@ -19,7 +19,11 @@ from litellm.proxy.openai_files_endpoints.file_content_streaming_handler import FileContentStreamingHandler, ) from litellm.proxy.proxy_server import app -from litellm.types.llms.openai import HttpxBinaryResponseContent, OpenAIFileObject +from litellm.types.llms.openai import ( + FileListPage, + HttpxBinaryResponseContent, + OpenAIFileObject, +) client = TestClient(app) from litellm.caching.caching import DualCache @@ -325,7 +329,15 @@ def test_mock_create_audio_file(mocker: MockerFixture, monkeypatch, llm_router: async def afile_retrieve(self, file_id, litellm_parent_otel_span, llm_router): raise NotImplementedError("Not implemented for test") - async def afile_list(self, purpose, litellm_parent_otel_span): + async def afile_list( + self, + purpose, + litellm_parent_otel_span, + user_api_key_dict, + limit=None, + after=None, + **data, + ): raise NotImplementedError("Not implemented for test") async def afile_delete( @@ -899,7 +911,15 @@ def test_create_file_with_expires_after( async def afile_retrieve(self, file_id, litellm_parent_otel_span, llm_router): raise NotImplementedError("Not implemented for test") - async def afile_list(self, purpose, litellm_parent_otel_span): + async def afile_list( + self, + purpose, + litellm_parent_otel_span, + user_api_key_dict, + limit=None, + after=None, + **data, + ): raise NotImplementedError("Not implemented for test") async def afile_delete( @@ -1062,7 +1082,15 @@ def test_create_file_with_expires_after_valid_values( async def afile_retrieve(self, file_id, litellm_parent_otel_span, llm_router): raise NotImplementedError("Not implemented for test") - async def afile_list(self, purpose, litellm_parent_otel_span): + async def afile_list( + self, + purpose, + litellm_parent_otel_span, + user_api_key_dict, + limit=None, + after=None, + **data, + ): raise NotImplementedError("Not implemented for test") async def afile_delete( @@ -1150,7 +1178,15 @@ def test_create_file_without_expires_after( async def afile_retrieve(self, file_id, litellm_parent_otel_span, llm_router): raise NotImplementedError("Not implemented for test") - async def afile_list(self, purpose, litellm_parent_otel_span): + async def afile_list( + self, + purpose, + litellm_parent_otel_span, + user_api_key_dict, + limit=None, + after=None, + **data, + ): raise NotImplementedError("Not implemented for test") async def afile_delete( @@ -1247,7 +1283,15 @@ def test_managed_files_with_loadbalancing( async def afile_retrieve(self, file_id, litellm_parent_otel_span, llm_router): raise NotImplementedError("Not implemented for test") - async def afile_list(self, purpose, litellm_parent_otel_span): + async def afile_list( + self, + purpose, + litellm_parent_otel_span, + user_api_key_dict, + limit=None, + after=None, + **data, + ): raise NotImplementedError("Not implemented for test") async def afile_delete( @@ -1364,7 +1408,15 @@ def test_create_file_with_nested_litellm_metadata( async def afile_retrieve(self, file_id, litellm_parent_otel_span, llm_router): raise NotImplementedError("Not implemented for test") - async def afile_list(self, purpose, litellm_parent_otel_span): + async def afile_list( + self, + purpose, + litellm_parent_otel_span, + user_api_key_dict, + limit=None, + after=None, + **data, + ): raise NotImplementedError("Not implemented for test") async def afile_delete( @@ -1468,7 +1520,15 @@ def test_create_file_with_deep_nested_litellm_metadata( async def afile_retrieve(self, file_id, litellm_parent_otel_span, llm_router): raise NotImplementedError("Not implemented for test") - async def afile_list(self, purpose, litellm_parent_otel_span): + async def afile_list( + self, + purpose, + litellm_parent_otel_span, + user_api_key_dict, + limit=None, + after=None, + **data, + ): raise NotImplementedError("Not implemented for test") async def afile_delete( @@ -1564,7 +1624,15 @@ def _make_capturing_managed_files(): async def afile_retrieve(self, file_id, litellm_parent_otel_span, llm_router): raise NotImplementedError - async def afile_list(self, purpose, litellm_parent_otel_span): + async def afile_list( + self, + purpose, + litellm_parent_otel_span, + user_api_key_dict, + limit=None, + after=None, + **data, + ): raise NotImplementedError async def afile_delete( @@ -2047,7 +2115,15 @@ def test_require_managed_files_allows_managed_file_upload( async def afile_retrieve(self, file_id, litellm_parent_otel_span, llm_router): raise NotImplementedError - async def afile_list(self, purpose, litellm_parent_otel_span): + async def afile_list( + self, + purpose, + litellm_parent_otel_span, + user_api_key_dict, + limit=None, + after=None, + **data, + ): raise NotImplementedError async def afile_delete( @@ -2171,7 +2247,15 @@ def test_require_managed_files_accepts_target_model_names_bracket_form( async def afile_retrieve(self, file_id, litellm_parent_otel_span, llm_router): raise NotImplementedError - async def afile_list(self, purpose, litellm_parent_otel_span): + async def afile_list( + self, + purpose, + litellm_parent_otel_span, + user_api_key_dict, + limit=None, + after=None, + **data, + ): raise NotImplementedError async def afile_delete( @@ -2251,7 +2335,15 @@ def test_require_managed_files_accepts_repeated_target_model_names_bracket_form( async def afile_retrieve(self, file_id, litellm_parent_otel_span, llm_router): raise NotImplementedError - async def afile_list(self, purpose, litellm_parent_otel_span): + async def afile_list( + self, + purpose, + litellm_parent_otel_span, + user_api_key_dict, + limit=None, + after=None, + **data, + ): raise NotImplementedError async def afile_delete( @@ -2463,6 +2555,403 @@ def test_list_files_without_target_model_names_uses_team_openai_deployment( proxy_logging_obj.post_call_failure_hook.assert_not_called() +def test_unscoped_list_files_uses_managed_file_store( + mocker: MockerFixture, monkeypatch, llm_router: Router +): + import litellm.proxy.proxy_server as ps + from litellm.llms.base_llm.files.transformation import BaseFileEndpoints + from litellm.proxy._types import LitellmUserRoles + + managed_file = OpenAIFileObject( + id="unified-file-id", + object="file", + bytes=100, + created_at=1700000000, + filename="output.jsonl", + purpose="batch_output", + status="processed", + ) + + proxy_logging_obj = setup_proxy_logging_object(monkeypatch, llm_router) + managed_files = mocker.MagicMock(spec=BaseFileEndpoints) + managed_files.afile_list = mocker.AsyncMock( + return_value={ + "object": "list", + "data": [managed_file], + "first_id": managed_file.id, + "last_id": managed_file.id, + "has_more": False, + } + ) + proxy_logging_obj.proxy_hook_mapping["managed_files"] = managed_files + proxy_logging_obj.update_request_status = mocker.AsyncMock() + proxy_logging_obj.post_call_success_hook = mocker.AsyncMock(return_value=None) + proxy_logging_obj.post_call_failure_hook = mocker.AsyncMock() + monkeypatch.setattr("litellm.proxy.proxy_server.master_key", None) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", llm_router) + provider_list = mocker.patch.object(litellm, "afile_list", new=mocker.AsyncMock()) + + app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth( + api_key="test-key", + user_role=LitellmUserRoles.INTERNAL_USER, + user_id="test-user", + ) + + try: + response = client.get( + "/v1/files", + headers={"Authorization": "Bearer test-key"}, + ) + finally: + app.dependency_overrides.pop(ps.user_api_key_auth, None) + + assert response.status_code == 200, response.text + assert response.json()["data"][0]["id"] == "unified-file-id" + managed_files.afile_list.assert_awaited_once() + assert managed_files.afile_list.await_args.kwargs["user_api_key_dict"].user_id == "test-user" + assert managed_files.afile_list.await_args.kwargs["limit"] is None + assert managed_files.afile_list.await_args.kwargs["after"] is None + provider_list.assert_not_awaited() + proxy_logging_obj.post_call_failure_hook.assert_not_called() + + +def test_unscoped_list_files_forwards_limit_and_after_to_the_managed_file_store( + mocker: MockerFixture, monkeypatch, llm_router: Router +): + import litellm.proxy.proxy_server as ps + from litellm.llms.base_llm.files.transformation import BaseFileEndpoints + from litellm.proxy._types import LitellmUserRoles + + second_page_file = OpenAIFileObject( + id="unified-file-id-2", + object="file", + bytes=100, + created_at=1700000000, + filename="output.jsonl", + purpose="batch", + status="processed", + ) + + proxy_logging_obj = setup_proxy_logging_object(monkeypatch, llm_router) + managed_files = mocker.MagicMock(spec=BaseFileEndpoints) + managed_files.afile_list = mocker.AsyncMock( + return_value={ + "object": "list", + "data": [second_page_file], + "first_id": second_page_file.id, + "last_id": second_page_file.id, + "has_more": True, + } + ) + proxy_logging_obj.proxy_hook_mapping["managed_files"] = managed_files + proxy_logging_obj.update_request_status = mocker.AsyncMock() + proxy_logging_obj.post_call_success_hook = mocker.AsyncMock(return_value=None) + proxy_logging_obj.post_call_failure_hook = mocker.AsyncMock() + monkeypatch.setattr("litellm.proxy.proxy_server.master_key", None) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", llm_router) + provider_list = mocker.patch.object(litellm, "afile_list", new=mocker.AsyncMock()) + + app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth( + api_key="test-key", + user_role=LitellmUserRoles.INTERNAL_USER, + user_id="test-user", + ) + + try: + response = client.get( + "/v1/files?limit=2&after=unified-file-id-1&purpose=batch", + headers={"Authorization": "Bearer test-key"}, + ) + finally: + app.dependency_overrides.pop(ps.user_api_key_auth, None) + + assert response.status_code == 200, response.text + assert response.json()["data"][0]["id"] == "unified-file-id-2" + assert response.json()["has_more"] is True + call_kwargs = managed_files.afile_list.await_args.kwargs + assert call_kwargs["limit"] == 2 + assert call_kwargs["after"] == "unified-file-id-1" + assert call_kwargs["purpose"] == "batch" + provider_list.assert_not_awaited() + proxy_logging_obj.post_call_failure_hook.assert_not_called() + + +def _setup_unscoped_list_files_route(mocker, monkeypatch, llm_router: Router, afile_list): + """Wire GET /v1/files to the managed file store, with afile_list as the store.""" + import litellm.proxy.proxy_server as ps + from litellm.llms.base_llm.files.transformation import BaseFileEndpoints + from litellm.proxy._types import LitellmUserRoles + + proxy_logging_obj = setup_proxy_logging_object(monkeypatch, llm_router) + managed_files = mocker.MagicMock(spec=BaseFileEndpoints) + managed_files.afile_list = mocker.AsyncMock(side_effect=afile_list) + proxy_logging_obj.proxy_hook_mapping["managed_files"] = managed_files + proxy_logging_obj.update_request_status = mocker.AsyncMock() + proxy_logging_obj.post_call_success_hook = mocker.AsyncMock(return_value=None) + proxy_logging_obj.post_call_failure_hook = mocker.AsyncMock() + monkeypatch.setattr("litellm.proxy.proxy_server.master_key", None) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", llm_router) + mocker.patch.object(litellm, "afile_list", new=mocker.AsyncMock()) + + app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth( + api_key="test-key", + user_role=LitellmUserRoles.INTERNAL_USER, + user_id="test-user", + ) + return managed_files + + +def _get_list_files(path: str): + try: + return client.get(path, headers={"Authorization": "Bearer test-key"}) + finally: + import litellm.proxy.proxy_server as ps + + app.dependency_overrides.pop(ps.user_api_key_auth, None) + + +def _get_unscoped_list_files(query: str): + return _get_list_files(f"/v1/files{query}") + + +_EMPTY_FILE_LIST_PAGE: Final = { + "object": "list", + "data": [], + "first_id": None, + "last_id": None, + "has_more": False, +} + + +async def _validating_afile_list(**kwargs): + """Stand in for the managed file store, applying the real request validation.""" + from litellm.proxy.openai_files_endpoints.common_utils import ( + validate_file_list_limit, + validate_file_list_purpose, + ) + + validate_file_list_limit(kwargs.get("limit")) + validate_file_list_purpose(kwargs.get("purpose")) + return FileListPage(**_EMPTY_FILE_LIST_PAGE) + + +async def _permissive_afile_list(**kwargs): + """Stand in for a file store that validates nothing, so only the route can reject.""" + return FileListPage(**_EMPTY_FILE_LIST_PAGE) + + +@pytest.mark.parametrize( + "limit, bound, expected_range", + [ + (0, "below minimum", ">= 1"), + (-1, "below minimum", ">= 1"), + (10001, "above maximum", "<= 10000"), + ], +) +def test_unscoped_list_files_returns_400_for_a_limit_outside_the_openai_range( + mocker: MockerFixture, monkeypatch, llm_router: Router, limit, bound, expected_range +): + """An out-of-range limit is the caller's mistake, so it must not read as a 500 the SDK retries.""" + _setup_unscoped_list_files_route(mocker, monkeypatch, llm_router, _validating_afile_list) + + response = _get_unscoped_list_files(f"?limit={limit}") + + assert response.status_code == 400, response.text + assert response.json() == { + "error": { + "message": ( + f"Invalid 'limit': integer {bound} value. " + f"Expected a value {expected_range}, but got {limit} instead." + ), + "type": "invalid_request_error", + "param": "limit", + "code": "400", + } + } + + +@pytest.mark.parametrize("limit", [1, 10000]) +def test_unscoped_list_files_accepts_the_ends_of_the_openai_limit_range( + mocker: MockerFixture, monkeypatch, llm_router: Router, limit +): + managed_files = _setup_unscoped_list_files_route(mocker, monkeypatch, llm_router, _validating_afile_list) + + response = _get_unscoped_list_files(f"?limit={limit}") + + assert response.status_code == 200, response.text + assert response.json()["data"] == [] + assert managed_files.afile_list.await_args.kwargs["limit"] == limit + + +@pytest.mark.parametrize( + "path", + [ + "/v1/files?limit=0", + "/v1/files?limit=0&target_model_names=gpt-3.5-turbo", + "/openai/v1/files?limit=0", + ], + ids=["managed-file-store", "target-model-names", "provider-route"], +) +def test_list_files_validates_the_limit_on_every_branch( + mocker: MockerFixture, monkeypatch, llm_router: Router, path +): + """The limit is a route-level contract, so the scoped and provider branches reject it too.""" + _setup_unscoped_list_files_route(mocker, monkeypatch, llm_router, _permissive_afile_list) + + response = _get_list_files(path) + + assert response.status_code == 400, response.text + assert response.json()["error"]["param"] == "limit" + assert response.json()["error"]["message"] == ( + "Invalid 'limit': integer below minimum value. Expected a value >= 1, but got 0 instead." + ) + + +def test_unscoped_list_files_returns_400_for_an_unknown_after_cursor( + mocker: MockerFixture, monkeypatch, llm_router: Router +): + from litellm.proxy._types import ProxyException + + async def _unknown_cursor(**kwargs): + raise ProxyException( + message=f"Invalid 'after' cursor: no file found with id '{kwargs['after']}'.", + type="invalid_request_error", + param="after", + code=400, + openai_code="invalid_value", + ) + + _setup_unscoped_list_files_route(mocker, monkeypatch, llm_router, _unknown_cursor) + + response = _get_unscoped_list_files("?after=file-does-not-exist-xyz") + + assert response.status_code == 400, response.text + assert response.json() == { + "error": { + "message": "Invalid 'after' cursor: no file found with id 'file-does-not-exist-xyz'.", + "type": "invalid_request_error", + "param": "after", + "code": "400", + } + } + + +def _managed_file(file_id: str) -> OpenAIFileObject: + return OpenAIFileObject( + id=file_id, + bytes=17, + created_at=1700000000, + filename="batch_input.jsonl", + object="file", + purpose="batch", + status="uploaded", + ) + + +def test_unscoped_list_files_hands_post_call_hooks_a_page_object( + mocker: MockerFixture, monkeypatch, llm_router: Router +): + """Logging callbacks read ``response.data`` off a listing, so the managed + branch has to hand them the same page shape the provider branch does. A bare + mapping turns every registered callback into a 500 on this route.""" + import litellm.proxy.proxy_server as ps + + seen_by_callback: list[list[str]] = [] + + async def _reads_response_data(data, user_api_key_dict, response): + seen_by_callback.append([file.id for file in response.data]) + return None + + async def _one_managed_file(**kwargs): + return FileListPage( + data=[_managed_file("unified-file-id")], + first_id="unified-file-id", + last_id="unified-file-id", + ) + + _setup_unscoped_list_files_route(mocker, monkeypatch, llm_router, _one_managed_file) + ps.proxy_logging_obj.post_call_success_hook = _reads_response_data + + response = _get_unscoped_list_files("") + + assert response.status_code == 200, response.text + assert seen_by_callback == [["unified-file-id"]] + body = response.json() + assert list(body) == ["object", "data", "first_id", "last_id", "has_more"] + assert body["object"] == "list" + assert [file["id"] for file in body["data"]] == ["unified-file-id"] + assert body["has_more"] is False + + +@pytest.mark.parametrize("purpose", ["nonexistent_purpose", "EVALS", "batch "]) +def test_unscoped_list_files_returns_400_for_a_purpose_the_api_never_accepts( + mocker: MockerFixture, monkeypatch, llm_router: Router, purpose +): + """An unknown purpose matches nothing, so reporting an empty page would dress + a bad request up as a successful one. The provider-backed branches reject the + same values, and so does the upload route.""" + from urllib.parse import quote + + _setup_unscoped_list_files_route(mocker, monkeypatch, llm_router, _validating_afile_list) + + response = _get_list_files(f"/v1/files?purpose={quote(purpose)}") + + assert response.status_code == 400, response.text + assert response.json()["error"]["param"] == "purpose" + assert response.json()["error"]["type"] == "invalid_request_error" + assert response.json()["error"]["message"].startswith(f"Invalid purpose: {purpose}. Must be one of: ") + + +@pytest.mark.parametrize("purpose", ["batch", "assistants", "fine-tune"]) +def test_unscoped_list_files_accepts_every_documented_purpose( + mocker: MockerFixture, monkeypatch, llm_router: Router, purpose +): + managed_files = _setup_unscoped_list_files_route(mocker, monkeypatch, llm_router, _validating_afile_list) + + response = _get_list_files(f"/v1/files?purpose={purpose}") + + assert response.status_code == 200, response.text + assert managed_files.afile_list.await_args.kwargs["purpose"] == purpose + + +def test_list_files_reports_a_bad_target_model_names_as_a_400( + mocker: MockerFixture, monkeypatch, llm_router: Router +): + """The exception tail reports an HTTPException with its own status and error + type rather than relabelling it, so a client that branches on either keeps + reading the same thing off a bad request.""" + _setup_unscoped_list_files_route(mocker, monkeypatch, llm_router, _permissive_afile_list) + + response = _get_list_files("/v1/files?target_model_names=gpt-3.5-turbo,gpt-4o") + + assert response.status_code == 400, response.text + assert response.json() == { + "error": { + "message": "target_model_names on list files must be a list of one model name. Example: ['gpt-4o']", + "type": "None", + "param": "None", + "code": "400", + } + } + + +def test_list_files_reports_an_unexpected_file_store_error_as_a_500( + mocker: MockerFixture, monkeypatch, llm_router: Router +): + async def _blows_up(**kwargs): + raise RuntimeError("managed file table is unreachable") + + _setup_unscoped_list_files_route(mocker, monkeypatch, llm_router, _blows_up) + + response = _get_unscoped_list_files("") + + assert response.status_code == 500, response.text + assert response.json()["error"]["message"] == "managed file table is unreachable" + + def test_list_files_restricted_team_does_not_leak_global_openai_credentials( mocker: MockerFixture, monkeypatch ): @@ -3677,6 +4166,66 @@ def test_batch_upload_redacts_per_record(monkeypatch, llm_router: Router): ProxyLogging._callback_capabilities_cache.clear() +PLAIN_UPLOAD_RESPONSE_BODY = { + "id": "dummy-id", + "object": "file", + "bytes": 0, + "created_at": 1234567890, + "filename": "batch.jsonl", + "purpose": "batch", + "status": "uploaded", + "expires_at": None, + "status_details": None, +} + + +def test_create_file_omits_batch_guardrail_field_when_no_guardrail_configured(monkeypatch, llm_router: Router): + """An upload no guardrail is configured for serialises the plain OpenAI file shape.""" + forwarded_calls = _setup_batch_upload_endpoint(monkeypatch, llm_router) + try: + response = client.post( + "/v1/files", + files={"file": ("batch.jsonl", VALID_BATCH_LINE, "application/jsonl")}, + data={"purpose": "batch"}, + headers={"Authorization": "Bearer test-key"}, + ) + finally: + _teardown_batch_upload_endpoint() + + assert response.status_code == 200, response.text + assert len(forwarded_calls) == 1 + assert response.json() == PLAIN_UPLOAD_RESPONSE_BODY + + +def test_create_file_omits_batch_guardrail_field_when_guardrail_made_no_changes(monkeypatch, llm_router: Router): + """A guardrail that runs and changes nothing leaves the response the plain OpenAI file shape.""" + import litellm + from litellm.integrations.custom_guardrail import CustomGuardrail + from litellm.proxy.utils import ProxyLogging + + class _Passthrough(CustomGuardrail): + async def async_pre_call_hook(self, user_api_key_dict, cache, data, call_type): + return data + + forwarded_calls = _setup_batch_upload_endpoint(monkeypatch, llm_router) + monkeypatch.setattr(litellm, "callbacks", [_Passthrough(guardrail_name="noop", default_on=True)]) + ProxyLogging._callback_capabilities_cache.clear() + try: + response = client.post( + "/v1/files", + files={"file": ("batch.jsonl", VALID_BATCH_LINE, "application/jsonl")}, + data={"purpose": "batch"}, + headers={"Authorization": "Bearer test-key"}, + ) + finally: + _teardown_batch_upload_endpoint() + ProxyLogging._callback_capabilities_cache.clear() + + assert response.status_code == 200, response.text + assert len(forwarded_calls) == 1 + assert response.json() == PLAIN_UPLOAD_RESPONSE_BODY + + def test_batch_upload_closes_the_spools_it_opened(monkeypatch, llm_router: Router): """The scan and the rewrite each open a spool; the request owns both and must not leak them.""" import json as _json diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_managed_id_rewriter.py b/tests/test_litellm/proxy/pass_through_endpoints/test_managed_id_rewriter.py index 9c16c52f589..f5bec4a2585 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_managed_id_rewriter.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_managed_id_rewriter.py @@ -126,3 +126,24 @@ async def test_list_files_limit_above_batch_cap_still_served(): assert result is not None assert [item["id"] for item in result["data"]] == [managed_id] + + +@pytest.mark.asyncio +async def test_list_files_drops_batch_guardrail_key_persisted_by_an_older_proxy(): + """Rows written before the response serializer dropped the key still carry an explicit null.""" + managed_id = new_managed_id("openai", "file-abc") + row = _file_row(managed_id) + row.file_object = {**row.file_object, "litellm_batch_guardrail": None} + pc = _prisma_client(file_rows=[row]) + + result = await list_passthrough_ids_from_db( + provider="openai", + route="/openai/v1/files", + user_api_key_dict=_user(), + prisma_client=pc, + query_params={}, + ) + + assert result is not None + assert "litellm_batch_guardrail" not in result["data"][0] + assert result["data"][0]["filename"] == "test.jsonl" diff --git a/tests/test_litellm/proxy/proxy_server/test_proxy_config.py b/tests/test_litellm/proxy/proxy_server/test_proxy_config.py index b0b2c68e30d..ee0de8840f6 100644 --- a/tests/test_litellm/proxy/proxy_server/test_proxy_config.py +++ b/tests/test_litellm/proxy/proxy_server/test_proxy_config.py @@ -26,6 +26,7 @@ from litellm.proxy.proxy_server import ( _scrub_guardrail_inner, resolve_complexity_router_plugins, resolve_routing_plugins, + validate_deployment_max_agentic_loops, ) from .conftest import normalize @@ -153,6 +154,71 @@ def test_resolve_complexity_router_plugins_resolves_dotted_path_to_live_instance assert type(config["plugins"][0]).__name__ == "_Plugin" +def test_validate_deployment_max_agentic_loops_allows_a_deployment_without_the_key(): + model = {"model_name": "gpt-4o", "litellm_params": {"model": "gpt-4o"}} + + validate_deployment_max_agentic_loops(model) + + assert "max_agentic_loops" not in model["litellm_params"] + + +def test_validate_deployment_max_agentic_loops_leaves_a_valid_ceiling_alone(): + model = {"model_name": "gpt-4o", "litellm_params": {"model": "gpt-4o", "max_agentic_loops": 5}} + + validate_deployment_max_agentic_loops(model) + + assert model["litellm_params"]["max_agentic_loops"] == 5 + + +def test_validate_deployment_max_agentic_loops_rejects_zero(): + """ + A per-deployment 0 used to be swallowed by an `or 3` and read as the default + ceiling of 3, handing the loosest setting to whoever asked for the tightest. + """ + with pytest.raises(ValueError, match="must be at least 1, got 0"): + validate_deployment_max_agentic_loops( + {"model_name": "gpt-4o", "litellm_params": {"model": "gpt-4o", "max_agentic_loops": 0}} + ) + + +def test_validate_deployment_max_agentic_loops_rejects_a_non_integer(): + """ + A per-deployment non-integer used to let the proxy boot and then fail every + request to that model with `invalid literal for int() with base 10`. + """ + with pytest.raises(TypeError, match="must be an integer"): + validate_deployment_max_agentic_loops( + {"model_name": "gpt-4o", "litellm_params": {"model": "gpt-4o", "max_agentic_loops": "three"}} + ) + + +def test_validate_deployment_max_agentic_loops_rejects_a_bool(): + with pytest.raises(TypeError, match="must be an integer"): + validate_deployment_max_agentic_loops( + {"model_name": "gpt-4o", "litellm_params": {"model": "gpt-4o", "max_agentic_loops": True}} + ) + + +def test_validate_deployment_max_agentic_loops_accepts_a_ceiling_from_an_env_var(): + """ + `max_agentic_loops: os.environ/MAX_AGENTIC_LOOPS` is resolved to a string + before this check runs, and the old `int(... or 3)` accepted that, so + refusing it here would stop an already working proxy from booting. + """ + model = {"model_name": "gpt-4o", "litellm_params": {"model": "gpt-4o", "max_agentic_loops": "5"}} + + validate_deployment_max_agentic_loops(model) + + assert model["litellm_params"]["max_agentic_loops"] == "5" + + +def test_validate_deployment_max_agentic_loops_names_the_offending_model(): + with pytest.raises(ValueError, match="on model 'claude-sonnet-4-5'"): + validate_deployment_max_agentic_loops( + {"model_name": "claude-sonnet-4-5", "litellm_params": {"max_agentic_loops": -1}} + ) + + def test_resolve_complexity_router_plugins_rejects_non_routing_plugin_object(tmp_path): plugin_file = tmp_path / "bad_plugin.py" plugin_file.write_text("not_a_plugin = object()\n") diff --git a/tests/test_litellm/proxy/test_common_request_processing.py b/tests/test_litellm/proxy/test_common_request_processing.py index 3c738aa164c..58714a5e319 100644 --- a/tests/test_litellm/proxy/test_common_request_processing.py +++ b/tests/test_litellm/proxy/test_common_request_processing.py @@ -323,6 +323,62 @@ class TestProxyBaseLLMRequestProcessing: pytest.fail("litellm_call_id is not a valid UUID") assert data_passed["litellm_call_id"] == returned_data["litellm_call_id"] + @pytest.mark.asyncio + async def test_common_processing_pre_call_logic_refreshes_proxy_server_request_body_after_guardrails( + self, monkeypatch + ): + """ + A guardrail (e.g. Presidio PII masking) mutates data["messages"] in place inside + pre_call_hook. The proxy_server_request.body snapshot is taken before that hook + runs, so it must be refreshed afterward or SpendLogs (when store_prompts_in_spend_logs + is enabled) persists the raw pre-guardrail body, bypassing the masking entirely. + """ + processing_obj = ProxyBaseLLMRequestProcessing(data={}) + mock_request = MagicMock(spec=Request) + mock_request.headers = {} + + raw_messages = [{"role": "user", "content": "my ssn is 123-45-6789"}] + + async def mock_add_litellm_data_to_request(*args, **kwargs): + return { + "messages": raw_messages, + "proxy_server_request": { + "url": "http://testserver/chat/completions", + "method": "POST", + "body": {"messages": raw_messages}, + }, + } + + async def mock_pre_call_hook(user_api_key_dict, data, call_type): + data["messages"] = [{"role": "user", "content": "my ssn is "}] + return data + + mock_proxy_logging_obj = MagicMock(spec=ProxyLogging) + mock_proxy_logging_obj.pre_call_hook = AsyncMock(side_effect=mock_pre_call_hook) + monkeypatch.setattr( + litellm.proxy.common_request_processing, + "add_litellm_data_to_request", + mock_add_litellm_data_to_request, + ) + + returned_data, _ = await processing_obj.common_processing_pre_call_logic( + request=mock_request, + general_settings={}, + user_api_key_dict=MagicMock(spec=UserAPIKeyAuth), + proxy_logging_obj=mock_proxy_logging_obj, + proxy_config=MagicMock(spec=ProxyConfig), + route_type="acompletion", + ) + + persisted_body = returned_data["proxy_server_request"]["body"] + assert persisted_body["messages"] == returned_data["messages"] + assert "123-45-6789" not in json.dumps(persisted_body["messages"]) + # litellm_logging_obj is stamped onto `data` by function_setup between the + # initial snapshot and pre_call_hook; it must never leak into the persisted + # audit body, which needs to stay plain-JSON-serializable end to end. + assert "litellm_logging_obj" not in persisted_body + json.dumps(persisted_body) + def test_add_dd_apm_tags_for_litellm_call_id_uses_dd_tracing_helper(self, monkeypatch): mock_set_active_span_tag = MagicMock(return_value=True) import litellm.proxy.dd_span_tagger @@ -1336,11 +1392,25 @@ class TestProxyBaseLLMRequestProcessing: route_type=route_type, ) - # Verify queue_time_seconds is set and non-negative + # Verify queue_time_seconds is set and non-negative. Ends at start_time + # (captured before this mock runs, so it can precede the mock's own + # time.time() by a handful of microseconds) rather than a freshly + # captured time.time(), so a tiny tolerance below 0.5 is expected and + # correct -- see LIT-6012. metadata = returned_data.get("metadata", {}) assert "queue_time_seconds" in metadata, "queue_time_seconds should be set in metadata" - assert metadata["queue_time_seconds"] >= 0.5, ( - f"queue_time_seconds should be at least 0.5, got {metadata['queue_time_seconds']}" + assert metadata["queue_time_seconds"] >= 0.49, ( + f"queue_time_seconds should be at least ~0.5, got {metadata['queue_time_seconds']}" + ) + + # queue_time_seconds must end exactly where logging_obj.start_time begins + # (the same start_time litellm_request_total_latency_metric's window + # starts from) so the two windows share a boundary, not an overlap. + # A mutant that reintroduces a separately-captured processing_start_time + # would make this assertion fail. + arrival_time = returned_data["proxy_server_request"]["arrival_time"] + assert arrival_time + metadata["queue_time_seconds"] == pytest.approx( + logging_obj.start_time.timestamp(), abs=1e-6 ) diff --git a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py index 111f11f85ba..81a97a70efa 100644 --- a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py +++ b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py @@ -2,6 +2,9 @@ import asyncio import copy import json import os +import time +from datetime import datetime, timezone +from types import SimpleNamespace from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -268,6 +271,75 @@ async def test_stamped_auth_object_reflects_header_derived_identity(): assert stamped.end_user_id == "end-user-from-header" +@pytest.mark.asyncio +async def test_arrival_time_prefers_litellm_received_at_over_time_time(): + """LIT-6012: by the time this function runs, auth has already completed, so + time.time() here would silently exclude the whole auth phase from the + queue-time window. request.state.litellm_received_at (stamped at the top of + user_api_key_auth, before auth work) must win when present.""" + from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request + + request_mock = MagicMock(spec=Request) + request_mock.url = MagicMock() + request_mock.url.path = "/v1/chat/completions" + request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions" + request_mock.method = "POST" + request_mock.query_params = {} + request_mock.headers = {"Content-Type": "application/json"} + request_mock.client = MagicMock() + request_mock.client.host = "127.0.0.1" + received_at = datetime(2024, 1, 1, tzinfo=timezone.utc) + request_mock.state = SimpleNamespace(litellm_received_at=received_at) + + user_api_key_dict = UserAPIKeyAuth(api_key="hashed-key", metadata={}, team_metadata={}) + + updated_data = await add_litellm_data_to_request( + data={"model": "gpt-3.5-turbo"}, + request=request_mock, + user_api_key_dict=user_api_key_dict, + proxy_config=MagicMock(), + general_settings={}, + version="test-version", + ) + + assert updated_data["proxy_server_request"]["arrival_time"] == received_at.timestamp() + + +@pytest.mark.asyncio +async def test_arrival_time_falls_back_to_time_time_without_litellm_received_at(): + """Callers that never went through user_api_key_auth (no stamp on request.state) + must still get a usable arrival_time instead of erroring.""" + from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request + + request_mock = MagicMock(spec=Request) + request_mock.url = MagicMock() + request_mock.url.path = "/v1/chat/completions" + request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions" + request_mock.method = "POST" + request_mock.query_params = {} + request_mock.headers = {"Content-Type": "application/json"} + request_mock.client = MagicMock() + request_mock.client.host = "127.0.0.1" + request_mock.state = SimpleNamespace() # no litellm_received_at attribute + + user_api_key_dict = UserAPIKeyAuth(api_key="hashed-key", metadata={}, team_metadata={}) + + before = time.time() + updated_data = await add_litellm_data_to_request( + data={"model": "gpt-3.5-turbo"}, + request=request_mock, + user_api_key_dict=user_api_key_dict, + proxy_config=MagicMock(), + general_settings={}, + version="test-version", + ) + after = time.time() + + arrival_time = updated_data["proxy_server_request"]["arrival_time"] + assert isinstance(arrival_time, float) + assert before <= arrival_time <= after + + @pytest.mark.asyncio async def test_add_litellm_data_to_request_strips_admin_injection_slots(): """User-supplied user_api_key_metadata / user_api_key_team_metadata / @@ -710,6 +782,54 @@ async def test_add_litellm_data_to_request_body_snapshot_excludes_proxy_server_r ) +def test_refresh_proxy_server_request_body_snapshot_picks_up_guardrail_masking(): + """ + Regression: proxy_server_request['body'] is snapshotted by + add_litellm_data_to_request BEFORE guardrails (e.g. Presidio PII masking) run + in pre_call_hook. Without a refresh after pre_call_hook, the persisted body + silently bypasses whatever masking the guardrail applied, so raw PII/PCI + lands in SpendLogs when store_prompts_in_spend_logs is enabled. + """ + from litellm.proxy.litellm_pre_call_utils import ( + refresh_proxy_server_request_body_snapshot, + ) + + class _FakeLoggingObj: + """Stands in for the live, non-JSON-serializable Logging instance that + litellm.utils.function_setup stamps onto `data` between the initial + snapshot and pre_call_hook.""" + + data = { + "model": "gpt-3.5-turbo", + "messages": [{"role": "user", "content": "my ssn is 123-45-6789"}], + "secret_fields": {"raw_headers": {"authorization": "Bearer sk-secret"}}, + "litellm_logging_obj": _FakeLoggingObj(), + "proxy_server_request": { + "url": "http://localhost/v1/chat/completions", + "method": "POST", + "body": { + "model": "gpt-3.5-turbo", + "messages": [{"role": "user", "content": "my ssn is 123-45-6789"}], + }, + }, + } + + # Simulate a PII-masking guardrail mutating `messages` in place, like Presidio's + # async_pre_call_hook does, after the initial snapshot was already taken. + data["messages"] = [{"role": "user", "content": "my ssn is "}] + + refresh_proxy_server_request_body_snapshot(data) + + refreshed_body = data["proxy_server_request"]["body"] + assert refreshed_body["messages"] == data["messages"] + # Still excludes secrets, self-reference, and the live logging object, same as + # the initial snapshot -- and proves the persisted body stays JSON-serializable. + assert "secret_fields" not in refreshed_body + assert "proxy_server_request" not in refreshed_body + assert "litellm_logging_obj" not in refreshed_body + assert "123-45-6789" not in json.dumps(refreshed_body) + + @pytest.mark.asyncio async def test_add_litellm_data_to_request_strips_string_encoded_admin_injection(): """Regression: metadata arriving as a JSON string (multipart/form-data or @@ -2738,7 +2858,6 @@ def test_add_headers_to_llm_call_by_model_group_existing_headers_in_data(): litellm.model_group_settings = original_model_group_settings -import time from typing import Optional from fastapi.responses import Response diff --git a/tests/test_litellm/proxy/test_prisma_migration.py b/tests/test_litellm/proxy/test_prisma_migration.py index 01b768ea8dc..729adcfb9e0 100644 --- a/tests/test_litellm/proxy/test_prisma_migration.py +++ b/tests/test_litellm/proxy/test_prisma_migration.py @@ -34,24 +34,22 @@ class TestPrismaMigration: mock_run_server.assert_called_once_with(("--skip_server_startup",), standalone_mode=False) + @pytest.mark.parametrize("env", [{}, {"ENFORCE_PRISMA_MIGRATION_CHECK": "false"}]) @patch("litellm.proxy.prisma_migration.subprocess.run") @patch("litellm.proxy.prisma_migration.run_server") - def test_main_returns_prisma_generate_exit_code_when_enforced( - self, mock_run_server: MagicMock, mock_subprocess_run: MagicMock + def test_main_exits_zero_when_only_prisma_generate_fails( + self, + mock_run_server: MagicMock, + mock_subprocess_run: MagicMock, + env: dict[str, str], ) -> None: - mock_subprocess_run.return_value = MagicMock(returncode=7, stdout="", stderr="") + mock_subprocess_run.return_value = MagicMock( + returncode=1, + stdout="", + stderr="PermissionError: [Errno 13] Permission denied: '/app/.venv/lib/python3.13/site-packages/prisma/schema.prisma'", + ) - with patch.dict(os.environ, {}, clear=True): - assert prisma_migration.main() == 7 - - @patch("litellm.proxy.prisma_migration.subprocess.run") - @patch("litellm.proxy.prisma_migration.run_server") - def test_main_ignores_prisma_generate_exit_code_when_disabled( - self, mock_run_server: MagicMock, mock_subprocess_run: MagicMock - ) -> None: - mock_subprocess_run.return_value = MagicMock(returncode=7, stdout="", stderr="") - - with patch.dict(os.environ, {"ENFORCE_PRISMA_MIGRATION_CHECK": "false"}, clear=True): + with patch.dict(os.environ, env, clear=True): assert prisma_migration.main() == 0 @patch("litellm.proxy.prisma_migration.subprocess.run") diff --git a/tests/test_litellm/responses/litellm_completion_transformation/test_image_generation_output.py b/tests/test_litellm/responses/litellm_completion_transformation/test_image_generation_output.py index ed7a3f63a8e..80e5335b4d4 100644 --- a/tests/test_litellm/responses/litellm_completion_transformation/test_image_generation_output.py +++ b/tests/test_litellm/responses/litellm_completion_transformation/test_image_generation_output.py @@ -57,9 +57,6 @@ class TestExtractImageGenerationOutputItems: def test_extracts_images_correctly(self): """Should extract OutputImageGenerationCall objects from images""" - mock_response = Mock(spec=ModelResponse) - mock_response.id = "test_123" - mock_message = Mock(spec=Message) mock_message.images = [ { @@ -80,7 +77,6 @@ class TestExtractImageGenerationOutputItems: result = ( LiteLLMCompletionResponsesConfig._extract_image_generation_output_items( - chat_completion_response=mock_response, choice=mock_choice, ) ) @@ -89,13 +85,13 @@ class TestExtractImageGenerationOutputItems: assert result[0].type == "image_generation_call" assert result[0].result == "IMG1" assert result[1].result == "IMG2" - assert result[0].id == "test_123_img_0" - assert result[1].id == "test_123_img_1" + assert result[0].id.startswith("ig_") + assert result[1].id.startswith("ig_") + assert result[0].id != result[1].id assert result[0].status == "completed" def test_returns_empty_for_no_images(self): """Should return empty list if no images""" - mock_response = Mock(spec=ModelResponse) mock_message = Mock(spec=Message) mock_message.images = [] @@ -105,7 +101,6 @@ class TestExtractImageGenerationOutputItems: result = ( LiteLLMCompletionResponsesConfig._extract_image_generation_output_items( - chat_completion_response=mock_response, choice=mock_choice, ) ) @@ -114,9 +109,6 @@ class TestExtractImageGenerationOutputItems: def test_maps_finish_reason_to_status(self): """Should correctly map finish_reason to status""" - mock_response = Mock(spec=ModelResponse) - mock_response.id = "test_finish" - mock_message = Mock(spec=Message) mock_message.images = [ { @@ -132,7 +124,6 @@ class TestExtractImageGenerationOutputItems: result = ( LiteLLMCompletionResponsesConfig._extract_image_generation_output_items( - chat_completion_response=mock_response, choice=mock_choice, ) ) @@ -198,3 +189,40 @@ class TestExtractMessageOutputItemsIntegration: assert len(result) == 1 assert isinstance(result[0], GenericResponseOutputItem) assert result[0].type == "message" + + +class TestImageGenerationOutputItemIds: + """Image generation call IDs must use the ig_ prefix (issue #27333). + + Native OpenAI Responses validates the prefix before it looks the item up, so a + replayed chatcmpl-*_img_N ID is rejected outright. + """ + + def _choice_with_images(self, count): + mock_message = Mock(spec=Message) + mock_message.images = [ + {"image_url": {"url": f"data:image/png;base64,IMG{idx}"}} + for idx in range(count) + ] + mock_choice = Mock(spec=Choices) + mock_choice.message = mock_message + mock_choice.finish_reason = "stop" + return mock_choice + + def test_image_generation_item_id_uses_ig_prefix(self): + result = LiteLLMCompletionResponsesConfig._extract_image_generation_output_items( + choice=self._choice_with_images(2), + ) + + assert len(result) == 2 + for item in result: + assert item.id.startswith("ig_") + assert "chatcmpl-" not in item.id + assert "_img_" not in item.id + + def test_image_generation_item_ids_are_unique(self): + result = LiteLLMCompletionResponsesConfig._extract_image_generation_output_items( + choice=self._choice_with_images(3), + ) + + assert len({item.id for item in result}) == 3 diff --git a/tests/test_litellm/responses/litellm_completion_transformation/test_litellm_completion_responses.py b/tests/test_litellm/responses/litellm_completion_transformation/test_litellm_completion_responses.py index 5efabed4b8d..2273f23b1cc 100644 --- a/tests/test_litellm/responses/litellm_completion_transformation/test_litellm_completion_responses.py +++ b/tests/test_litellm/responses/litellm_completion_transformation/test_litellm_completion_responses.py @@ -2841,9 +2841,9 @@ class TestStreamingIDConsistency: # Verify the cached ID is set and matches assert iterator._cached_item_id is not None, "Iterator should cache the item_id" assert iterator._cached_item_id == item_id_1, "Cached ID should match event IDs" - assert ( - iterator._cached_item_id == "chatcmpl-first-id" - ), "Should use the first chunk's ID" + assert iterator._cached_item_id.startswith( + "msg_" + ), "Message item IDs must use the Responses API msg_ prefix (issue #27333)" def test_streaming_iterator_initial_events_use_cached_id(self): """ @@ -3771,3 +3771,234 @@ def test_function_call_tool_id_falls_back_to_unique_id_for_degenerate_call_id(): id="fc_2", call_id="call_tokyo", name="get_weather", arguments="{}" ) assert convert(openai)["id"] == "call_tokyo" + + +BRIDGED_CHAT_COMPLETION_ID = "chatcmpl-dfa2da3a-1586-4ff7-b64e-f59c692a5d11" + + +def _bridged_chat_completion_response(**overrides): + defaults = dict( + id=BRIDGED_CHAT_COMPLETION_ID, + created=1717000000, + model="claude-sonnet-4-5", + object="chat.completion", + choices=[ + Choices( + index=0, + finish_reason="stop", + message=Message(role="assistant", content="apple"), + ) + ], + usage=Usage(prompt_tokens=10, completion_tokens=5, total_tokens=15), + ) + defaults.update(overrides) + return ModelResponse(**defaults) + + +def _bridged_output_items(response, item_type): + return [item for item in response.output if getattr(item, "type", None) == item_type] + + +class TestBridgedOutputItemIdPrefixes: + """Bridged output items must carry Responses API ID prefixes (issue #27333). + + Native OpenAI Responses rejects a replayed history whose message item ID does not + begin with "msg", so leaking the upstream chatcmpl-* ID makes the conversation + impossible to hand off from a bridged provider to OpenAI. + """ + + def _transform(self, chat_completion_response): + return LiteLLMCompletionResponsesConfig.transform_chat_completion_response_to_responses_api_response( + request_input="Say the single word: apple", + responses_api_request={}, + chat_completion_response=chat_completion_response, + ) + + def test_message_item_id_uses_msg_prefix(self): + response = self._transform(_bridged_chat_completion_response()) + + message_items = _bridged_output_items(response, "message") + assert len(message_items) == 1 + assert message_items[0].id.startswith("msg_") + + def test_message_item_id_does_not_leak_chat_completion_id(self): + response = self._transform(_bridged_chat_completion_response()) + + for item in _bridged_output_items(response, "message"): + assert item.id != BRIDGED_CHAT_COMPLETION_ID + assert not item.id.startswith("chatcmpl-") + + def test_message_item_ids_are_unique_across_responses(self): + first = self._transform(_bridged_chat_completion_response()) + second = self._transform(_bridged_chat_completion_response()) + + first_id = _bridged_output_items(first, "message")[0].id + second_id = _bridged_output_items(second, "message")[0].id + assert first_id != second_id + + def _reasoning_items(self): + message = Message(role="assistant", content="apple") + message.reasoning_content = "thinking about fruit" + choice = Choices(index=0, finish_reason="stop", message=message) + return LiteLLMCompletionResponsesConfig._extract_reasoning_output_items( + chat_completion_response=_bridged_chat_completion_response(), + choices=[choice], + ) + + def test_reasoning_item_id_uses_rs_prefix(self): + items = self._reasoning_items() + + assert len(items) == 1 + assert items[0].id.startswith("rs_") + + def test_reasoning_item_id_is_not_a_salted_hash(self): + """Python's hash() is salted per process, so the old rs_{hash(...)} ID for the + same reasoning text differed between workers and across restarts.""" + suffix = self._reasoning_items()[0].id.removeprefix("rs_") + + assert not suffix.lstrip("-").isdigit() + assert not suffix.startswith("-") + + +class TestStreamingSnapshotItemIds: + """The response.completed snapshot must reuse the streamed item ID (issue #27333). + + The incremental events already minted msg_* IDs while the final snapshot went back + through the non-streaming transform, so a streaming client replaying the snapshot + sent back an ID it had never been shown. + """ + + def _make_iterator(self): + from unittest.mock import Mock + + import litellm + from litellm.responses.litellm_completion_transformation.streaming_iterator import ( + LiteLLMCompletionStreamingIterator, + ) + + mock_stream_wrapper = Mock(spec=litellm.CustomStreamWrapper) + mock_stream_wrapper.logging_obj = Mock() + return LiteLLMCompletionStreamingIterator( + model="anthropic/claude-sonnet-4-5", + litellm_custom_stream_wrapper=mock_stream_wrapper, + request_input="Say the single word: apple", + responses_api_request={}, + custom_llm_provider="anthropic", + ) + + def _make_chunk(self, content): + from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices + + return ModelResponseStream( + id=BRIDGED_CHAT_COMPLETION_ID, + choices=[ + StreamingChoices( + index=0, + delta=Delta(content=content, role="assistant"), + finish_reason=None, + ) + ], + created=1717000000, + model="claude-sonnet-4-5", + object="chat.completion.chunk", + ) + + def test_incremental_item_id_uses_msg_prefix(self): + iterator = self._make_iterator() + + event = iterator._transform_chat_completion_chunk_to_response_api_chunk( + self._make_chunk("apple") + ) + + assert event is not None + assert event.item_id.startswith("msg_") + assert event.item_id != BRIDGED_CHAT_COMPLETION_ID + + def test_completed_snapshot_reuses_streamed_item_id(self): + iterator = self._make_iterator() + + streamed_event = iterator._transform_chat_completion_chunk_to_response_api_chunk( + self._make_chunk("apple") + ) + assert streamed_event is not None + + completed_event = iterator._emit_response_completed_event( + _bridged_chat_completion_response() + ) + + assert completed_event is not None + message_items = _bridged_output_items(completed_event.response, "message") + assert len(message_items) == 1 + assert message_items[0].id == streamed_event.item_id + + def test_completed_snapshot_item_id_is_replayable(self): + iterator = self._make_iterator() + iterator._transform_chat_completion_chunk_to_response_api_chunk( + self._make_chunk("apple") + ) + + completed_event = iterator._emit_response_completed_event( + _bridged_chat_completion_response() + ) + + assert completed_event is not None + for item in _bridged_output_items(completed_event.response, "message"): + assert item.id.startswith("msg_") + assert not item.id.startswith("chatcmpl-") + + def _make_reasoning_chunk(self, reasoning_content): + from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices + + return ModelResponseStream( + id=BRIDGED_CHAT_COMPLETION_ID, + choices=[ + StreamingChoices( + index=0, + delta=Delta(role="assistant", reasoning_content=reasoning_content), + finish_reason=None, + ) + ], + created=1717000000, + model="claude-sonnet-4-5", + object="chat.completion.chunk", + ) + + def _reasoning_chat_completion_response(self): + message = Message(role="assistant", content="apple") + message.reasoning_content = "thinking about fruit" + return _bridged_chat_completion_response( + choices=[Choices(index=0, finish_reason="stop", message=message)] + ) + + def test_reasoning_delta_events_share_one_item_id(self): + """The old rs_{hash(text)} ID changed with every delta, so a client accumulating + reasoning by item ID saw a new item per chunk.""" + iterator = self._make_iterator() + + first = iterator._transform_chat_completion_chunk_to_response_api_chunk( + self._make_reasoning_chunk("thinking ") + ) + second = iterator._transform_chat_completion_chunk_to_response_api_chunk( + self._make_reasoning_chunk("about fruit") + ) + + assert first is not None and second is not None + assert first.item_id.startswith("rs_") + assert first.item_id == second.item_id + + def test_completed_snapshot_reuses_streamed_reasoning_item_id(self): + iterator = self._make_iterator() + + streamed_event = iterator._transform_chat_completion_chunk_to_response_api_chunk( + self._make_reasoning_chunk("thinking about fruit") + ) + assert streamed_event is not None + + completed_event = iterator._emit_response_completed_event( + self._reasoning_chat_completion_response() + ) + + assert completed_event is not None + reasoning_items = _bridged_output_items(completed_event.response, "reasoning") + assert len(reasoning_items) == 1 + assert reasoning_items[0].id == streamed_event.item_id diff --git a/tests/test_litellm/responses/litellm_completion_transformation/test_reasoning_input_item_preservation.py b/tests/test_litellm/responses/litellm_completion_transformation/test_reasoning_input_item_preservation.py new file mode 100644 index 00000000000..5e001bdbbbb --- /dev/null +++ b/tests/test_litellm/responses/litellm_completion_transformation/test_reasoning_input_item_preservation.py @@ -0,0 +1,384 @@ +""" +Unit tests for preserving prior-turn ``reasoning`` input items when the +Responses API is bridged to chat completions. + +Without this handling, a ``ResponseReasoningItemParam`` falls through to the +generic message branch, polluting the prompt as visible assistant ``content`` +or being silently dropped. Chat-completions providers such as DeepSeek V4 and +Kimi K2.6 require the chain-of-thought to be replayed as ``reasoning_content`` +on an assistant message. + +Providers whose reasoning is signed (Anthropic, Bedrock converse) get their +blocks back through ``encrypted_content``, which LiteLLM itself writes as a +JSON array of thinking blocks on the response side. +""" + +import json + +import pytest + +from litellm.responses.litellm_completion_transformation.transformation import ( + LiteLLMCompletionResponsesConfig, +) +from litellm.types.utils import Message + + +def _transform_item(item): + return LiteLLMCompletionResponsesConfig._transform_responses_api_input_item_to_chat_completion_message( + input_item=item, replay_reasoning=True + ) + + +def _transform_input(input_items): + return LiteLLMCompletionResponsesConfig._transform_response_input_param_to_chat_completion_message( + input=input_items, replay_reasoning=True + ) + + +def _inspect_input(input_items): + return LiteLLMCompletionResponsesConfig.transform_responses_api_input_to_messages( + input=input_items, responses_api_request={} + ) + + +class TestReasoningInputItemHandler: + """Reasoning input items map to assistant ``reasoning_content``.""" + + def test_reasoning_item_with_output_text_content(self): + """Standard Responses-API reasoning item with output_text blocks.""" + item = { + "type": "reasoning", + "id": "rs_abc", + "summary": [], + "content": [{"type": "output_text", "text": "step 1: think about X"}], + } + messages = _transform_item(item) + assert len(messages) == 1 + assert messages[0]["role"] == "assistant" + assert messages[0]["content"] is None + assert messages[0]["reasoning_content"] == "step 1: think about X" + + def test_reasoning_item_with_string_content(self): + """Variant: reasoning content as a plain string.""" + item = {"type": "reasoning", "id": "rs_1", "content": "step 1: ..."} + messages = _transform_item(item) + assert messages[0]["reasoning_content"] == "step 1: ..." + + def test_reasoning_item_with_summary_only(self): + """SDK form: reasoning carried in summary list, no content.""" + item = { + "type": "reasoning", + "id": "rs_2", + "summary": [{"type": "summary_text", "text": "..."}], + } + messages = _transform_item(item) + assert messages[0]["reasoning_content"] == "..." + + def test_reasoning_item_with_opaque_encrypted_content_dropped(self): + """An encrypted blob LiteLLM did not write cannot be forwarded.""" + item = {"type": "reasoning", "id": "rs_3", "encrypted_content": "opaque-blob"} + assert _transform_item(item) == [] + + def test_reasoning_item_empty_dropped(self): + """Reasoning item with neither content nor summary drops cleanly.""" + assert _transform_item({"type": "reasoning", "id": "rs_4"}) == [] + + +class TestReasoningInputItemMerging: + """Standalone reasoning messages merge into the following assistant turn.""" + + def test_reasoning_merged_into_following_assistant_message(self): + """Reasoning + assistant answer become one assistant message.""" + messages = _transform_input( + [ + { + "type": "reasoning", + "id": "rs_1", + "content": [{"type": "output_text", "text": "secret reasoning"}], + }, + {"type": "message", "role": "assistant", "content": "The answer."}, + ] + ) + assert len(messages) == 1 + assert messages[0]["role"] == "assistant" + assert messages[0]["content"] == "The answer." + assert messages[0]["reasoning_content"] == "secret reasoning" + + def test_reasoning_preserved_when_followed_by_user_message(self): + """Stateless chain: reasoning + user prompt keeps the reasoning turn.""" + messages = _transform_input( + [ + { + "type": "reasoning", + "id": "rs_1", + "content": [{"type": "output_text", "text": "secret BLUEBERRY"}], + }, + {"role": "user", "content": "What is the secret word?"}, + ] + ) + assert len(messages) == 2 + assert messages[0]["role"] == "assistant" + assert messages[0]["content"] is None + assert messages[0]["reasoning_content"] == "secret BLUEBERRY" + assert messages[1]["role"] == "user" + + def test_reasoning_merged_into_function_call_assistant(self): + """Reasoning + function_call becomes one assistant tool-call message.""" + messages = _transform_input( + [ + { + "type": "reasoning", + "id": "rs_1", + "content": [{"type": "output_text", "text": "I should look this up"}], + }, + { + "type": "function_call", + "call_id": "call_1", + "name": "lookup", + "arguments": '{"cwe": "79"}', + }, + ] + ) + assert len(messages) == 1 + assert messages[0]["role"] == "assistant" + assert messages[0]["reasoning_content"] == "I should look this up" + assert len(messages[0]["tool_calls"]) == 1 + + def test_reasoning_merged_into_assistant_with_existing_reasoning_content(self): + """Old reasoning precedes existing reasoning on the target assistant turn.""" + messages = LiteLLMCompletionResponsesConfig._merge_reasoning_only_assistant_messages( + [ + {"role": "assistant", "content": None, "reasoning_content": "old reasoning"}, + {"role": "assistant", "content": "The answer.", "reasoning_content": "new reasoning"}, + ] + ) + assert len(messages) == 1 + assert messages[0]["content"] == "The answer." + assert messages[0]["reasoning_content"] == "old reasoning\nnew reasoning" + + +class TestEncryptedReasoningRoundTrip: + """``encrypted_content`` LiteLLM wrote decodes back into thinking blocks.""" + + def test_encoded_thinking_blocks_decode_back(self): + """The decoder is the inverse of the encoder the response side uses.""" + blocks = [ + {"type": "thinking", "thinking": "step one", "signature": "sig-one"}, + {"type": "redacted_thinking", "data": "redacted-payload"}, + ] + message = Message(role="assistant", content="answer", thinking_blocks=blocks) + encoded = LiteLLMCompletionResponsesConfig._encode_thinking_blocks(message) + decoded = LiteLLMCompletionResponsesConfig._decode_thinking_blocks_from_input_item( + {"type": "reasoning", "encrypted_content": encoded} + ) + assert list(decoded) == blocks + + def test_signed_thinking_blocks_replayed_on_assistant_message(self): + """A signed block survives the bridge instead of vanishing.""" + item = { + "type": "reasoning", + "id": "rs_1", + "encrypted_content": json.dumps( + [{"type": "thinking", "thinking": "hidden", "signature": "sig-one"}] + ), + } + messages = _transform_item(item) + assert len(messages) == 1 + assert messages[0]["content"] is None + assert messages[0]["thinking_blocks"] == [ + {"type": "thinking", "thinking": "hidden", "signature": "sig-one"} + ] + + def test_unsigned_blocks_dropped(self): + """Blocks without a signature or redacted payload are not replayed.""" + item = { + "type": "reasoning", + "id": "rs_2", + "encrypted_content": json.dumps([{"type": "thinking", "thinking": "unsigned"}]), + } + assert _transform_item(item) == [] + + def test_json_object_encrypted_content_dropped(self): + """A JSON payload that is not a block array is treated as opaque.""" + item = { + "type": "reasoning", + "id": "rs_3", + "encrypted_content": json.dumps({"ciphertext": "abc"}), + } + assert _transform_item(item) == [] + + def test_thinking_blocks_merged_onto_tool_call_assistant(self): + """Signed reasoning lands on the assistant turn carrying the tool call.""" + messages = _transform_input( + [ + { + "type": "reasoning", + "id": "rs_1", + "summary": [{"type": "summary_text", "text": "look it up"}], + "encrypted_content": json.dumps( + [{"type": "thinking", "thinking": "hidden", "signature": "sig-one"}] + ), + }, + { + "type": "function_call", + "call_id": "call_1", + "name": "lookup", + "arguments": '{"cwe": "79"}', + }, + ] + ) + assert len(messages) == 1 + assert messages[0]["reasoning_content"] == "look it up" + assert messages[0]["thinking_blocks"] == [ + {"type": "thinking", "thinking": "hidden", "signature": "sig-one"} + ] + assert len(messages[0]["tool_calls"]) == 1 + + def test_replayed_blocks_precede_existing_blocks(self): + """Signature verification depends on the original block order.""" + messages = LiteLLMCompletionResponsesConfig._merge_reasoning_only_assistant_messages( + [ + { + "role": "assistant", + "content": None, + "thinking_blocks": [{"type": "thinking", "thinking": "older", "signature": "a"}], + }, + { + "role": "assistant", + "content": "answer", + "thinking_blocks": [{"type": "thinking", "thinking": "newer", "signature": "b"}], + }, + ] + ) + assert len(messages) == 1 + assert [block["thinking"] for block in messages[0]["thinking_blocks"]] == ["older", "newer"] + + def test_encrypted_only_reasoning_preserved_before_user_turn(self): + """A signed item with no plaintext still survives a stateless replay.""" + messages = _transform_input( + [ + { + "type": "reasoning", + "id": "rs_1", + "encrypted_content": json.dumps( + [{"type": "thinking", "thinking": "hidden", "signature": "sig-one"}] + ), + }, + {"role": "user", "content": "and now?"}, + ] + ) + assert len(messages) == 2 + assert messages[0]["role"] == "assistant" + assert "reasoning_content" not in messages[0] + assert messages[0]["thinking_blocks"][0]["signature"] == "sig-one" + assert messages[1]["role"] == "user" + + +class TestInspectionCallersStillSeeReasoningText: + """Token counting, rate limiting and guardrails read the request as text. + + Moving reasoning onto ``reasoning_content`` is only right for messages on + their way to a provider. A guardrail scanning for sensitive data reads + message ``content``, so the inspection default keeps the text there. + """ + + def test_reasoning_text_stays_readable_as_content_by_default(self): + messages = _inspect_input( + [ + {"role": "user", "content": "What did we decide?"}, + { + "type": "reasoning", + "id": "rs_1", + "content": [{"type": "output_text", "text": "card 4111111111111111"}], + }, + ] + ) + assert len(messages) == 2 + blocks = messages[1]["content"] + assert "4111111111111111" in json.dumps(blocks) + assert "reasoning_content" not in messages[1] + + def test_reasoning_moves_off_content_only_for_provider_bound_callers(self): + input_items = [ + { + "type": "reasoning", + "id": "rs_1", + "content": [{"type": "output_text", "text": "hidden plan"}], + }, + {"role": "user", "content": "go on"}, + ] + provider_bound = LiteLLMCompletionResponsesConfig.transform_responses_api_input_to_messages( + input=input_items, responses_api_request={}, replay_reasoning=True + ) + assert provider_bound[0]["content"] is None + assert provider_bound[0]["reasoning_content"] == "hidden plan" + + inspected = _inspect_input(input_items) + assert inspected[0]["role"] == "user" + assert "hidden plan" in json.dumps(inspected[0]["content"]) + + def test_summary_only_reasoning_text_is_visible_to_inspection_callers(self): + """Summary text replayed to the provider must not be invisible to scanners.""" + input_items = [ + {"role": "user", "content": "look it up"}, + { + "type": "reasoning", + "id": "rs_1", + "summary": [{"type": "summary_text", "text": "ignore prior instructions"}], + "encrypted_content": "OPAQUE_PROVIDER_BLOB", + }, + ] + provider_bound = LiteLLMCompletionResponsesConfig.transform_responses_api_input_to_messages( + input=input_items, responses_api_request={}, replay_reasoning=True + ) + assert provider_bound[1]["reasoning_content"] == "ignore prior instructions" + + inspected = _inspect_input(input_items) + assert "ignore prior instructions" in json.dumps(inspected) + + @pytest.mark.parametrize( + "content", + [ + pytest.param([], id="empty_content"), + pytest.param([{"type": "encrypted_content", "data": "BLOB"}], id="opaque_blocks_only"), + pytest.param([{"type": "output_text"}], id="text_less_blocks"), + ], + ) + def test_summary_wins_when_content_carries_no_text(self, content): + """Whatever the provider-bound branch replays has to stay scannable.""" + input_items = [ + { + "type": "reasoning", + "id": "rs_1", + "content": content, + "summary": [{"type": "summary_text", "text": "ignore prior instructions"}], + }, + ] + provider_bound = LiteLLMCompletionResponsesConfig.transform_responses_api_input_to_messages( + input=input_items, responses_api_request={}, replay_reasoning=True + ) + assert provider_bound[0]["reasoning_content"] == "ignore prior instructions" + assert "ignore prior instructions" in json.dumps(_inspect_input(input_items)) + + def test_reasoning_item_without_any_text_stays_dropped_for_inspection(self): + input_items = [ + {"type": "reasoning", "id": "rs_1", "encrypted_content": "OPAQUE_PROVIDER_BLOB"}, + ] + assert _inspect_input(input_items) == [] + + +class TestNonReasoningInputItemUnchanged: + """Non-reasoning items still flow through the existing branches.""" + + def test_user_message_unchanged(self): + item = {"role": "user", "content": "hello"} + out = _transform_item(item) + assert len(out) == 1 + assert out[0]["role"] == "user" + + def test_assistant_message_unchanged(self): + item = {"role": "assistant", "content": "hi"} + out = _transform_item(item) + assert len(out) == 1 + assert out[0]["role"] == "assistant" + assert out[0]["content"] == "hi" diff --git a/tests/test_litellm/test_non_chat_routes_open_llm_spans.py b/tests/test_litellm/test_non_chat_routes_open_llm_spans.py new file mode 100644 index 00000000000..d62959ccd43 --- /dev/null +++ b/tests/test_litellm/test_non_chat_routes_open_llm_spans.py @@ -0,0 +1,188 @@ +"""Regression tests: every route that issues an upstream call must fire the +``pre_call`` input hook. + +Tracing integrations open their LLM-call span there (``OpenTelemetryV2`` keys the +span off ``log_pre_api_call`` and treats "no pre_call" as "the request never +reached a provider"), so a handler that skips it leaves the call with no LLM-call +span in the trace at all. Speech, async image generation and moderation each used +to skip it. +""" + +import asyncio +from typing import Any, Final + +import httpx +import pytest +from openai import AsyncAzureOpenAI, AsyncOpenAI + +import litellm +from litellm.integrations.custom_logger import CustomLogger + + +class _PreCallRecorder(CustomLogger): + def __init__(self) -> None: + super().__init__() + self.call_types: list[str] = [] # mutable-ok: test recorder of hook calls + self.api_bases: list[str] = [] # mutable-ok: test recorder of hook calls + self.request_bodies: list[Any] = [] # mutable-ok: test recorder of hook calls + + def log_pre_api_call(self, model, messages, kwargs) -> None: + self.call_types.append(str(kwargs.get("call_type"))) + self.api_bases.append(str(kwargs.get("litellm_params", {}).get("api_base"))) + self.request_bodies.append(kwargs.get("additional_args", {}).get("complete_input_dict")) + + +class _FakeSpeech: + def __init__(self) -> None: + self.calls: list[dict[str, Any]] = [] # mutable-ok: test recorder of SDK calls + + async def create(self, **kwargs: Any) -> Any: + self.calls.append(kwargs) + request: Final = httpx.Request("POST", "https://api.openai.com/v1/audio/speech") + return type( + "_Speech", + (), + {"response": httpx.Response(200, content=b"audio-bytes", request=request)}, + )() + + +class _FakeImages: + async def generate(self, **kwargs: Any) -> Any: + return type( + "_Images", + (), + { + "model_dump": lambda self: { + "created": 1, + "data": [{"url": "https://example.com/img.png"}], + } + }, + )() + + +class _FakeModerations: + async def create(self, **kwargs: Any) -> Any: + return type( + "_Moderations", + (), + { + "model_dump": lambda self: { + "id": "modr-1", + "model": "omni-moderation-latest", + "results": [ + { + "flagged": False, + "categories": {}, + "category_scores": {}, + "category_applied_input_types": {}, + } + ], + } + }, + )() + + +class _FakeAsyncOpenAI(AsyncOpenAI): + """Stands in for the injected client: a real ``AsyncOpenAI`` (``amoderation`` + type-checks it) whose resource namespaces answer without a network call.""" + + def __init__(self, base_url: str = "https://api.openai.com/v1") -> None: + super().__init__(api_key="sk-test", base_url=base_url) + self.speech = _FakeSpeech() + self.audio = type("_Audio", (), {"speech": self.speech})() + self.images = _FakeImages() + self.moderations = _FakeModerations() + + +class _FakeAsyncAzureOpenAI(AsyncAzureOpenAI): + """Same idea for the Azure entrypoint, which resolves no default endpoint of + its own when ``AZURE_API_BASE`` is unset.""" + + def __init__(self) -> None: + super().__init__( + api_key="sk-test", + api_version="2024-02-01", + azure_endpoint="https://unit-test.openai.azure.com", + ) + self.speech = _FakeSpeech() + self.audio = type("_Audio", (), {"speech": self.speech})() + + +@pytest.fixture +def recorder(monkeypatch): + recorder: Final = _PreCallRecorder() + monkeypatch.setattr(litellm, "callbacks", [recorder]) + monkeypatch.setattr(litellm, "success_callback", []) + return recorder + + +def test_async_speech_opens_an_llm_span(recorder): + asyncio.run( + litellm.aspeech( + model="openai/tts-1", + input="hello", + voice="alloy", + client=_FakeAsyncOpenAI(), + ) + ) + assert recorder.call_types == ["aspeech"] + + +def test_azure_async_speech_opens_an_llm_span_without_api_base(recorder, monkeypatch): + """Azure resolves no default endpoint, so a missing ``api_base`` used to reach + ``_get_masked_api_base`` as ``None``; the ``TypeError`` was swallowed and the + whole callback dispatch was skipped.""" + monkeypatch.delenv("AZURE_API_BASE", raising=False) + asyncio.run( + litellm.aspeech( + model="azure/tts-deployment", + input="hello", + voice="alloy", + client=_FakeAsyncAzureOpenAI(), + ) + ) + assert recorder.call_types == ["aspeech"] + assert recorder.api_bases == ["https://unit-test.openai.azure.com/openai/"] + + +def test_azure_async_speech_keeps_caller_headers_out_of_the_logged_body(recorder): + """The Azure entrypoint carries caller headers in ``optional_params``, so they reach + the provider as a request kwarg; telemetry reads the logged body, which must stay + free of them.""" + headers: Final = {"authorization": "Bearer caller-secret"} + client: Final = _FakeAsyncAzureOpenAI() + asyncio.run( + litellm.aspeech( + model="azure/tts-deployment", + input="hello", + voice="alloy", + extra_headers=headers, + client=client, + ) + ) + assert recorder.call_types == ["aspeech"] + assert "extra_headers" not in recorder.request_bodies[0] + assert client.speech.calls[0]["extra_headers"] == headers + + +def test_async_image_generation_opens_an_llm_span(recorder): + asyncio.run( + litellm.aimage_generation( + model="openai/dall-e-3", + prompt="a cat", + client=_FakeAsyncOpenAI(), + ) + ) + assert recorder.call_types == ["aimage_generation"] + + +def test_async_moderation_opens_an_llm_span(recorder): + asyncio.run( + litellm.amoderation( + model="omni-moderation-latest", + input="hello", + client=_FakeAsyncOpenAI(base_url="https://gateway.example/v1"), + ) + ) + assert recorder.call_types == ["amoderation"] + assert recorder.api_bases == ["https://gateway.example/v1/"] diff --git a/tests/test_litellm/types/llms/test_types_llms_openai.py b/tests/test_litellm/types/llms/test_types_llms_openai.py index e5e5c0183a0..3966677e928 100644 --- a/tests/test_litellm/types/llms/test_types_llms_openai.py +++ b/tests/test_litellm/types/llms/test_types_llms_openai.py @@ -450,3 +450,75 @@ def test_openai_file_object_accepts_pending_status(): status="pending", ) assert file_obj.status == "pending" + + +class TestOpenAIFileObjectBatchGuardrailSerialization: + """The proxy-only `litellm_batch_guardrail` key must reach the wire only when something set it.""" + + @staticmethod + def _file_object(**overrides): + from litellm.types.llms.openai import OpenAIFileObject + + return OpenAIFileObject( + id="file-123", + object="file", + bytes=1024, + created_at=1677610602, + filename="batch.jsonl", + purpose="batch", + status="uploaded", + **overrides, + ) + + @staticmethod + def _report(): + from litellm.types.llms.openai import BatchGuardrailRecord, BatchGuardrailReport + + return BatchGuardrailReport( + submitted_records=3, + modified_records=(BatchGuardrailRecord(line=2, custom_id="dirty", action="redacted"),), + ) + + @pytest.mark.parametrize("mode", ["python", "json"]) + def test_key_absent_when_unset(self, mode): + assert "litellm_batch_guardrail" not in self._file_object().model_dump(mode=mode) + + @pytest.mark.parametrize("mode", ["python", "json"]) + def test_key_present_when_set(self, mode): + dumped = self._file_object(litellm_batch_guardrail=self._report()).model_dump(mode=mode) + assert dumped["litellm_batch_guardrail"]["submitted_records"] == 3 + + def test_nested_nulls_of_a_set_report_survive(self): + """`exclude_none=True` was rejected as the fix because it would strip these.""" + dumped = self._file_object(litellm_batch_guardrail=self._report()).model_dump(mode="json") + assert dumped["litellm_batch_guardrail"]["modified_records"] == [ + {"line": 2, "custom_id": "dirty", "action": "redacted", "guardrail": None} + ] + + def test_by_alias_dump_also_omits_the_key(self): + """Tripwire: the serializer filters a literal key name, which an added alias would bypass.""" + assert "litellm_batch_guardrail" not in self._file_object().model_dump(mode="json", by_alias=True) + + def test_other_optional_fields_still_serialize_as_null(self): + dumped = self._file_object().model_dump(mode="json") + assert dumped["expires_at"] is None + assert dumped["status_details"] is None + + def test_round_trip_of_a_set_report_is_lossless(self): + from litellm.types.llms.openai import OpenAIFileObject + + original = self._file_object(litellm_batch_guardrail=self._report()) + assert OpenAIFileObject(**original.model_dump()) == original + + def test_serialization_json_schema_still_describes_the_model(self): + """A return annotation on the wrap serializer would collapse this to a bare object.""" + from litellm.types.llms.openai import OpenAIFileObject + + schema = OpenAIFileObject.model_json_schema(mode="serialization") + assert "litellm_batch_guardrail" in schema["properties"] + + def test_key_omitted_inside_a_file_list_page(self): + from litellm.types.llms.openai import FileListPage + + page = FileListPage(object="list", data=[self._file_object()], has_more=False) + assert "litellm_batch_guardrail" not in page.model_dump(mode="json")["data"][0] diff --git a/ui/litellm-dashboard/src/app/(dashboard)/playground/components/chat_ui/EndpointUtils.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/playground/components/chat_ui/EndpointUtils.test.tsx index 778effdaea6..c528071fa8a 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/playground/components/chat_ui/EndpointUtils.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/playground/components/chat_ui/EndpointUtils.test.tsx @@ -247,6 +247,13 @@ describe("isModelCompatibleWithEndpoint / filterModelsForEndpoint", () => { expect(isModelCompatibleWithEndpoint(batchModel, EndpointType.REALTIME)).toBe(false); }); + it("keeps completion-mode models for the chat endpoint", () => { + const completionModel: ModelGroup = { model_group: "davinci-002", mode: "completion" }; + expect(isModelCompatibleWithEndpoint(completionModel, EndpointType.CHAT)).toBe(true); + expect(isModelCompatibleWithEndpoint(completionModel, EndpointType.RESPONSES)).toBe(true); + expect(isModelCompatibleWithEndpoint(completionModel, EndpointType.SPEECH)).toBe(false); + }); + it("keeps image-edit models for the image-edits endpoint using the mode the backend sends", () => { const imageEditModel: ModelGroup = { model_group: "gpt-image-1", mode: "image_edit" }; const imageModel: ModelGroup = { model_group: "dall-e-3", mode: "image_generation" }; diff --git a/ui/litellm-dashboard/src/components/chat_ui/mode_endpoint_mapping.tsx b/ui/litellm-dashboard/src/components/chat_ui/mode_endpoint_mapping.tsx index 18e44e06efe..930ded5d1a5 100644 --- a/ui/litellm-dashboard/src/components/chat_ui/mode_endpoint_mapping.tsx +++ b/ui/litellm-dashboard/src/components/chat_ui/mode_endpoint_mapping.tsx @@ -7,6 +7,7 @@ export enum ModelMode { IMAGE_GENERATION = "image_generation", VIDEO_GENERATION = "video_generation", CHAT = "chat", + COMPLETION = "completion", RESPONSES = "responses", IMAGE_EDITS = "image_edit", ANTHROPIC_MESSAGES = "anthropic_messages", @@ -36,6 +37,7 @@ export const litellmModeMapping: Record = { [ModelMode.IMAGE_GENERATION]: EndpointType.IMAGE, [ModelMode.VIDEO_GENERATION]: EndpointType.VIDEO, [ModelMode.CHAT]: EndpointType.CHAT, + [ModelMode.COMPLETION]: EndpointType.CHAT, [ModelMode.RESPONSES]: EndpointType.RESPONSES, [ModelMode.IMAGE_EDITS]: EndpointType.IMAGE_EDITS, [ModelMode.ANTHROPIC_MESSAGES]: EndpointType.ANTHROPIC_MESSAGES, diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index cf55dc69e86..d311bfa3cbc 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -43144,6 +43144,8 @@ export interface operations { provider?: string | null; target_model_names?: string | null; purpose?: string | null; + limit?: number | null; + after?: string | null; }; header?: never; path?: never; @@ -58015,6 +58017,8 @@ export interface operations { provider?: string | null; target_model_names?: string | null; purpose?: string | null; + limit?: number | null; + after?: string | null; }; header?: never; path?: never; @@ -64268,6 +64272,8 @@ export interface operations { query?: { target_model_names?: string | null; purpose?: string | null; + limit?: number | null; + after?: string | null; }; header?: never; path: {