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:
Lucky-Lodhi2004 2026-01-21 09:15:55 +05:30 committed by GitHub
parent 76433f9f04
commit efbdf9d60d
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 66 additions and 14 deletions

View file

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

View file

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