mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix: address Bedrock guardrail review feedback
This commit is contained in:
parent
08cfef6791
commit
34e041e9f7
3 changed files with 150 additions and 85 deletions
|
|
@ -667,16 +667,13 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
|
||||
return merged_messages
|
||||
|
||||
# NOTE: Consider moving these helpers to CustomGuardrail when the filtering
|
||||
# logic becomes shared across providers.
|
||||
|
||||
#### CALL HOOKS - proxy only ####
|
||||
@staticmethod
|
||||
def _resolve_model_provider(model: str) -> str | None:
|
||||
try:
|
||||
_, custom_llm_provider, _, _ = litellm.get_llm_provider(model=model)
|
||||
return custom_llm_provider
|
||||
except Exception:
|
||||
except Exception: # noqa: BLE001 # provider resolution has a safe prefix fallback
|
||||
return model.partition("/")[0]
|
||||
|
||||
@staticmethod
|
||||
|
|
@ -709,9 +706,29 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
)
|
||||
or []
|
||||
)
|
||||
model_group_aliases: Final = getattr(llm_router, "model_group_alias", None)
|
||||
team_filtered_deployments: Final = (
|
||||
[
|
||||
deployment
|
||||
for deployment in listed_deployments
|
||||
if (
|
||||
(
|
||||
deployment.get("model_info")
|
||||
if isinstance(deployment, Mapping)
|
||||
else getattr(deployment, "model_info", None)
|
||||
)
|
||||
or {}
|
||||
).get("team_id")
|
||||
in (None, resolved_team_id)
|
||||
]
|
||||
if resolved_team_id is not None
|
||||
and isinstance(model_group_aliases, Mapping)
|
||||
and model in model_group_aliases
|
||||
else listed_deployments
|
||||
)
|
||||
model_id_deployment: Final = (
|
||||
llm_router.get_deployment(model_id=model)
|
||||
if not listed_deployments and llm_router.has_model_id(model) is True
|
||||
if not team_filtered_deployments and llm_router.has_model_id(model) is True
|
||||
else None
|
||||
)
|
||||
model_id_deployment_row: Final = (
|
||||
|
|
@ -720,7 +737,9 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
else model_id_deployment
|
||||
)
|
||||
candidate_deployments: Final = (
|
||||
[model_id_deployment_row] if model_id_deployment_row is not None else listed_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)
|
||||
|
|
@ -737,7 +756,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
deployments: Final = (
|
||||
filtered_deployments if isinstance(filtered_deployments, list) else candidate_deployments
|
||||
)
|
||||
except Exception:
|
||||
except Exception: # noqa: BLE001 # optional router state must not break guardrail auth
|
||||
return False
|
||||
active_deployments = []
|
||||
for deployment in deployments:
|
||||
|
|
@ -793,6 +812,10 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
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
|
||||
|
||||
model: Final[object | None] = request_data.get("model")
|
||||
model_provider: Final[str | None] = (
|
||||
BedrockGuardrail._resolve_model_provider(model) if isinstance(model, str) else None
|
||||
|
|
|
|||
|
|
@ -10,84 +10,6 @@ from litellm.proxy._types import UserAPIKeyAuth
|
|||
from litellm.caching import DualCache
|
||||
from unittest.mock import MagicMock, AsyncMock, patch
|
||||
|
||||
|
||||
def test_bedrock_guardrail_uses_active_metadata_bucket_for_team_id():
|
||||
router = MagicMock()
|
||||
router.get_model_list.return_value = [
|
||||
{"litellm_params": {"custom_llm_provider": "bedrock"}, "model_info": {}}
|
||||
]
|
||||
request_data = {
|
||||
"model": "team-alias",
|
||||
"metadata": {"user_api_key_team_id": "legacy-team"},
|
||||
"litellm_metadata": {"user_api_key_team_id": "active-team"},
|
||||
}
|
||||
|
||||
with patch("litellm.proxy.proxy_server.llm_router", router):
|
||||
assert BedrockGuardrail._router_allows_bedrock(request_data) is True
|
||||
|
||||
router.get_model_list.assert_called_once_with(model_name="team-alias", team_id="active-team")
|
||||
|
||||
|
||||
def test_bedrock_guardrail_uses_proxy_team_when_alternate_metadata_is_empty():
|
||||
router = MagicMock()
|
||||
router.get_model_list.return_value = [
|
||||
{"litellm_params": {"custom_llm_provider": "bedrock"}, "model_info": {}}
|
||||
]
|
||||
request_data = {
|
||||
"model": "team-alias",
|
||||
"metadata": {"user_api_key_team_id": "proxy-team"},
|
||||
"litellm_metadata": {},
|
||||
}
|
||||
|
||||
with patch("litellm.proxy.proxy_server.llm_router", router):
|
||||
assert BedrockGuardrail._router_allows_bedrock(request_data) is True
|
||||
|
||||
router.get_model_list.assert_called_once_with(model_name="team-alias", team_id="proxy-team")
|
||||
|
||||
|
||||
def test_bedrock_guardrail_resolves_router_model_id():
|
||||
router = MagicMock()
|
||||
router.get_model_list.return_value = []
|
||||
router.has_model_id.return_value = True
|
||||
deployment = MagicMock()
|
||||
deployment.model_dump.return_value = {
|
||||
"litellm_params": {"custom_llm_provider": "bedrock"},
|
||||
"model_info": {},
|
||||
}
|
||||
router.get_deployment.return_value = deployment
|
||||
|
||||
with patch("litellm.proxy.proxy_server.llm_router", router):
|
||||
assert BedrockGuardrail._router_allows_bedrock({"model": "deployment-id"}) is True
|
||||
|
||||
router.get_deployment.assert_called_once_with(model_id="deployment-id")
|
||||
deployment.model_dump.assert_called_once_with(exclude_none=True)
|
||||
|
||||
|
||||
def test_bedrock_guardrail_filters_access_group_deployments():
|
||||
router = MagicMock()
|
||||
router.get_model_list.return_value = [
|
||||
{"litellm_params": {"custom_llm_provider": "openai"}, "model_info": {}},
|
||||
{"litellm_params": {"custom_llm_provider": "bedrock"}, "model_info": {}},
|
||||
]
|
||||
router._filter_deployments_by_model_access_groups.return_value = [
|
||||
{"litellm_params": {"custom_llm_provider": "bedrock"}, "model_info": {}}
|
||||
]
|
||||
|
||||
with patch("litellm.proxy.proxy_server.llm_router", router):
|
||||
assert BedrockGuardrail._router_allows_bedrock({"model": "scoped-alias"}) is True
|
||||
|
||||
|
||||
def test_bedrock_guardrail_ignores_blocked_deployments():
|
||||
router = MagicMock()
|
||||
router.get_model_list.return_value = [
|
||||
{"litellm_params": {"custom_llm_provider": "openai"}, "model_info": {"blocked": True}},
|
||||
{"litellm_params": {"custom_llm_provider": "bedrock"}, "model_info": {}},
|
||||
]
|
||||
|
||||
with patch("litellm.proxy.proxy_server.llm_router", router):
|
||||
assert BedrockGuardrail._router_allows_bedrock({"model": "alias"}) is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bedrock_guardrails_pii_masking():
|
||||
# Create proper mock objects
|
||||
|
|
|
|||
|
|
@ -30,6 +30,126 @@ from litellm.types.proxy.guardrails.guardrail_hooks.bedrock_guardrails import (
|
|||
from litellm.types.utils import CallTypes, ModelResponse
|
||||
|
||||
|
||||
def test_bedrock_guardrail_uses_active_metadata_bucket_for_team_id():
|
||||
router = MagicMock()
|
||||
router.get_model_list.return_value = [
|
||||
{"litellm_params": {"custom_llm_provider": "bedrock"}, "model_info": {}}
|
||||
]
|
||||
request_data = {
|
||||
"model": "team-alias",
|
||||
"metadata": {"user_api_key_team_id": "legacy-team"},
|
||||
"litellm_metadata": {"user_api_key_team_id": "active-team"},
|
||||
}
|
||||
|
||||
with patch("litellm.proxy.proxy_server.llm_router", router):
|
||||
assert BedrockGuardrail._router_allows_bedrock(request_data) is True
|
||||
|
||||
router.get_model_list.assert_called_once_with(model_name="team-alias", team_id="active-team")
|
||||
|
||||
|
||||
def test_bedrock_guardrail_uses_proxy_team_when_alternate_metadata_is_empty():
|
||||
router = MagicMock()
|
||||
router.get_model_list.return_value = [
|
||||
{"litellm_params": {"custom_llm_provider": "bedrock"}, "model_info": {}}
|
||||
]
|
||||
request_data = {
|
||||
"model": "team-alias",
|
||||
"metadata": {"user_api_key_team_id": "proxy-team"},
|
||||
"litellm_metadata": {},
|
||||
}
|
||||
|
||||
with patch("litellm.proxy.proxy_server.llm_router", router):
|
||||
assert BedrockGuardrail._router_allows_bedrock(request_data) is True
|
||||
|
||||
router.get_model_list.assert_called_once_with(model_name="team-alias", team_id="proxy-team")
|
||||
|
||||
|
||||
def test_bedrock_guardrail_resolves_router_model_id():
|
||||
router = MagicMock()
|
||||
router.get_model_list.return_value = []
|
||||
router.has_model_id.return_value = True
|
||||
deployment = MagicMock()
|
||||
deployment.model_dump.return_value = {
|
||||
"litellm_params": {"custom_llm_provider": "bedrock"},
|
||||
"model_info": {},
|
||||
}
|
||||
router.get_deployment.return_value = deployment
|
||||
|
||||
with patch("litellm.proxy.proxy_server.llm_router", router):
|
||||
assert BedrockGuardrail._router_allows_bedrock({"model": "deployment-id"}) is True
|
||||
|
||||
router.get_deployment.assert_called_once_with(model_id="deployment-id")
|
||||
deployment.model_dump.assert_called_once_with(exclude_none=True)
|
||||
|
||||
|
||||
def test_bedrock_guardrail_filters_access_group_deployments():
|
||||
router = MagicMock()
|
||||
router.get_model_list.return_value = [
|
||||
{"litellm_params": {"custom_llm_provider": "openai"}, "model_info": {}},
|
||||
{"litellm_params": {"custom_llm_provider": "bedrock"}, "model_info": {}},
|
||||
]
|
||||
router._filter_deployments_by_model_access_groups.return_value = [
|
||||
{"litellm_params": {"custom_llm_provider": "bedrock"}, "model_info": {}}
|
||||
]
|
||||
|
||||
with patch("litellm.proxy.proxy_server.llm_router", router):
|
||||
assert BedrockGuardrail._router_allows_bedrock({"model": "scoped-alias"}) is True
|
||||
|
||||
|
||||
def test_bedrock_guardrail_ignores_blocked_deployments():
|
||||
router = MagicMock()
|
||||
router.get_model_list.return_value = [
|
||||
{"litellm_params": {"custom_llm_provider": "openai"}, "model_info": {"blocked": True}},
|
||||
{"litellm_params": {"custom_llm_provider": "bedrock"}, "model_info": {}},
|
||||
]
|
||||
|
||||
with patch("litellm.proxy.proxy_server.llm_router", router):
|
||||
assert BedrockGuardrail._router_allows_bedrock({"model": "alias"}) is True
|
||||
|
||||
|
||||
def test_bedrock_guardrail_filters_alias_deployments_by_team():
|
||||
router = MagicMock()
|
||||
router.model_group_alias = {"team-alias": "shared-group"}
|
||||
router.get_model_list.return_value = [
|
||||
{
|
||||
"litellm_params": {"custom_llm_provider": "openai"},
|
||||
"model_info": {"team_id": "other-team"},
|
||||
},
|
||||
{
|
||||
"litellm_params": {"custom_llm_provider": "bedrock"},
|
||||
"model_info": {"team_id": "active-team"},
|
||||
},
|
||||
]
|
||||
|
||||
with patch("litellm.proxy.proxy_server.llm_router", router):
|
||||
assert (
|
||||
BedrockGuardrail._router_allows_bedrock(
|
||||
{
|
||||
"model": "team-alias",
|
||||
"litellm_metadata": {"user_api_key_team_id": "active-team"},
|
||||
}
|
||||
)
|
||||
is True
|
||||
)
|
||||
|
||||
|
||||
def test_bedrock_guardrail_honors_explicit_provider_without_router(monkeypatch: pytest.MonkeyPatch):
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
monkeypatch.setattr(proxy_server, "llm_router", None)
|
||||
|
||||
assert (
|
||||
BedrockGuardrail._get_bedrock_api_key(
|
||||
{
|
||||
"model": "bedrock/anthropic.claude-3-haiku",
|
||||
"custom_llm_provider": "openai",
|
||||
"api_key": "openai-key",
|
||||
}
|
||||
)
|
||||
is None
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test__redact_pii_matches_function():
|
||||
"""Test the _redact_pii_matches function directly"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue