mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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:
parent
e8d7537b57
commit
dd92b1a0ac
15 changed files with 681 additions and 113 deletions
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
103
litellm/llms/base_llm/passthrough/transformation.py
Normal file
103
litellm/llms/base_llm/passthrough/transformation.py
Normal 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
|
||||
)
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
"""
|
||||
|
|
|
|||
53
litellm/llms/bedrock/passthrough/transformation.py
Normal file
53
litellm/llms/bedrock/passthrough/transformation.py
Normal 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,
|
||||
)
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
32
litellm/llms/vllm/passthrough/transformation.py
Normal file
32
litellm/llms/vllm/passthrough/transformation.py
Normal 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,
|
||||
)
|
||||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue