mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-07 08:26:10 +00:00
fix #19414 - [Bug]: get_model_info on Bedrock suffixed models doesn't return proper information (#19421)
* fixed get_model_info for bedrock models * fixed import
This commit is contained in:
parent
76433f9f04
commit
efbdf9d60d
2 changed files with 66 additions and 14 deletions
|
|
@ -1,3 +1,5 @@
|
|||
from __future__ import annotations
|
||||
|
||||
"""
|
||||
Common utilities used across bedrock chat/embedding/image generation
|
||||
"""
|
||||
|
|
@ -8,6 +10,7 @@ from typing import TYPE_CHECKING, Dict, List, Literal, Optional, Union
|
|||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.types.llms.bedrock import BedrockCreateBatchRequest
|
||||
from litellm.types.utils import ProviderSpecificModelInfo
|
||||
|
||||
import httpx
|
||||
|
||||
|
|
@ -34,7 +37,7 @@ _get_model_info = None
|
|||
def get_cached_model_info():
|
||||
"""
|
||||
Lazy import and cache get_model_info to avoid circular imports.
|
||||
|
||||
|
||||
This function is used by bedrock transformation classes that need get_model_info
|
||||
but cannot import it at module level due to circular import issues.
|
||||
The function is cached after first use to avoid performance impact.
|
||||
|
|
@ -42,6 +45,7 @@ def get_cached_model_info():
|
|||
global _get_model_info
|
||||
if _get_model_info is None:
|
||||
from litellm import get_model_info
|
||||
|
||||
_get_model_info = get_model_info
|
||||
return _get_model_info
|
||||
|
||||
|
|
@ -135,16 +139,16 @@ def add_custom_header(headers):
|
|||
def _get_bedrock_client_ssl_verify() -> Union[bool, str]:
|
||||
"""
|
||||
Get SSL verification setting for Bedrock client.
|
||||
|
||||
|
||||
Returns the SSL verification setting which can be:
|
||||
- True: Use default SSL verification
|
||||
- False: Disable SSL verification
|
||||
- str: Path to a custom CA bundle file
|
||||
"""
|
||||
from litellm.secret_managers.main import str_to_bool
|
||||
|
||||
|
||||
ssl_verify: Union[bool, str, None] = os.getenv("SSL_VERIFY", litellm.ssl_verify)
|
||||
|
||||
|
||||
# Convert string "False"/"True" to boolean
|
||||
if isinstance(ssl_verify, str):
|
||||
# Check if it's a file path
|
||||
|
|
@ -154,13 +158,13 @@ def _get_bedrock_client_ssl_verify() -> Union[bool, str]:
|
|||
ssl_verify_bool = str_to_bool(ssl_verify)
|
||||
if ssl_verify_bool is not None:
|
||||
ssl_verify = ssl_verify_bool
|
||||
|
||||
|
||||
# Check SSL_CERT_FILE environment variable for custom CA bundle
|
||||
if ssl_verify is True or ssl_verify == "True":
|
||||
ssl_cert_file = os.getenv("SSL_CERT_FILE")
|
||||
if ssl_cert_file and os.path.exists(ssl_cert_file):
|
||||
return ssl_cert_file
|
||||
|
||||
|
||||
return ssl_verify if ssl_verify is not None else True
|
||||
|
||||
|
||||
|
|
@ -287,7 +291,7 @@ def init_bedrock_client(
|
|||
"sts",
|
||||
aws_access_key_id=aws_access_key_id,
|
||||
aws_secret_access_key=aws_secret_access_key,
|
||||
verify=ssl_verify
|
||||
verify=ssl_verify,
|
||||
)
|
||||
|
||||
sts_response = sts_client.assume_role(
|
||||
|
|
@ -426,7 +430,7 @@ def strip_bedrock_routing_prefix(model: str) -> str:
|
|||
|
||||
|
||||
def strip_bedrock_throughput_suffix(model: str) -> str:
|
||||
""" Strip throughput tier suffixes from Bedrock model names. """
|
||||
"""Strip throughput tier suffixes from Bedrock model names."""
|
||||
import re
|
||||
|
||||
# Pattern matches model:version:throughput where throughput is like 51k, 18k, etc.
|
||||
|
|
@ -500,6 +504,22 @@ class BedrockModelInfo(BaseLLMModelInfo):
|
|||
) -> List[str]:
|
||||
return []
|
||||
|
||||
def get_provider_info(self, model: str) -> Optional[ProviderSpecificModelInfo]:
|
||||
"""
|
||||
Handles Bedrock throughput suffixes like ":28k", ":51k".
|
||||
"""
|
||||
import re
|
||||
|
||||
overrides: ProviderSpecificModelInfo = {}
|
||||
|
||||
# Parse context window suffix (e.g., :28k, :51k)
|
||||
match = re.search(r":(\d+)k$", model)
|
||||
if match:
|
||||
throughput_value = int(match.group(1)) * 1000
|
||||
overrides["max_input_tokens"] = throughput_value
|
||||
|
||||
return overrides if overrides else None
|
||||
|
||||
def get_token_counter(self) -> Optional[BaseTokenCounter]:
|
||||
"""
|
||||
Factory method to create a Bedrock token counter.
|
||||
|
|
@ -532,12 +552,29 @@ class BedrockModelInfo(BaseLLMModelInfo):
|
|||
@staticmethod
|
||||
def get_bedrock_route(
|
||||
model: str,
|
||||
) -> Literal["converse", "invoke", "converse_like", "agent", "agentcore", "async_invoke", "openai"]:
|
||||
) -> Literal[
|
||||
"converse",
|
||||
"invoke",
|
||||
"converse_like",
|
||||
"agent",
|
||||
"agentcore",
|
||||
"async_invoke",
|
||||
"openai",
|
||||
]:
|
||||
"""
|
||||
Get the bedrock route for the given model.
|
||||
"""
|
||||
route_mappings: Dict[
|
||||
str, Literal["invoke", "converse_like", "converse", "agent", "agentcore", "async_invoke", "openai"]
|
||||
str,
|
||||
Literal[
|
||||
"invoke",
|
||||
"converse_like",
|
||||
"converse",
|
||||
"agent",
|
||||
"agentcore",
|
||||
"async_invoke",
|
||||
"openai",
|
||||
],
|
||||
] = {
|
||||
"invoke/": "invoke",
|
||||
"converse_like/": "converse_like",
|
||||
|
|
@ -645,10 +682,10 @@ class BedrockModelInfo(BaseLLMModelInfo):
|
|||
def get_bedrock_chat_config(model: str):
|
||||
"""
|
||||
Helper function to get the appropriate Bedrock chat config based on model and route.
|
||||
|
||||
|
||||
Args:
|
||||
model: The model name/identifier
|
||||
|
||||
|
||||
Returns:
|
||||
The appropriate Bedrock config class instance
|
||||
"""
|
||||
|
|
@ -667,11 +704,13 @@ def get_bedrock_chat_config(model: str):
|
|||
from litellm.llms.bedrock.chat.invoke_agent.transformation import (
|
||||
AmazonInvokeAgentConfig,
|
||||
)
|
||||
|
||||
return AmazonInvokeAgentConfig()
|
||||
elif bedrock_route == "agentcore":
|
||||
from litellm.llms.bedrock.chat.agentcore.transformation import (
|
||||
AmazonAgentCoreConfig,
|
||||
)
|
||||
|
||||
return AmazonAgentCoreConfig()
|
||||
|
||||
# Handle provider-specific configs
|
||||
|
|
|
|||
|
|
@ -4647,7 +4647,9 @@ def add_provider_specific_params_to_optional_params(
|
|||
else:
|
||||
for k in passed_params.keys():
|
||||
if k not in openai_params and passed_params[k] is not None:
|
||||
if _should_drop_param(k=k, additional_drop_params=additional_drop_params):
|
||||
if _should_drop_param(
|
||||
k=k, additional_drop_params=additional_drop_params
|
||||
):
|
||||
continue
|
||||
optional_params[k] = passed_params[k]
|
||||
return optional_params
|
||||
|
|
@ -5775,6 +5777,14 @@ def get_model_info(model: str, custom_llm_provider: Optional[str] = None) -> Mod
|
|||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
|
||||
provider_info = get_provider_info(
|
||||
model=model, custom_llm_provider=custom_llm_provider
|
||||
)
|
||||
if provider_info:
|
||||
for key, value in provider_info.items():
|
||||
if value is not None:
|
||||
_model_info[key] = value # type: ignore
|
||||
|
||||
verbose_logger.debug(f"model_info: {_model_info}")
|
||||
|
||||
returned_model_info = ModelInfo(
|
||||
|
|
@ -8151,7 +8161,10 @@ class ProviderConfigManager:
|
|||
# Note: GPT models (gpt-3.5, gpt-4, gpt-5, etc.) support temperature parameter
|
||||
# O-series models (o1, o3) do not contain "gpt" and have different parameter restrictions
|
||||
is_gpt_model = model and "gpt" in model.lower()
|
||||
is_o_series = model and ("o_series" in model.lower() or (supports_reasoning(model) and not is_gpt_model))
|
||||
is_o_series = model and (
|
||||
"o_series" in model.lower()
|
||||
or (supports_reasoning(model) and not is_gpt_model)
|
||||
)
|
||||
|
||||
is_o_series = model and (
|
||||
"o_series" in model.lower()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue