mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-02 02:11:58 +00:00
test(bedrock): accept regional aliases that inherit Converse routing (#43785)
* test(bedrock): accept regional aliases that inherit Converse routing * test(bedrock): cover regional alias metadata independently of catalog
This commit is contained in:
parent
ba6d6d1a95
commit
d02ff435bf
1 changed files with 38 additions and 8 deletions
|
|
@ -1,14 +1,18 @@
|
|||
# What is this?
|
||||
## Unit testing for the 'get_model_info()' function
|
||||
import os
|
||||
import re
|
||||
from collections.abc import Collection, Mapping
|
||||
|
||||
|
||||
from typing import List, Dict, Any
|
||||
from typing import List, Dict, Any, Final, Literal
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm import get_model_info
|
||||
from litellm.llms.bedrock.common_utils import BedrockModelInfo
|
||||
from litellm.types.utils import ModelInfoBase
|
||||
from litellm.utils import _invalidate_model_cost_lowercase_map
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
|
|
@ -116,26 +120,31 @@ def test_get_model_info_ft_model_with_provider_prefix():
|
|||
|
||||
|
||||
def _enforce_bedrock_converse_models(
|
||||
model_cost: List[Dict[str, Any]], whitelist_models: List[str]
|
||||
):
|
||||
model_cost: Mapping[str, ModelInfoBase], whitelist_models: Collection[str]
|
||||
) -> None:
|
||||
"""
|
||||
Assert all new bedrock chat models are added as `bedrock_converse` unless explicitly whitelisted.
|
||||
Assert unlisted Bedrock chat models declare or inherit Converse routing.
|
||||
"""
|
||||
# Check for unwhitelisted models
|
||||
for model, info in litellm.model_cost.items():
|
||||
for model, info in model_cost.items():
|
||||
if (
|
||||
info["litellm_provider"] == "bedrock"
|
||||
and info["mode"] == "chat"
|
||||
and model not in whitelist_models
|
||||
and not (
|
||||
(base_model := BedrockModelInfo.get_base_model(model)) != model
|
||||
and model_cost.get(base_model, {}).get("litellm_provider") == "bedrock_converse"
|
||||
and BedrockModelInfo.get_bedrock_route(model) == "converse"
|
||||
)
|
||||
):
|
||||
raise AssertionError(
|
||||
f"New bedrock chat model detected: {model}. Please set `litellm_provider='bedrock_converse'` for this model."
|
||||
f"Unlisted Bedrock chat model does not route to Converse: {model}"
|
||||
)
|
||||
|
||||
|
||||
def test_model_info_bedrock_converse(monkeypatch):
|
||||
"""
|
||||
Assert all new bedrock chat models are added as `bedrock_converse` unless explicitly whitelisted.
|
||||
Assert unlisted Bedrock chat models declare or inherit Converse routing.
|
||||
|
||||
This ensures they are automatically routed to the converse endpoint.
|
||||
"""
|
||||
|
|
@ -173,7 +182,7 @@ def test_model_info_bedrock_converse_enforcement(monkeypatch):
|
|||
whitelist_models = [line.strip() for line in file.readlines()]
|
||||
|
||||
# Check for unwhitelisted models
|
||||
with pytest.raises(AssertionError):
|
||||
with pytest.raises(AssertionError, match=r"fake\.bedrock-chat-model"):
|
||||
_enforce_bedrock_converse_models(
|
||||
model_cost=litellm.model_cost, whitelist_models=whitelist_models
|
||||
)
|
||||
|
|
@ -181,6 +190,27 @@ def test_model_info_bedrock_converse_enforcement(monkeypatch):
|
|||
pytest.skip("whitelisted_bedrock_models.txt not found")
|
||||
|
||||
|
||||
@pytest.mark.parametrize("region", ("us-gov-east-1", "us-gov-west-1"))
|
||||
@pytest.mark.parametrize("base_provider", ("bedrock_converse", "bedrock"))
|
||||
def test_regional_bedrock_alias_requires_canonical_converse_metadata(
|
||||
region: str, base_provider: Literal["bedrock_converse", "bedrock"]
|
||||
) -> None:
|
||||
base_model: Final = next(
|
||||
model for model in sorted(litellm.bedrock_converse_models) if BedrockModelInfo.get_base_model(model) == model
|
||||
)
|
||||
model: Final = f"bedrock/{region}/{base_model}"
|
||||
model_cost: Final[Mapping[str, ModelInfoBase]] = {
|
||||
model: {"litellm_provider": "bedrock", "mode": "chat"},
|
||||
base_model: {"litellm_provider": base_provider, "mode": "chat"},
|
||||
}
|
||||
assert BedrockModelInfo.get_bedrock_route(model) == "converse"
|
||||
if base_provider == "bedrock":
|
||||
with pytest.raises(AssertionError, match=re.escape(model)):
|
||||
_enforce_bedrock_converse_models(model_cost, ())
|
||||
return
|
||||
_enforce_bedrock_converse_models(model_cost, ())
|
||||
|
||||
|
||||
def test_get_model_info_custom_provider():
|
||||
# Custom provider example copied from https://docs.litellm.ai/docs/providers/custom_llm_server:
|
||||
import litellm
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue