fix: align Bedrock tag fallback

This commit is contained in:
aiedwardyi 2026-08-26 16:18:59 +09:00
parent 6183bc25a1
commit 053a26a57a
No known key found for this signature in database
2 changed files with 120 additions and 14 deletions

View file

@ -57,10 +57,12 @@ from litellm.proxy.guardrails.anthropic_sse import (
from litellm.router_strategy.tag_based_routing import (
_chain_tag_filtering_override,
_get_tags_from_request_kwargs,
_inherited_constraint_sets,
_match_deployment,
_request_tags_after_router_consumption,
_split_tags,
_strip_routing_prefix,
_unknown_required_tag_hides_an_answer,
)
from litellm.router_utils.common_utils import filter_team_based_models, filter_web_search_deployments
from litellm.router_utils.cooldown_handlers import _get_cooldown_deployments
@ -744,10 +746,13 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
request_tags = _request_tags_after_router_consumption(metadata, model) or ()
routing_prefix: Final[object] = getattr(router, "tag_routing_prefix", "")
resolved_prefix: Final[str] = routing_prefix if isinstance(routing_prefix, str) else ""
rewritten_tags, _ = _strip_routing_prefix(request_tags, resolved_prefix)
rewritten_tags, routing_confirmed = _strip_routing_prefix(request_tags, resolved_prefix)
required_tags, positive_tags, excluded_tags = _split_tags(rewritten_tags)
required_set: Final = frozenset(required_tags)
excluded_set: Final = frozenset(excluded_tags)
inherited_required_set, inherited_excluded_set = _inherited_constraint_sets(
metadata.get("inherited_tags"), resolved_prefix
)
allowed_deployments: Final = [
deployment for deployment in deployments if not excluded_set.intersection(_deployment_tags(deployment))
]
@ -764,9 +769,41 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
has_positive_filter: Final = bool(positive_tags) or (
bool(header_strings) and has_regex_deployments and not required_set
)
def _fail_open_deployments() -> list[object]:
if _unknown_required_tag_hides_an_answer(
deployments,
excluded_set,
required_set,
routing_confirmed,
):
return []
if not any(
isinstance(deployment, Mapping)
and (deployment.get("model_info") or {}).get("allow_fail_open") is True
for deployment in deployments
):
return []
trusted_excluded: Final = (
frozenset() if inherited_excluded_set is None else inherited_excluded_set & excluded_set
)
trusted_required: Final = (
frozenset() if inherited_required_set is None else inherited_required_set & required_set
)
trusted_deployments: Final = [
deployment
for deployment in deployments
if not trusted_excluded.intersection(_deployment_tags(deployment))
and trusted_required.issubset(_deployment_tags(deployment))
]
default_deployments: Final = [
deployment for deployment in trusted_deployments if "default" in _deployment_tags(deployment)
]
return default_deployments or trusted_deployments
if not has_positive_filter:
if required_set or excluded_set:
return candidate_deployments
return candidate_deployments or _fail_open_deployments()
default_deployments: Final = [
deployment for deployment in candidate_deployments if "default" in _deployment_tags(deployment)
]
@ -791,12 +828,10 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
]
if matched_deployments:
return matched_deployments
if any(
isinstance(deployment, Mapping) and (deployment.get("model_info") or {}).get("allow_fail_open") is True
for deployment in allowed_deployments
):
return candidate_deployments or allowed_deployments
return []
default_deployments: Final = [
deployment for deployment in candidate_deployments if "default" in _deployment_tags(deployment)
]
return default_deployments or _fail_open_deployments()
@staticmethod
def _get_trusted_router_request_kwargs(request_data: Mapping[str, object]) -> dict[str, object]:
@ -1250,12 +1285,6 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
if provider is None:
return None
providers.append(provider)
router_allows_bedrock: Final = BedrockGuardrail._router_allows_bedrock(
request_data,
cooldown_deployments=[],
)
if router_allows_bedrock is not None:
return api_key if router_allows_bedrock else None
return (
api_key
if all(provider in ("bedrock", "bedrock_converse") for provider in providers)

View file

@ -349,6 +349,61 @@ def test_bedrock_guardrail_keeps_all_required_tag_matches(monkeypatch: pytest.Mo
)
def test_bedrock_guardrail_mirrors_router_fail_open_default(monkeypatch: pytest.MonkeyPatch):
from litellm.proxy import proxy_server
router = MagicMock()
router.enable_tag_filtering = True
router.get_model_list.return_value = [
{
"litellm_params": {"custom_llm_provider": "bedrock", "tags": ["caller-only"]},
"model_info": {"allow_fail_open": True},
},
{
"litellm_params": {"custom_llm_provider": "openai", "tags": ["default"]},
"model_info": {},
},
]
router._get_all_deployments.return_value = router.get_model_list.return_value
monkeypatch.setattr(proxy_server, "llm_router", router)
assert (
BedrockGuardrail._router_allows_bedrock(
{
"model": "shared-alias",
"metadata": {"tags": ["&caller-only", "unmatched"], "inherited_tags": []},
}
)
is False
)
def test_bedrock_guardrail_preserves_default_for_unknown_tag(monkeypatch: pytest.MonkeyPatch):
from litellm.proxy import proxy_server
router = MagicMock()
router.enable_tag_filtering = True
router.get_model_list.return_value = [
{
"litellm_params": {"custom_llm_provider": "openai", "tags": ["other"]},
"model_info": {},
},
{
"litellm_params": {"custom_llm_provider": "bedrock", "tags": ["default"]},
"model_info": {},
},
]
router._get_all_deployments.return_value = router.get_model_list.return_value
monkeypatch.setattr(proxy_server, "llm_router", router)
assert (
BedrockGuardrail._router_allows_bedrock(
{"model": "shared-alias", "metadata": {"tags": ["unknown"]}}
)
is True
)
@pytest.mark.asyncio
async def test_bedrock_guardrail_async_sync_strategy_uses_unfiltered_pool(monkeypatch: pytest.MonkeyPatch):
from litellm.proxy import proxy_server
@ -389,6 +444,28 @@ async def test_bedrock_guardrail_async_sync_strategy_uses_unfiltered_pool(monkey
router.async_get_healthy_deployments.assert_not_awaited()
@pytest.mark.asyncio
async def test_bedrock_guardrail_async_uses_callback_filtered_pool(monkeypatch: pytest.MonkeyPatch):
from litellm.proxy import proxy_server
router = MagicMock()
router.routing_strategy = "usage-based-routing-v2"
router.async_get_healthy_deployments = AsyncMock(
return_value=[{"litellm_params": {"custom_llm_provider": "bedrock"}, "model_info": {}}]
)
monkeypatch.setattr(proxy_server, "llm_router", router)
with patch.object(BedrockGuardrail, "_router_allows_bedrock", return_value=False) as router_allows:
assert (
await BedrockGuardrail._async_get_bedrock_api_key(
{"model": "shared-alias", "api_key": "bedrock-key"}
)
== "bedrock-key"
)
router_allows.assert_not_called()
def test_bedrock_guardrail_matches_regex_tag_pool(monkeypatch: pytest.MonkeyPatch):
from litellm.proxy import proxy_server