Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py: 53%
1349 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
1import ast
2import asyncio
3import copy
4import json
5import posixpath
6import traceback
7from base64 import b64encode
8from collections.abc import AsyncGenerator, AsyncIterator, Callable, Iterable, Mapping, Sequence
9from dataclasses import dataclass
10from datetime import datetime
11from itertools import count, groupby
12from types import MappingProxyType
13from typing import TYPE_CHECKING, Any, Final, TypedDict, cast
14from urllib.parse import urlencode, urlparse
16import httpx
17from fastapi import (
18 APIRouter,
19 Depends,
20 FastAPI,
21 HTTPException,
22 Request,
23 Response,
24 UploadFile,
25 WebSocket,
26 status,
27)
28from fastapi.responses import StreamingResponse
29from starlette.datastructures import UploadFile as StarletteUploadFile
30from starlette.routing import BaseRoute, Route
31from starlette.websockets import WebSocketState
32from websockets.asyncio.client import connect
33from websockets.exceptions import (
34 ConnectionClosedError,
35 ConnectionClosedOK,
36 InvalidStatus,
37)
38from websockets.frames import Close, CloseCode
40import litellm
41from litellm._logging import verbose_proxy_logger
42from litellm._uuid import uuid
43from litellm.constants import (
44 MAXIMUM_TRACEBACK_LINES_TO_LOG,
45 PASSTHROUGH_UPSTREAM_ERROR_BODY_MAX_LOG_CHARS,
46 REDACTED_BY_LITELLM,
47 SESSION_ID_OMITTED_METADATA_KEY,
48 WEBSOCKET_CLOSE_REASON_MAX_BYTES,
49)
50from litellm.integrations.custom_guardrail import CustomGuardrail
51from litellm.integrations.custom_logger import CustomLogger
52from litellm.litellm_core_utils.core_helpers import (
53 bind_budget_reservation_to_callbacks,
54 get_metadata_variable_name_from_kwargs,
55 get_or_create_metadata_bucket,
56)
57from litellm.litellm_core_utils.initialize_dynamic_callback_params import validate_no_callback_env_reference
58from litellm.litellm_core_utils.internal_call_metadata import MODEL_ACCESS_GROUP_METADATA_KEY
59from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
60from litellm.litellm_core_utils.litellm_logging import _get_masked_values
61from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
62from litellm.litellm_core_utils.redact_messages import should_redact_message_logging
63from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
64from litellm.llms.base_llm.managed_resources.utils import (
65 resolve_passthrough_managed_id_provider,
66)
67from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
68from litellm.passthrough import BasePassthroughUtils
69from litellm.proxy._lazy_features import lazy_owned_routes
70from litellm.proxy._types import (
71 ConfigFieldInfo,
72 ConfigFieldUpdate,
73 LiteLLMRoutes,
74 PassThroughEndpointResponse,
75 PassThroughGenericEndpoint,
76 ProxyException,
77 UserAPIKeyAuth,
78)
79from litellm.proxy.auth.auth_utils import request_dispatched_to_pass_through_endpoint
80from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
81from litellm.proxy.common_request_processing import (
82 ProxyBaseLLMRequestProcessing,
83 log_llm_api_exception,
84 open_sse_before_first_byte,
85 resolve_litellm_call_id,
86)
87from litellm.proxy.common_utils.error_body_call_id import JSON_OBJECT, error_body_call_id, with_call_id
88from litellm.proxy.common_utils.http_parsing_utils import (
89 _read_request_body,
90 _safe_get_request_headers,
91)
92from litellm.proxy.common_utils.openai_error_payload import (
93 LITELLM_CALL_ID_HEADER,
94 error_status_code,
95 litellm_call_id_headers,
96 openai_error_param,
97 openai_error_type,
98)
99from litellm.proxy.common_utils.sse_keepalive import (
100 wrap_passthrough_sse_bytes_with_keepalive_pings,
101)
102from litellm.proxy.litellm_pre_call_utils import (
103 LiteLLMProxyRequestSetup,
104 _get_dynamic_logging_metadata, # pyright: ignore[reportPrivateUsage] # shared proxy helper, same import style as _read_request_body above
105)
106from litellm.proxy.route_llm_request import ProxyModelNotFoundError
107from litellm.proxy.utils import normalize_route_for_root_path
108from litellm.repositories.team_repository import TeamRepository
109from litellm.secret_managers.main import get_secret_str
110from litellm.types import utils as types_utils
111from litellm.types.litellm_params import ProxyRequestState, wire_names
112from litellm.types.llms.custom_http import httpxSpecialProvider
113from litellm.types.passthrough_endpoints.pass_through_endpoints import (
114 LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY,
115 LITELLM_PASS_THROUGH_DEPLOYMENT_MODEL_INFO_STATE_KEY,
116 LITELLM_PASS_THROUGH_ENDPOINT_MARKER,
117 LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY,
118 EndpointType,
119 PassthroughStandardLoggingPayload,
120)
121from litellm.types.utils import TRUSTED_CALLBACK_VARS_FIELD, Usage
123from .llm_provider_handlers.tinyfish_passthrough_logging_handler import (
124 is_tinyfish_agent_url,
125)
126from .streaming_handler import PassThroughStreamingHandler
127from .success_handler import PassThroughEndpointLogging
128from .upstream_usage_headers import (
129 UpstreamReportedUsage,
130 apply_upstream_reported_usage,
131)
133if TYPE_CHECKING: 133 ↛ 134line 133 didn't jump to line 134 because the condition on line 133 was never true
134 from litellm.proxy.proxy_server import ProxyConfig
136router: Final = APIRouter()
138pass_through_endpoint_logging: Final = PassThroughEndpointLogging()
140_METADATA_KEYS: Final = frozenset(("litellm_metadata", "metadata"))
141_KEPT_OUT_OF_LITELLM_PARAMS: Final = _METADATA_KEYS | frozenset(wire_names(ProxyRequestState))
143# Global registry to track registered pass-through routes and prevent memory leaks
144_registered_pass_through_routes: Final[dict[str, dict[str, str | bool | list[str] | Mapping[str, object]]]] = {}
147def get_response_body(response: httpx.Response) -> dict | None:
148 try:
149 return response.json()
150 except Exception:
151 return None
154async def set_env_variables_in_header(custom_headers: dict | None) -> dict | None:
155 """
156 checks if any headers on config.yaml are defined as os.environ/COHERE_API_KEY etc
158 only runs for headers defined on config.yaml
160 example header can be
162 {"Authorization": "Bearer os.environ/COHERE_API_KEY"}
163 """
164 if custom_headers is None: 164 ↛ 165line 164 didn't jump to line 165 because the condition on line 164 was never true
165 return None
166 headers: Final = {}
167 for key, value in custom_headers.items():
168 # langfuse Api requires base64 encoded headers - it's simpleer to just ask litellm users to set their langfuse public and secret keys
169 # we can then get the b64 encoded keys here
170 if key == "LANGFUSE_PUBLIC_KEY" or key == "LANGFUSE_SECRET_KEY": 170 ↛ 172line 170 didn't jump to line 172 because the condition on line 170 was never true
171 # langfuse requires b64 encoded headers - we construct that here
172 _langfuse_public_key = custom_headers["LANGFUSE_PUBLIC_KEY"]
173 _langfuse_secret_key = custom_headers["LANGFUSE_SECRET_KEY"]
174 if isinstance(_langfuse_public_key, str) and _langfuse_public_key.startswith("os.environ/"):
175 _langfuse_public_key = get_secret_str(_langfuse_public_key)
176 if isinstance(_langfuse_secret_key, str) and _langfuse_secret_key.startswith("os.environ/"):
177 _langfuse_secret_key = get_secret_str(_langfuse_secret_key)
178 headers["Authorization"] = "Basic " + b64encode(
179 f"{_langfuse_public_key}:{_langfuse_secret_key}".encode()
180 ).decode("ascii")
181 else:
182 # for all other headers
183 headers[key] = value
184 if isinstance(value, str) and "os.environ/" in value: 184 ↛ 185line 184 didn't jump to line 185 because the condition on line 184 was never true
185 verbose_proxy_logger.debug("pass through endpoint - looking up 'os.environ/' variable")
186 # get string section that is os.environ/
187 start_index = value.find("os.environ/")
188 _variable_name = value[start_index:]
190 verbose_proxy_logger.debug(
191 "pass through endpoint - getting secret for variable name: %s",
192 _variable_name,
193 )
194 _secret_value = get_secret_str(_variable_name)
195 if _secret_value is not None:
196 new_value = value.replace(_variable_name, _secret_value)
197 headers[key] = new_value
198 return headers
201async def chat_completion_pass_through_endpoint(
202 fastapi_response: Response,
203 request: Request,
204 adapter_id: str,
205 user_api_key_dict: UserAPIKeyAuth,
206):
207 from litellm.proxy.proxy_server import (
208 add_litellm_data_to_request,
209 general_settings,
210 llm_router,
211 proxy_config,
212 proxy_logging_obj,
213 user_api_base,
214 user_max_tokens,
215 user_model,
216 user_request_timeout,
217 user_temperature,
218 version,
219 )
221 litellm_call_id: Final = resolve_litellm_call_id(request.headers.get("x-litellm-call-id"))
222 data = {"litellm_call_id": litellm_call_id}
223 try:
224 body: Final = await request.body()
225 body_str: Final = body.decode()
226 try:
227 data = ast.literal_eval(body_str) | data
228 except Exception:
229 data = json.loads(body_str) | data
231 data["adapter_id"] = adapter_id
233 verbose_proxy_logger.debug("Request received by LiteLLM:\n%s", data)
234 data["model"] = (
235 general_settings.get("completion_model", None) # server default
236 or user_model # model name passed via cli args
237 or data.get("model", None) # default passed in http request
238 )
239 if user_model:
240 data["model"] = user_model
242 data = await add_litellm_data_to_request(
243 data=data,
244 request=request,
245 general_settings=general_settings,
246 user_api_key_dict=user_api_key_dict,
247 version=version,
248 proxy_config=proxy_config,
249 )
251 # override with user settings, these are params passed via cli
252 if user_temperature:
253 data["temperature"] = user_temperature
254 if user_request_timeout:
255 data["request_timeout"] = user_request_timeout
256 if user_max_tokens:
257 data["max_tokens"] = user_max_tokens
258 if user_api_base:
259 data["api_base"] = user_api_base
261 ### MODEL ALIAS MAPPING ###
262 # check if model name in model alias map
263 # get the actual model name
264 if data["model"] in litellm.model_alias_map:
265 data["model"] = litellm.model_alias_map[data["model"]]
267 # Check key-specific aliases
268 if (
269 isinstance(data["model"], str)
270 and user_api_key_dict.aliases
271 and isinstance(user_api_key_dict.aliases, dict)
272 and data["model"] in user_api_key_dict.aliases
273 ):
274 data["model"] = user_api_key_dict.aliases[data["model"]]
276 ### CALL HOOKS ### - modify incoming data before calling the model
277 data = await proxy_logging_obj.pre_call_hook(
278 user_api_key_dict=user_api_key_dict, data=data, call_type="text_completion"
279 )
281 ### ROUTE THE REQUESTs ###
282 router_model_names: Final = llm_router.model_names if llm_router is not None else []
283 # skip router if user passed their key
284 if "api_key" in data:
285 llm_response = asyncio.create_task(litellm.aadapter_completion(**data))
286 elif llm_router is not None and llm_router.is_recognized_model(data["model"]):
287 llm_response = asyncio.create_task(llm_router.aadapter_completion(**data))
288 elif (
289 llm_router is not None
290 and data["model"] not in router_model_names
291 and (llm_router.default_deployment is not None or len(llm_router.pattern_router.patterns) > 0)
292 ): # check for wildcard routes or default deployment before checking deployment_names
293 llm_response = asyncio.create_task(llm_router.aadapter_completion(**data))
294 elif (
295 llm_router is not None and data["model"] in llm_router.deployment_names
296 ): # model in router deployments, calling a specific deployment on the router (lowest priority)
297 llm_response = asyncio.create_task(llm_router.aadapter_completion(**data, specific_deployment=True))
298 elif user_model is not None: # `litellm --model <your-model-name>`
299 llm_response = asyncio.create_task(litellm.aadapter_completion(**data))
300 else:
301 raise ProxyModelNotFoundError(
302 route="completion", model_name=data.get("model", ""), retryable_with_model_read_through=False
303 )
305 # Await the llm_response task
306 response: Final = await llm_response
308 hidden_params: Final = getattr(response, "_hidden_params", {}) or {}
309 model_id: Final = hidden_params.get("model_id", None) or ""
310 cache_key: Final = hidden_params.get("cache_key", None) or ""
311 api_base: Final = hidden_params.get("api_base", None) or ""
312 response_cost: Final = hidden_params.get("response_cost", None) or ""
314 ### ALERTING ###
315 asyncio.create_task(proxy_logging_obj.update_request_status(litellm_call_id=litellm_call_id, status="success"))
317 verbose_proxy_logger.debug("final response: %s", response)
319 fastapi_response.headers.update(
320 ProxyBaseLLMRequestProcessing.get_custom_headers(
321 user_api_key_dict=user_api_key_dict,
322 model_id=model_id,
323 cache_key=cache_key,
324 api_base=api_base,
325 version=version,
326 response_cost=response_cost,
327 )
328 )
330 verbose_proxy_logger.debug("\nResponse from Litellm:\n%s", response)
331 return response
332 except Exception as e:
333 await proxy_logging_obj.post_call_failure_hook(
334 user_api_key_dict=user_api_key_dict, original_exception=e, request_data=data
335 )
336 log_llm_api_exception(e, litellm_call_id)
337 error_msg: Final = f"{e}"
338 raise ProxyException(
339 message=getattr(e, "message", error_msg),
340 type=openai_error_type(e, error_status_code(e, 500)),
341 param=openai_error_param(e),
342 headers=litellm_call_id_headers(litellm_call_id),
343 code=error_status_code(e, 500),
344 )
347class HttpPassThroughEndpointHelpers(BasePassthroughUtils):
348 @staticmethod
349 def get_response_headers(
350 headers: httpx.Headers,
351 litellm_call_id: str | None = None,
352 custom_headers: Mapping[str, str] | None = None,
353 ) -> dict:
354 # Exclude headers that uvicorn writes itself (server, date) and
355 # encoding/length headers that don't survive re-serialization.
356 # If we forward the upstream's Server header, uvicorn adds its
357 # own and strict HTTP parsers (e.g. aiohttp) reject the
358 # response with "Duplicate 'Server' header found".
359 excluded_headers: Final = {
360 "transfer-encoding",
361 "content-encoding",
362 "content-length",
363 "server",
364 "date",
365 "connection",
366 "keep-alive",
367 }
369 return_headers: Final = {key: value for key, value in headers.items() if key.lower() not in excluded_headers}
370 if litellm_call_id: 370 ↛ 371line 370 didn't jump to line 371 because the condition on line 370 was never true
371 return_headers["x-litellm-call-id"] = litellm_call_id
372 if custom_headers: 372 ↛ 379line 372 didn't jump to line 379 because the condition on line 372 was always true
373 # Ensure custom headers don't override actual upstream response headers or let framework defaults (like content-length: 0) interfere.
374 sanitized_custom_headers: Final = {
375 key: value for key, value in custom_headers.items() if key.lower() not in excluded_headers
376 }
377 return_headers.update(sanitized_custom_headers)
379 return return_headers
381 @staticmethod
382 def get_endpoint_type(url: str) -> EndpointType:
383 parsed_url: Final = urlparse(url)
384 if (
385 ("generateContent") in url
386 or ("streamGenerateContent") in url
387 or ("rawPredict") in url
388 or ("streamRawPredict") in url
389 ):
390 return EndpointType.VERTEX_AI
391 elif parsed_url.hostname == "api.anthropic.com": 391 ↛ 392line 391 didn't jump to line 392 because the condition on line 391 was never true
392 return EndpointType.ANTHROPIC
393 elif ( 393 ↛ 398line 393 didn't jump to line 398 because the condition on line 393 was never true
394 parsed_url.hostname == "api.openai.com"
395 or parsed_url.hostname == "openai.azure.com"
396 or (parsed_url.hostname and "openai.com" in parsed_url.hostname)
397 ):
398 return EndpointType.OPENAI
399 elif is_tinyfish_agent_url(url): 399 ↛ 400line 399 didn't jump to line 400 because the condition on line 399 was never true
400 return EndpointType.TINYFISH
401 return EndpointType.GENERIC
403 @staticmethod
404 async def _make_non_streaming_http_request(
405 request: Request,
406 async_client: httpx.AsyncClient,
407 url: str,
408 headers: dict,
409 requested_query_params: dict | None = None,
410 custom_body: dict | None = None,
411 ) -> httpx.Response:
412 """
413 Make a non-streaming HTTP request
415 If request is GET, don't include a JSON body
416 """
417 if request.method == "GET":
418 response = await async_client.request(
419 method=request.method,
420 url=url,
421 headers=headers,
422 params=requested_query_params,
423 )
424 else:
425 response = await async_client.request(
426 method=request.method,
427 url=url,
428 headers=headers,
429 params=requested_query_params,
430 json=custom_body,
431 )
432 return response
434 @staticmethod
435 async def non_streaming_http_request_handler(
436 request: Request,
437 async_client: httpx.AsyncClient,
438 url: httpx.URL,
439 headers: dict,
440 requested_query_params: dict | None = None,
441 _parsed_body: dict | None = None,
442 forward_multipart: bool = False,
443 ) -> httpx.Response:
444 """
445 Handle non-SSE HTTP requests
447 Handles special cases when GET requests, multipart/form-data requests, and generic httpx requests.
449 GET and generic requests are sent with httpx stream semantics so the caller can
450 decide from the response headers whether to buffer the body (JSON, inspected for
451 logging/guardrails) or relay it to the client without materializing it in memory
452 (LIT-4009: large batch results files must not be buffered in proxy RSS).
453 """
454 if request.method == "GET":
455 get_request: Final = async_client.build_request(
456 request.method,
457 url,
458 headers=headers,
459 params=requested_query_params,
460 )
461 return await async_client.send(get_request, stream=True)
462 if HttpPassThroughEndpointHelpers.is_multipart(request) is True and forward_multipart: 462 ↛ 467line 462 didn't jump to line 467 because the condition on line 462 was never true
463 # Forward multipart via make_multipart_http_request even when _parsed_body is
464 # non-empty (pass_through_request always injects litellm_logging_obj, etc.).
465 # forward_multipart is False when custom_body was supplied (JSON body despite
466 # multipart content-type) — those requests use the generic json= path.
467 return await HttpPassThroughEndpointHelpers.make_multipart_http_request(
468 request=request,
469 async_client=async_client,
470 url=url,
471 headers=headers,
472 requested_query_params=requested_query_params,
473 )
474 generic_request: Final = async_client.build_request(
475 request.method,
476 url,
477 headers=headers,
478 params=requested_query_params,
479 json=_parsed_body,
480 )
481 return await async_client.send(generic_request, stream=True)
483 @staticmethod
484 def is_multipart(request: Request) -> bool:
485 """Check if the request is a multipart/form-data request"""
486 return "multipart/form-data" in request.headers.get("content-type", "")
488 @staticmethod
489 async def _build_request_files_from_upload_file(
490 upload_file: UploadFile | StarletteUploadFile,
491 ) -> tuple[str | None, bytes, str | None]:
492 """Build a request files dict from an UploadFile object"""
493 file_content: Final = await upload_file.read()
494 return (upload_file.filename, file_content, upload_file.content_type)
496 @staticmethod
497 async def make_multipart_http_request(
498 request: Request,
499 async_client: httpx.AsyncClient,
500 url: httpx.URL,
501 headers: dict,
502 requested_query_params: dict | None = None,
503 stream: bool = False,
504 ) -> httpx.Response:
505 """Process multipart/form-data requests, handling both files and form fields.
507 Iterates ``form.multi_items()`` rather than ``form.items()`` so repeated
508 field names (e.g. several ``-F file=@...`` parts) are all forwarded;
509 ``items()`` collapses duplicate keys to the last value. Files go out as a
510 list of ``(field_name, (filename, content, content_type))`` tuples and
511 repeated non-file fields are grouped into list values, both of which httpx
512 encodes as separate multipart parts. A form with no file parts is sent
513 entirely through ``files`` as ``(field_name, (None, value))`` tuples,
514 because httpx downgrades a file-less ``data=`` payload to
515 application/x-www-form-urlencoded.
516 """
517 form_items: Final = (await request.form()).multi_items()
519 files: Final = [
520 (
521 field_name,
522 await HttpPassThroughEndpointHelpers._build_request_files_from_upload_file(upload_file=field_value),
523 )
524 for field_name, field_value in form_items
525 if isinstance(field_value, (StarletteUploadFile, UploadFile))
526 ]
528 non_file_items: Final = tuple(
529 (field_name, field_value)
530 for field_name, field_value in form_items
531 if not isinstance(field_value, (StarletteUploadFile, UploadFile))
532 )
533 field_order: Final = {
534 field_name: index
535 for index, field_name in enumerate(dict.fromkeys(field_name for field_name, _ in non_file_items))
536 }
537 form_data_dict: Final = {
538 field_name: [value for _, value in group]
539 for field_name, group in groupby(
540 sorted(non_file_items, key=lambda item: field_order[item[0]]),
541 key=lambda item: item[0],
542 )
543 }
545 multipart_files: Final = (
546 files if files else tuple((field_name, (None, field_value)) for field_name, field_value in non_file_items)
547 )
548 multipart_data: Final = form_data_dict if files else None
550 # Remove content-type header - httpx will set it correctly with the new boundary
551 # when it creates the multipart body from files/data parameters
552 headers_copy: Final = headers.copy()
553 headers_copy.pop("content-type", None)
555 # httpx.AsyncClient.request() does not accept stream=; use send() for streaming.
556 if stream:
557 req: Final = async_client.build_request(
558 request.method,
559 url,
560 headers=headers_copy,
561 params=requested_query_params,
562 files=multipart_files,
563 data=multipart_data,
564 )
565 return await async_client.send(req, stream=True)
567 return await async_client.request(
568 method=request.method,
569 url=url,
570 headers=headers_copy,
571 params=requested_query_params,
572 files=multipart_files,
573 data=multipart_data,
574 )
576 @staticmethod
577 def _init_kwargs_for_pass_through_endpoint(
578 request: Request,
579 user_api_key_dict: UserAPIKeyAuth,
580 passthrough_logging_payload: PassthroughStandardLoggingPayload,
581 logging_obj: LiteLLMLoggingObj,
582 _parsed_body: dict | None = None,
583 litellm_call_id: str | None = None,
584 ) -> dict:
585 """
586 Filter out litellm params from the request body
587 """
588 _parsed_body = _parsed_body or {}
590 litellm_keys_in_body: Final = MappingProxyType(
591 {k: _parsed_body.pop(k) for k in types_utils.all_litellm_params if k in _parsed_body}
592 )
593 litellm_params_in_body: Final = MappingProxyType(
594 {k: v for k, v in litellm_keys_in_body.items() if k not in _KEPT_OUT_OF_LITELLM_PARAMS}
595 )
597 _metadata = dict(
598 LiteLLMProxyRequestSetup.get_sanitized_user_information_from_key(user_api_key_dict=user_api_key_dict)
599 )
601 litellm_metadata: Final = litellm_keys_in_body.get("litellm_metadata")
602 metadata: Final = litellm_keys_in_body.get("metadata")
603 if litellm_metadata: 603 ↛ 604line 603 didn't jump to line 604 because the condition on line 603 was never true
604 _metadata.update(litellm_metadata)
605 if metadata: 605 ↛ 606line 605 didn't jump to line 606 because the condition on line 605 was never true
606 _metadata.update(metadata)
608 _metadata = _update_metadata_with_tags_in_header(
609 request=request,
610 metadata=_metadata,
611 )
613 # Set internal keys after merging client-supplied metadata so a request
614 # body that mirrors them cannot clobber the authenticated key, the real
615 # parent span, or the proxy's own session-id decision.
616 _metadata.pop(SESSION_ID_OMITTED_METADATA_KEY, None)
617 _metadata["user_api_key"] = LiteLLMProxyRequestSetup.get_logged_api_key(user_api_key_dict)
618 _metadata["litellm_parent_otel_span"] = user_api_key_dict.parent_otel_span
619 _metadata["user_api_key_budget_reservation"] = user_api_key_dict.budget_reservation
620 _metadata[MODEL_ACCESS_GROUP_METADATA_KEY] = user_api_key_dict.matched_model_access_groups
621 # The per-model budget counters are keyed off these. get_sanitized_user_information_from_key
622 # returns StandardLoggingUserAPIKeyMetadata, which carries no budget field, so without this
623 # the post-call increment finds nothing and every passthrough request goes untracked and
624 # unenforced. Set after the client merge so a request body cannot supply its own budget.
625 #
626 # Only for the built-in provider routes. `get_model_from_request` returns
627 # None for a user-defined pass-through, deliberately: its body is forwarded
628 # verbatim, so `model` there names an UPSTREAM model rather than a
629 # LiteLLM-managed one. Enforcement is therefore skipped on those routes, and
630 # charging a counter anyway would track spend that nothing can refuse, and
631 # would attribute it to a budget the operator scoped to a LiteLLM model that
632 # merely shares the name.
633 if not request_dispatched_to_pass_through_endpoint(request):
634 _metadata["user_api_key_model_max_budget"] = user_api_key_dict.model_max_budget
635 _metadata["user_api_key_team_model_max_budget"] = user_api_key_dict.team_model_max_budget
636 _metadata["user_api_key_user_model_max_budget"] = user_api_key_dict.user_model_max_budget
637 _metadata["user_api_key_end_user_model_max_budget"] = user_api_key_dict.end_user_model_max_budget
638 _metadata.update(
639 LiteLLMProxyRequestSetup.get_sanitized_user_information_from_key(user_api_key_dict=user_api_key_dict)
640 )
641 _request_state: Final = getattr(request, "state", None)
642 deployment_model_info: Final = getattr(
643 _request_state, LITELLM_PASS_THROUGH_DEPLOYMENT_MODEL_INFO_STATE_KEY, None
644 )
645 if isinstance(deployment_model_info, Mapping): 645 ↛ 646line 645 didn't jump to line 646 because the condition on line 645 was never true
646 _metadata["model_info"] = dict(deployment_model_info)
648 kwargs: Final = {
649 "litellm_params": {
650 **litellm_params_in_body,
651 "metadata": _metadata,
652 "proxy_server_request": {
653 "url": str(request.url),
654 "method": request.method,
655 "body": copy.copy(_parsed_body), # use copy instead of deepcopy
656 "headers": request.headers,
657 },
658 },
659 "call_type": "pass_through_endpoint",
660 "litellm_call_id": litellm_call_id,
661 "passthrough_logging_payload": passthrough_logging_payload,
662 }
664 logging_obj.model_call_details["passthrough_logging_payload"] = passthrough_logging_payload
666 return kwargs
668 @staticmethod
669 def construct_target_url_with_subpath(base_target: str, subpath: str, include_subpath: bool | None) -> str:
670 """
671 Helper function to construct the full target URL with subpath handling.
673 Args:
674 base_target: The base target URL
675 subpath: The captured subpath from the request
676 include_subpath: Whether to include the subpath in the target URL
678 Returns:
679 The constructed full target URL
680 """
681 if not include_subpath:
682 return base_target
684 if not subpath: 684 ↛ 685line 684 didn't jump to line 685 because the condition on line 684 was never true
685 return base_target
687 # Ensure base_target ends with / and subpath doesn't start with /
688 if not base_target.endswith("/"): 688 ↛ 690line 688 didn't jump to line 690 because the condition on line 688 was always true
689 base_target = base_target + "/"
690 subpath = subpath.removeprefix("/")
692 # Resolve any '..' segments in the subpath so it cannot climb above
693 # the base_target prefix that the operator configured. Preserve a
694 # trailing slash on the original subpath since some upstreams treat
695 # `/foo` and `/foo/` as different resources.
696 trailing_slash: Final = subpath.endswith("/")
697 safe_subpath = posixpath.normpath("/" + subpath).lstrip("/")
698 if safe_subpath == ".": 698 ↛ 699line 698 didn't jump to line 699 because the condition on line 698 was never true
699 safe_subpath = ""
700 if trailing_slash and safe_subpath and not safe_subpath.endswith("/"): 700 ↛ 701line 700 didn't jump to line 701 because the condition on line 700 was never true
701 safe_subpath += "/"
703 return base_target + safe_subpath
705 @staticmethod
706 def join_base_and_endpoint_path(base_url: httpx.URL, endpoint_path: str) -> str:
707 """
708 Combine the path component of ``base_url`` with ``endpoint_path``.
710 Preserves any path prefix configured on the base URL and resolves
711 ``..`` segments in the endpoint so the result stays within the base
712 path. A trailing slash on ``endpoint_path`` is preserved.
713 """
714 trailing_slash: Final = endpoint_path.endswith("/")
715 base_path = base_url.path or ""
716 if not base_path or base_path == "/":
717 normalized_endpoint = posixpath.normpath("/" + endpoint_path.lstrip("/"))
718 if trailing_slash and normalized_endpoint != "/": 718 ↛ 719line 718 didn't jump to line 719 because the condition on line 718 was never true
719 normalized_endpoint += "/"
720 return normalized_endpoint
722 base_path = base_path.rstrip("/")
723 clean_endpoint: Final = endpoint_path.lstrip("/")
724 combined = posixpath.normpath(base_path + "/" + clean_endpoint)
725 # If normalization climbs out of the base path, fall back to base.
726 if combined != base_path and not combined.startswith(base_path + "/"): 726 ↛ 727line 726 didn't jump to line 727 because the condition on line 726 was never true
727 return base_path + "/"
728 if trailing_slash and not combined.endswith("/"):
729 combined += "/"
730 return combined
732 @staticmethod
733 def _update_stream_param_based_on_request_body(
734 parsed_body: dict,
735 stream: bool | None = None,
736 ) -> bool | None:
737 """
738 If stream is provided in the request body, use it.
739 Otherwise, use the stream parameter passed to the `pass_through_request` function
740 """
741 if "stream" in parsed_body: 741 ↛ 742line 741 didn't jump to line 742 because the condition on line 741 was never true
742 return parsed_body.get("stream", stream)
743 return stream
746def _carry_guardrail_logging_info(request_data: dict, guardrail_data: dict | None) -> None:
747 """Copy guardrail logging entries from ``guardrail_data`` onto ``request_data``.
749 Post-call guardrails run against a throwaway ``hook_data`` dict (its
750 ``metadata`` is what ``_init_kwargs_for_pass_through_endpoint`` already
751 stripped off ``_parsed_body``), so a block records the
752 ``standard_logging_guardrail_information`` there and not on the dict the
753 failure handler forwards to ``post_call_failure_hook``. Without this the
754 otel guardrail span is emitted on allow but missing on block. Carry the
755 entries over so the failure path matches the unified path.
756 """
757 if guardrail_data is None: 757 ↛ 759line 757 didn't jump to line 759 because the condition on line 757 was always true
758 return
759 source_key: Final = get_metadata_variable_name_from_kwargs(guardrail_data)
760 source_metadata: Final = guardrail_data.get(source_key) or {}
761 entries: Final = source_metadata.get("standard_logging_guardrail_information")
762 if not entries:
763 return
765 _, metadata = get_or_create_metadata_bucket(request_data)
766 metadata.setdefault("standard_logging_guardrail_information", list(entries))
769def _build_passthrough_failure_request_payload(
770 parsed_body: dict | None,
771 kwargs: dict | None,
772 logging_obj: LiteLLMLoggingObj | None,
773 custom_llm_provider: str | None,
774 upstream_usage: UpstreamReportedUsage | None = None,
775) -> dict:
776 """Build the ``request_data`` dict passed to ``post_call_failure_hook``.
778 Shared by the outer exception handler (LiteLLM-internal failures) and
779 upstream HTTP error logging, so both failure paths report the same shape
780 of request data (model, custom_llm_provider, litellm_logging_obj, ...).
782 ``upstream_usage`` carries the cost and tokens an upstream reported on an
783 error response. Spend tracking only attributes a recovered cost when it
784 comes paired with a usage object, so both keys are written together.
785 """
786 request_payload: Final[dict] = dict(parsed_body or {})
787 if kwargs: 787 ↛ 789line 787 didn't jump to line 789 because the condition on line 787 was always true
788 request_payload.update(kwargs)
789 if logging_obj is not None: 789 ↛ 791line 789 didn't jump to line 791 because the condition on line 789 was always true
790 request_payload["litellm_logging_obj"] = logging_obj
791 if "model" not in request_payload and parsed_body and isinstance(parsed_body, dict): 791 ↛ 792line 791 didn't jump to line 792 because the condition on line 791 was never true
792 request_payload["model"] = parsed_body.get("model", "")
793 if "custom_llm_provider" not in request_payload and custom_llm_provider:
794 request_payload["custom_llm_provider"] = custom_llm_provider
795 if upstream_usage is not None: 795 ↛ 796line 795 didn't jump to line 796 because the condition on line 795 was never true
796 request_payload["response_cost"] = upstream_usage.response_cost or 0.0
797 request_payload["combined_usage_object"] = Usage(total_tokens=upstream_usage.total_tokens or 0)
798 return request_payload
801@dataclass(frozen=True, slots=True)
802class _TeamCallbackWiring:
803 success_callbacks: "list[str | Callable | CustomLogger] | None" = None # mutable-ok: Logging.__init__ arg
804 failure_callbacks: "list[str | Callable | CustomLogger] | None" = None # mutable-ok: Logging.__init__ arg
805 logging_kwargs: dict[str, str | dict[str, str]] | None = None # mutable-ok: Logging.__init__ arg
808def _resolve_team_callback_wiring(
809 user_api_key_dict: UserAPIKeyAuth,
810 proxy_config: "ProxyConfig",
811 route_description: str,
812) -> _TeamCallbackWiring:
813 """Resolve key/team dynamic logging callbacks for a passthrough request.
815 Mirrors add_litellm_data_to_request: callback_vars are unpacked top-level
816 (read by initialize_standard_callback_dynamic_params) and also stamped on
817 the proxy-owned trusted-vars field (read by get_trusted_callback_params).
819 Fails open: a callback resolution or validation error is logged at error
820 level and the request proceeds without dynamic callbacks, since a broken
821 logging config must not fail the customer's upstream call (and the
822 websocket is already accepted by the time this runs on that path). The
823 env-reference check runs here because the deprecated callback_settings
824 branch skips AddTeamCallback validation, and Logging.__init__ would
825 otherwise reject the vars mid-request.
826 """
827 try:
828 callback_settings_obj: Final = _get_dynamic_logging_metadata(
829 user_api_key_dict=user_api_key_dict, proxy_config=proxy_config
830 )
831 if callback_settings_obj and callback_settings_obj.callback_vars: 831 ↛ 832line 831 didn't jump to line 832 because the condition on line 831 was never true
832 for item in callback_settings_obj.callback_vars.items():
833 validate_no_callback_env_reference(item[0], item[1], source="key/team callback metadata")
834 except Exception: # noqa: BLE001 - a broken logging config must never fail the passthrough request
835 verbose_proxy_logger.exception(
836 "%s: failed to resolve team logging callbacks, continuing without them",
837 route_description,
838 )
839 return _TeamCallbackWiring()
840 if callback_settings_obj is None: 840 ↛ 842line 840 didn't jump to line 842 because the condition on line 840 was always true
841 return _TeamCallbackWiring()
842 callback_vars: Final = callback_settings_obj.callback_vars
843 success_callbacks: Final = callback_settings_obj.success_callback
844 failure_callbacks: Final = callback_settings_obj.failure_callback
845 logging_kwargs: Final = (
846 None
847 if not callback_vars
848 else { # mutable-ok: Logging arg
849 **callback_vars,
850 TRUSTED_CALLBACK_VARS_FIELD: callback_vars,
851 "metadata": {}, # mutable-ok: Logging arg
852 "model_info": {}, # mutable-ok: Logging arg
853 }
854 )
855 return _TeamCallbackWiring(
856 success_callbacks=None if success_callbacks is None else [*success_callbacks], # mutable-ok: Logging arg
857 failure_callbacks=None if failure_callbacks is None else [*failure_callbacks], # mutable-ok: Logging arg
858 logging_kwargs=logging_kwargs,
859 )
862def _truncate_upstream_error_body(body: str) -> str:
863 if len(body) <= PASSTHROUGH_UPSTREAM_ERROR_BODY_MAX_LOG_CHARS: 863 ↛ 865line 863 didn't jump to line 865 because the condition on line 863 was always true
864 return body
865 return (
866 f"{body[:PASSTHROUGH_UPSTREAM_ERROR_BODY_MAX_LOG_CHARS]}... "
867 f"(truncated at {PASSTHROUGH_UPSTREAM_ERROR_BODY_MAX_LOG_CHARS} chars)"
868 )
871def _sanitize_upstream_error_body(body: str) -> str:
872 return " ".join("".join(char if char.isprintable() else " " for char in body).split())
875class _PrefixReplayStream(httpx.AsyncByteStream):
876 def __init__(self, prefix: bytes, rest: AsyncIterator[bytes], upstream: httpx.Response) -> None:
877 self._prefix: Final = prefix
878 self._rest: Final = rest
879 self._upstream: Final = upstream
881 async def __aiter__(self) -> AsyncIterator[bytes]:
882 if self._prefix:
883 yield self._prefix
884 async for chunk in self._rest:
885 yield chunk
887 async def aclose(self) -> None:
888 await self._upstream.aclose()
891async def _no_more_chunks() -> AsyncIterator[bytes]:
892 return
893 yield b""
896async def _read_error_body_preview(
897 stream: AsyncIterator[bytes],
898) -> tuple[bytes, AsyncIterator[bytes]]:
899 collected: Final[list[bytes]] = [] # mutable-ok: accumulated until the preview byte budget, then joined once
900 total = 0 # rebind-ok: running byte count against the preview budget
901 try:
902 async for chunk in stream:
903 collected.append(chunk)
904 total += len(chunk)
905 if total > PASSTHROUGH_UPSTREAM_ERROR_BODY_MAX_LOG_CHARS:
906 break
907 except httpx.HTTPError as err:
908 partial: Final = b"".join(collected)
909 verbose_proxy_logger.warning(
910 "pass_through_endpoint: upstream error body read failed after %d bytes: %s",
911 len(partial),
912 type(err).__name__,
913 )
914 return partial, _no_more_chunks()
915 return b"".join(collected), stream
918def _headers_without_body_framing(headers: httpx.Headers) -> httpx.Headers:
919 return httpx.Headers(
920 [(name, value) for name, value in headers.raw if name.lower() not in (b"content-encoding", b"content-length")]
921 )
924async def _error_body_preview_and_relay(response: httpx.Response) -> tuple[str, httpx.Response]:
925 if response.is_stream_consumed: 925 ↛ 927line 925 didn't jump to line 927 because the condition on line 925 was always true
926 return response.text, response
927 body_iter: Final = response.aiter_bytes()
928 prefix, rest = await _read_error_body_preview(body_iter)
929 preview_text: Final = prefix.decode(response.encoding or "utf-8", errors="replace")
930 return preview_text, httpx.Response(
931 status_code=response.status_code,
932 headers=_headers_without_body_framing(response.headers),
933 stream=_PrefixReplayStream(prefix=prefix, rest=rest, upstream=response),
934 request=response.request,
935 extensions=response.extensions,
936 )
939async def _log_passthrough_upstream_failure(
940 response: httpx.Response,
941 user_api_key_dict: UserAPIKeyAuth,
942 request_payload: dict,
943 logging_obj: LiteLLMLoggingObj,
944) -> httpx.Response:
945 if response.status_code < 400: 945 ↛ 946line 945 didn't jump to line 946 because the condition on line 945 was never true
946 return response
947 from litellm.proxy.proxy_server import proxy_logging_obj
949 preview_text, relay_response = await _error_body_preview_and_relay(response)
950 upstream_error_body: Final = (
951 REDACTED_BY_LITELLM
952 if should_redact_message_logging(logging_obj.model_call_details)
953 else _truncate_upstream_error_body(_sanitize_upstream_error_body(preview_text))
954 )
955 verbose_proxy_logger.warning(
956 "pass_through_endpoint: upstream %s %s returned %s: %s",
957 response.request.method,
958 response.url.copy_with(query=None, fragment=None),
959 response.status_code,
960 upstream_error_body,
961 )
962 try:
963 response.raise_for_status()
964 except httpx.HTTPStatusError:
965 # Reported as an HTTPException, not the raw httpx error: ProxyLogging's
966 # alerting path only excludes HTTPException/ProxyException from its
967 # "High" severity llm_exceptions alert, treating everything else as an
968 # operational LLM-API failure. An upstream 4xx/5xx returned unchanged
969 # to the client is a user-facing error like any other, not something
970 # ops needs paged for, so it must be excluded the same way auth and
971 # rate-limit errors already are.
972 synthetic_exception: Final = HTTPException(
973 status_code=response.status_code,
974 detail=f"Upstream passthrough request failed with status {response.status_code}: {upstream_error_body}",
975 )
976 try:
977 await proxy_logging_obj.post_call_failure_hook(
978 user_api_key_dict=user_api_key_dict,
979 original_exception=synthetic_exception,
980 request_data=request_payload,
981 traceback_str=traceback.format_exc(limit=MAXIMUM_TRACEBACK_LINES_TO_LOG),
982 )
983 except Exception: # noqa: BLE001 - a failing logging callback must never break the passthrough response
984 verbose_proxy_logger.warning(
985 "pass_through_endpoint: post_call_failure_hook raised for upstream error",
986 exc_info=True,
987 )
988 return relay_response
991async def _relay_reporting_failures(
992 stream: AsyncGenerator[bytes, None],
993 upstream_status: int,
994 user_api_key_dict: UserAPIKeyAuth,
995 request_payload: dict, # mutable-ok: post_call_failure_hook lifts fields onto request_data in place
996) -> AsyncGenerator[bytes, None]:
997 from litellm.proxy.proxy_server import proxy_logging_obj
999 try:
1000 async for chunk in stream:
1001 yield chunk
1002 except Exception as e:
1003 if upstream_status >= 400:
1004 raise
1005 try:
1006 await proxy_logging_obj.post_call_failure_hook(
1007 user_api_key_dict=user_api_key_dict,
1008 original_exception=e,
1009 request_data=request_payload,
1010 traceback_str=traceback.format_exc(limit=MAXIMUM_TRACEBACK_LINES_TO_LOG),
1011 )
1012 except Exception: # noqa: BLE001 - a failing logging callback must never mask the upstream error
1013 verbose_proxy_logger.warning(
1014 "pass_through_endpoint: post_call_failure_hook raised for a mid-stream upstream error",
1015 exc_info=True,
1016 )
1017 raise
1020from litellm.passthrough.timeout_utils import (
1021 DEFAULT_PASS_THROUGH_REQUEST_TIMEOUT_SECONDS, # noqa: F401 - re-exported for backward compat
1022 resolve_llm_passthrough_timeout, # noqa: F401 - re-exported for backward compat
1023 resolve_pass_through_request_timeout,
1024)
1027async def pass_through_request(
1028 request: Request,
1029 target: str,
1030 custom_headers: dict,
1031 user_api_key_dict: UserAPIKeyAuth,
1032 custom_body: dict | None = None,
1033 forward_headers: bool | None = False,
1034 merge_query_params: bool | None = False,
1035 query_params: dict | None = None,
1036 default_query_params: dict | None = None,
1037 stream: bool | None = None,
1038 cost_per_request: float | None = None,
1039 custom_llm_provider: str | None = None,
1040 guardrails_config: dict | None = None,
1041 timeout: float | None = None,
1042):
1043 """
1044 Pass through endpoint handler, makes the httpx request for pass-through endpoints and ensures logging hooks are called
1046 Args:
1047 request: The incoming request
1048 target: The target URL
1049 custom_headers: The custom headers
1050 user_api_key_dict: The user API key dictionary
1051 custom_body: The custom body
1052 forward_headers: Whether to forward headers
1053 merge_query_params: Whether to merge query params
1054 query_params: The query params
1055 default_query_params: The default query params to be applied if not overridden by client
1056 stream: Whether to stream the response
1057 cost_per_request: Optional field - cost per request to the target endpoint
1058 custom_llm_provider: Optional field - custom LLM provider for the endpoint
1059 guardrails_config: Optional field - guardrails configuration for passthrough endpoint
1060 timeout: Optional per-endpoint timeout in seconds. Falls back to
1061 general_settings.pass_through_request_timeout, then 600s.
1062 """
1063 from litellm.exceptions import ModifyResponseException
1064 from litellm.litellm_core_utils.litellm_logging import Logging
1065 from litellm.proxy.pass_through_endpoints.passthrough_guardrails import (
1066 PassthroughGuardrailHandler,
1067 )
1068 from litellm.proxy.proxy_server import proxy_config, proxy_logging_obj
1070 #########################################################
1071 # Initialize variables
1072 #########################################################
1073 litellm_call_id: Final = str(uuid.uuid4())
1074 url: httpx.URL | None = None
1076 # parsed request body
1077 _parsed_body: dict | None = None
1078 # kwargs for pass through endpoint, contains metadata, litellm_params, call_type, litellm_call_id, passthrough_logging_payload
1079 kwargs: dict | None = None
1080 logging_obj: Logging | None = None
1081 # the dict post-call guardrails wrote their logging info into; the failure
1082 # handler reuses it so a guardrail block still surfaces its span/logs
1083 post_call_guardrail_data: dict | None = None
1085 #########################################################
1086 try:
1087 url = httpx.URL(target)
1088 headers = custom_headers
1089 headers = HttpPassThroughEndpointHelpers.forward_headers_from_request(
1090 request_headers=_safe_get_request_headers(request).copy(),
1091 headers=headers,
1092 forward_headers=forward_headers,
1093 )
1094 upstream_headers: Final = _with_trace_context(headers, parent_span=user_api_key_dict.parent_otel_span)
1096 requested_query_params: dict | None = query_params or dict(request.query_params) or None
1098 endpoint_type: Final[EndpointType] = HttpPassThroughEndpointHelpers.get_endpoint_type(str(url))
1100 # SigV4-signed callers (e.g. Bedrock) attach the exact bytes that were
1101 # signed via request.state; we must send those instead of re-encoding the
1102 # parsed dict (hooks mutate it, breaking the signature / Content-Length).
1103 # Tolerate request objects without `state` (test fixtures) and only honor
1104 # values httpx accepts for `content=`.
1105 _request_state: Final = getattr(request, "state", None)
1106 state_raw_body: str | bytes | None = (
1107 getattr(_request_state, LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY, None)
1108 if _request_state is not None
1109 else None
1110 )
1111 if state_raw_body is not None and not isinstance(state_raw_body, (str, bytes, bytearray)): 1111 ↛ 1112line 1111 didn't jump to line 1112 because the condition on line 1111 was never true
1112 state_raw_body = None
1114 # Skip body parsing for multipart requests - make_multipart_http_request will handle it
1115 # But if custom_body is provided (e.g., JSON parsed despite multipart content-type), use it
1116 is_multipart: Final = HttpPassThroughEndpointHelpers.is_multipart(request) and not custom_body
1118 if custom_body: 1118 ↛ 1119line 1118 didn't jump to line 1119 because the condition on line 1118 was never true
1119 _parsed_body = custom_body
1120 elif is_multipart: 1120 ↛ 1122line 1120 didn't jump to line 1122 because the condition on line 1120 was never true
1121 # Don't parse multipart body here - it will be handled by make_multipart_http_request
1122 _parsed_body = {}
1123 else:
1124 _parsed_body = await _read_request_body(request)
1125 verbose_proxy_logger.debug(
1126 "Pass through endpoint sending request to \nURL %s\nheaders: %s\nbody: %s\n",
1127 url,
1128 _get_masked_values(upstream_headers),
1129 _parsed_body,
1130 )
1132 ### COLLECT GUARDRAILS FOR PASSTHROUGH ENDPOINT ###
1133 # Passthrough endpoints are opt-in only for guardrails
1134 # When enabled, collect guardrails from org/team/key levels + passthrough-specific
1135 guardrails_to_run: Final = PassthroughGuardrailHandler.collect_guardrails(
1136 user_api_key_dict=user_api_key_dict,
1137 passthrough_guardrails_config=guardrails_config,
1138 )
1140 # Add guardrails to metadata if any should run
1141 if guardrails_to_run and len(guardrails_to_run) > 0: 1141 ↛ 1142line 1141 didn't jump to line 1142 because the condition on line 1141 was never true
1142 if _parsed_body is None:
1143 _parsed_body = {}
1144 if "metadata" not in _parsed_body:
1145 _parsed_body["metadata"] = {}
1146 _parsed_body["metadata"]["guardrails"] = guardrails_to_run
1147 verbose_proxy_logger.debug("Added guardrails to passthrough request metadata: %s", guardrails_to_run)
1149 ## LOGGING OBJECT ## - initialize before pre_call_hook so guardrails can access it
1150 # Surface the requested model (when the body carries one) so logging/spans
1151 # read e.g. ``chat gpt-4o`` instead of ``chat unknown``.
1152 passthrough_model: Final = (_parsed_body.get("model") if isinstance(_parsed_body, dict) else None) or "unknown"
1153 start_time: Final = datetime.now()
1154 team_callbacks: Final = _resolve_team_callback_wiring(
1155 user_api_key_dict=user_api_key_dict,
1156 proxy_config=proxy_config,
1157 route_description="pass_through_endpoint",
1158 )
1159 logging_obj = Logging(
1160 model=passthrough_model,
1161 messages=[{"role": "user", "content": safe_dumps(_parsed_body)}],
1162 stream=False,
1163 call_type="pass_through_endpoint",
1164 start_time=start_time,
1165 litellm_call_id=litellm_call_id,
1166 function_id="1245",
1167 dynamic_success_callbacks=team_callbacks.success_callbacks,
1168 dynamic_failure_callbacks=team_callbacks.failure_callbacks,
1169 kwargs=team_callbacks.logging_kwargs,
1170 )
1172 # Store passthrough guardrails config on logging_obj for field targeting
1173 logging_obj.passthrough_guardrails_config = guardrails_config
1175 # Store logging_obj in data so guardrails can access it
1176 if _parsed_body is None: 1176 ↛ 1177line 1176 didn't jump to line 1177 because the condition on line 1176 was never true
1177 _parsed_body = {}
1178 _parsed_body["litellm_logging_obj"] = logging_obj
1180 ### CALL HOOKS ### - modify incoming data / reject request before calling the model
1181 _parsed_body = await proxy_logging_obj.pre_call_hook(
1182 user_api_key_dict=user_api_key_dict,
1183 data=_parsed_body,
1184 call_type="pass_through_endpoint",
1185 )
1186 resolved_timeout: Final = resolve_pass_through_request_timeout(timeout)
1187 async_client_obj: Final = get_async_httpx_client(
1188 llm_provider=httpxSpecialProvider.PassThroughEndpoint,
1189 params={"timeout": resolved_timeout},
1190 )
1191 async_client: Final = async_client_obj.client
1192 passthrough_logging_payload: Final = PassthroughStandardLoggingPayload(
1193 url=str(url),
1194 request_body=_parsed_body,
1195 request_method=getattr(request, "method", None),
1196 cost_per_request=cost_per_request,
1197 )
1198 kwargs = HttpPassThroughEndpointHelpers._init_kwargs_for_pass_through_endpoint(
1199 user_api_key_dict=user_api_key_dict,
1200 _parsed_body=_parsed_body,
1201 passthrough_logging_payload=passthrough_logging_payload,
1202 litellm_call_id=litellm_call_id,
1203 request=request,
1204 logging_obj=logging_obj,
1205 )
1207 # Store custom_llm_provider in kwargs and logging object if provided
1208 if custom_llm_provider:
1209 logging_obj.model_call_details["custom_llm_provider"] = custom_llm_provider
1210 logging_obj.model_call_details["litellm_params"] = kwargs.get("litellm_params", {})
1212 # done for supporting 'parallel_request_limiter.py' with pass-through endpoints
1213 logging_obj.update_environment_variables(
1214 model=passthrough_model,
1215 user="unknown",
1216 optional_params={},
1217 litellm_params=kwargs["litellm_params"],
1218 call_type="pass_through_endpoint",
1219 )
1220 logging_obj.model_call_details["litellm_call_id"] = litellm_call_id
1222 ## PASSTHROUGH MANAGED ID RESOLUTION (INPUT) ##
1223 # Resolve managed IDs in path, query params, and body back to raw
1224 # provider IDs before forwarding upstream. Gated by feature flag and
1225 # enterprise managed-files hook. Runs after pre_call_hook so
1226 # guardrails have already seen the managed IDs.
1227 from litellm.proxy.proxy_server import (
1228 general_settings as proxy_general_settings,
1229 )
1230 from litellm.proxy.proxy_server import (
1231 general_settings_view,
1232 )
1234 _managed_id_provider: Final = resolve_passthrough_managed_id_provider(custom_llm_provider)
1236 if proxy_general_settings.get("passthrough_managed_object_ids", False) and _managed_id_provider is not None: 1236 ↛ 1237line 1236 didn't jump to line 1237 because the condition on line 1236 was never true
1237 verbose_proxy_logger.debug(
1238 "pass_through_endpoint: managed-id input rewrite enabled for route=%s method=%s",
1239 request.url.path,
1240 request.method,
1241 )
1242 _passthrough_managed_hook = proxy_logging_obj.get_proxy_hook("managed_files")
1243 if _passthrough_managed_hook is not None:
1244 from litellm.proxy.pass_through_endpoints.managed_id_rewriter import (
1245 rewrite_body_ids,
1246 rewrite_path_ids,
1247 rewrite_query_ids,
1248 )
1249 from litellm.proxy.proxy_server import (
1250 prisma_client as _passthrough_prisma,
1251 )
1253 _original_path: Final = url.path
1254 _original_query_params: Final = requested_query_params
1255 _original_body: Final = _parsed_body
1256 _new_path: Final = await rewrite_path_ids(
1257 url.path,
1258 _managed_id_provider,
1259 user_api_key_dict,
1260 _passthrough_prisma,
1261 _passthrough_managed_hook,
1262 )
1263 if _new_path != url.path:
1264 url = url.copy_with(path=_new_path)
1265 requested_query_params = await rewrite_query_ids(
1266 requested_query_params,
1267 _managed_id_provider,
1268 user_api_key_dict,
1269 _passthrough_prisma,
1270 _passthrough_managed_hook,
1271 )
1272 _parsed_body = await rewrite_body_ids(
1273 _parsed_body,
1274 _managed_id_provider,
1275 user_api_key_dict,
1276 _passthrough_prisma,
1277 _passthrough_managed_hook,
1278 )
1279 verbose_proxy_logger.debug(
1280 "pass_through_endpoint: managed-id input rewrite results path_changed=%s query_changed=%s body_changed=%s route=%s method=%s",
1281 _new_path != _original_path,
1282 requested_query_params is not _original_query_params,
1283 _parsed_body is not _original_body,
1284 request.url.path,
1285 request.method,
1286 )
1287 else:
1288 verbose_proxy_logger.debug(
1289 "pass_through_endpoint: managed-id input rewrite skipped (managed_files hook not available) route=%s method=%s",
1290 request.url.path,
1291 request.method,
1292 )
1294 # Apply default query parameters if provided, regardless of merge_query_params setting
1295 if default_query_params or merge_query_params:
1296 # Create a new URL with the merged query params
1297 url = url.copy_with(
1298 query=urlencode(
1299 HttpPassThroughEndpointHelpers.get_merged_query_parameters(
1300 existing_url=url,
1301 request_query_params=requested_query_params or MappingProxyType({}),
1302 default_query_params=default_query_params,
1303 )
1304 ).encode("ascii")
1305 )
1306 requested_query_params = None
1308 ## PASSTHROUGH MANAGED LIST (DB-only response) ##
1309 # For GET /v1/files and GET /v1/batches passthrough routes, serve the
1310 # listing entirely from our DB so each caller only sees their own IDs.
1311 # Admins / master-key callers see all rows. Gated on the same
1312 # conditions as INPUT/OUTPUT rewrite: feature flag, provider, AND
1313 # the managed_files hook must be present. Without the hook no managed
1314 # IDs are ever minted or stored, so the DB is empty and intercepting
1315 # the list would silently hide the caller's real upstream files/batches.
1316 if ( 1316 ↛ 1322line 1316 didn't jump to line 1322 because the condition on line 1316 was never true
1317 proxy_general_settings.get("passthrough_managed_object_ids", False)
1318 and _managed_id_provider is not None
1319 and request.method == "GET"
1320 and proxy_logging_obj.get_proxy_hook("managed_files") is not None
1321 ):
1322 from litellm.proxy.auth.auth_utils import get_request_route
1323 from litellm.proxy.pass_through_endpoints.managed_id_rewriter import (
1324 is_passthrough_list_route,
1325 list_passthrough_ids_from_db,
1326 )
1327 from litellm.proxy.proxy_server import prisma_client as _list_prisma
1329 if (
1330 is_passthrough_list_route(_managed_id_provider, request.method, get_request_route(request))
1331 and _list_prisma is not None
1332 ):
1333 _list_result: Final = await list_passthrough_ids_from_db(
1334 provider=_managed_id_provider,
1335 route=get_request_route(request),
1336 user_api_key_dict=user_api_key_dict,
1337 prisma_client=_list_prisma,
1338 query_params=dict(request.query_params),
1339 )
1340 if _list_result is not None:
1341 verbose_proxy_logger.debug(
1342 "pass_through_endpoint: list served from DB route=%s count=%d",
1343 request.url.path,
1344 len(_list_result.get("data", [])),
1345 )
1346 return Response(
1347 content=json.dumps(_list_result),
1348 status_code=200,
1349 media_type="application/json",
1350 )
1352 requested_query_params_str = None
1353 if requested_query_params:
1354 requested_query_params_str = "&".join(f"{k}={v}" for k, v in requested_query_params.items())
1356 logging_url = str(url)
1357 if requested_query_params_str:
1358 if "?" in str(url): 1358 ↛ 1359line 1358 didn't jump to line 1359 because the condition on line 1358 was never true
1359 logging_url = str(url) + "&" + requested_query_params_str
1360 else:
1361 logging_url = str(url) + "?" + requested_query_params_str
1363 logging_obj.pre_call(
1364 input=[{"role": "user", "content": safe_dumps(_parsed_body)}],
1365 api_key="",
1366 additional_args={
1367 "complete_input_dict": _parsed_body,
1368 "api_base": str(logging_url),
1369 "headers": upstream_headers,
1370 },
1371 )
1372 stream = HttpPassThroughEndpointHelpers._update_stream_param_based_on_request_body(
1373 parsed_body=_parsed_body or {},
1374 stream=stream,
1375 )
1377 if stream: 1377 ↛ 1378line 1377 didn't jump to line 1378 because the condition on line 1377 was never true
1378 logging_obj.stream = True
1379 logging_obj.model_call_details["stream"] = True
1381 if is_multipart:
1382 response = await HttpPassThroughEndpointHelpers.make_multipart_http_request(
1383 request=request,
1384 async_client=async_client,
1385 url=url,
1386 headers=upstream_headers,
1387 requested_query_params=requested_query_params,
1388 stream=True,
1389 )
1390 else:
1391 # SigV4-signed callers (Bedrock) supply the exact pre-signed bytes;
1392 # otherwise httpx encodes the parsed JSON dict as before.
1393 req: Final = (
1394 async_client.build_request(
1395 request.method,
1396 url,
1397 params=requested_query_params,
1398 headers=upstream_headers,
1399 content=state_raw_body,
1400 )
1401 if state_raw_body is not None
1402 else async_client.build_request(
1403 request.method,
1404 url,
1405 params=requested_query_params,
1406 headers=upstream_headers,
1407 json=_parsed_body,
1408 )
1409 )
1411 response = await async_client.send(req, stream=stream)
1413 upstream_usage = apply_upstream_reported_usage(
1414 logging_obj=logging_obj,
1415 headers=response.headers,
1416 )
1418 relay_response: Final = await _log_passthrough_upstream_failure(
1419 response=response,
1420 user_api_key_dict=user_api_key_dict,
1421 request_payload=_build_passthrough_failure_request_payload(
1422 parsed_body=_parsed_body,
1423 kwargs=kwargs,
1424 logging_obj=logging_obj,
1425 custom_llm_provider=custom_llm_provider,
1426 upstream_usage=upstream_usage,
1427 ),
1428 logging_obj=logging_obj,
1429 )
1431 # Call response headers hook for streaming pass-through
1432 _response_headers = HttpPassThroughEndpointHelpers.get_response_headers(
1433 headers=relay_response.headers,
1434 litellm_call_id=litellm_call_id,
1435 )
1436 callback_headers = await proxy_logging_obj.post_call_response_headers_hook(
1437 data=_parsed_body or {},
1438 user_api_key_dict=user_api_key_dict,
1439 response=relay_response,
1440 request_headers=dict(request.headers),
1441 )
1442 if callback_headers:
1443 _response_headers.update(callback_headers)
1445 return StreamingResponse(
1446 wrap_passthrough_sse_bytes_with_keepalive_pings(
1447 stream=_own_streamed_managed_ids(
1448 stream=_relay_reporting_failures(
1449 stream=PassThroughStreamingHandler.chunk_processor(
1450 response=relay_response,
1451 request_body=_parsed_body,
1452 litellm_logging_obj=logging_obj,
1453 endpoint_type=endpoint_type,
1454 start_time=start_time,
1455 passthrough_success_handler_obj=pass_through_endpoint_logging,
1456 url_route=str(url),
1457 ),
1458 upstream_status=relay_response.status_code,
1459 user_api_key_dict=user_api_key_dict,
1460 request_payload=_build_passthrough_failure_request_payload(
1461 parsed_body=_parsed_body,
1462 kwargs=kwargs,
1463 logging_obj=logging_obj,
1464 custom_llm_provider=custom_llm_provider,
1465 ),
1466 ),
1467 managed_id_provider=_managed_id_provider,
1468 request=request,
1469 user_api_key_dict=user_api_key_dict,
1470 ),
1471 ping_interval_seconds=litellm.sse_keepalive_ping_interval_seconds,
1472 upstream_headers=relay_response.headers,
1473 ),
1474 headers=_response_headers,
1475 status_code=relay_response.status_code,
1476 )
1478 if state_raw_body is not None: 1478 ↛ 1481line 1478 didn't jump to line 1481 because the condition on line 1478 was never true
1479 # SigV4-signed callers (Bedrock) require the exact pre-signed bytes
1480 # to be forwarded so the signature/Content-Length stay valid.
1481 raw_body_request: Final = async_client.build_request(
1482 request.method,
1483 url,
1484 headers=upstream_headers,
1485 params=requested_query_params,
1486 content=state_raw_body,
1487 )
1488 response = await async_client.send(raw_body_request, stream=True)
1489 else:
1490 response = await HttpPassThroughEndpointHelpers.non_streaming_http_request_handler(
1491 request=request,
1492 async_client=async_client,
1493 url=url,
1494 headers=upstream_headers,
1495 requested_query_params=requested_query_params,
1496 _parsed_body=_parsed_body,
1497 forward_multipart=is_multipart,
1498 )
1499 verbose_proxy_logger.debug("response.headers= %s", response.headers)
1501 upstream_usage = apply_upstream_reported_usage(
1502 logging_obj=logging_obj,
1503 headers=response.headers,
1504 )
1506 if _is_streaming_response(response) is True: 1506 ↛ 1507line 1506 didn't jump to line 1507 because the condition on line 1506 was never true
1507 logging_obj.stream = True
1508 logging_obj.model_call_details["stream"] = True
1510 detected_relay_response: Final = await _log_passthrough_upstream_failure(
1511 response=response,
1512 user_api_key_dict=user_api_key_dict,
1513 request_payload=_build_passthrough_failure_request_payload(
1514 parsed_body=_parsed_body,
1515 kwargs=kwargs,
1516 logging_obj=logging_obj,
1517 custom_llm_provider=custom_llm_provider,
1518 upstream_usage=upstream_usage,
1519 ),
1520 logging_obj=logging_obj,
1521 )
1523 # Call response headers hook for detected streaming pass-through
1524 _response_headers = HttpPassThroughEndpointHelpers.get_response_headers(
1525 headers=detected_relay_response.headers,
1526 litellm_call_id=litellm_call_id,
1527 )
1528 callback_headers = await proxy_logging_obj.post_call_response_headers_hook(
1529 data=_parsed_body or {},
1530 user_api_key_dict=user_api_key_dict,
1531 response=detected_relay_response,
1532 request_headers=dict(request.headers),
1533 )
1534 if callback_headers:
1535 _response_headers.update(callback_headers)
1537 return StreamingResponse(
1538 wrap_passthrough_sse_bytes_with_keepalive_pings(
1539 stream=_own_streamed_managed_ids(
1540 stream=_relay_reporting_failures(
1541 stream=PassThroughStreamingHandler.chunk_processor(
1542 response=detected_relay_response,
1543 request_body=_parsed_body,
1544 litellm_logging_obj=logging_obj,
1545 endpoint_type=endpoint_type,
1546 start_time=start_time,
1547 passthrough_success_handler_obj=pass_through_endpoint_logging,
1548 url_route=str(url),
1549 ),
1550 upstream_status=detected_relay_response.status_code,
1551 user_api_key_dict=user_api_key_dict,
1552 request_payload=_build_passthrough_failure_request_payload(
1553 parsed_body=_parsed_body,
1554 kwargs=kwargs,
1555 logging_obj=logging_obj,
1556 custom_llm_provider=custom_llm_provider,
1557 ),
1558 ),
1559 managed_id_provider=_managed_id_provider,
1560 request=request,
1561 user_api_key_dict=user_api_key_dict,
1562 ),
1563 ping_interval_seconds=litellm.sse_keepalive_ping_interval_seconds,
1564 upstream_headers=detected_relay_response.headers,
1565 ),
1566 headers=_response_headers,
1567 status_code=detected_relay_response.status_code,
1568 )
1570 if not _should_buffer_passthrough_response(response): 1570 ↛ 1571line 1570 didn't jump to line 1571 because the condition on line 1570 was never true
1571 relay_custom_headers: Final = ProxyBaseLLMRequestProcessing.get_custom_headers(
1572 user_api_key_dict=user_api_key_dict,
1573 call_id=litellm_call_id,
1574 model_id=None,
1575 cache_key=None,
1576 api_base=str(url._uri_reference),
1577 )
1578 relay_callback_headers: Final = await proxy_logging_obj.post_call_response_headers_hook(
1579 data=_parsed_body or {},
1580 user_api_key_dict=user_api_key_dict,
1581 response=response,
1582 request_headers=dict(request.headers),
1583 )
1584 if relay_callback_headers:
1585 relay_custom_headers.update(relay_callback_headers)
1587 return StreamingResponse(
1588 _relay_passthrough_response_bytes(
1589 response=response,
1590 request_body=_parsed_body or {},
1591 url_route=str(url),
1592 start_time=start_time,
1593 logging_obj=logging_obj,
1594 custom_llm_provider=custom_llm_provider,
1595 success_handler_kwargs=kwargs,
1596 ),
1597 status_code=response.status_code,
1598 headers=HttpPassThroughEndpointHelpers.get_response_headers(
1599 headers=response.headers,
1600 custom_headers=relay_custom_headers,
1601 ),
1602 )
1604 content = await response.aread()
1606 ## POST-CALL GUARDRAILS ##
1607 # Guardrails and managed-id rewriting only apply to successful upstream
1608 # responses; response_body itself is parsed unconditionally so the
1609 # failure-hook log payload below still reflects upstream error bodies.
1610 _content_modified = False
1611 response_body: dict | None = get_response_body(response)
1613 failure_request_payload: Final = _build_passthrough_failure_request_payload(
1614 parsed_body=_parsed_body,
1615 kwargs=kwargs,
1616 logging_obj=logging_obj,
1617 custom_llm_provider=custom_llm_provider,
1618 upstream_usage=upstream_usage,
1619 )
1620 failure_request_payload["response_body"] = response_body
1621 await _log_passthrough_upstream_failure(
1622 response=response,
1623 user_api_key_dict=user_api_key_dict,
1624 request_payload=failure_request_payload,
1625 logging_obj=logging_obj,
1626 )
1628 if response.status_code < 400 and response_body is not None and guardrails_to_run: 1628 ↛ 1633line 1628 didn't jump to line 1633 because the condition on line 1628 was never true
1629 # Build an enriched data dict: _parsed_body has been stripped of
1630 # `metadata` by both pre_call_hook and _init_kwargs_for_pass_through_endpoint,
1631 # so we re-attach the configured guardrails here so should_run_guardrail
1632 # sees them.
1633 hook_data: Final = dict(_parsed_body or {})
1634 existing_metadata = hook_data.get("metadata")
1635 if not isinstance(existing_metadata, dict):
1636 existing_metadata = {}
1637 hook_data["metadata"] = {
1638 **existing_metadata,
1639 "guardrails": guardrails_to_run,
1640 }
1641 post_call_guardrail_data = hook_data
1642 response_body = await proxy_logging_obj.post_call_success_hook(
1643 data=hook_data,
1644 user_api_key_dict=user_api_key_dict,
1645 response=response_body,
1646 )
1647 if isinstance(response_body, dict):
1648 content = json.dumps(response_body).encode("utf-8")
1649 _content_modified = True
1650 else:
1651 verbose_proxy_logger.debug(
1652 "pass_through_endpoint: post_call_success_hook returned %s, expected dict — using original response",
1653 type(response_body).__name__,
1654 )
1655 elif response_body is None:
1656 verbose_proxy_logger.debug(
1657 "pass_through_endpoint: response body not JSON-parseable, skipping post-call guardrails"
1658 )
1660 ## PASSTHROUGH MANAGED ID MINTING (OUTPUT) ##
1661 # Mint managed IDs for raw provider IDs in the response body and swap
1662 # them before the response reaches the client. Runs after guardrails
1663 # so guardrails see the raw IDs (cleaner) and the client receives the
1664 # managed IDs. Gated by feature flag and enterprise managed-files hook.
1665 if ( 1665 ↛ 1671line 1665 didn't jump to line 1671 because the condition on line 1665 was never true
1666 proxy_general_settings.get("passthrough_managed_object_ids", False)
1667 and _managed_id_provider is not None
1668 and isinstance(response_body, dict)
1669 and response.status_code < 300
1670 ):
1671 verbose_proxy_logger.debug(
1672 "pass_through_endpoint: managed-id output rewrite enabled for route=%s method=%s status=%s",
1673 request.url.path,
1674 request.method,
1675 response.status_code,
1676 )
1677 _passthrough_managed_hook = proxy_logging_obj.get_proxy_hook("managed_files")
1678 if _passthrough_managed_hook is not None:
1679 from litellm.proxy.auth.auth_utils import get_request_route
1680 from litellm.proxy.pass_through_endpoints.managed_id_rewriter import (
1681 rewrite_response_ids,
1682 )
1683 from litellm.proxy.proxy_server import (
1684 prisma_client as _passthrough_prisma,
1685 )
1687 _new_body: Final = await rewrite_response_ids(
1688 provider=_managed_id_provider,
1689 method=request.method,
1690 route=get_request_route(request),
1691 body=response_body,
1692 user_api_key_dict=user_api_key_dict,
1693 prisma_client=_passthrough_prisma,
1694 managed_files_hook=_passthrough_managed_hook,
1695 )
1696 if _new_body is not response_body:
1697 response_body = _new_body
1698 content = json.dumps(response_body).encode("utf-8")
1699 _content_modified = True
1700 verbose_proxy_logger.debug(
1701 "pass_through_endpoint: managed-id output rewrite applied route=%s method=%s",
1702 request.url.path,
1703 request.method,
1704 )
1705 else:
1706 verbose_proxy_logger.debug(
1707 "pass_through_endpoint: managed-id output rewrite no-op route=%s method=%s",
1708 request.url.path,
1709 request.method,
1710 )
1711 else:
1712 verbose_proxy_logger.debug(
1713 "pass_through_endpoint: managed-id output rewrite skipped (managed_files hook not available) route=%s method=%s",
1714 request.url.path,
1715 request.method,
1716 )
1718 ## LOG SUCCESS
1719 # Upstream errors are already logged via _log_passthrough_upstream_failure
1720 # above; the success handler has no status-code awareness of its own; so
1721 # calling it here for a 4xx/5xx would double-log the same request as both
1722 # a failure and a success (corrupting spend tracking).
1723 passthrough_logging_payload["response_body"] = response_body
1724 end_time: Final = datetime.now()
1725 if response.status_code < 400: 1725 ↛ 1726line 1725 didn't jump to line 1726 because the condition on line 1725 was never true
1726 GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue(
1727 async_coroutine=pass_through_endpoint_logging.pass_through_async_success_handler(
1728 httpx_response=response,
1729 response_body=response_body,
1730 url_route=str(url),
1731 result="",
1732 start_time=start_time,
1733 end_time=end_time,
1734 logging_obj=logging_obj,
1735 cache_hit=False,
1736 request_body=_parsed_body or {},
1737 custom_llm_provider=custom_llm_provider,
1738 **kwargs,
1739 )
1740 )
1741 bind_budget_reservation_to_callbacks(logging_obj.litellm_params)
1743 ## CUSTOM HEADERS - `x-litellm-*`
1744 custom_headers = ProxyBaseLLMRequestProcessing.get_custom_headers(
1745 user_api_key_dict=user_api_key_dict,
1746 call_id=litellm_call_id,
1747 model_id=None,
1748 cache_key=None,
1749 api_base=str(url._uri_reference),
1750 )
1752 # Call response headers hook
1753 callback_headers = await proxy_logging_obj.post_call_response_headers_hook(
1754 data=_parsed_body or {},
1755 user_api_key_dict=user_api_key_dict,
1756 response=response,
1757 request_headers=dict(request.headers),
1758 )
1759 if callback_headers: 1759 ↛ 1760line 1759 didn't jump to line 1760 because the condition on line 1759 was never true
1760 custom_headers.update(callback_headers)
1762 response_headers: Final = HttpPassThroughEndpointHelpers.get_response_headers(
1763 headers=response.headers,
1764 custom_headers=custom_headers,
1765 )
1766 emitted_call_id: Final = (
1767 JSON_OBJECT.validate_python(response_headers).get(LITELLM_CALL_ID_HEADER)
1768 if response.status_code >= 400
1769 else None
1770 )
1771 error_call_id: Final = (
1772 error_body_call_id(general_settings_view(), emitted_call_id) if isinstance(emitted_call_id, str) else None
1773 )
1774 relayed_content: Final = (
1775 json.dumps(with_call_id(JSON_OBJECT.validate_python(response_body), error_call_id)).encode("utf-8")
1776 if error_call_id is not None and isinstance(response_body, dict)
1777 else content
1778 )
1779 if _content_modified: 1779 ↛ 1780line 1779 didn't jump to line 1780 because the condition on line 1779 was never true
1780 response_headers.pop("content-length", None)
1782 return Response(
1783 content=relayed_content,
1784 status_code=response.status_code,
1785 headers=response_headers,
1786 )
1787 except ModifyResponseException as e:
1788 verbose_proxy_logger.info(
1789 "pass_through_endpoint: Guardrail %s modified response: %s",
1790 e.guardrail_name,
1791 str(e.message or "")[:200],
1792 )
1793 try:
1794 await proxy_logging_obj.post_call_failure_hook(
1795 user_api_key_dict=user_api_key_dict,
1796 original_exception=e,
1797 request_data=e.request_data,
1798 )
1799 except Exception:
1800 verbose_proxy_logger.warning(
1801 "pass_through_endpoint: post_call_failure_hook raised during guardrail block",
1802 exc_info=True,
1803 )
1804 error_body: Final = {
1805 "error": {
1806 "message": e.message or "Response blocked by guardrail",
1807 "type": "content_filter",
1808 "guardrail_name": e.guardrail_name,
1809 "model": e.model,
1810 }
1811 }
1812 return Response(
1813 content=json.dumps(error_body),
1814 status_code=200,
1815 media_type="application/json",
1816 )
1817 except Exception as e:
1818 custom_headers = ProxyBaseLLMRequestProcessing.get_custom_headers(
1819 user_api_key_dict=user_api_key_dict,
1820 call_id=litellm_call_id,
1821 model_id=None,
1822 cache_key=None,
1823 api_base=str(url._uri_reference) if url else None,
1824 )
1825 if CustomGuardrail._is_guardrail_intervention(e): 1825 ↛ 1826line 1825 didn't jump to line 1826 because the condition on line 1825 was never true
1826 verbose_proxy_logger.warning(
1827 "pass_through_endpoint: request blocked by guardrail - %s",
1828 str(e),
1829 )
1830 else:
1831 verbose_proxy_logger.exception(
1832 "litellm.proxy.proxy_server.pass_through_endpoint(): Exception occured - %s", e
1833 )
1835 #########################################################
1836 # Monitoring: Trigger post_call_failure_hook
1837 # for pass through endpoint failure
1838 #########################################################
1839 request_payload: Final[dict] = _parsed_body or {}
1840 # add user_api_key_dict, litellm_call_id, passthrough_logging_payloa for logging
1841 if kwargs: 1841 ↛ 1844line 1841 didn't jump to line 1844 because the condition on line 1841 was always true
1842 for key, value in kwargs.items():
1843 request_payload[key] = value
1844 if logging_obj is not None: 1844 ↛ 1847line 1844 didn't jump to line 1847 because the condition on line 1844 was always true
1845 request_payload["litellm_logging_obj"] = logging_obj
1847 if "model" not in request_payload and _parsed_body and isinstance(_parsed_body, dict):
1848 request_payload["model"] = _parsed_body.get("model", "")
1849 if "custom_llm_provider" not in request_payload and custom_llm_provider:
1850 request_payload["custom_llm_provider"] = custom_llm_provider
1852 _carry_guardrail_logging_info(request_payload, post_call_guardrail_data)
1854 await proxy_logging_obj.post_call_failure_hook(
1855 user_api_key_dict=user_api_key_dict,
1856 original_exception=e,
1857 request_data=request_payload,
1858 traceback_str=traceback.format_exc(
1859 limit=MAXIMUM_TRACEBACK_LINES_TO_LOG,
1860 ),
1861 )
1863 #########################################################
1865 if isinstance(e, ProxyException): 1865 ↛ 1866line 1865 didn't jump to line 1866 because the condition on line 1865 was never true
1866 raise
1867 if isinstance(e, HTTPException): 1867 ↛ 1868line 1867 didn't jump to line 1868 because the condition on line 1867 was never true
1868 raise ProxyException(
1869 message=getattr(e, "message", str(getattr(e, "detail", str(e)))),
1870 type=openai_error_type(e, error_status_code(e, status.HTTP_400_BAD_REQUEST)),
1871 param=openai_error_param(e),
1872 code=error_status_code(e, status.HTTP_400_BAD_REQUEST),
1873 headers=custom_headers,
1874 )
1875 else:
1876 error_msg: Final = f"{e}"
1877 raise ProxyException(
1878 message=getattr(e, "message", error_msg),
1879 type=openai_error_type(e, error_status_code(e, 500)),
1880 param=openai_error_param(e),
1881 code=error_status_code(e, 500),
1882 headers=custom_headers,
1883 )
1886def _update_metadata_with_tags_in_header(request: Request, metadata: dict) -> dict:
1887 """
1888 If tags are in the request headers, add them to the metadata
1890 Used for google and vertex JS SDKs, and Azure passthrough
1891 Checks both 'tags' and 'x-litellm-tags' headers
1892 """
1893 tags_to_add: Final = []
1895 # Check for 'tags' header first
1896 _tags = request.headers.get("tags")
1897 if _tags: 1897 ↛ 1898line 1897 didn't jump to line 1898 because the condition on line 1897 was never true
1898 tags_to_add.extend([tag.strip() for tag in _tags.split(",")])
1900 _tags = request.headers.get("x-litellm-tags")
1901 if _tags: 1901 ↛ 1902line 1901 didn't jump to line 1902 because the condition on line 1901 was never true
1902 tags_to_add.extend([tag.strip() for tag in _tags.split(",")])
1904 # Only add tags key if there are tags to add
1905 if tags_to_add: 1905 ↛ 1906line 1905 didn't jump to line 1906 because the condition on line 1905 was never true
1906 if "tags" not in metadata:
1907 metadata["tags"] = []
1908 metadata["tags"].extend(tags_to_add)
1910 return metadata
1913class _PassThroughRequestEnvelope(TypedDict, total=False):
1914 query_params: Mapping[str, object] | None
1915 custom_body: Mapping[str, object] | None
1916 stream: bool | None
1919async def _parse_request_data_by_content_type(
1920 request: Request,
1921) -> tuple[object, object, None, bool | None]:
1922 """
1923 Parse request data based on content type.
1925 Handles JSON, multipart/form-data, and URL-encoded form data.
1927 Returns:
1928 Tuple of (query_params_data, custom_body_data, file_data, stream)
1929 """
1930 content_type: Final = request.headers.get("content-type", "")
1932 query_params_data = None
1933 custom_body_data = None
1934 file_data: Final = None
1935 stream = None
1937 if "application/json" in content_type:
1938 # ✅ Handle JSON
1939 try:
1940 body: _PassThroughRequestEnvelope = await request.json()
1941 query_params_data = body.get("query_params")
1942 custom_body_data = body.get("custom_body")
1943 stream = body.get("stream")
1944 except json.JSONDecodeError:
1945 # Handle requests with no body (e.g., DELETE requests)
1946 pass
1947 elif "multipart/form-data" in content_type: 1947 ↛ 1950line 1947 didn't jump to line 1950 because the condition on line 1947 was never true
1948 # ✅ Try to parse as JSON first (handles misconfigured clients sending JSON with multipart content-type)
1949 # If that fails, skip parsing - pass_through_request will handle actual multipart
1950 try:
1951 body = await request.json()
1952 # Successfully parsed as JSON - treat as JSON body
1953 query_params_data = body.get("query_params")
1954 custom_body_data = body.get("custom_body")
1955 stream = body.get("stream")
1956 # If custom_body is not set, use the entire body
1957 if custom_body_data is None and body:
1958 custom_body_data = body
1959 except (json.JSONDecodeError, Exception):
1960 # Not JSON - this is actual multipart data
1961 # Skip parsing here to avoid consuming the request body stream
1962 # make_multipart_http_request will handle it
1963 pass
1965 elif "application/x-www-form-urlencoded" in content_type: 1965 ↛ 1967line 1965 didn't jump to line 1967 because the condition on line 1965 was never true
1966 # ✅ Handle URL-encoded form data
1967 form: Final = await request.form()
1968 query_params_data = form.get("query_params")
1969 custom_body_data = form.get("custom_body")
1971 else:
1972 # ✅ Fallback: maybe no body, just query params
1973 query_params_data = dict(request.query_params) or None
1975 return query_params_data, custom_body_data, file_data, stream
1978def create_pass_through_route(
1979 endpoint,
1980 target: str,
1981 custom_headers: Mapping[str, object] | None = None,
1982 _forward_headers: bool | None = False,
1983 _merge_query_params: bool | None = False,
1984 dependencies: list | None = None,
1985 include_subpath: bool | None = False,
1986 cost_per_request: float | None = None,
1987 custom_llm_provider: str | None = None,
1988 is_streaming_request: bool | None = False,
1989 query_params: dict | None = None,
1990 default_query_params: dict | None = None,
1991 guardrails: dict[str, object] | None = None,
1992 config_file_path: str | None = None,
1993 timeout: float | None = None,
1994):
1995 # check if target is an adapter.py or a url
1996 from litellm._uuid import uuid
1997 from litellm.proxy.types_utils.utils import get_instance_fn
1999 try:
2000 if isinstance(target, CustomLogger): 2000 ↛ 2001line 2000 didn't jump to line 2001 because the condition on line 2000 was never true
2001 adapter = target
2002 else:
2003 adapter = get_instance_fn(value=target, config_file_path=config_file_path)
2004 adapter_id: Final = str(uuid.uuid4())
2005 litellm.adapters = [{"id": adapter_id, "adapter": adapter}]
2007 async def endpoint_func(
2008 request: Request,
2009 fastapi_response: Response,
2010 user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
2011 subpath: str = "", # captures sub-paths when include_subpath=True
2012 ):
2013 return await chat_completion_pass_through_endpoint(
2014 fastapi_response=fastapi_response,
2015 request=request,
2016 adapter_id=adapter_id,
2017 user_api_key_dict=user_api_key_dict,
2018 )
2020 except Exception:
2021 verbose_proxy_logger.debug("Defaulting to target being a url.")
2023 async def endpoint_func(
2024 request: Request,
2025 fastapi_response: Response,
2026 user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
2027 subpath: str = "", # captures sub-paths when include_subpath=True
2028 ):
2029 from litellm.proxy.auth.auth_utils import ( # noqa: PLC0415
2030 get_request_route,
2031 )
2032 from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
2033 InitPassThroughEndpointHelpers,
2034 )
2036 path: Final = get_request_route(request)
2038 # Parse request data based on content type
2039 (
2040 query_params_data,
2041 custom_body_data,
2042 file_data,
2043 stream,
2044 ) = await _parse_request_data_by_content_type(request)
2046 if not InitPassThroughEndpointHelpers.is_registered_pass_through_route(route=path): 2046 ↛ 2047line 2046 didn't jump to line 2047 because the condition on line 2046 was never true
2047 raise HTTPException(
2048 status_code=404,
2049 detail=f"Pass-through endpoint {endpoint} not found. This could have been deleted or not yet added to the proxy.",
2050 )
2052 passthrough_params: Final = InitPassThroughEndpointHelpers.get_registered_pass_through_route(
2053 route=path, method=request.method
2054 )
2055 if ( 2055 ↛ 2059line 2055 didn't jump to line 2059 because the condition on line 2055 was never true
2056 passthrough_params is None
2057 and InitPassThroughEndpointHelpers.get_registered_pass_through_route(route=path) is not None
2058 ):
2059 raise HTTPException(
2060 status_code=status.HTTP_405_METHOD_NOT_ALLOWED,
2061 detail=f"Method {request.method} is not allowed for pass-through endpoint {path}.",
2062 )
2063 target_params: Final = {
2064 "target": target,
2065 "custom_headers": custom_headers,
2066 "forward_headers": _forward_headers,
2067 "merge_query_params": _merge_query_params,
2068 "cost_per_request": cost_per_request,
2069 "guardrails": None,
2070 "timeout": timeout,
2071 }
2073 if passthrough_params is not None:
2074 target_params.update(passthrough_params.get("passthrough_params", {}))
2076 # Extract and cast parameters with proper types
2077 param_target: Final = target_params.get("target") or target
2078 param_custom_headers: Final = target_params.get("custom_headers", custom_headers)
2079 param_forward_headers: Final = target_params.get("forward_headers", _forward_headers)
2080 param_merge_query_params: Final = target_params.get("merge_query_params", _merge_query_params)
2081 param_cost_per_request: Final = target_params.get("cost_per_request", cost_per_request)
2082 param_guardrails: Final = target_params.get("guardrails", None)
2083 param_default_query_params: Final = target_params.get("default_query_params", None)
2084 param_timeout: Final = target_params.get("timeout", timeout)
2086 # Construct the full target URL with subpath if needed
2087 full_target: Final = HttpPassThroughEndpointHelpers.construct_target_url_with_subpath(
2088 base_target=cast(str, param_target),
2089 subpath=subpath,
2090 include_subpath=include_subpath,
2091 )
2093 # Ensure custom_headers is a dict. Botocore returns a HeadersDict
2094 # for SigV4-prepared requests, which is a Mapping but not a dict.
2095 headers_dict: Final = dict(param_custom_headers) if isinstance(param_custom_headers, Mapping) else {}
2097 # Ensure query_params and custom_body are dicts or None
2098 final_query_params: Final = query_params_data if isinstance(query_params_data, dict) else {}
2099 if query_params: 2099 ↛ 2100line 2099 didn't jump to line 2100 because the condition on line 2099 was never true
2100 final_query_params.update(query_params)
2101 # Programmatic callers set LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY on
2102 # request.state (see Bedrock proxy). Parsed JSON envelope otherwise.
2103 state_custom_body: Final[dict | None] = getattr(
2104 request.state,
2105 LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY,
2106 None,
2107 )
2108 final_custom_body: dict | None = None
2109 if isinstance(state_custom_body, dict): 2109 ↛ 2110line 2109 didn't jump to line 2110 because the condition on line 2109 was never true
2110 final_custom_body = state_custom_body
2111 elif isinstance(custom_body_data, dict): 2111 ↛ 2112line 2111 didn't jump to line 2112 because the condition on line 2111 was never true
2112 final_custom_body = custom_body_data
2114 is_stream: Final = bool(is_streaming_request or stream)
2116 async def _relay() -> Response:
2117 try:
2118 return await pass_through_request(
2119 request=request,
2120 target=full_target,
2121 custom_headers=headers_dict,
2122 user_api_key_dict=user_api_key_dict,
2123 forward_headers=cast(bool | None, param_forward_headers),
2124 merge_query_params=cast(bool | None, param_merge_query_params),
2125 query_params=final_query_params,
2126 default_query_params=cast(dict | None, param_default_query_params),
2127 stream=is_stream,
2128 custom_body=final_custom_body,
2129 cost_per_request=cast(float | None, param_cost_per_request),
2130 custom_llm_provider=custom_llm_provider,
2131 guardrails_config=cast(dict | None, param_guardrails),
2132 timeout=cast(float | None, param_timeout),
2133 )
2134 finally:
2135 if hasattr(request.state, LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY): 2135 ↛ 2136line 2135 didn't jump to line 2136 because the condition on line 2135 was never true
2136 delattr(request.state, LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY)
2137 if hasattr(request.state, LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY): 2137 ↛ 2138line 2137 didn't jump to line 2138 because the condition on line 2137 was never true
2138 delattr(request.state, LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY)
2139 if hasattr(request.state, LITELLM_PASS_THROUGH_DEPLOYMENT_MODEL_INFO_STATE_KEY):
2140 delattr(request.state, LITELLM_PASS_THROUGH_DEPLOYMENT_MODEL_INFO_STATE_KEY)
2142 # The upstream withholds its response headers until its first token, so
2143 # the whole time-to-first-token is spent inside _relay with nothing on
2144 # the wire. Off unless an operator sets an interval.
2145 return await open_sse_before_first_byte(
2146 _relay(),
2147 ping_interval_seconds=(litellm.sse_keepalive_ping_interval_seconds if is_stream else None),
2148 )
2150 setattr(endpoint_func, LITELLM_PASS_THROUGH_ENDPOINT_MARKER, True)
2151 return endpoint_func
2154def create_websocket_passthrough_route(
2155 endpoint: str,
2156 target: str,
2157 custom_headers: dict | None = None,
2158 _forward_headers: bool | None = False,
2159 dependencies: list | None = None,
2160 cost_per_request: float | None = None,
2161):
2162 """
2163 Create a WebSocket passthrough route function.
2165 Args:
2166 endpoint: The endpoint path (for logging purposes)
2167 target: The target WebSocket URL (e.g., "wss://api.example.com/ws")
2168 custom_headers: Custom headers to include in the WebSocket connection
2169 _forward_headers: Whether to forward incoming headers
2170 dependencies: FastAPI dependencies to inject
2172 Returns:
2173 A WebSocket passthrough function that can be registered with app.websocket()
2174 """
2175 from litellm.proxy.auth.user_api_key_auth import user_api_key_auth_websocket
2177 async def websocket_endpoint_func(
2178 websocket: WebSocket,
2179 user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth_websocket),
2180 **kwargs, # For additional query parameters
2181 ):
2182 """
2183 WebSocket passthrough endpoint function.
2185 This function handles the WebSocket connection by:
2186 1. Accepting the incoming WebSocket connection
2187 2. Establishing a connection to the target WebSocket
2188 3. Forwarding messages bidirectionally
2189 4. Handling connection cleanup
2190 """
2191 return await websocket_passthrough_request(
2192 websocket=websocket,
2193 target=target,
2194 custom_headers=custom_headers or {},
2195 user_api_key_dict=user_api_key_dict,
2196 forward_headers=_forward_headers,
2197 endpoint=endpoint,
2198 cost_per_request=cost_per_request,
2199 accept_websocket=True, # Generic usage should accept the WebSocket
2200 )
2202 return websocket_endpoint_func
2205def _rewrite_vertex_live_setup_model(text_data: str, setup_model_rewriter: Callable[[str], str] | None) -> str:
2206 """
2207 Rewrite the model of a Vertex AI Live ``setup`` frame, leaving every other frame byte-identical
2208 """
2209 if setup_model_rewriter is None:
2210 return text_data
2211 try:
2212 message: Final = json.loads(text_data)
2213 except json.JSONDecodeError:
2214 return text_data
2215 if not isinstance(message, dict):
2216 return text_data
2217 setup: Final = message.get("setup")
2218 if not isinstance(setup, dict):
2219 return text_data
2220 setup_model: Final = setup.get("model")
2221 if not isinstance(setup_model, str):
2222 return text_data
2223 rewritten_model: Final = setup_model_rewriter(setup_model)
2224 if rewritten_model == setup_model:
2225 return text_data
2226 return json.dumps({**message, "setup": {**setup, "model": rewritten_model}}) # mutable-ok: one-shot json payload
2229def _resolved_vertex_live_setup(
2230 setup_data: Mapping[str, object], setup_model_rewriter: Callable[[str], str] | None
2231) -> Mapping[str, object]:
2232 """
2233 Give the model extractor the same fully qualified path the upstream will receive.
2235 Clients may name a bare gateway alias, which the rewriter turns into a ``projects/...`` path before
2236 it reaches Vertex. The extractor only reads a path containing ``/models/``, so running it on the raw
2237 frame logs the session as ``unknown`` at no cost, which is precisely the supported client form
2238 """
2239 setup_model: Final = setup_data.get("model")
2240 if setup_model_rewriter is None or not isinstance(setup_model, str):
2241 return setup_data
2242 return {**setup_data, "model": setup_model_rewriter(setup_model)}
2245def _json_object_frame(frame: str | bytes) -> dict[str, object] | None:
2246 try:
2247 decoded: Final = json.loads(frame if isinstance(frame, str) else frame.decode("utf-8"))
2248 except (json.JSONDecodeError, UnicodeDecodeError):
2249 return None
2250 return decoded if isinstance(decoded, dict) else None
2253def _truncated_close_reason(reason: str) -> str:
2254 """
2255 Fit a close reason inside the byte budget a WebSocket close frame allows, without splitting a character
2256 """
2257 encoded: Final = reason.encode("utf-8")
2258 if len(encoded) <= WEBSOCKET_CLOSE_REASON_MAX_BYTES:
2259 return reason
2260 return encoded[:WEBSOCKET_CLOSE_REASON_MAX_BYTES].decode("utf-8", errors="ignore")
2263SENDABLE_CLOSE_CODES: Final = frozenset(CloseCode) - frozenset(
2264 {CloseCode.NO_STATUS_RCVD, CloseCode.ABNORMAL_CLOSURE, CloseCode.TLS_HANDSHAKE}
2265)
2268def _client_socket_is_open(websocket: WebSocket) -> bool:
2269 """
2270 Starlette tracks the two halves separately and raises on a second close, so both have to still be live
2271 """
2272 return (
2273 websocket.client_state != WebSocketState.DISCONNECTED
2274 and websocket.application_state != WebSocketState.DISCONNECTED
2275 )
2278def _upstream_close_to_relay(task_results: Iterable[object]) -> Close | None:
2279 """
2280 The upstream close worth telling the client about: anything other than a plain, reasonless normal close.
2282 Codes outside ``SENDABLE_CLOSE_CODES`` and the private range never travel on the wire (1006 for a socket that
2283 died without a close frame, 1005 for one that sent no code), so relaying them would build an invalid frame
2284 """
2285 upstream_close: Final = next((result for result in task_results if isinstance(result, Close)), None)
2286 if upstream_close is None:
2287 return None
2288 if upstream_close.code == 1000 and upstream_close.reason == "":
2289 return None
2290 if upstream_close.code not in SENDABLE_CLOSE_CODES and not 3000 <= upstream_close.code < 5000:
2291 return None
2292 return upstream_close
2295_WEBSOCKET_FORWARDED_HEADERS: Final = frozenset(("authorization", "x-api-key", "x-goog-user-project"))
2298def _with_trace_context(headers: Mapping[str, str], parent_span: object) -> dict[str, str]:
2299 try:
2300 from litellm.integrations.otel.plumbing.context import inject_trace_context
2301 except ImportError:
2302 return dict(headers) # mutable-ok: matches inject_trace_context's carrier return type
2303 return inject_trace_context(headers, parent_span=parent_span)
2306async def websocket_passthrough_request(
2307 websocket: WebSocket,
2308 target: str,
2309 custom_headers: dict,
2310 user_api_key_dict: UserAPIKeyAuth,
2311 forward_headers: bool | None = False,
2312 endpoint: str | None = None,
2313 cost_per_request: float | None = None,
2314 accept_websocket: bool = True,
2315 setup_model_rewriter: Callable[[str], str] | None = None,
2316):
2317 """
2318 WebSocket passthrough request handler.
2320 Args:
2321 websocket: The incoming WebSocket connection
2322 target: The target WebSocket URL
2323 custom_headers: Custom headers to include in the connection
2324 user_api_key_dict: The user API key dictionary
2325 forward_headers: Whether to forward incoming headers
2326 endpoint: The endpoint path (for logging purposes)
2327 cost_per_request: Optional field - cost per request to the target endpoint
2328 setup_model_rewriter: Optional rewrite of the setup frame's model before it reaches the upstream
2329 """
2330 from litellm.litellm_core_utils.litellm_logging import Logging
2331 from litellm.proxy.proxy_server import proxy_config, proxy_logging_obj
2332 from litellm.types.passthrough_endpoints.pass_through_endpoints import (
2333 PassthroughStandardLoggingPayload,
2334 )
2336 # Initialize tracking variables
2337 start_time: Final = datetime.now()
2338 websocket_messages: Final[list[dict[str, object]]] = []
2339 litellm_call_id: Final = str(uuid.uuid4())
2341 verbose_proxy_logger.info("WebSocket passthrough (%s): Starting WebSocket connection to %s", endpoint, target)
2343 # Only accept the WebSocket if requested (for generic usage)
2344 if accept_websocket:
2345 await websocket.accept()
2346 verbose_proxy_logger.debug("WebSocket passthrough (%s): WebSocket connection accepted", endpoint)
2348 forwarded_headers: Final = { # mutable-ok: one-shot upstream header dict, read as a Mapping
2349 **custom_headers,
2350 **{
2351 header_name: header_value
2352 for header_name, header_value in websocket.headers.items()
2353 if forward_headers and header_name.lower() in _WEBSOCKET_FORWARDED_HEADERS
2354 },
2355 }
2356 upstream_headers: Final = _with_trace_context(forwarded_headers, parent_span=user_api_key_dict.parent_otel_span)
2358 # Initialize logging object similar to HTTP passthrough
2359 team_callbacks: Final = _resolve_team_callback_wiring(
2360 user_api_key_dict=user_api_key_dict,
2361 proxy_config=proxy_config,
2362 route_description="websocket_passthrough",
2363 )
2364 logging_obj: Final = Logging(
2365 model="unknown",
2366 messages=[{"role": "user", "content": "WebSocket connection"}],
2367 stream=True, # WebSockets are inherently streaming
2368 call_type="pass_through_endpoint",
2369 start_time=start_time,
2370 litellm_call_id=litellm_call_id,
2371 function_id="websocket_passthrough",
2372 dynamic_success_callbacks=team_callbacks.success_callbacks,
2373 dynamic_failure_callbacks=team_callbacks.failure_callbacks,
2374 kwargs=team_callbacks.logging_kwargs,
2375 )
2377 # Create passthrough logging payload
2378 passthrough_logging_payload: Final = PassthroughStandardLoggingPayload(
2379 url=target,
2380 request_body={}, # WebSocket doesn't have a traditional request body
2381 request_method="WEBSOCKET",
2382 cost_per_request=cost_per_request,
2383 )
2385 # Create a dummy request object for WebSocket connections to maintain compatibility
2386 # with the existing _init_kwargs_for_pass_through_endpoint function
2387 class DummyRequest:
2388 def __init__(self, url: str, method: str = "WEBSOCKET", headers: dict | None = None):
2389 self.url = url
2390 self.method = method
2391 self.headers = headers or {}
2393 def __str__(self):
2394 return f"DummyRequest(url={self.url}, method={self.method})"
2396 dummy_request: Final = DummyRequest(
2397 url=target,
2398 method="WEBSOCKET",
2399 headers=dict(websocket.headers) if hasattr(websocket, "headers") else {},
2400 )
2402 # Initialize kwargs for logging using the same pattern as HTTP passthrough
2403 kwargs: Final = HttpPassThroughEndpointHelpers._init_kwargs_for_pass_through_endpoint(
2404 user_api_key_dict=user_api_key_dict,
2405 _parsed_body={}, # WebSocket doesn't have a traditional request body
2406 passthrough_logging_payload=passthrough_logging_payload,
2407 litellm_call_id=litellm_call_id,
2408 request=dummy_request,
2409 logging_obj=logging_obj,
2410 )
2412 # Update logging environment variables
2413 logging_obj.update_environment_variables(
2414 model="unknown",
2415 user="unknown",
2416 optional_params={},
2417 litellm_params=dict(kwargs.get("litellm_params", {})),
2418 call_type="pass_through_endpoint",
2419 )
2420 logging_obj.model_call_details["litellm_call_id"] = litellm_call_id
2422 # Pre-call logging
2423 logging_obj.pre_call(
2424 input=[{"role": "user", "content": "WebSocket connection"}],
2425 api_key="",
2426 additional_args={
2427 "complete_input_dict": {},
2428 "api_base": target,
2429 "headers": upstream_headers,
2430 },
2431 )
2433 ### CALL HOOKS ### - modify incoming data / reject request before calling the model
2434 websocket_data: dict[str, object] = {}
2435 websocket_data = await proxy_logging_obj.pre_call_hook(
2436 user_api_key_dict=user_api_key_dict,
2437 data=websocket_data,
2438 call_type="pass_through_endpoint",
2439 )
2441 try:
2442 verbose_proxy_logger.debug(
2443 "WebSocket passthrough (%s): Establishing upstream connection to %s", endpoint, target
2444 )
2445 async with connect(
2446 target,
2447 additional_headers=upstream_headers,
2448 ) as upstream_ws:
2449 verbose_proxy_logger.info(
2450 "WebSocket passthrough (%s): Upstream connection established successfully", endpoint
2451 )
2453 async def forward_client_to_upstream() -> None:
2454 """Forward messages from client to upstream WebSocket"""
2455 try:
2456 while True:
2457 message = await websocket.receive()
2458 message_type = message.get("type")
2459 if message_type == "websocket.disconnect":
2460 await upstream_ws.close()
2461 break
2463 text_data: str | None = message.get("text")
2464 bytes_data: bytes | None = message.get("bytes")
2466 if text_data is not None:
2467 # Try to extract model from client setup message for Vertex AI Live
2468 if endpoint and "/vertex_ai/live" in endpoint:
2469 verbose_proxy_logger.debug(
2470 "WebSocket passthrough (%s): Processing client message for model extraction",
2471 endpoint,
2472 )
2473 try:
2474 client_message = json.loads(text_data)
2475 if isinstance(client_message, dict) and "setup" in client_message:
2476 setup_data = client_message["setup"]
2477 verbose_proxy_logger.debug(
2478 "WebSocket passthrough (%s): Found setup data in client message: %s",
2479 endpoint,
2480 setup_data,
2481 )
2482 if isinstance(setup_data, dict) and "model" in setup_data:
2483 extracted_model = _extract_model_from_vertex_ai_setup(
2484 _resolved_vertex_live_setup(setup_data, setup_model_rewriter)
2485 )
2486 if extracted_model:
2487 kwargs["model"] = extracted_model
2488 kwargs["custom_llm_provider"] = "vertex_ai-language-models"
2489 # Update logging object with correct model
2490 logging_obj.model = extracted_model
2491 logging_obj.model_call_details["model"] = extracted_model
2492 logging_obj.model_call_details["custom_llm_provider"] = "vertex_ai"
2493 verbose_proxy_logger.info(
2494 "WebSocket passthrough (%s): Successfully extracted model '%s' and set provider to 'vertex_ai' from client setup message",
2495 endpoint,
2496 extracted_model,
2497 )
2498 else:
2499 verbose_proxy_logger.warning(
2500 "WebSocket passthrough (%s): Failed to extract model from client setup data: %s",
2501 endpoint,
2502 setup_data,
2503 )
2504 else:
2505 verbose_proxy_logger.debug(
2506 "WebSocket passthrough (%s): Setup data does not contain model field: %s",
2507 endpoint,
2508 setup_data,
2509 )
2510 else:
2511 verbose_proxy_logger.debug(
2512 "WebSocket passthrough (%s): Client message does not contain setup data",
2513 endpoint,
2514 )
2515 except (json.JSONDecodeError, KeyError, TypeError) as e:
2516 verbose_proxy_logger.debug(
2517 "WebSocket passthrough (%s): Client message is not a valid setup message: %s",
2518 endpoint,
2519 e,
2520 )
2521 # Not a JSON message or doesn't contain setup data
2523 await upstream_ws.send(_rewrite_vertex_live_setup_model(text_data, setup_model_rewriter))
2524 elif bytes_data is not None:
2525 await upstream_ws.send(bytes_data)
2526 except asyncio.CancelledError:
2527 raise
2528 except Exception:
2529 verbose_proxy_logger.exception(
2530 "WebSocket passthrough (%s): error forwarding client message", endpoint
2531 )
2532 await upstream_ws.close()
2534 def _extract_vertex_live_model_from_setup_response(setup_response: Mapping[str, object]) -> None:
2535 extracted_model: Final = _extract_model_from_vertex_ai_setup(setup_response)
2536 if not extracted_model:
2537 verbose_proxy_logger.warning(
2538 "WebSocket passthrough (%s): Failed to extract model from server setup response: %s",
2539 endpoint,
2540 setup_response,
2541 )
2542 return
2543 kwargs["model"] = extracted_model
2544 kwargs["custom_llm_provider"] = "vertex_ai_language_models"
2545 logging_obj.model = extracted_model
2546 logging_obj.model_call_details["model"] = extracted_model
2547 logging_obj.model_call_details["custom_llm_provider"] = "vertex_ai_language_models"
2549 is_vertex_live: Final = bool(endpoint and "/vertex_ai/live" in endpoint)
2550 json_frame_ordinal: Final = count()
2552 async def relay_upstream_frame(upstream_message: str | bytes) -> None:
2553 if isinstance(upstream_message, bytes):
2554 await websocket.send_bytes(upstream_message)
2555 else:
2556 await websocket.send_text(upstream_message)
2557 message_data: Final = _json_object_frame(upstream_message)
2558 if message_data is None:
2559 return
2560 if is_vertex_live and next(json_frame_ordinal) == 0:
2561 _extract_vertex_live_model_from_setup_response(message_data)
2562 return
2563 websocket_messages.append(message_data)
2565 async def forward_upstream_to_client() -> Close | None:
2566 try:
2567 while True:
2568 await relay_upstream_frame(await upstream_ws.recv())
2569 except (ConnectionClosedOK, ConnectionClosedError) as e:
2570 verbose_proxy_logger.debug("Upstream WebSocket connection closed: %s", e)
2571 return e.rcvd
2572 except asyncio.CancelledError:
2573 verbose_proxy_logger.debug("asyncio.CancelledError in forward_upstream_to_client")
2574 raise
2575 except Exception as e:
2576 verbose_proxy_logger.debug("Exception in forward_upstream_to_client: %s", e)
2577 verbose_proxy_logger.exception(
2578 "WebSocket passthrough (%s): error forwarding upstream message", endpoint
2579 )
2580 raise
2582 # Create tasks for bidirectional message forwarding
2583 tasks: Final = [
2584 asyncio.create_task(forward_client_to_upstream()),
2585 asyncio.create_task(forward_upstream_to_client()),
2586 ]
2588 done, pending = await asyncio.wait(tasks, return_when=asyncio.FIRST_COMPLETED)
2590 # Cancel remaining tasks
2591 for task in pending:
2592 task.cancel()
2593 try:
2594 await task
2595 except asyncio.CancelledError:
2596 pass
2598 # Check for exceptions in completed tasks
2599 for task in done:
2600 exception = task.exception()
2601 if exception is not None:
2602 raise exception
2604 upstream_close: Final = _upstream_close_to_relay(task.result() for task in done)
2605 if upstream_close is not None and _client_socket_is_open(websocket):
2606 await websocket.close(
2607 code=upstream_close.code,
2608 reason=_truncated_close_reason(upstream_close.reason),
2609 )
2611 end_time: Final = datetime.now()
2613 # Update passthrough logging payload with response data
2614 passthrough_logging_payload["response_body"] = websocket_messages
2615 passthrough_logging_payload["end_time"] = end_time
2617 # Remove logging_obj from kwargs to avoid duplicate keyword argument
2618 success_kwargs: Final = kwargs.copy()
2619 success_kwargs.pop("logging_obj", None)
2621 # # Add user authentication context for database logging
2622 # if user_api_key_dict:
2623 # success_kwargs.setdefault('litellm_params', {})
2624 # success_kwargs['litellm_params'].update({
2625 # 'proxy_server_request': {
2626 # 'body': {
2627 # 'user': user_api_key_dict.user_id,
2628 # 'team_id': user_api_key_dict.team_id,
2629 # 'end_user_id': user_api_key_dict.end_user_id,
2630 # }
2631 # }
2632 # })
2633 # # Also add the user_api_key for direct access
2634 # success_kwargs['user_api_key'] = user_api_key_dict.api_key
2636 # Create a dummy httpx.Response for WebSocket connections
2637 class MockWebSocketResponse:
2638 def __init__(self, target_url: str):
2639 self.status_code = 200
2640 self.text = "WebSocket connection successful"
2641 self.headers: dict[str, str] = {}
2642 self.request = MockWebSocketRequest(target_url)
2644 class MockWebSocketRequest:
2645 def __init__(self, target_url: str):
2646 self.method = "WEBSOCKET"
2647 self.url = target_url
2649 mock_response: Final = MockWebSocketResponse(target)
2651 # Use the same success handler as HTTP passthrough endpoints
2652 GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue(
2653 async_coroutine=pass_through_endpoint_logging.pass_through_async_success_handler(
2654 httpx_response=mock_response,
2655 response_body=websocket_messages,
2656 url_route=endpoint or "",
2657 result="websocket_connection_successful",
2658 start_time=start_time,
2659 end_time=end_time,
2660 logging_obj=logging_obj,
2661 cache_hit=False,
2662 request_body={},
2663 **success_kwargs,
2664 )
2665 )
2666 bind_budget_reservation_to_callbacks(logging_obj.litellm_params)
2668 # Call the proxy logging success hook
2669 if proxy_logging_obj:
2670 await proxy_logging_obj.post_call_success_hook(
2671 data={},
2672 user_api_key_dict=user_api_key_dict,
2673 response={"status": "websocket_connection_successful"},
2674 )
2676 except InvalidStatus as exc:
2677 verbose_proxy_logger.exception("WebSocket passthrough (%s): upstream rejected WebSocket connection", endpoint)
2679 # Prepare request payload for logging
2680 request_payload = {}
2681 if kwargs:
2682 for key, value in kwargs.items():
2683 request_payload[key] = value
2684 if logging_obj is not None:
2685 request_payload["litellm_logging_obj"] = logging_obj
2687 # Log the connection failure using the same pattern as HTTP
2688 await proxy_logging_obj.post_call_failure_hook(
2689 user_api_key_dict=user_api_key_dict,
2690 original_exception=exc,
2691 request_data=request_payload,
2692 traceback_str=traceback.format_exc(
2693 limit=MAXIMUM_TRACEBACK_LINES_TO_LOG,
2694 ),
2695 )
2697 if _client_socket_is_open(websocket):
2698 await websocket.close(
2699 code=getattr(exc, "status_code", 1011),
2700 reason="Upstream connection rejected",
2701 )
2702 except Exception as e:
2703 verbose_proxy_logger.exception(
2704 "WebSocket passthrough (%s): unexpected error while proxying WebSocket", endpoint
2705 )
2707 # Prepare request payload for logging
2708 request_payload = {}
2709 if kwargs:
2710 for key, value in kwargs.items():
2711 request_payload[key] = value
2712 if logging_obj is not None:
2713 request_payload["litellm_logging_obj"] = logging_obj
2715 # Log the unexpected error using the same pattern as HTTP
2716 await proxy_logging_obj.post_call_failure_hook(
2717 user_api_key_dict=user_api_key_dict,
2718 original_exception=e,
2719 request_data=request_payload,
2720 traceback_str=traceback.format_exc(
2721 limit=MAXIMUM_TRACEBACK_LINES_TO_LOG,
2722 ),
2723 )
2725 if _client_socket_is_open(websocket):
2726 await websocket.close(code=1011, reason="WebSocket passthrough error")
2727 finally:
2728 if _client_socket_is_open(websocket):
2729 await websocket.close()
2732def _is_streaming_response(response: httpx.Response) -> bool:
2733 _content_type: Final = response.headers.get("content-type")
2734 if _content_type is not None and "text/event-stream" in _content_type: 2734 ↛ 2735line 2734 didn't jump to line 2735 because the condition on line 2734 was never true
2735 return True
2736 return False
2739def _own_streamed_managed_ids(
2740 stream: AsyncGenerator[bytes, None],
2741 managed_id_provider: str | None,
2742 request: Request,
2743 user_api_key_dict: UserAPIKeyAuth,
2744) -> AsyncGenerator[bytes, None]:
2745 from litellm.proxy.proxy_server import general_settings, prisma_client, proxy_logging_obj
2747 if (
2748 managed_id_provider is None
2749 or not general_settings.get("passthrough_managed_object_ids", False)
2750 or prisma_client is None
2751 or proxy_logging_obj.get_proxy_hook("managed_files") is None
2752 ):
2753 return stream
2754 from litellm.proxy.auth.auth_utils import get_request_route
2755 from litellm.proxy.pass_through_endpoints.managed_id_rewriter import (
2756 rewrite_streamed_response_ids,
2757 )
2759 return rewrite_streamed_response_ids(
2760 stream=stream,
2761 provider=managed_id_provider,
2762 method=request.method,
2763 route=get_request_route(request),
2764 user_api_key_dict=user_api_key_dict,
2765 prisma_client=prisma_client,
2766 )
2769def _should_buffer_passthrough_response(response: httpx.Response) -> bool:
2770 """
2771 Decide from the response headers whether the body must be read into memory.
2773 JSON bodies (including the AWS JSON protocol media types) and upstream errors
2774 stay buffered: spend logging, guardrails and managed-id rewriting inspect them,
2775 and they are small in practice. Everything else (jsonl batch results,
2776 octet-stream files, ...) is relayed to the client chunk by chunk so a large
2777 body is never resident in full (LIT-4009). A missing content-type is buffered
2778 because the body cannot be classified.
2779 """
2780 if response.status_code >= 400: 2780 ↛ 2782line 2780 didn't jump to line 2782 because the condition on line 2780 was always true
2781 return True
2782 content_type_header: Final[str] = response.headers.get("content-type", "")
2783 media_type: Final = content_type_header.split(";")[0].strip().lower()
2784 return (
2785 media_type in ("", "application/json")
2786 or media_type.endswith("+json")
2787 or media_type.startswith("application/x-amz-json")
2788 )
2791async def _relay_passthrough_response_bytes(
2792 response: httpx.Response,
2793 request_body: dict,
2794 url_route: str,
2795 start_time: datetime,
2796 logging_obj: LiteLLMLoggingObj,
2797 custom_llm_provider: str | None,
2798 success_handler_kwargs: dict,
2799) -> AsyncGenerator[bytes, None]:
2800 """
2801 Yield upstream bytes to the client without accumulating them, then fire the
2802 passthrough success handler with response_body=None (uninspected body). The
2803 finally block also runs on client disconnect (GeneratorExit) so partial
2804 downloads still produce a spend-log row, mirroring chunk_processor; a
2805 disconnect additionally logs a warning with the number of bytes relayed so
2806 partial deliveries are distinguishable from complete ones in proxy logs.
2807 """
2808 bytes_relayed = 0
2809 upstream_fully_relayed = False
2810 try:
2811 async for chunk in response.aiter_bytes():
2812 bytes_relayed += len(chunk)
2813 yield chunk
2814 upstream_fully_relayed = True
2815 finally:
2816 if not upstream_fully_relayed:
2817 verbose_proxy_logger.warning(
2818 "Passthrough stream for %s ended before upstream body was fully relayed; %s bytes were sent to the client",
2819 url_route,
2820 bytes_relayed,
2821 )
2822 await response.aclose()
2823 GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue(
2824 async_coroutine=pass_through_endpoint_logging.pass_through_async_success_handler(
2825 httpx_response=response,
2826 response_body=None,
2827 url_route=url_route,
2828 result="",
2829 start_time=start_time,
2830 end_time=datetime.now(),
2831 logging_obj=logging_obj,
2832 cache_hit=False,
2833 request_body=request_body,
2834 custom_llm_provider=custom_llm_provider,
2835 **success_handler_kwargs,
2836 )
2837 )
2838 bind_budget_reservation_to_callbacks(logging_obj.litellm_params)
2841def _extract_model_from_vertex_ai_setup(setup_response: Mapping[str, object]) -> str | None:
2842 """
2843 Extract the model name from Vertex AI Live setup response.
2845 The setup response can contain a model field in two formats:
2846 1. Direct: {"model": "projects/.../models/gemini-2.0-flash-live-preview-04-09"}
2847 2. Nested: {"setup": {"model": "projects/.../models/gemini-2.0-flash-live-preview-04-09"}}
2849 We extract just the model name: "gemini-2.0-flash-live-preview-04-09"
2850 """
2851 try:
2852 # Handle both direct model field and nested setup.model field
2853 model_path = None
2854 if isinstance(setup_response, dict):
2855 if "model" in setup_response:
2856 model_path = setup_response["model"]
2857 elif (
2858 "setup" in setup_response
2859 and isinstance(setup_response["setup"], dict)
2860 and "model" in setup_response["setup"]
2861 ):
2862 model_path = setup_response["setup"]["model"]
2864 if isinstance(model_path, str) and "/models/" in model_path:
2865 # Extract the model name after the last "/models/"
2866 model_name: Final = model_path.split("/models/")[-1]
2867 return model_name
2868 except Exception as e:
2869 verbose_proxy_logger.debug("Error extracting model from setup response: %s", e)
2870 return None
2873def _placed_ahead(routes: Sequence[BaseRoute], moving: BaseRoute, before: BaseRoute) -> tuple[BaseRoute, ...]:
2874 kept: Final = tuple(route for route in routes if route is not moving)
2875 at: Final = next(index for index, route in enumerate(kept) if route is before)
2876 return (*kept[:at], moving, *kept[at:])
2879class SafeRouteAdder:
2880 """
2881 Wrapper class for adding routes to FastAPI app.
2882 Only adds routes if they don't already exist on the app. A route a lazy feature registered
2883 does not count: a route added at its path goes ahead of it, the precedence a config
2884 pass-through at /v1/decisions gets in lazy mode, where the feature has not loaded yet.
2885 """
2887 @staticmethod
2888 def _colliding_routes(app: FastAPI, path: str, methods: Sequence[str]) -> tuple[Route, ...]:
2889 wanted: Final = frozenset(methods)
2890 return tuple(
2891 route
2892 for route in app.routes
2893 if isinstance(route, Route) and route.path == path and not wanted.isdisjoint(route.methods or ())
2894 )
2896 @staticmethod
2897 def _is_path_registered(app: FastAPI, path: str, methods: list[str]) -> bool:
2898 """True when a route the app itself defines already serves the path with one of the methods."""
2899 lazy_owned: Final = lazy_owned_routes(app)
2900 return any(id(route) not in lazy_owned for route in SafeRouteAdder._colliding_routes(app, path, methods))
2902 @staticmethod
2903 def add_api_route_if_not_exists(
2904 app: FastAPI,
2905 path: str,
2906 endpoint: Callable[..., object],
2907 methods: list[str],
2908 dependencies: list | None = None,
2909 ) -> bool:
2910 """
2911 Add an API route to the app only if it doesn't already exist.
2913 Args:
2914 app: The FastAPI application instance
2915 path: The path for the route
2916 endpoint: The endpoint function/callable
2917 methods: List of HTTP methods
2918 dependencies: Optional list of dependencies
2920 Returns:
2921 True if route was added, False if it already existed
2922 """
2923 if SafeRouteAdder._is_path_registered(app=app, path=path, methods=methods):
2924 verbose_proxy_logger.debug(
2925 "Skipping route registration - path %s with methods %s already registered on app",
2926 path,
2927 methods,
2928 )
2929 return False
2931 shadowed: Final = SafeRouteAdder._colliding_routes(app, path, methods)
2932 app.add_api_route(
2933 path=path,
2934 endpoint=endpoint,
2935 methods=methods,
2936 dependencies=dependencies,
2937 )
2938 if shadowed: 2938 ↛ 2939line 2938 didn't jump to line 2939 because the condition on line 2938 was never true
2939 app.router.routes[:] = _placed_ahead( # rebind-ok: the app owns its route table
2940 app.router.routes, app.router.routes[-1], shadowed[0]
2941 )
2942 verbose_proxy_logger.debug(
2943 "Successfully added route: %s with methods %s",
2944 path,
2945 methods,
2946 )
2947 return True
2950class InitPassThroughEndpointHelpers:
2951 @staticmethod
2952 def add_exact_path_route(
2953 app: FastAPI,
2954 path: str,
2955 target: str,
2956 custom_headers: dict | None,
2957 forward_headers: bool | None,
2958 merge_query_params: bool | None,
2959 dependencies: list | None,
2960 cost_per_request: float | None,
2961 endpoint_id: str,
2962 guardrails: dict | None = None,
2963 methods: list[str] | None = None,
2964 default_query_params: dict | None = None,
2965 config_file_path: str | None = None,
2966 auth: bool = False,
2967 timeout: float | None = None,
2968 ):
2969 """Add exact path route for pass-through endpoint"""
2970 # Default to all methods if none specified (backward compatibility)
2971 if methods is None or len(methods) == 0:
2972 methods = ["GET", "POST", "PUT", "DELETE", "PATCH"]
2974 # Create route key that includes methods for uniqueness
2975 methods_str: Final = ",".join(sorted(methods))
2976 route_key: Final = f"{endpoint_id}:exact:{path}:{methods_str}"
2978 # Check if this exact route is already registered
2979 if route_key in _registered_pass_through_routes:
2980 verbose_proxy_logger.debug(
2981 "Updating duplicate exact pass through endpoint: %s with methods %s (already registered)",
2982 path,
2983 methods,
2984 )
2986 verbose_proxy_logger.debug(
2987 "adding exact pass through endpoint: %s, methods: %s, dependencies: %s",
2988 path,
2989 methods,
2990 dependencies,
2991 )
2993 # Use SafeRouteAdder to only add route if it doesn't exist on the app
2994 SafeRouteAdder.add_api_route_if_not_exists(
2995 app=app,
2996 path=path,
2997 endpoint=create_pass_through_route(
2998 path,
2999 target,
3000 custom_headers,
3001 forward_headers,
3002 merge_query_params,
3003 dependencies,
3004 cost_per_request=cost_per_request,
3005 default_query_params=default_query_params,
3006 guardrails=guardrails,
3007 config_file_path=config_file_path,
3008 timeout=timeout,
3009 ),
3010 methods=methods,
3011 dependencies=dependencies,
3012 )
3014 # Always register/update the route metadata (headers, target) even if FastAPI route exists
3015 _registered_pass_through_routes[route_key] = {
3016 "endpoint_id": endpoint_id,
3017 "path": path,
3018 "type": "exact",
3019 "methods": methods,
3020 "auth": auth,
3021 "passthrough_params": {
3022 "target": target,
3023 "custom_headers": custom_headers,
3024 "forward_headers": forward_headers,
3025 "merge_query_params": merge_query_params,
3026 "default_query_params": default_query_params,
3027 "dependencies": dependencies,
3028 "cost_per_request": cost_per_request,
3029 "guardrails": guardrails,
3030 "timeout": timeout,
3031 },
3032 }
3034 @staticmethod
3035 def add_subpath_route(
3036 app: FastAPI,
3037 path: str,
3038 target: str,
3039 custom_headers: dict | None,
3040 forward_headers: bool | None,
3041 merge_query_params: bool | None,
3042 dependencies: list | None,
3043 cost_per_request: float | None,
3044 endpoint_id: str,
3045 guardrails: dict | None = None,
3046 methods: list[str] | None = None,
3047 default_query_params: dict | None = None,
3048 config_file_path: str | None = None,
3049 auth: bool = False,
3050 timeout: float | None = None,
3051 ):
3052 """Add wildcard route for sub-paths"""
3053 # Default to all methods if none specified (backward compatibility)
3054 if methods is None or len(methods) == 0:
3055 methods = ["GET", "POST", "PUT", "DELETE", "PATCH"]
3057 wildcard_path: Final = f"{path}/{{subpath:path}}"
3058 methods_str: Final = ",".join(sorted(methods))
3059 route_key: Final = f"{endpoint_id}:subpath:{path}:{methods_str}"
3061 # Check if this subpath route is already registered
3062 if route_key in _registered_pass_through_routes:
3063 verbose_proxy_logger.debug(
3064 "Updating duplicate wildcard pass through endpoint: %s with methods %s (already registered)",
3065 wildcard_path,
3066 methods,
3067 )
3069 verbose_proxy_logger.debug(
3070 "adding wildcard pass through endpoint: %s, methods: %s, dependencies: %s",
3071 wildcard_path,
3072 methods,
3073 dependencies,
3074 )
3076 # Use SafeRouteAdder to only add route if it doesn't exist on the app
3077 SafeRouteAdder.add_api_route_if_not_exists(
3078 app=app,
3079 path=wildcard_path,
3080 endpoint=create_pass_through_route(
3081 path,
3082 target,
3083 custom_headers,
3084 forward_headers,
3085 merge_query_params,
3086 dependencies,
3087 include_subpath=True,
3088 cost_per_request=cost_per_request,
3089 default_query_params=default_query_params,
3090 guardrails=guardrails,
3091 config_file_path=config_file_path,
3092 timeout=timeout,
3093 ),
3094 methods=methods,
3095 dependencies=dependencies,
3096 )
3098 # Register the route to prevent duplicates only if it was added
3099 _registered_pass_through_routes[route_key] = {
3100 "endpoint_id": endpoint_id,
3101 "path": path,
3102 "type": "subpath",
3103 "methods": methods,
3104 "auth": auth,
3105 "passthrough_params": {
3106 "target": target,
3107 "custom_headers": custom_headers,
3108 "forward_headers": forward_headers,
3109 "merge_query_params": merge_query_params,
3110 "default_query_params": default_query_params,
3111 "dependencies": dependencies,
3112 "cost_per_request": cost_per_request,
3113 "guardrails": guardrails,
3114 "timeout": timeout,
3115 },
3116 }
3118 @staticmethod
3119 def remove_endpoint_routes(endpoint_id: str):
3120 """Remove all routes for a specific endpoint ID from the registry
3121 and clean up corresponding entries from LiteLLMRoutes.openai_routes."""
3122 keys_to_remove: Final = [
3123 key for key, value in _registered_pass_through_routes.items() if value["endpoint_id"] == endpoint_id
3124 ]
3125 for key in keys_to_remove:
3126 route_info = _registered_pass_through_routes[key]
3127 path = route_info.get("path")
3128 if isinstance(path, str): 3128 ↛ 3136line 3128 didn't jump to line 3136 because the condition on line 3128 was always true
3129 openai_routes = LiteLLMRoutes.openai_routes.value
3130 if path in openai_routes:
3131 openai_routes.remove(path)
3132 if route_info.get("type") == "subpath":
3133 wildcard_path = path.rstrip("/") + "/*"
3134 if wildcard_path in openai_routes:
3135 openai_routes.remove(wildcard_path)
3136 del _registered_pass_through_routes[key]
3137 verbose_proxy_logger.debug("Removed pass-through route from registry: %s", key)
3139 @staticmethod
3140 def clear_all_pass_through_routes():
3141 """Clear all pass-through routes from the registry"""
3142 _registered_pass_through_routes.clear()
3144 @staticmethod
3145 def get_all_registered_pass_through_routes() -> list[str]:
3146 """Get all registered pass-through endpoints from the registry"""
3147 return list(_registered_pass_through_routes.keys())
3149 @staticmethod
3150 def _route_for_registry_lookup(route: str) -> str:
3151 """
3152 Normalize an incoming route to the bare path stored in the registry.
3154 Registry keys store root-stripped paths. Callers should pass routes from
3155 ``get_request_route()`` (already stripped); prefixed ``request.url.path``
3156 values are stripped via ``normalize_route_for_root_path``.
3157 """
3158 normalized_route: Final = normalize_route_for_root_path(route)
3159 return normalized_route if normalized_route is not None else route
3161 @staticmethod
3162 def is_registered_pass_through_route(route: str) -> bool:
3163 """
3164 Check if route is a registered pass-through endpoint from DB
3166 Uses the in-memory registry to avoid additional DB queries
3167 Optimized for minimal latency
3169 Args:
3170 route: The route to check
3172 Returns:
3173 bool: True if route is a registered pass-through endpoint, False otherwise
3174 """
3175 ## CHECK IF MAPPED PASS THROUGH ENDPOINT
3176 normalized_route: Final = normalize_route_for_root_path(route)
3177 if normalized_route is not None: 3177 ↛ 3182line 3177 didn't jump to line 3182 because the condition on line 3177 was always true
3178 for mapped_route in LiteLLMRoutes.mapped_pass_through_routes.value:
3179 if normalized_route.startswith(mapped_route):
3180 return True
3182 comparison_route: Final = InitPassThroughEndpointHelpers._route_for_registry_lookup(route)
3184 # Fast path: check if any registered route key contains this path
3185 # Keys are in format: "{endpoint_id}:exact:{path}:{methods}" or "{endpoint_id}:subpath:{path}:{methods}"
3186 # For backward compatibility, also support old format: "{endpoint_id}:exact:{path}" or "{endpoint_id}:subpath:{path}"
3187 # Extract unique paths from keys for quick checking
3188 for key in _registered_pass_through_routes: 3188 ↛ 3199line 3188 didn't jump to line 3199 because the loop on line 3188 didn't complete
3189 parts = key.split(":", 3) # Split into [endpoint_id, type, path, methods?]
3190 if len(parts) >= 3: 3190 ↛ 3188line 3190 didn't jump to line 3188 because the condition on line 3190 was always true
3191 route_type = parts[1]
3192 registered_path = parts[2]
3193 if route_type == "exact" and comparison_route == registered_path: 3193 ↛ 3194line 3193 didn't jump to line 3194 because the condition on line 3193 was never true
3194 return True
3195 elif route_type == "subpath":
3196 if comparison_route == registered_path or comparison_route.startswith(registered_path + "/"):
3197 return True
3199 return False
3201 @staticmethod
3202 def get_registered_pass_through_route(route: str, method: str | None = None) -> dict[str, Any] | None:
3203 """Get passthrough params for a given route and optionally filter by HTTP method"""
3204 comparison_route: Final = InitPassThroughEndpointHelpers._route_for_registry_lookup(route)
3205 for key in _registered_pass_through_routes:
3206 parts = key.split(":", 3) # Split into [endpoint_id, type, path, methods?]
3207 if len(parts) >= 3: 3207 ↛ 3205line 3207 didn't jump to line 3205 because the condition on line 3207 was always true
3208 route_type = parts[1]
3209 registered_path = parts[2]
3211 # Get the methods for this route. Prefer the registered metadata,
3212 # but keep supporting test fixtures / older registry entries that
3213 # only encoded methods in the route key.
3214 methods_entry = _registered_pass_through_routes[key].get("methods", [])
3215 route_methods: list[str] = methods_entry if isinstance(methods_entry, list) else []
3216 if not route_methods and len(parts) == 4: 3216 ↛ 3217line 3216 didn't jump to line 3217 because the condition on line 3216 was never true
3217 route_methods = parts[3].split(",")
3219 # Check if path matches
3220 path_matches = False
3221 if route_type == "exact" and comparison_route == registered_path: 3221 ↛ 3222line 3221 didn't jump to line 3222 because the condition on line 3221 was never true
3222 path_matches = True
3223 elif route_type == "subpath":
3224 if comparison_route == registered_path or comparison_route.startswith(registered_path + "/"):
3225 path_matches = True
3227 # If path matches and method filter is provided, check if method is allowed
3228 if path_matches:
3229 if method is None or not route_methods or method in route_methods: 3229 ↛ 3205line 3229 didn't jump to line 3205 because the condition on line 3229 was always true
3230 return _registered_pass_through_routes[key]
3232 return None
3235def _get_combined_pass_through_endpoints(
3236 pass_through_endpoints: list[dict] | list[PassThroughGenericEndpoint],
3237 config_pass_through_endpoints: list[dict],
3238):
3239 """Get combined pass-through endpoints from db + config"""
3240 return pass_through_endpoints + config_pass_through_endpoints
3243async def _register_pass_through_endpoint(
3244 endpoint: dict[str, object] | PassThroughGenericEndpoint,
3245 app: FastAPI,
3246 premium_user: bool,
3247 visited_endpoints: set[str],
3248 config_file_path: str | None = None,
3249) -> None:
3250 endpoint_data: dict[str, Any]
3251 if isinstance(endpoint, PassThroughGenericEndpoint):
3252 endpoint_data = endpoint.model_dump()
3253 else:
3254 endpoint_data = endpoint
3256 if endpoint_data.get("id") is None: 3256 ↛ 3257line 3256 didn't jump to line 3257 because the condition on line 3256 was never true
3257 endpoint_data["id"] = str(uuid.uuid4())
3258 endpoint_id: Final = cast(str, endpoint_data["id"])
3260 target: Final[str | None] = endpoint_data.get("target")
3261 path: Final[str | None] = endpoint_data.get("path")
3262 if path is None: 3262 ↛ 3263line 3262 didn't jump to line 3263 because the condition on line 3262 was never true
3263 raise ValueError("Path is required for pass-through endpoint")
3265 custom_headers: Final = await set_env_variables_in_header(custom_headers=endpoint_data.get("headers"))
3266 forward_headers: Final = endpoint_data.get("forward_headers")
3267 merge_query_params: Final = endpoint_data.get("merge_query_params")
3268 default_query_params: Final = endpoint_data.get("default_query_params")
3269 auth: Final[bool | str | None] = endpoint_data.get("auth")
3270 dependencies = None
3271 auth_enforced: Final = auth is not None and str(auth).lower() == "true"
3273 if auth_enforced:
3274 # Authentication on a pass-through endpoint used to be enterprise-only.
3275 # That left OSS with no safe configuration: auth=True raised at startup
3276 # unless the operator had a license. The safe option must always be free,
3277 # and unauthenticated forwarding should require explicit opt-in.
3278 dependencies = [Depends(user_api_key_auth)]
3279 if path not in LiteLLMRoutes.openai_routes.value:
3280 LiteLLMRoutes.openai_routes.value.append(path)
3282 if target is None: 3282 ↛ 3283line 3282 didn't jump to line 3283 because the condition on line 3282 was never true
3283 return
3285 guardrails: Final = endpoint_data.get("guardrails")
3286 methods: Final = endpoint_data.get("methods")
3287 cost_per_request: Final = endpoint_data.get("cost_per_request")
3288 timeout: Final = endpoint_data.get("timeout")
3290 verbose_proxy_logger.debug("Initializing pass through endpoint: %s (ID: %s)", path, endpoint_id)
3291 InitPassThroughEndpointHelpers.add_exact_path_route(
3292 app=app,
3293 path=path,
3294 target=target,
3295 custom_headers=custom_headers,
3296 forward_headers=forward_headers,
3297 merge_query_params=merge_query_params,
3298 dependencies=dependencies,
3299 cost_per_request=cost_per_request,
3300 endpoint_id=endpoint_id,
3301 guardrails=guardrails,
3302 methods=methods,
3303 default_query_params=default_query_params,
3304 config_file_path=config_file_path,
3305 auth=auth_enforced,
3306 timeout=timeout,
3307 )
3309 methods_for_key: Final = methods if methods else ["GET", "POST", "PUT", "DELETE", "PATCH"]
3310 methods_str: Final = ",".join(sorted(methods_for_key))
3311 visited_endpoints.add(f"{endpoint_id}:exact:{path}:{methods_str}")
3313 if endpoint_data.get("include_subpath", False) is True:
3314 if auth is not None and str(auth).lower() == "true":
3315 wildcard_path: Final = path.rstrip("/") + "/*"
3316 if wildcard_path not in LiteLLMRoutes.openai_routes.value:
3317 LiteLLMRoutes.openai_routes.value.append(wildcard_path)
3318 InitPassThroughEndpointHelpers.add_subpath_route(
3319 app=app,
3320 path=path,
3321 target=target,
3322 custom_headers=custom_headers,
3323 forward_headers=forward_headers,
3324 merge_query_params=merge_query_params,
3325 dependencies=dependencies,
3326 cost_per_request=cost_per_request,
3327 endpoint_id=endpoint_id,
3328 guardrails=guardrails,
3329 methods=methods,
3330 default_query_params=default_query_params,
3331 config_file_path=config_file_path,
3332 auth=auth_enforced,
3333 timeout=timeout,
3334 )
3335 visited_endpoints.add(f"{endpoint_id}:subpath:{path}:{methods_str}")
3337 verbose_proxy_logger.debug("Added new pass through endpoint: %s (ID: %s)", path, endpoint_id)
3340async def initialize_pass_through_endpoints(
3341 pass_through_endpoints: list[dict] | list[PassThroughGenericEndpoint],
3342 config_file_path: str | None = None,
3343):
3344 """
3345 1. Create a global list of pass-through endpoints (db + config)
3346 2. Clear all existing pass-through endpoints from the FastAPI app routes
3347 3. Add new endpoints to the in-memory registry
3349 Initialize a list of pass-through endpoints by adding them to the FastAPI app routes
3351 Args:
3352 pass_through_endpoints: List of pass-through endpoints to initialize
3353 config_file_path: Path to the operator's config.yaml when this call
3354 originates from a YAML-load. Threaded through to
3355 ``create_pass_through_route`` so an operator using
3356 ``s3://``/``gcs://`` ``custom_handler`` in their config still
3357 loads. Callers from the DB-overlay / runtime API path must leave
3358 this ``None`` so the runtime gate in ``get_instance_fn`` fires.
3360 Returns:
3361 None
3362 """
3363 verbose_proxy_logger.debug("initializing pass through endpoints")
3364 from litellm.proxy.proxy_server import (
3365 app,
3366 config_passthrough_endpoints,
3367 premium_user,
3368 )
3370 ## get combined pass-through endpoints from db + config
3371 combined_pass_through_endpoints: list[dict | PassThroughGenericEndpoint]
3373 if config_passthrough_endpoints is not None: 3373 ↛ 3374line 3373 didn't jump to line 3374 because the condition on line 3373 was never true
3374 combined_pass_through_endpoints = _get_combined_pass_through_endpoints(
3375 pass_through_endpoints, config_passthrough_endpoints
3376 )
3377 else:
3378 combined_pass_through_endpoints = pass_through_endpoints
3380 ## clear all existing pass-through endpoints from the FastAPI app routes
3381 # InitPassThroughEndpointHelpers.clear_all_pass_through_routes()
3383 # get a list of all registered pass-through endpoints
3384 # mark the ones that are visited in the list
3385 # remove the ones that are not visited from the list
3386 registered_pass_through_endpoints: Final = InitPassThroughEndpointHelpers.get_all_registered_pass_through_routes()
3388 visited_endpoints: Final[set[str]] = set()
3390 for endpoint in combined_pass_through_endpoints:
3391 await _register_pass_through_endpoint(
3392 endpoint=endpoint,
3393 app=app,
3394 premium_user=premium_user,
3395 visited_endpoints=visited_endpoints,
3396 config_file_path=config_file_path,
3397 )
3399 # Drop stale registry entries by their exact route key. registered_pass_through_endpoints
3400 # holds route keys ("{id}:{type}:{path}:{methods}"), not endpoint ids, so remove_endpoint_routes
3401 # (which matches on endpoint_id) never matched and left the registry growing every reload cycle.
3402 # We pop the key directly and leave openai_routes alone: its append is path-deduped, and the path
3403 # is still owned by the live endpoint that was just re-registered under a new id this same cycle.
3404 for endpoint_key in registered_pass_through_endpoints:
3405 if endpoint_key not in visited_endpoints: 3405 ↛ 3406line 3405 didn't jump to line 3406 because the condition on line 3405 was never true
3406 _registered_pass_through_routes.pop(endpoint_key, None)
3409def _get_pass_through_endpoints_from_config() -> list[PassThroughGenericEndpoint]:
3410 """
3411 Get pass-through endpoints defined in the config file.
3412 These are read-only and cannot be edited via the UI.
3413 Malformed endpoints are logged and skipped; they do not crash the function.
3414 """
3415 from pydantic import ValidationError
3417 from litellm.proxy.proxy_server import config_passthrough_endpoints
3419 if config_passthrough_endpoints is None or len(config_passthrough_endpoints) == 0: 3419 ↛ 3422line 3419 didn't jump to line 3422 because the condition on line 3419 was always true
3420 return []
3422 returned_endpoints: Final[list[PassThroughGenericEndpoint]] = []
3423 for endpoint in config_passthrough_endpoints:
3424 try:
3425 if isinstance(endpoint, dict):
3426 endpoint_dict = dict(endpoint)
3427 endpoint_dict["is_from_config"] = True
3428 returned_endpoints.append(PassThroughGenericEndpoint.model_validate(endpoint_dict))
3429 elif isinstance(endpoint, PassThroughGenericEndpoint):
3430 # Create a copy with is_from_config=True
3431 endpoint_dict = endpoint.model_dump()
3432 endpoint_dict["is_from_config"] = True
3433 returned_endpoints.append(PassThroughGenericEndpoint.model_validate(endpoint_dict))
3434 except ValidationError as e:
3435 verbose_proxy_logger.warning(
3436 "Skipping malformed pass-through endpoint from config: %s",
3437 e,
3438 exc_info=False,
3439 )
3441 return returned_endpoints
3444def _config_field_endpoints(response: ConfigFieldInfo) -> list[object] | None:
3445 return response.field_value
3448def _request_app(request: Request) -> FastAPI:
3449 return request.app
3452async def _get_pass_through_endpoints_from_db(
3453 endpoint_id: str | None = None,
3454 user_api_key_dict: UserAPIKeyAuth | None = None,
3455) -> list[PassThroughGenericEndpoint]:
3456 from litellm.proxy._types import LitellmUserRoles
3457 from litellm.proxy.proxy_server import get_config_general_settings
3459 try:
3460 if user_api_key_dict is None:
3461 user_api_key_dict = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN)
3462 response: Final[ConfigFieldInfo] = await get_config_general_settings(
3463 field_name="pass_through_endpoints", user_api_key_dict=user_api_key_dict
3464 )
3465 except Exception:
3466 return []
3468 pass_through_endpoint_data: Final = _config_field_endpoints(response)
3469 if pass_through_endpoint_data is None: 3469 ↛ 3470line 3469 didn't jump to line 3470 because the condition on line 3469 was never true
3470 return []
3472 returned_endpoints: Final[list[PassThroughGenericEndpoint]] = []
3473 if endpoint_id is None:
3474 # Return all endpoints from DB, mark as not from config
3475 for endpoint in pass_through_endpoint_data:
3476 if isinstance(endpoint, dict): 3476 ↛ 3480line 3476 didn't jump to line 3480 because the condition on line 3476 was always true
3477 endpoint_dict = dict(endpoint)
3478 endpoint_dict["is_from_config"] = False
3479 returned_endpoints.append(PassThroughGenericEndpoint.model_validate(endpoint_dict))
3480 elif isinstance(endpoint, PassThroughGenericEndpoint):
3481 endpoint_dict = endpoint.model_dump()
3482 endpoint_dict["is_from_config"] = False
3483 returned_endpoints.append(PassThroughGenericEndpoint.model_validate(endpoint_dict))
3484 else:
3485 # Find specific endpoint by ID
3486 found_endpoint: Final = _find_endpoint_by_id(pass_through_endpoint_data, endpoint_id)
3487 if found_endpoint is not None:
3488 endpoint_dict = (
3489 found_endpoint.model_dump()
3490 if isinstance(found_endpoint, PassThroughGenericEndpoint)
3491 else dict(found_endpoint)
3492 )
3493 endpoint_dict["is_from_config"] = False
3494 returned_endpoints.append(PassThroughGenericEndpoint.model_validate(endpoint_dict))
3496 return returned_endpoints
3499async def _filter_endpoints_by_team_allowed_routes(
3500 team_id: str,
3501 pass_through_endpoints: list[PassThroughGenericEndpoint],
3502 prisma_client,
3503) -> list[PassThroughGenericEndpoint]:
3504 """
3505 Filter pass-through endpoints based on team's allowed_passthrough_routes metadata.
3507 Args:
3508 team_id: The team ID to check permissions for
3509 pass_through_endpoints: List of endpoints to filter
3510 prisma_client: Database client
3512 Returns:
3513 Filtered list of endpoints based on team permissions
3515 Raises:
3516 HTTPException: If team is not found
3517 """
3518 # retrieve team from db
3519 team: Final = await TeamRepository(prisma_client).table.find_unique(
3520 where={"team_id": team_id},
3521 )
3522 if team is None:
3523 raise HTTPException(
3524 status_code=404,
3525 detail={"error": "Team not found"},
3526 )
3528 # retrieve team metadata
3529 team_metadata: Final = cast( # cast-ok: prisma types the Json column as str; reads hand back the decoded value
3530 "Mapping[str, object] | None", team.metadata
3531 )
3532 if team_metadata is not None and team_metadata.get("allowed_passthrough_routes") is not None: 3532 ↛ 3534line 3532 didn't jump to line 3534 because the condition on line 3532 was never true
3533 ## FILTER pass_through_endpoints by allowed_passthrough_routes
3534 pass_through_endpoints = [
3535 endpoint
3536 for endpoint in pass_through_endpoints
3537 if endpoint.path
3538 in cast( # cast-ok: guarded above; team metadata stores this key as a list of route paths
3539 Sequence[str], team_metadata.get("allowed_passthrough_routes")
3540 )
3541 ]
3543 return pass_through_endpoints
3546@router.get(
3547 "/config/pass_through_endpoint",
3548 dependencies=[Depends(user_api_key_auth)],
3549 response_model=PassThroughEndpointResponse,
3550)
3551@router.get(
3552 "/config/pass_through_endpoint/team/{team_id}",
3553 dependencies=[Depends(user_api_key_auth)],
3554 response_model=PassThroughEndpointResponse,
3555)
3556async def get_pass_through_endpoints(
3557 endpoint_id: str | None = None,
3558 user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
3559 team_id: str | None = None,
3560):
3561 """
3562 GET configured pass through endpoint.
3564 If no endpoint_id given, return all configured endpoints.
3565 """ ## Get existing pass-through endpoint field value
3566 from litellm.proxy._types import CommonProxyErrors
3567 from litellm.proxy.proxy_server import prisma_client
3569 if prisma_client is None: 3569 ↛ 3570line 3569 didn't jump to line 3570 because the condition on line 3569 was never true
3570 raise HTTPException(
3571 status_code=500,
3572 detail={"error": CommonProxyErrors.db_not_connected_error.value},
3573 )
3575 # Get endpoints from DB (editable via UI)
3576 db_endpoints: Final = await _get_pass_through_endpoints_from_db(
3577 endpoint_id=endpoint_id, user_api_key_dict=user_api_key_dict
3578 )
3580 # Get endpoints from config file (read-only, not editable via UI)
3581 config_endpoints: Final = _get_pass_through_endpoints_from_config()
3583 # Merge: config endpoints not in DB + all DB endpoints (DB overrides config for same path)
3584 db_paths: Final = {ep.path for ep in db_endpoints}
3585 config_only_endpoints: Final = [ep for ep in config_endpoints if ep.path not in db_paths]
3586 if endpoint_id is not None:
3587 # When filtering by endpoint_id, only return if found in DB (config endpoints use generated IDs)
3588 pass_through_endpoints = db_endpoints
3589 else:
3590 pass_through_endpoints = config_only_endpoints + db_endpoints
3592 if team_id is not None:
3593 pass_through_endpoints = await _filter_endpoints_by_team_allowed_routes(
3594 team_id=team_id,
3595 pass_through_endpoints=pass_through_endpoints,
3596 prisma_client=prisma_client,
3597 )
3599 return PassThroughEndpointResponse(endpoints=pass_through_endpoints)
3602@router.post(
3603 "/config/pass_through_endpoint/{endpoint_id}",
3604 dependencies=[Depends(user_api_key_auth)],
3605)
3606async def update_pass_through_endpoints(
3607 endpoint_id: str,
3608 data: PassThroughGenericEndpoint,
3609 request: Request,
3610 user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
3611):
3612 """
3613 Update a pass-through endpoint by ID.
3614 """
3615 from litellm.proxy.proxy_server import (
3616 get_config_general_settings,
3617 update_config_general_settings,
3618 )
3620 ## Get existing pass-through endpoint field value
3621 try:
3622 response: Final[ConfigFieldInfo] = await get_config_general_settings(
3623 field_name="pass_through_endpoints", user_api_key_dict=user_api_key_dict
3624 )
3625 except Exception:
3626 raise HTTPException(
3627 status_code=404,
3628 detail={"error": "No pass-through endpoints found"},
3629 )
3631 pass_through_endpoint_data: Final[list | None] = _config_field_endpoints(response)
3632 if pass_through_endpoint_data is None: 3632 ↛ 3633line 3632 didn't jump to line 3633 because the condition on line 3632 was never true
3633 raise HTTPException(
3634 status_code=404,
3635 detail={"error": "No pass-through endpoints found"},
3636 )
3638 # Find the endpoint to update
3639 found_endpoint: Final = _find_endpoint_by_id(pass_through_endpoint_data, endpoint_id)
3641 if found_endpoint is None:
3642 raise HTTPException(
3643 status_code=404,
3644 detail={"error": f"Endpoint with ID '{endpoint_id}' not found"},
3645 )
3647 # Find the index for updating the list
3648 endpoint_index = None
3649 for idx, endpoint in enumerate(pass_through_endpoint_data): 3649 ↛ 3655line 3649 didn't jump to line 3655 because the loop on line 3649 didn't complete
3650 _endpoint = PassThroughGenericEndpoint(**endpoint) if isinstance(endpoint, dict) else endpoint
3651 if _endpoint.id == endpoint_id:
3652 endpoint_index = idx
3653 break
3655 if endpoint_index is None: 3655 ↛ 3656line 3655 didn't jump to line 3656 because the condition on line 3655 was never true
3656 raise HTTPException(
3657 status_code=404,
3658 detail={"error": f"Could not find index for endpoint with ID '{endpoint_id}'"},
3659 )
3661 # Only merge fields the caller explicitly sent so omitted fields keep their
3662 # stored value. Without exclude_unset, defaults like auth=True would overwrite
3663 # an existing auth=false entry on any unrelated edit.
3664 # Exclude is_from_config as it's a response-only field (computed at read time)
3665 update_data: Final = data.model_dump(exclude_unset=True, exclude_none=True, exclude={"is_from_config"})
3667 # Start with existing endpoint data
3668 endpoint_dict: Final = found_endpoint.model_dump()
3670 # Update with new data (only explicitly provided values)
3671 endpoint_dict.update(update_data)
3673 # Preserve existing ID if not provided in update and endpoint has ID
3674 if "id" not in update_data and found_endpoint.id is not None:
3675 endpoint_dict["id"] = found_endpoint.id
3677 # Remove is_from_config before saving - it's a response-only field (computed at read time)
3678 endpoint_dict.pop("is_from_config", None)
3680 # Create updated endpoint object
3681 updated_endpoint: Final = PassThroughGenericEndpoint.model_validate(endpoint_dict)
3683 # Update the list
3684 pass_through_endpoint_data[endpoint_index] = endpoint_dict
3686 # Remove old routes from registry before they get re-registered
3687 InitPassThroughEndpointHelpers.remove_endpoint_routes(endpoint_id)
3689 ## Update db
3690 updated_data: Final = ConfigFieldUpdate(
3691 field_name="pass_through_endpoints",
3692 field_value=pass_through_endpoint_data,
3693 config_type="general_settings",
3694 )
3696 await update_config_general_settings(data=updated_data, user_api_key_dict=user_api_key_dict)
3698 # Re-register the route with updated headers
3699 _custom_headers: dict | None = updated_endpoint.headers or {}
3700 _custom_headers = await set_env_variables_in_header(custom_headers=_custom_headers)
3702 route_app: Final = _request_app(request)
3703 if updated_endpoint.include_subpath:
3704 InitPassThroughEndpointHelpers.add_subpath_route(
3705 app=route_app,
3706 path=updated_endpoint.path,
3707 target=updated_endpoint.target,
3708 custom_headers=_custom_headers,
3709 forward_headers=None, # Defaults not available in model? assuming None logic handles it
3710 merge_query_params=None,
3711 dependencies=None,
3712 cost_per_request=updated_endpoint.cost_per_request,
3713 endpoint_id=updated_endpoint.id or endpoint_id or "",
3714 guardrails=getattr(updated_endpoint, "guardrails", None),
3715 methods=updated_endpoint.methods,
3716 default_query_params=updated_endpoint.default_query_params,
3717 auth=updated_endpoint.auth,
3718 timeout=updated_endpoint.timeout,
3719 )
3720 else:
3721 InitPassThroughEndpointHelpers.add_exact_path_route(
3722 app=route_app,
3723 path=updated_endpoint.path,
3724 target=updated_endpoint.target,
3725 custom_headers=_custom_headers,
3726 forward_headers=None,
3727 merge_query_params=None,
3728 dependencies=None,
3729 cost_per_request=updated_endpoint.cost_per_request,
3730 endpoint_id=updated_endpoint.id or endpoint_id or "",
3731 guardrails=getattr(updated_endpoint, "guardrails", None),
3732 methods=updated_endpoint.methods,
3733 default_query_params=updated_endpoint.default_query_params,
3734 auth=updated_endpoint.auth,
3735 timeout=updated_endpoint.timeout,
3736 )
3738 return PassThroughEndpointResponse(endpoints=[updated_endpoint] if updated_endpoint else [])
3741@router.post(
3742 "/config/pass_through_endpoint",
3743 dependencies=[Depends(user_api_key_auth)],
3744)
3745async def create_pass_through_endpoints(
3746 data: PassThroughGenericEndpoint,
3747 request: Request,
3748 user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
3749):
3750 """
3751 Create new pass-through endpoint
3752 """
3753 from litellm._uuid import uuid
3754 from litellm.proxy.proxy_server import (
3755 get_config_general_settings,
3756 update_config_general_settings,
3757 )
3759 ## Get existing pass-through endpoint field value
3761 try:
3762 response: ConfigFieldInfo = await get_config_general_settings(
3763 field_name="pass_through_endpoints", user_api_key_dict=user_api_key_dict
3764 )
3765 except Exception:
3766 response = ConfigFieldInfo(field_name="pass_through_endpoints", field_value=None)
3768 ## Auto-generate ID if not provided
3769 # Exclude is_from_config as it's a response-only field (computed at read time)
3770 data_dict: Final = data.model_dump(exclude={"is_from_config"})
3771 if data_dict.get("id") is None:
3772 data_dict["id"] = str(uuid.uuid4())
3774 if response.field_value is None:
3775 response.field_value = [data_dict]
3776 elif isinstance(response.field_value, list): 3776 ↛ 3780line 3776 didn't jump to line 3780 because the condition on line 3776 was always true
3777 response.field_value.append(data_dict)
3779 ## Update db
3780 updated_data: Final = ConfigFieldUpdate(
3781 field_name="pass_through_endpoints",
3782 field_value=response.field_value,
3783 config_type="general_settings",
3784 )
3785 await update_config_general_settings(data=updated_data, user_api_key_dict=user_api_key_dict)
3787 # Return the created endpoint with the generated ID
3788 created_endpoint: Final = PassThroughGenericEndpoint.model_validate(data_dict)
3790 # Register the new route
3791 _custom_headers: dict | None = created_endpoint.headers or {}
3792 _custom_headers = await set_env_variables_in_header(custom_headers=_custom_headers)
3794 route_app: Final = _request_app(request)
3795 if created_endpoint.include_subpath:
3796 InitPassThroughEndpointHelpers.add_subpath_route(
3797 app=route_app,
3798 path=created_endpoint.path,
3799 target=created_endpoint.target,
3800 custom_headers=_custom_headers,
3801 forward_headers=None,
3802 merge_query_params=None,
3803 dependencies=None,
3804 cost_per_request=created_endpoint.cost_per_request,
3805 endpoint_id=created_endpoint.id or "",
3806 guardrails=getattr(created_endpoint, "guardrails", None),
3807 methods=created_endpoint.methods,
3808 default_query_params=created_endpoint.default_query_params,
3809 auth=created_endpoint.auth,
3810 timeout=created_endpoint.timeout,
3811 )
3812 else:
3813 InitPassThroughEndpointHelpers.add_exact_path_route(
3814 app=route_app,
3815 path=created_endpoint.path,
3816 target=created_endpoint.target,
3817 custom_headers=_custom_headers,
3818 forward_headers=None,
3819 merge_query_params=None,
3820 dependencies=None,
3821 cost_per_request=created_endpoint.cost_per_request,
3822 endpoint_id=created_endpoint.id or "",
3823 guardrails=getattr(created_endpoint, "guardrails", None),
3824 methods=created_endpoint.methods,
3825 default_query_params=created_endpoint.default_query_params,
3826 auth=created_endpoint.auth,
3827 timeout=created_endpoint.timeout,
3828 )
3830 return PassThroughEndpointResponse(endpoints=[created_endpoint])
3833@router.delete(
3834 "/config/pass_through_endpoint",
3835 dependencies=[Depends(user_api_key_auth)],
3836 response_model=PassThroughEndpointResponse,
3837)
3838async def delete_pass_through_endpoints(
3839 endpoint_id: str,
3840 user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
3841):
3842 """
3843 Delete a pass-through endpoint by ID.
3845 Returns - the deleted endpoint
3846 """
3847 from litellm.proxy.proxy_server import (
3848 get_config_general_settings,
3849 update_config_general_settings,
3850 )
3852 ## Get existing pass-through endpoint field value
3854 try:
3855 response: ConfigFieldInfo = await get_config_general_settings(
3856 field_name="pass_through_endpoints", user_api_key_dict=user_api_key_dict
3857 )
3858 except Exception:
3859 response = ConfigFieldInfo(field_name="pass_through_endpoints", field_value=None)
3861 ## Update field by removing endpoint
3862 pass_through_endpoint_data: Final[list | None] = _config_field_endpoints(response)
3863 if response.field_value is None or pass_through_endpoint_data is None: 3863 ↛ 3864line 3863 didn't jump to line 3864 because the condition on line 3863 was never true
3864 raise HTTPException(
3865 status_code=400,
3866 detail={"error": "There are no pass-through endpoints setup."},
3867 )
3869 # Find the endpoint to delete
3870 found_endpoint: Final = _find_endpoint_by_id(pass_through_endpoint_data, endpoint_id)
3872 if found_endpoint is None:
3873 raise HTTPException(
3874 status_code=400,
3875 detail={"error": f"Endpoint with ID '{endpoint_id}' was not found in pass-through endpoint list."},
3876 )
3878 # Find the index for deleting from the list
3879 endpoint_index = None
3880 for idx, endpoint in enumerate(pass_through_endpoint_data): 3880 ↛ 3886line 3880 didn't jump to line 3886 because the loop on line 3880 didn't complete
3881 _endpoint = PassThroughGenericEndpoint(**endpoint) if isinstance(endpoint, dict) else endpoint
3882 if _endpoint.id == endpoint_id:
3883 endpoint_index = idx
3884 break
3886 if endpoint_index is None: 3886 ↛ 3887line 3886 didn't jump to line 3887 because the condition on line 3886 was never true
3887 raise HTTPException(
3888 status_code=400,
3889 detail={"error": f"Could not find index for endpoint with ID '{endpoint_id}'"},
3890 )
3892 # Remove the endpoint
3893 pass_through_endpoint_data.pop(endpoint_index)
3894 response_obj: Final = found_endpoint
3896 # Remove routes from registry
3897 InitPassThroughEndpointHelpers.remove_endpoint_routes(endpoint_id)
3899 ## Update db
3900 updated_data: Final = ConfigFieldUpdate(
3901 field_name="pass_through_endpoints",
3902 field_value=pass_through_endpoint_data,
3903 config_type="general_settings",
3904 )
3905 await update_config_general_settings(data=updated_data, user_api_key_dict=user_api_key_dict)
3907 return PassThroughEndpointResponse(endpoints=[response_obj])
3910def _find_endpoint_by_id(
3911 endpoints_data: list,
3912 endpoint_id: str,
3913) -> PassThroughGenericEndpoint | None:
3914 """
3915 Find an endpoint by ID.
3917 Args:
3918 endpoints_data: List of endpoint data (dicts or PassThroughGenericEndpoint objects)
3919 endpoint_id: ID to search for
3921 Returns:
3922 Found endpoint or None if not found
3923 """
3924 for endpoint in endpoints_data:
3925 _endpoint: PassThroughGenericEndpoint | None = None
3926 if isinstance(endpoint, dict): 3926 ↛ 3928line 3926 didn't jump to line 3928 because the condition on line 3926 was always true
3927 _endpoint = PassThroughGenericEndpoint(**endpoint)
3928 elif isinstance(endpoint, PassThroughGenericEndpoint):
3929 _endpoint = endpoint
3931 # Only compare IDs to IDs
3932 if _endpoint is not None and _endpoint.id == endpoint_id:
3933 return _endpoint
3935 return None
3938async def initialize_pass_through_endpoints_in_db():
3939 """
3940 Gets all pass-through endpoints from db and initializes them in the proxy server.
3941 """
3942 pass_through_endpoints: Final = await _get_pass_through_endpoints_from_db()
3943 await initialize_pass_through_endpoints(pass_through_endpoints=pass_through_endpoints)