diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 7e32c5c438b..713c3358b51 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -1278,10 +1278,12 @@ class Logging(LiteLLMLoggingBaseClass): if self.call_type == CallTypes.anthropic_messages.value: result = self._handle_anthropic_messages_response_logging(result=result) elif ( - self.call_type == CallTypes.generate_content.value or - self.call_type == CallTypes.agenerate_content.value + self.call_type == CallTypes.generate_content.value + or self.call_type == CallTypes.agenerate_content.value ): - result = self._handle_non_streaming_google_genai_generate_content_response_logging(result=result) + result = self._handle_non_streaming_google_genai_generate_content_response_logging( + result=result + ) ## if model in model cost map - log the response cost ## else set cost to None @@ -2757,12 +2759,15 @@ class Logging(LiteLLMLoggingBaseClass): json_mode=None, ) return result - - def _handle_non_streaming_google_genai_generate_content_response_logging(self, result: Any) -> ModelResponse: + + def _handle_non_streaming_google_genai_generate_content_response_logging( + self, result: Any + ) -> ModelResponse: """ Handles logging for Google GenAI generate content responses. """ import httpx + httpx_response = self.model_call_details.get("httpx_response", None) if httpx_response is None: raise ValueError("Google GenAI Generate Content: httpx_response is None") @@ -2775,7 +2780,7 @@ class Logging(LiteLLMLoggingBaseClass): raw_response=httpx.Response( status_code=200, headers={}, - ), + ), ) return result diff --git a/litellm/llms/base_llm/base_utils.py b/litellm/llms/base_llm/base_utils.py index 712f5de8cc0..35959f0d083 100644 --- a/litellm/llms/base_llm/base_utils.py +++ b/litellm/llms/base_llm/base_utils.py @@ -41,7 +41,9 @@ class BaseLLMModelInfo(ABC): @staticmethod @abstractmethod - def get_api_base(api_base: Optional[str] = None) -> Optional[str]: + def get_api_base( + api_base: Optional[str] = None, + ) -> Optional[str]: pass @abstractmethod diff --git a/litellm/llms/base_llm/passthrough/transformation.py b/litellm/llms/base_llm/passthrough/transformation.py new file mode 100644 index 00000000000..e157a8ffde0 --- /dev/null +++ b/litellm/llms/base_llm/passthrough/transformation.py @@ -0,0 +1,103 @@ +from abc import abstractmethod +from typing import TYPE_CHECKING, Optional, Tuple, Union + +from ..base_utils import BaseLLMModelInfo + +if TYPE_CHECKING: + from httpx import URL, Headers + + from ..chat.transformation import BaseLLMException + + +class BasePassthroughConfig(BaseLLMModelInfo): + @abstractmethod + def is_streaming_request(self, endpoint: str, request_data: dict) -> bool: + """ + Check if the request is a streaming request + """ + pass + + def format_url( + self, + endpoint: str, + base_target_url: str, + request_query_params: Optional[dict], + ) -> "URL": + """ + Helper function to add query params to the url + Args: + endpoint: str - the endpoint to add to the url + base_target_url: str - the base url to add the endpoint to + request_query_params: dict - the query params to add to the url + Returns: + str - the formatted url + """ + from urllib.parse import urlencode + + import httpx + + encoded_endpoint = httpx.URL(endpoint).path + + # Ensure endpoint starts with '/' for proper URL construction + if not encoded_endpoint.startswith("/"): + encoded_endpoint = "/" + encoded_endpoint + + # Construct the full target URL using httpx + base_url = httpx.URL(base_target_url) + updated_url = base_url.copy_with(path=encoded_endpoint) + + if request_query_params: + # Create a new URL with the merged query params + updated_url = updated_url.copy_with( + query=urlencode(request_query_params).encode("ascii") + ) + return updated_url + + @abstractmethod + def get_complete_url( + self, + api_base: Optional[str], + api_key: Optional[str], + model: str, + endpoint: str, + request_query_params: Optional[dict], + litellm_params: dict, + ) -> Tuple["URL", str]: + """ + Get the complete url for the request + Returns: + - complete_url: URL - the complete url for the request + - base_target_url: str - the base url to add the endpoint to. Useful for auth headers. + """ + pass + + def sign_request( + self, + headers: dict, + litellm_params: dict, + request_data: Optional[dict], + api_base: str, + model: Optional[str] = None, + ) -> Tuple[dict, Optional[bytes]]: + """ + Some providers like Bedrock require signing the request. The sign request funtion needs access to `request_data` and `complete_url` + Args: + headers: dict + optional_params: dict + request_data: dict - the request body being sent in http request + api_base: str - the complete url being sent in http request + Returns: + dict - the signed headers + + Update the headers with the signed headers in this function. The return values will be sent as headers in the http request. + """ + return headers, None + + def get_error_class( + self, error_message: str, status_code: int, headers: Union[dict, "Headers"] + ) -> "BaseLLMException": + from litellm.llms.base_llm.chat.transformation import BaseLLMException + + return BaseLLMException( + status_code=status_code, message=error_message, headers=headers + ) diff --git a/litellm/llms/bedrock/base_aws_llm.py b/litellm/llms/bedrock/base_aws_llm.py index 20587080e46..ce3f66339e2 100644 --- a/litellm/llms/bedrock/base_aws_llm.py +++ b/litellm/llms/bedrock/base_aws_llm.py @@ -113,7 +113,7 @@ class BaseAWSLLM: elif param is None: # check if uppercase value in env key = self.aws_authentication_params[i] if key.upper() in os.environ: - params_to_check[i] = os.getenv(key) + params_to_check[i] = os.getenv(key.upper()) # Assign updated values back to parameters ( @@ -710,6 +710,7 @@ class BaseAWSLLM: Returns: Tuple[dict, Optional[str]]: A tuple containing the headers and the json str body of the request """ + try: from botocore.auth import SigV4Auth from botocore.awsrequest import AWSRequest @@ -762,4 +763,5 @@ class BaseAWSLLM: headers is not None and "Authorization" in headers ): # prevent sigv4 from overwriting the auth header request_headers_dict["Authorization"] = headers["Authorization"] + return request_headers_dict, request.body diff --git a/litellm/llms/bedrock/common_utils.py b/litellm/llms/bedrock/common_utils.py index fc6f52233e1..ee2a6cda8dc 100644 --- a/litellm/llms/bedrock/common_utils.py +++ b/litellm/llms/bedrock/common_utils.py @@ -3,7 +3,7 @@ Common utilities used across bedrock chat/embedding/image generation """ import os -from typing import List, Literal, Optional, Union +from typing import TYPE_CHECKING, List, Literal, Optional, Union import httpx @@ -12,6 +12,9 @@ from litellm.llms.base_llm.base_utils import BaseLLMModelInfo from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.secret_managers.main import get_secret +if TYPE_CHECKING: + from litellm.types.llms.openai import AllMessageValues + class BedrockError(BaseLLMException): pass @@ -333,6 +336,37 @@ class BedrockModelInfo(BaseLLMModelInfo): global_config = AmazonBedrockGlobalConfig() all_global_regions = global_config.get_all_regions() + @staticmethod + def get_api_base(api_base: Optional[str] = None) -> Optional[str]: + """ + Get the API base for the given model. + """ + return api_base + + @staticmethod + def get_api_key(api_key: Optional[str] = None) -> Optional[str]: + """ + Get the API key for the given model. + """ + return api_key + + def validate_environment( + self, + headers: dict, + model: str, + messages: List["AllMessageValues"], + optional_params: dict, + litellm_params: dict, + api_key: Optional[str] = None, + api_base: Optional[str] = None, + ) -> dict: + return headers + + def get_models( + self, api_key: Optional[str] = None, api_base: Optional[str] = None + ) -> List[str]: + return [] + @staticmethod def extract_model_name_from_arn(model: str) -> str: """ diff --git a/litellm/llms/bedrock/passthrough/transformation.py b/litellm/llms/bedrock/passthrough/transformation.py new file mode 100644 index 00000000000..6ed6d7829ec --- /dev/null +++ b/litellm/llms/bedrock/passthrough/transformation.py @@ -0,0 +1,53 @@ +from typing import TYPE_CHECKING, Optional, Tuple + +from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConfig + +from ..base_aws_llm import BaseAWSLLM +from ..common_utils import BedrockModelInfo + +if TYPE_CHECKING: + from httpx import URL + + +class BedrockPassthroughConfig(BaseAWSLLM, BedrockModelInfo, BasePassthroughConfig): + def is_streaming_request(self, endpoint: str, request_data: dict) -> bool: + return "stream" in endpoint + + def get_complete_url( + self, + api_base: Optional[str], + api_key: Optional[str], + model: str, + endpoint: str, + request_query_params: Optional[dict], + litellm_params: dict, + ) -> Tuple["URL", str]: + optional_params = litellm_params.copy() + + aws_region_name = self._get_aws_region_name( + optional_params=optional_params, + model=model, + model_id=None, + ) + + api_base = f"https://bedrock-runtime.{aws_region_name}.amazonaws.com" + + return self.format_url(endpoint, api_base, request_query_params or {}), api_base + + def sign_request( + self, + headers: dict, + litellm_params: dict, + request_data: Optional[dict], + api_base: str, + model: Optional[str] = None, + ) -> Tuple[dict, Optional[bytes]]: + optional_params = litellm_params.copy() + return self._sign_request( + service_name="bedrock", + headers=headers, + optional_params=optional_params, + request_data=request_data or {}, + api_base=api_base, + model=model, + ) diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index bc3e293a452..dbabdaa1490 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -79,6 +79,7 @@ from litellm.utils import ( if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj + from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConfig LiteLLMLoggingObj = _LiteLLMLoggingObj else: @@ -2364,6 +2365,7 @@ class BaseLLMHTTPHandler: BaseVectorStoreConfig, BaseGoogleGenAIGenerateContentConfig, BaseAnthropicMessagesConfig, + "BasePassthroughConfig", ], ): status_code = getattr(e, "status_code", 500) @@ -2383,6 +2385,15 @@ class BaseLLMHTTPHandler: else: error_headers = {} + if provider_config is None: + from litellm.llms.base_llm.chat.transformation import BaseLLMException + + raise BaseLLMException( + status_code=status_code, + message=error_text, + headers=error_headers, + ) + raise provider_config.get_error_class( error_message=error_text, status_code=status_code, diff --git a/litellm/llms/vllm/passthrough/transformation.py b/litellm/llms/vllm/passthrough/transformation.py new file mode 100644 index 00000000000..cc8a78fb50d --- /dev/null +++ b/litellm/llms/vllm/passthrough/transformation.py @@ -0,0 +1,32 @@ +from typing import TYPE_CHECKING, Optional, Tuple + +from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConfig + +from ..common_utils import VLLMModelInfo + +if TYPE_CHECKING: + from httpx import URL + + +class VLLMPassthroughConfig(VLLMModelInfo, BasePassthroughConfig): + def is_streaming_request(self, endpoint: str, request_data: dict) -> bool: + return "stream" in request_data + + def get_complete_url( + self, + api_base: Optional[str], + api_key: Optional[str], + model: str, + endpoint: str, + request_query_params: Optional[dict], + litellm_params: dict, + ) -> Tuple["URL", str]: + base_target_url = self.get_api_base(api_base) + + if base_target_url is None: + raise Exception("VLLM api base not found") + + return ( + self.format_url(endpoint, base_target_url, request_query_params), + base_target_url, + ) diff --git a/litellm/passthrough/main.py b/litellm/passthrough/main.py index 208d0dbbaf9..a0d09ef7ce6 100644 --- a/litellm/passthrough/main.py +++ b/litellm/passthrough/main.py @@ -5,8 +5,7 @@ This module is used to pass through requests to the LLM APIs. import asyncio import contextvars from functools import partial -from typing import Any, Coroutine, Optional, Union -from urllib.parse import urlencode +from typing import TYPE_CHECKING, Any, Coroutine, Optional, Union, cast import httpx from httpx._types import CookieTypes, QueryParamTypes, RequestFiles @@ -14,22 +13,27 @@ from httpx._types import CookieTypes, QueryParamTypes, RequestFiles import litellm from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler +from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler from litellm.utils import client +base_llm_http_handler = BaseLLMHTTPHandler() from .utils import BasePassthroughUtils +if TYPE_CHECKING: + from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConfig + @client async def allm_passthrough_route( *, method: str, endpoint: str, + model: str, custom_llm_provider: Optional[str] = None, api_base: Optional[str] = None, api_key: Optional[str] = None, request_query_params: Optional[dict] = None, request_headers: Optional[dict] = None, - stream: bool = False, content: Optional[Any] = None, data: Optional[dict] = None, files: Optional[RequestFiles] = None, @@ -50,12 +54,12 @@ async def allm_passthrough_route( llm_passthrough_route, method=method, endpoint=endpoint, + model=model, custom_llm_provider=custom_llm_provider, api_base=api_base, api_key=api_key, request_query_params=request_query_params, request_headers=request_headers, - stream=stream, content=content, data=data, files=files, @@ -72,11 +76,20 @@ async def allm_passthrough_route( if asyncio.iscoroutine(init_response): response = await init_response + try: + response.raise_for_status() + except httpx.HTTPStatusError as e: + error_text = await e.response.aread() + error_text_str = error_text.decode("utf-8") + raise Exception(error_text_str) else: response = init_response return response except Exception as e: - raise e + raise base_llm_http_handler._handle_error( + e=e, + provider_config=None, + ) @client @@ -91,7 +104,6 @@ def llm_passthrough_route( request_query_params: Optional[dict] = None, request_headers: Optional[dict] = None, allm_passthrough_route: bool = False, - stream: bool = False, content: Optional[Any] = None, data: Optional[dict] = None, files: Optional[RequestFiles] = None, @@ -123,37 +135,29 @@ def llm_passthrough_route( api_key=api_key, ) + from litellm.litellm_core_utils.get_litellm_params import get_litellm_params from litellm.types.utils import LlmProviders from litellm.utils import ProviderConfigManager - provider_config = ProviderConfigManager.get_provider_model_info( + litellm_params_dict = get_litellm_params(**kwargs) + + provider_config = cast( + Optional["BasePassthroughConfig"], kwargs.get("provider_config") + ) or ProviderConfigManager.get_provider_passthrough_config( provider=LlmProviders(custom_llm_provider), model=model, ) if provider_config is None: raise Exception(f"Provider {custom_llm_provider} not found") - base_target_url = provider_config.get_api_base(api_base) - - if base_target_url is None: - raise Exception(f"Provider {custom_llm_provider} api base not found") - - encoded_endpoint = httpx.URL(endpoint).path - - # Ensure endpoint starts with '/' for proper URL construction - if not encoded_endpoint.startswith("/"): - encoded_endpoint = "/" + encoded_endpoint - - # Construct the full target URL using httpx - base_url = httpx.URL(base_target_url) - updated_url = base_url.copy_with(path=encoded_endpoint) - - if request_query_params: - # Create a new URL with the merged query params - updated_url = updated_url.copy_with( - query=urlencode(request_query_params).encode("ascii") - ) - + updated_url, base_target_url = provider_config.get_complete_url( + api_base=api_base, + api_key=api_key, + model=model, + endpoint=endpoint, + request_query_params=request_query_params, + litellm_params=litellm_params_dict, + ) # Add or update query parameters provider_api_key = provider_config.get_api_key(api_key) @@ -173,6 +177,14 @@ def llm_passthrough_route( forward_headers=False, ) + headers, signed_json_body = provider_config.sign_request( + headers=headers, + litellm_params=litellm_params_dict, + request_data=data if data else json, + api_base=str(updated_url), + model=model, + ) + ## SWAP MODEL IN JSON BODY if json and isinstance(json, dict) and "model" in json: json["model"] = model @@ -180,14 +192,31 @@ def llm_passthrough_route( request = client.client.build_request( method=method, url=updated_url, - content=content, - data=data, + content=signed_json_body, + data=data if signed_json_body is None else None, files=files, - json=json, + json=json if signed_json_body is None else None, params=params, headers=headers, cookies=cookies, ) - response = client.client.send(request=request, stream=stream) - return response + ## IS STREAMING REQUEST + is_streaming_request = provider_config.is_streaming_request( + endpoint=endpoint, + request_data=data or json or {}, + ) + + try: + response = client.client.send(request=request, stream=is_streaming_request) + if asyncio.iscoroutine(response): + return response + response.raise_for_status() + return response + except Exception as e: + if provider_config is None: + raise e + raise base_llm_http_handler._handle_error( + e=e, + provider_config=provider_config, + ) diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index afe596b03d5..f943549a412 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -259,6 +259,7 @@ class ProxyBaseLLMRequestProcessing: "aimage_edit", "agenerate_content", "agenerate_content_stream", + "allm_passthrough_route", ], version: Optional[str] = None, user_model: Optional[str] = None, @@ -340,6 +341,7 @@ class ProxyBaseLLMRequestProcessing: "alist_input_items", "agenerate_content", "agenerate_content_stream", + "allm_passthrough_route", ], proxy_logging_obj: ProxyLogging, general_settings: dict, @@ -428,7 +430,9 @@ class ProxyBaseLLMRequestProcessing: litellm_call_id=self.data.get("litellm_call_id", ""), status="success" ) ) - if self._is_streaming_request(data=self.data, is_streaming_request=is_streaming_request): # use generate_responses to stream responses + if self._is_streaming_request( + data=self.data, is_streaming_request=is_streaming_request + ): # use generate_responses to stream responses custom_headers = ProxyBaseLLMRequestProcessing.get_custom_headers( user_api_key_dict=user_api_key_dict, call_id=logging_obj.litellm_call_id, @@ -443,16 +447,23 @@ class ProxyBaseLLMRequestProcessing: hidden_params=hidden_params, **additional_headers, ) - selected_data_generator = select_data_generator( - response=response, - user_api_key_dict=user_api_key_dict, - request_data=self.data, - ) - return await create_streaming_response( - generator=selected_data_generator, - media_type="text/event-stream", - headers=custom_headers, - ) + if route_type == "allm_passthrough_route": + return StreamingResponse( + content=response.aiter_bytes(), + status_code=response.status_code, + headers=custom_headers, + ) + else: + selected_data_generator = select_data_generator( + response=response, + user_api_key_dict=user_api_key_dict, + request_data=self.data, + ) + return await create_streaming_response( + generator=selected_data_generator, + media_type="text/event-stream", + headers=custom_headers, + ) ### CALL HOOKS ### - modify outgoing data response = await proxy_logging_obj.post_call_success_hook( @@ -484,6 +495,96 @@ class ProxyBaseLLMRequestProcessing: return response + async def base_passthrough_process_llm_request( + self, + request: Request, + fastapi_response: Response, + user_api_key_dict: UserAPIKeyAuth, + proxy_logging_obj: ProxyLogging, + general_settings: dict, + proxy_config: ProxyConfig, + select_data_generator: Callable, + llm_router: Optional[Router] = None, + model: Optional[str] = None, + user_model: Optional[str] = None, + user_temperature: Optional[float] = None, + user_request_timeout: Optional[float] = None, + user_max_tokens: Optional[int] = None, + user_api_base: Optional[str] = None, + version: Optional[str] = None, + ): + from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( + HttpPassThroughEndpointHelpers, + ) + + result = await self.base_process_llm_request( + request=request, + fastapi_response=fastapi_response, + user_api_key_dict=user_api_key_dict, + route_type="allm_passthrough_route", + proxy_logging_obj=proxy_logging_obj, + llm_router=llm_router, + general_settings=general_settings, + proxy_config=proxy_config, + select_data_generator=select_data_generator, + model=model, + user_model=user_model, + user_temperature=user_temperature, + user_request_timeout=user_request_timeout, + user_max_tokens=user_max_tokens, + user_api_base=user_api_base, + version=version, + ) + + # Check if result is actually a streaming response by inspecting its type + if isinstance(result, StreamingResponse): + return result + + content = await result.aread() + return Response( + content=content, + status_code=result.status_code, + headers=HttpPassThroughEndpointHelpers.get_response_headers( + headers=result.headers, + custom_headers=None, + ), + ) + + def _is_streaming_response(self, response: Any) -> bool: + """ + Check if the response object is actually a streaming response by inspecting its type. + + This uses standard Python inspection to detect streaming/async iterator objects + rather than relying on specific wrapper classes. + """ + import asyncio + import inspect + from collections.abc import AsyncGenerator, AsyncIterator + + # Check if it's an async generator (most reliable) + if inspect.isasyncgen(response): + return True + + # Check if it implements the async iterator protocol + if isinstance(response, (AsyncIterator, AsyncGenerator)): + return True + + # Check for __aiter__ method (async iterator protocol) + if hasattr(response, "__aiter__") and callable(getattr(response, "__aiter__")): + return True + + # Check if it's a coroutine that might yield an async generator + if asyncio.iscoroutine(response): + return True + + # Check for common streaming HTTP response patterns + if hasattr(response, "aiter_bytes") and callable( + getattr(response, "aiter_bytes") + ): + return True + + return False + def _is_streaming_request( self, data: dict, is_streaming_request: Optional[bool] = False ) -> bool: @@ -499,7 +600,6 @@ class ProxyBaseLLMRequestProcessing: return True return False - async def _handle_llm_api_exception( self, e: Exception, @@ -569,7 +669,7 @@ class ProxyBaseLLMRequestProcessing: return "completion" elif route_type == "aresponses": return "responses" - + ######################################################### # Proxy Level Streaming Data Generator ######################################################### @@ -644,5 +744,3 @@ class ProxyBaseLLMRequestProcessing: ) error_returned = json.dumps({"error": proxy_exception.to_dict()}) yield f"{STREAM_SSE_DATA_PREFIX}{error_returned}\n\n" - - diff --git a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py index f9c72304947..f85ecb47a2f 100644 --- a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py @@ -306,9 +306,11 @@ async def vllm_proxy_route( content=None, data=None, files=None, - json=request_body - if request.headers.get("content-type") == "application/json" - else None, + json=( + request_body + if request.headers.get("content-type") == "application/json" + else None + ), params=None, headers=None, cookies=None, @@ -462,6 +464,76 @@ async def anthropic_proxy_route( return received_value +async def bedrock_llm_proxy_route( + endpoint: str, + request: Request, + fastapi_response: Response, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): + """ + Handles Bedrock LLM API calls. + """ + from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing + from litellm.proxy.proxy_server import ( + general_settings, + llm_router, + proxy_config, + proxy_logging_obj, + select_data_generator, + user_api_base, + user_max_tokens, + user_model, + user_request_timeout, + user_temperature, + version, + ) + + request_body = await _read_request_body(request=request) + data: Dict[str, Any] = {} + base_llm_response_processor = ProxyBaseLLMRequestProcessing(data=data) + try: + model = endpoint.split("/")[1] + except Exception: + raise HTTPException( + status_code=400, + detail={ + "error": "Model missing from endpoint. Expected format: /model//. Got: " + + endpoint, + }, + ) + + data["method"] = request.method + data["endpoint"] = endpoint + data["data"] = request_body + + try: + result = await base_llm_response_processor.base_passthrough_process_llm_request( + request=request, + fastapi_response=fastapi_response, + user_api_key_dict=user_api_key_dict, + proxy_logging_obj=proxy_logging_obj, + llm_router=llm_router, + general_settings=general_settings, + proxy_config=proxy_config, + select_data_generator=select_data_generator, + model=model, + user_model=user_model, + user_temperature=user_temperature, + user_request_timeout=user_request_timeout, + user_max_tokens=user_max_tokens, + user_api_base=user_api_base, + version=version, + ) + + return result + except Exception as e: + raise await base_llm_response_processor._handle_llm_api_exception( + e=e, + user_api_key_dict=user_api_key_dict, + proxy_logging_obj=proxy_logging_obj, + ) + + @router.api_route( "/bedrock/{endpoint:path}", methods=["GET", "POST", "PUT", "DELETE", "PATCH"], @@ -474,6 +546,8 @@ async def bedrock_proxy_route( user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ + This is the v1 passthrough for Bedrock. + V2 is handled by the `/bedrock/v2` endpoint. [Docs](https://docs.litellm.ai/docs/pass_through/bedrock) """ create_request_copy(request) @@ -491,7 +565,12 @@ async def bedrock_proxy_route( f"https://bedrock-agent-runtime.{aws_region_name}.amazonaws.com" ) else: - base_target_url = f"https://bedrock-runtime.{aws_region_name}.amazonaws.com" + return await bedrock_llm_proxy_route( + endpoint=endpoint, + request=request, + fastapi_response=fastapi_response, + user_api_key_dict=user_api_key_dict, + ) encoded_endpoint = httpx.URL(endpoint).path # Ensure endpoint starts with '/' for proper URL construction @@ -705,6 +784,7 @@ class VertexAIPassThroughHandler(BaseVertexAIPassThroughHandler): ) -> str: return get_vertex_base_url(vertex_location) + def get_vertex_base_url(vertex_location: Optional[str]) -> str: """ Returns the base URL for Vertex AI based on the provided location. @@ -713,8 +793,9 @@ def get_vertex_base_url(vertex_location: Optional[str]) -> str: return "https://aiplatform.googleapis.com/" return f"https://{vertex_location}-aiplatform.googleapis.com/" + def get_vertex_pass_through_handler( - call_type: Literal["discovery", "aiplatform"] + call_type: Literal["discovery", "aiplatform"], ) -> BaseVertexAIPassThroughHandler: if call_type == "discovery": return VertexAIDiscoveryPassThroughHandler() diff --git a/litellm/proxy/route_llm_request.py b/litellm/proxy/route_llm_request.py index e920210a75e..3d7882740b7 100644 --- a/litellm/proxy/route_llm_request.py +++ b/litellm/proxy/route_llm_request.py @@ -75,6 +75,7 @@ async def route_request( "aimage_edit", "agenerate_content", "agenerate_content_stream", + "allm_passthrough_route", ], ): """ @@ -152,7 +153,8 @@ async def route_request( elif user_model is not None: return getattr(litellm, f"{route_type}")(**data) - + elif route_type == "allm_passthrough_route": + return getattr(litellm, f"{route_type}")(**data) # if no route found then it's a bad request route_name = ROUTE_ENDPOINT_MAPPING.get(route_type, route_type) raise ProxyModelNotFoundError( diff --git a/litellm/utils.py b/litellm/utils.py index 9b176dfc411..592c626aeeb 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -244,6 +244,7 @@ from litellm.llms.base_llm.image_generation.transformation import ( from litellm.llms.base_llm.image_variations.transformation import ( BaseImageVariationConfig, ) +from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConfig from litellm.llms.base_llm.realtime.transformation import BaseRealtimeConfig from litellm.llms.base_llm.rerank.transformation import BaseRerankConfig from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig @@ -538,9 +539,9 @@ def function_setup( # noqa: PLR0915 function_id: Optional[str] = kwargs["id"] if "id" in kwargs else None ## DYNAMIC CALLBACKS ## - dynamic_callbacks: Optional[ - List[Union[str, Callable, CustomLogger]] - ] = kwargs.pop("callbacks", None) + dynamic_callbacks: Optional[List[Union[str, Callable, CustomLogger]]] = ( + kwargs.pop("callbacks", None) + ) all_callbacks = get_dynamic_callbacks(dynamic_callbacks=dynamic_callbacks) if len(all_callbacks) > 0: @@ -1258,9 +1259,9 @@ def client(original_function): # noqa: PLR0915 exception=e, retry_policy=kwargs.get("retry_policy"), ) - kwargs[ - "retry_policy" - ] = reset_retry_policy() # prevent infinite loops + kwargs["retry_policy"] = ( + reset_retry_policy() + ) # prevent infinite loops litellm.num_retries = ( None # set retries to None to prevent infinite loops ) @@ -1577,10 +1578,13 @@ def _is_streaming_request( except ValueError: return False - if call_type == CallTypes.generate_content_stream or call_type == CallTypes.agenerate_content_stream: + if ( + call_type == CallTypes.generate_content_stream + or call_type == CallTypes.agenerate_content_stream + ): return True ######################################################### - + return False @@ -2941,10 +2945,10 @@ def pre_process_non_default_params( if "response_format" in non_default_params: if provider_config is not None: - non_default_params[ - "response_format" - ] = provider_config.get_json_schema_from_pydantic_object( - response_format=non_default_params["response_format"] + non_default_params["response_format"] = ( + provider_config.get_json_schema_from_pydantic_object( + response_format=non_default_params["response_format"] + ) ) else: non_default_params["response_format"] = type_to_response_format_param( @@ -3071,16 +3075,16 @@ def pre_process_optional_params( True # so that main.py adds the function call to the prompt ) if "tools" in non_default_params: - optional_params[ - "functions_unsupported_model" - ] = non_default_params.pop("tools") + optional_params["functions_unsupported_model"] = ( + non_default_params.pop("tools") + ) non_default_params.pop( "tool_choice", None ) # causes ollama requests to hang elif "functions" in non_default_params: - optional_params[ - "functions_unsupported_model" - ] = non_default_params.pop("functions") + optional_params["functions_unsupported_model"] = ( + non_default_params.pop("functions") + ) elif ( litellm.add_function_to_prompt ): # if user opts to add it to prompt instead @@ -4163,9 +4167,9 @@ def _count_characters(text: str) -> int: def get_response_string(response_obj: Union[ModelResponse, ModelResponseStream]) -> str: - _choices: Union[ - List[Union[Choices, StreamingChoices]], List[StreamingChoices] - ] = response_obj.choices + _choices: Union[List[Union[Choices, StreamingChoices]], List[StreamingChoices]] = ( + response_obj.choices + ) response_str = "" for choice in _choices: @@ -4699,7 +4703,9 @@ def _get_model_info_helper( # noqa: PLR0915 output_cost_per_second=_model_info.get("output_cost_per_second", None), output_cost_per_image=_model_info.get("output_cost_per_image", None), output_vector_size=_model_info.get("output_vector_size", None), - citation_cost_per_token=_model_info.get("citation_cost_per_token", None), + citation_cost_per_token=_model_info.get( + "citation_cost_per_token", None + ), litellm_provider=_model_info.get( "litellm_provider", custom_llm_provider ), @@ -6919,6 +6925,26 @@ class ProviderConfigManager: return VLLMModelInfo() return None + @staticmethod + def get_provider_passthrough_config( + model: str, + provider: LlmProviders, + ) -> Optional[BasePassthroughConfig]: + if LlmProviders.BEDROCK == provider: + from litellm.llms.bedrock.passthrough.transformation import ( + BedrockPassthroughConfig, + ) + + return BedrockPassthroughConfig() + elif LlmProviders.VLLM == provider: + from litellm.llms.vllm.passthrough.transformation import ( + VLLMPassthroughConfig, + ) + + return VLLMPassthroughConfig() + + return None + @staticmethod def get_provider_image_variation_config( model: str, @@ -6958,7 +6984,6 @@ class ProviderConfigManager: if LlmProviders.BEDROCK == provider: return BedrockVectorStore.get_initialized_custom_logger() return None - @staticmethod def get_provider_vector_stores_config( @@ -6971,11 +6996,13 @@ class ProviderConfigManager: from litellm.llms.openai.vector_stores.transformation import ( OpenAIVectorStoreConfig, ) + return OpenAIVectorStoreConfig() elif litellm.LlmProviders.AZURE == provider: from litellm.llms.azure.vector_stores.transformation import ( AzureOpenAIVectorStoreConfig, ) + return AzureOpenAIVectorStoreConfig() return None @@ -7027,7 +7054,7 @@ class ProviderConfigManager: return AzureImageEditConfig() return None - + @staticmethod def get_provider_google_genai_generate_content_config( model: str, diff --git a/tests/llm_translation/test_bedrock_completion.py b/tests/llm_translation/test_bedrock_completion.py index b147132cd07..7b7ddf0a3de 100644 --- a/tests/llm_translation/test_bedrock_completion.py +++ b/tests/llm_translation/test_bedrock_completion.py @@ -1266,6 +1266,7 @@ def test_bedrock_tools_pt_invalid_names(): def test_bedrock_tools_transformation_valid_params(): from litellm.types.llms.bedrock import ToolJsonSchemaBlock + tools = [ { "type": "function", @@ -1288,14 +1289,17 @@ def test_bedrock_tools_transformation_valid_params(): result = _bedrock_tools_pt(tools) print("bedrock tools after prompt formatting=", result) - # Ensure the keys for properties in the response is a subset of keys in ToolJsonSchemaBlock + # Ensure the keys for properties in the response is a subset of keys in ToolJsonSchemaBlock toolJsonSchema = result[0]["toolSpec"]["inputSchema"]["json"] assert toolJsonSchema is not None print("transformed toolJsonSchema keys=", toolJsonSchema.keys()) - print("allowed ToolJsonSchemaBlock keys=", ToolJsonSchemaBlock.__annotations__.keys()) - assert set(toolJsonSchema.keys()).issubset(set(ToolJsonSchemaBlock.__annotations__.keys())) + print( + "allowed ToolJsonSchemaBlock keys=", ToolJsonSchemaBlock.__annotations__.keys() + ) + assert set(toolJsonSchema.keys()).issubset( + set(ToolJsonSchemaBlock.__annotations__.keys()) + ) - assert isinstance(result, list) assert len(result) == 1 assert "toolSpec" in result[0] @@ -1303,11 +1307,13 @@ def test_bedrock_tools_transformation_valid_params(): assert result[0]["toolSpec"]["description"] == "Invalid name test" assert "inputSchema" in result[0]["toolSpec"] assert "json" in result[0]["toolSpec"]["inputSchema"] - assert result[0]["toolSpec"]["inputSchema"]["json"]["properties"]["test"]["type"] == "string" + assert ( + result[0]["toolSpec"]["inputSchema"]["json"]["properties"]["test"]["type"] + == "string" + ) assert "test" in result[0]["toolSpec"]["inputSchema"]["json"]["required"] - def test_not_found_error(): with pytest.raises(litellm.NotFoundError): completion( @@ -2269,7 +2275,6 @@ class TestBedrockConverseNovaTestSuite(BaseLLMChatTest): """Test that tool calls with no arguments is translated correctly. Relevant issue: https://github.com/BerriAI/litellm/issues/6833""" pass - def test_prompt_caching(self): """ TODO: Ensure this test passes our base llm test suite @@ -2479,7 +2484,9 @@ def test_bedrock_process_empty_text_blocks(): assert modified_message["content"][0]["text"] == "Please continue." -@pytest.mark.skip(reason="Skipping test due to bedrock changing their response schema support. Come back to this.") +@pytest.mark.skip( + reason="Skipping test due to bedrock changing their response schema support. Come back to this." +) def test_nova_optional_params_tool_choice(): try: litellm.drop_params = True @@ -2499,7 +2506,12 @@ def test_nova_optional_params_tool_choice(): "parameters": { "$defs": { "TurnDurationEnum": { - "enum": ["action", "encounter", "battle", "operation"], + "enum": [ + "action", + "encounter", + "battle", + "operation", + ], "title": "TurnDurationEnum", "type": "string", } @@ -2512,10 +2524,22 @@ def test_nova_optional_params_tool_choice(): }, "prompt": {"title": "Prompt", "type": "string"}, "name": {"title": "Name", "type": "string"}, - "description": {"title": "Description", "type": "string"}, - "competitve": {"title": "Competitve", "type": "boolean"}, - "players_min": {"title": "Players Min", "type": "integer"}, - "players_max": {"title": "Players Max", "type": "integer"}, + "description": { + "title": "Description", + "type": "string", + }, + "competitve": { + "title": "Competitve", + "type": "boolean", + }, + "players_min": { + "title": "Players Min", + "type": "integer", + }, + "players_max": { + "title": "Players Max", + "type": "integer", + }, "turn_duration": { "$ref": "#/$defs/TurnDurationEnum", "description": "how long the passing of a turn should represent for a game at this scale", @@ -2540,6 +2564,7 @@ def test_nova_optional_params_tool_choice(): except litellm.APIConnectionError: pass + class TestBedrockEmbedding(BaseLLMEmbeddingTest): def get_base_embedding_call_args(self) -> dict: return { @@ -3021,7 +3046,8 @@ def test_bedrock_application_inference_profile(): client = HTTPHandler() client2 = HTTPHandler() - tools = [{ + tools = [ + { "type": "function", "function": { "name": "get_current_weather", @@ -3039,12 +3065,11 @@ def test_bedrock_application_inference_profile(): }, }, "required": ["location"], - } - } + }, + }, } ] - with patch.object(client, "post") as mock_post, patch.object( client2, "post" ) as mock_post2: @@ -3054,7 +3079,7 @@ def test_bedrock_application_inference_profile(): messages=[{"role": "user", "content": "Hello, how are you?"}], model_id="arn:aws:bedrock:eu-central-1:000000000000:application-inference-profile/a0a0a0a0a0a0", client=client, - tools=tools + tools=tools, ) except Exception as e: print(e) @@ -3064,7 +3089,7 @@ def test_bedrock_application_inference_profile(): model="bedrock/converse/arn:aws:bedrock:eu-central-1:000000000000:application-inference-profile/a0a0a0a0a0a0", messages=[{"role": "user", "content": "Hello, how are you?"}], client=client2, - tools=tools + tools=tools, ) except Exception as e: print(e) @@ -3094,7 +3119,6 @@ def return_mocked_response(model: str): } - @pytest.mark.parametrize( "model", [ @@ -3141,7 +3165,6 @@ async def test_bedrock_max_completion_tokens(model: str): } - def test_bedrock_meta_llama_function_calling(): """ Tests that: @@ -3149,9 +3172,10 @@ def test_bedrock_meta_llama_function_calling(): """ from litellm.utils import return_raw_request from litellm.types.utils import CallTypes + tools = [ - { - "type": "function", + { + "type": "function", "function": { "name": "get_current_weather", "description": "Get the current weather in a given location", @@ -3191,4 +3215,69 @@ def test_bedrock_meta_llama_function_calling(): print(response) - \ No newline at end of file + +def test_bedrock_passthrough(): + import litellm + + litellm._turn_on_debug() + + data = { + "max_tokens": 512, + "messages": [{"role": "user", "content": "Hey"}], + "system": [ + { + "type": "text", + "text": "Analyze if this message indicates a new conversation topic. If it does, extract a 2-3 word title that captures the new topic. Format your response as a JSON object with two fields: 'isNewTopic' (boolean) and 'title' (string, or null if isNewTopic is false). Only include these fields, no other text.", + } + ], + "temperature": 0, + "metadata": { + "user_id": "5dd07c33da27e6d2968d94ea20bf47a7b090b6b158b82328d54da2909a108e84" + }, + "anthropic_version": "bedrock-2023-05-31", + "anthropic_beta": ["claude-code-20250219"], + } + + response = litellm.llm_passthrough_route( + model="bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0", + method="POST", + endpoint="/model/us.anthropic.claude-3-5-sonnet-20240620-v1:0/invoke", + data=data, + ) + + print(response.text) + + assert response.status_code == 200 + + +def test_bedrock_streaming_passthrough(): + import litellm + + litellm._turn_on_debug() + + data = { + "max_tokens": 512, + "messages": [{"role": "user", "content": "Hey"}], + "system": [ + { + "type": "text", + "text": "Analyze if this message indicates a new conversation topic. If it does, extract a 2-3 word title that captures the new topic. Format your response as a JSON object with two fields: 'isNewTopic' (boolean) and 'title' (string, or null if isNewTopic is false). Only include these fields, no other text.", + } + ], + "temperature": 0, + "metadata": { + "user_id": "5dd07c33da27e6d2968d94ea20bf47a7b090b6b158b82328d54da2909a108e84" + }, + "anthropic_version": "bedrock-2023-05-31", + "anthropic_beta": ["claude-code-20250219"], + } + + response = litellm.llm_passthrough_route( + model="bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0", + method="POST", + endpoint="/model/us.anthropic.claude-3-5-sonnet-20240620-v1:0/invoke-with-response-stream", + data=data, + stream=True, + ) + + assert response.status_code == 200 diff --git a/tests/test_litellm/passthrough/test_passthrough_main.py b/tests/test_litellm/passthrough/test_passthrough_main.py index c2c2bba2ced..afbaba60cf3 100644 --- a/tests/test_litellm/passthrough/test_passthrough_main.py +++ b/tests/test_litellm/passthrough/test_passthrough_main.py @@ -27,13 +27,13 @@ def test_llm_passthrough_route(): return_value=MagicMock(status_code=200, json={"message": "Hello, world!"}), ) as mock_post: response = llm_passthrough_route( - model="gpt-3.5-turbo", + model="vllm/anthropic.claude-3-5-sonnet-20240620-v1:0", endpoint="v1/chat/completions", method="POST", request_url="http://localhost:8000/v1/chat/completions", api_base="http://localhost:8090", json={ - "model": "gpt-3.5-turbo", + "model": "my-custom-model", "messages": [{"role": "user", "content": "Hello, world!"}], }, client=client,