Skip to content
7 changes: 5 additions & 2 deletions litellm/litellm_core_utils/get_litellm_params.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,7 @@

# Pre-define optional kwargs keys as frozenset for O(1) lookups
# These are extracted from kwargs only if present, avoiding unnecessary .get() calls
_OPTIONAL_KWARGS_KEYS = frozenset(
OPTIONAL_KWARGS_KEYS = frozenset(
{
"azure_ad_token",
"tenant_id",
Expand Down Expand Up @@ -39,6 +39,9 @@
}
)

# Backward-compatible alias for existing imports/tests.
_OPTIONAL_KWARGS_KEYS = OPTIONAL_KWARGS_KEYS


def _get_base_model_from_litellm_call_metadata(
metadata: Optional[dict],
Expand Down Expand Up @@ -166,7 +169,7 @@ def get_litellm_params(

# Sparse extraction: only add kwargs keys that are actually present
if kwargs:
for key in _OPTIONAL_KWARGS_KEYS:
for key in OPTIONAL_KWARGS_KEYS:
if key in kwargs:
litellm_params[key] = kwargs[key]

Expand Down
23 changes: 12 additions & 11 deletions litellm/litellm_core_utils/get_llm_provider_logic.py
Original file line number Diff line number Diff line change
@@ -1,13 +1,13 @@
import re
from typing import Optional, Tuple
from typing import Optional, Tuple, cast
from urllib.parse import urlparse

import litellm
from litellm.constants import REPLICATE_MODEL_NAME_WITH_ID_LENGTH
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
from litellm.secret_managers.main import get_secret, get_secret_str

from ..types.router import LiteLLM_Params
from ..types.router import GenericLiteLLMParams, LiteLLM_Params


def _endpoint_matches_api_base(endpoint: str, api_base: str) -> bool:
Expand Down Expand Up @@ -159,7 +159,7 @@ def get_llm_provider( # noqa: PLR0915
custom_llm_provider: Optional[str] = None,
api_base: Optional[str] = None,
api_key: Optional[str] = None,
litellm_params: Optional[LiteLLM_Params] = None,
litellm_params: Optional[GenericLiteLLMParams] = None,
) -> Tuple[str, str, Optional[str], Optional[str]]:
"""
Returns the provider for a given model name - e.g. 'azure/chatgpt-v-2' -> 'azure'
Expand All @@ -178,20 +178,18 @@ def get_llm_provider( # noqa: PLR0915
)

if litellm.LiteLLMProxyChatConfig._should_use_litellm_proxy_by_default(
litellm_params=litellm_params
litellm_params=cast(Optional[LiteLLM_Params], litellm_params)
):
return litellm.LiteLLMProxyChatConfig.litellm_proxy_get_custom_llm_provider_info(
model=model, api_base=api_base, api_key=api_key
)

## IF LITELLM PARAMS GIVEN ##
if litellm_params:
assert (
custom_llm_provider is None and api_base is None and api_key is None
), "Either pass in litellm_params or the custom_llm_provider/api_base/api_key. Otherwise, these values will be overriden."
custom_llm_provider = litellm_params.custom_llm_provider
api_base = litellm_params.api_base
api_key = litellm_params.api_key
if custom_llm_provider is None and api_base is None and api_key is None:
custom_llm_provider = litellm_params.custom_llm_provider
api_base = litellm_params.api_base
api_key = litellm_params.api_key

dynamic_api_key = None
# check if llm provider provided
Expand Down Expand Up @@ -235,6 +233,7 @@ def get_llm_provider( # noqa: PLR0915
api_base=api_base,
api_key=api_key,
dynamic_api_key=dynamic_api_key,
litellm_params=litellm_params,
)

# check if llm provider part of model name
Expand All @@ -250,6 +249,7 @@ def get_llm_provider( # noqa: PLR0915
api_base=api_base,
api_key=api_key,
dynamic_api_key=dynamic_api_key,
litellm_params=litellm_params,
)
elif model.split("/", 1)[0] in litellm.provider_list:
custom_llm_provider = model.split("/", 1)[0]
Expand Down Expand Up @@ -570,6 +570,7 @@ def _get_openai_compatible_provider_info( # noqa: PLR0915
api_base: Optional[str],
api_key: Optional[str],
dynamic_api_key: Optional[str],
litellm_params: Optional[GenericLiteLLMParams] = None,
) -> Tuple[str, str, Optional[str], Optional[str]]:
"""
Returns:
Expand Down Expand Up @@ -637,7 +638,7 @@ def _get_openai_compatible_provider_info( # noqa: PLR0915
api_base,
dynamic_api_key,
) = litellm.BedrockMantleChatConfig()._get_openai_compatible_provider_info(
api_base, api_key
api_base, api_key, litellm_params=litellm_params
)
elif custom_llm_provider == "nvidia_nim":
# nvidia_nim is openai compatible, we just need to set this to custom_openai and have the api_base be https://api.endpoints.anyscale.com/v1
Expand Down
12 changes: 10 additions & 2 deletions litellm/llms/bedrock_mantle/chat/transformation.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,8 +12,10 @@

import litellm
from litellm._logging import verbose_logger
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.openai import AllMessageValues
from litellm.types.router import GenericLiteLLMParams

from ...openai_like.chat.transformation import OpenAILikeChatConfig

Expand All @@ -34,13 +36,19 @@ def get_config(cls):
return super().get_config()

def _get_openai_compatible_provider_info(
self, api_base: Optional[str], api_key: Optional[str]
self,
api_base: Optional[str],
api_key: Optional[str],
litellm_params: Optional[GenericLiteLLMParams] = None,
) -> Tuple[Optional[str], Optional[str]]:
region = (
get_secret_str("BEDROCK_MANTLE_REGION")
(litellm_params.aws_region_name if litellm_params else None)
or get_secret_str("BEDROCK_MANTLE_REGION")
or get_secret_str("AWS_REGION_NAME")
or get_secret_str("AWS_REGION")
or BEDROCK_MANTLE_DEFAULT_REGION
)
Comment thread
greptile-apps[bot] marked this conversation as resolved.
BaseAWSLLM._validate_aws_region_name(region)
api_base = (
api_base
or get_secret_str("BEDROCK_MANTLE_API_BASE")
Expand Down
60 changes: 59 additions & 1 deletion litellm/llms/bedrock_mantle/responses/transformation.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,7 +16,7 @@
"""

import re
from typing import Optional, Tuple
from typing import Any, Dict, List, Optional, Tuple

from botocore.exceptions import (
CredentialRetrievalError,
Expand All @@ -25,9 +25,11 @@
ProfileNotFound,
)

from litellm._logging import verbose_logger
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig
from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.openai import ResponsesAPIOptionalRequestParams
from litellm.types.router import GenericLiteLLMParams
from litellm.types.utils import LlmProviders

Expand All @@ -48,6 +50,11 @@
r"^https?://bedrock-mantle\.([^/.]+)\.api\.aws", re.IGNORECASE
)

# Per Bedrock Mantle Responses API validation errors.
_BEDROCK_MANTLE_SUPPORTED_RESPONSE_TOOL_TYPES = frozenset(
{"function", "mcp", "custom", "namespace", "tool_search"}
)


class BedrockMantleResponsesAPIConfig(OpenAIResponsesAPIConfig):
def __init__(
Expand All @@ -67,6 +74,7 @@ def custom_llm_provider(self) -> LlmProviders:
def _resolve_region(params: dict) -> str:
region = params.get("aws_region_name")
if region:
BaseAWSLLM._validate_aws_region_name(region)
return region
base = params.get("api_base") or get_secret_str("BEDROCK_MANTLE_API_BASE")
if base:
Expand Down Expand Up @@ -125,6 +133,56 @@ def supports_native_file_search(self) -> bool:
def supports_native_websocket(self) -> bool:
return False

@staticmethod
def _filter_unsupported_tools(tools: List[Any]) -> List[Any]:
"""Keep only tool types Mantle's Responses API accepts."""
kept: List[Any] = []
dropped_types: List[str] = []
for tool in tools:
if not isinstance(tool, dict):
kept.append(tool)
continue
tool_type = tool.get("type")
if tool_type in _BEDROCK_MANTLE_SUPPORTED_RESPONSE_TOOL_TYPES:
kept.append(tool)
else:
dropped_types.append(str(tool_type))

if dropped_types:
verbose_logger.warning(
"Bedrock Mantle Responses API: dropping unsupported tool type(s) "
"%s (supported: %s).",
sorted(set(dropped_types)),
sorted(_BEDROCK_MANTLE_SUPPORTED_RESPONSE_TOOL_TYPES),
)

return kept

def map_openai_params(
self,
response_api_optional_params: ResponsesAPIOptionalRequestParams,
model: str,
drop_params: bool,
) -> Dict:
params = super().map_openai_params(
response_api_optional_params=response_api_optional_params,
model=model,
drop_params=drop_params,
)

tools = params.get("tools")
if not tools:
return params

tools_list = tools if isinstance(tools, list) else [tools]
filtered = self._filter_unsupported_tools(tools_list)
if filtered:
params["tools"] = filtered
else:
params.pop("tools", None)

return params

def sign_request(
self,
headers: dict,
Expand Down
9 changes: 9 additions & 0 deletions litellm/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -86,6 +86,7 @@
get_audio_file_for_health_check,
)
from litellm.litellm_core_utils.completion_timeout import CompletionTimeout
from litellm.litellm_core_utils.get_litellm_params import OPTIONAL_KWARGS_KEYS
from litellm.litellm_core_utils.dd_tracing import tracer
from litellm.litellm_core_utils.get_provider_specific_headers import (
ProviderSpecificHeaderUtils,
Expand Down Expand Up @@ -1407,11 +1408,19 @@ def completion( # type: ignore # noqa: PLR0915
if deployment_id is not None: # azure llms
model = deployment_id
custom_llm_provider = "azure"
_supplemental_provider_params = {
k: kwargs[k] for k in OPTIONAL_KWARGS_KEYS if k in kwargs
}
model, custom_llm_provider, dynamic_api_key, api_base = get_llm_provider(
model=model,
custom_llm_provider=custom_llm_provider,
api_base=api_base,
api_key=api_key,
litellm_params=(
GenericLiteLLMParams(**_supplemental_provider_params)
if _supplemental_provider_params
else None
),
)

## RESPONSES API BRIDGE LOGIC ## - check early and normalize model name
Expand Down
26 changes: 6 additions & 20 deletions litellm/responses/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -673,16 +673,16 @@ def _resolve_model_provider_for_responses(
litellm_params: GenericLiteLLMParams,
local_vars: Dict[str, Any],
) -> tuple[str, Optional[str]]:
if custom_llm_provider is not None and not litellm_params.custom_llm_provider:
litellm_params.custom_llm_provider = custom_llm_provider
(
model,
custom_llm_provider,
dynamic_api_key,
dynamic_api_base,
) = litellm.get_llm_provider(
model=model,
custom_llm_provider=custom_llm_provider,
api_base=litellm_params.api_base,
api_key=litellm_params.api_key,
litellm_params=litellm_params,
)
local_vars["custom_llm_provider"] = custom_llm_provider
if dynamic_api_key is not None:
Expand Down Expand Up @@ -1972,27 +1972,13 @@ def compact_responses(
# get llm provider logic
litellm_params = GenericLiteLLMParams(**kwargs)

(
model,
custom_llm_provider,
dynamic_api_key,
dynamic_api_base,
) = litellm.get_llm_provider(
model, custom_llm_provider = _resolve_model_provider_for_responses(
model=model,
custom_llm_provider=custom_llm_provider,
api_base=litellm_params.api_base,
api_key=litellm_params.api_key,
litellm_params=litellm_params,
local_vars=local_vars,
)

# Update local_vars with detected provider (fixes #19782)
local_vars["custom_llm_provider"] = custom_llm_provider

# Use dynamic credentials from get_llm_provider (e.g., when use_litellm_proxy=True)
if dynamic_api_key is not None:
litellm_params.api_key = dynamic_api_key
if dynamic_api_base is not None:
litellm_params.api_base = dynamic_api_base

if custom_llm_provider is None:
raise ValueError("custom_llm_provider is required but passed as None")

Expand Down
Loading
Loading