mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix: close second Bedrock review round
This commit is contained in:
parent
34e041e9f7
commit
aa3c21eab9
2 changed files with 87 additions and 4 deletions
|
|
@ -737,9 +737,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
else model_id_deployment
|
||||
)
|
||||
candidate_deployments: Final = (
|
||||
[model_id_deployment_row]
|
||||
if model_id_deployment_row is not None
|
||||
else team_filtered_deployments
|
||||
[model_id_deployment_row] if model_id_deployment_row is not None else team_filtered_deployments
|
||||
)
|
||||
|
||||
filter_deployments: Final = getattr(llm_router, "_filter_deployments_by_model_access_groups", None)
|
||||
|
|
@ -771,6 +769,37 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
if blocked is not True:
|
||||
active_deployments.append(deployment)
|
||||
if not active_deployments:
|
||||
router_settings = getattr(llm_router, "router_general_settings", None)
|
||||
if getattr(router_settings, "pass_through_all_models", False) is True:
|
||||
provider = request_data.get("custom_llm_provider")
|
||||
if not isinstance(provider, str):
|
||||
provider = BedrockGuardrail._resolve_model_provider(model) if isinstance(model, str) else None
|
||||
return provider in ("bedrock", "bedrock_converse") if isinstance(provider, str) else None
|
||||
|
||||
default_deployment = getattr(llm_router, "default_deployment", None)
|
||||
if default_deployment is not None:
|
||||
default_params: object = (
|
||||
default_deployment.get("litellm_params")
|
||||
if isinstance(default_deployment, Mapping)
|
||||
else getattr(default_deployment, "litellm_params", None)
|
||||
)
|
||||
provider: object = (
|
||||
default_params.get("custom_llm_provider")
|
||||
if isinstance(default_params, Mapping)
|
||||
else getattr(default_params, "custom_llm_provider", None)
|
||||
)
|
||||
if not isinstance(provider, str):
|
||||
default_model: object = (
|
||||
default_params.get("model")
|
||||
if isinstance(default_params, Mapping)
|
||||
else getattr(default_params, "model", None)
|
||||
)
|
||||
provider = (
|
||||
BedrockGuardrail._resolve_model_provider(default_model)
|
||||
if isinstance(default_model, str)
|
||||
else None
|
||||
)
|
||||
return provider in ("bedrock", "bedrock_converse") if isinstance(provider, str) else None
|
||||
return False
|
||||
|
||||
providers: list[str] = []
|
||||
|
|
@ -808,11 +837,14 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
if not isinstance(api_key, str):
|
||||
return None
|
||||
|
||||
explicit_provider: Final[object | None] = request_data.get("custom_llm_provider")
|
||||
if isinstance(explicit_provider, str) and explicit_provider not in ("bedrock", "bedrock_converse"):
|
||||
return None
|
||||
|
||||
router_allows_bedrock: Final = BedrockGuardrail._router_allows_bedrock(request_data)
|
||||
if router_allows_bedrock is not None:
|
||||
return api_key if router_allows_bedrock else None
|
||||
|
||||
explicit_provider: Final[object | None] = request_data.get("custom_llm_provider")
|
||||
if isinstance(explicit_provider, str):
|
||||
return api_key if explicit_provider in ("bedrock", "bedrock_converse") else None
|
||||
|
||||
|
|
|
|||
|
|
@ -82,6 +82,57 @@ def test_bedrock_guardrail_resolves_router_model_id():
|
|||
deployment.model_dump.assert_called_once_with(exclude_none=True)
|
||||
|
||||
|
||||
def test_bedrock_guardrail_accepts_pass_through_bedrock_provider(monkeypatch: pytest.MonkeyPatch):
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
router = MagicMock()
|
||||
router.get_model_list.return_value = []
|
||||
router.router_general_settings.pass_through_all_models = True
|
||||
router.default_deployment = None
|
||||
monkeypatch.setattr(proxy_server, "llm_router", router)
|
||||
|
||||
request_data = {
|
||||
"model": "amazon.nova-lite-v1:0",
|
||||
"custom_llm_provider": "bedrock",
|
||||
"api_key": "bedrock-key",
|
||||
}
|
||||
assert BedrockGuardrail._router_allows_bedrock(request_data) is True
|
||||
assert BedrockGuardrail._get_bedrock_api_key(request_data) == "bedrock-key"
|
||||
|
||||
|
||||
def test_bedrock_guardrail_accepts_bedrock_default_deployment(monkeypatch: pytest.MonkeyPatch):
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
router = MagicMock()
|
||||
router.get_model_list.return_value = []
|
||||
router.router_general_settings.pass_through_all_models = False
|
||||
router.default_deployment = {"litellm_params": {"custom_llm_provider": "bedrock"}}
|
||||
monkeypatch.setattr(proxy_server, "llm_router", router)
|
||||
|
||||
assert BedrockGuardrail._router_allows_bedrock({"model": "unlisted-model"}) is True
|
||||
|
||||
|
||||
def test_bedrock_guardrail_explicit_non_bedrock_provider_wins_alias(monkeypatch: pytest.MonkeyPatch):
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
router = MagicMock()
|
||||
router.get_model_list.return_value = [
|
||||
{"litellm_params": {"custom_llm_provider": "bedrock"}, "model_info": {}}
|
||||
]
|
||||
monkeypatch.setattr(proxy_server, "llm_router", router)
|
||||
|
||||
assert (
|
||||
BedrockGuardrail._get_bedrock_api_key(
|
||||
{
|
||||
"model": "bedrock-alias",
|
||||
"custom_llm_provider": "openai",
|
||||
"api_key": "openai-key",
|
||||
}
|
||||
)
|
||||
is None
|
||||
)
|
||||
|
||||
|
||||
def test_bedrock_guardrail_filters_access_group_deployments():
|
||||
router = MagicMock()
|
||||
router.get_model_list.return_value = [
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue