Refactor: bedrock passthrough fixes - migrate to Passthrough SDK (#12089)

* feat: initial commit adding bedrock support via the new sdk passthrough logic

ensures correct sequencing of tasks (pre call checks etc. can run before signing request)

* fix(route_llm_requests.py): passthrough to allm_passthrough_route if no model found

* feat(bedrock/passthrough): working bedrock passthrough via sdk support

* fix(passthrough/main.py): re-add data and json

* feat(passthrough/main): support async passthrough calls to bedrock

* feat(passthrough/main.py): async streaming + completion support

* feat(llm_passthrough_endpoints.py): migrate bedrock passthrough calls to to new bedrock passthrough sdk

Enables calls to work correctly

* fix: fix linting errors

* test: update test
This commit is contained in:
Krish Dholakia 2025-06-26 22:51:35 -07:00 • committed by GitHub
parent e8d7537b57
commit dd92b1a0ac
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
15 changed files with 681 additions and 113 deletions

View file

@ -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

View file

@ -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

View file

@ -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
)

View file

@ -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

View file

@ -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:
"""

View file

@ -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,
)

View file

@ -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,

View file

@ -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,
)

View file

@ -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,
)

View file

@ -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"

View file

@ -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/<Model>/<endpoint>. 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()

View file

@ -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(

View file

@ -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,

View file

@ -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)
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

View file

@ -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,