fix: align Bedrock tag and fallback routing

This commit is contained in:
aiedwardyi 2026-08-25 10:56:39 +09:00
parent 9fe37b1a3a
commit 46d9a2a67d
No known key found for this signature in database
2 changed files with 175 additions and 22 deletions

View file

@ -31,6 +31,7 @@ from litellm.constants import BEDROCK_APPLY_GUARDRAIL_CHUNK_BUDGET_CHARS
from litellm.exceptions import ModifyResponseException
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.litellm_core_utils.core_helpers import (
get_metadata_variable_name_from_kwargs,
redact_nested_match_and_regex_keys,
)
from litellm.litellm_core_utils.llm_cost_calc.guardrail_cost import bedrock_guardrail_cost
@ -53,7 +54,14 @@ from litellm.proxy.guardrails.anthropic_sse import (
is_raw_sse_stream,
model_response_text,
)
from litellm.router_strategy.tag_based_routing import _get_tags_from_request_kwargs, is_valid_deployment_tag
from litellm.router_strategy.tag_based_routing import (
_chain_tag_filtering_override,
_get_tags_from_request_kwargs,
_match_deployment,
_request_tags_after_router_consumption,
_split_tags,
_strip_routing_prefix,
)
from litellm.router_utils.common_utils import filter_team_based_models
from litellm.router_utils.cooldown_handlers import _get_cooldown_deployments
from litellm.secret_managers.main import get_secret_str
@ -685,16 +693,18 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
deployments: list[object],
request_data: Mapping[str, object],
) -> list[object]:
request_tags: Final = tuple(_get_tags_from_request_kwargs(request_data))
deployment_tag_filtering: Final = any(
BedrockGuardrail._router_deployment_field(deployment, "enable_tag_filtering") is True
for deployment in deployments
model: Final[object] = request_data.get("model")
chain_tag_filtering: Final[object] = (
_chain_tag_filtering_override(router, model, deployments)
if isinstance(model, str)
else None
)
tag_filtering_enabled: Final = (
request_data.get("enable_tag_filtering") is True
or getattr(router, "enable_tag_filtering", False) is True
or deployment_tag_filtering
effective_tag_filtering: Final = (
chain_tag_filtering
if isinstance(chain_tag_filtering, bool)
else getattr(router, "enable_tag_filtering", False)
)
tag_filtering_enabled: Final = request_data.get("enable_tag_filtering") is True or effective_tag_filtering is True
if not tag_filtering_enabled:
return deployments
@ -713,23 +723,37 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
else ()
)
if not request_tags:
metadata_name: Final = get_metadata_variable_name_from_kwargs(request_data)
metadata: Final[object] = request_data.get(metadata_name)
if not isinstance(metadata, Mapping):
default_deployments: Final = [
deployment for deployment in deployments if "default" in _deployment_tags(deployment)
]
return default_deployments or deployments
required_tags: Final = frozenset(tag[1:] for tag in request_tags if tag.startswith("&") and len(tag) > 1)
excluded_tags: Final = frozenset(tag[1:] for tag in request_tags if tag.startswith("!") and len(tag) > 1)
positive_tags: Final = tuple(tag for tag in request_tags if not tag.startswith(("&", "!")))
request_tags: Sequence[str] = tuple(_get_tags_from_request_kwargs(request_data))
if isinstance(model, str):
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)
required_tags, positive_tags, excluded_tags = _split_tags(rewritten_tags)
required_set: Final = frozenset(required_tags)
excluded_set: Final = frozenset(excluded_tags)
allowed_deployments: Final = [
deployment for deployment in deployments if not excluded_tags.intersection(_deployment_tags(deployment))
deployment for deployment in deployments if not excluded_set.intersection(_deployment_tags(deployment))
]
required_deployments: Final = [
deployment for deployment in allowed_deployments if required_tags.issubset(_deployment_tags(deployment))
candidate_deployments: Final = [
deployment for deployment in allowed_deployments if required_set.issubset(_deployment_tags(deployment))
]
if not positive_tags:
return required_deployments
user_agent: Final[object] = metadata.get("user_agent")
header_strings: Final = [f"User-Agent: {user_agent}"] if isinstance(user_agent, str) and user_agent else []
if not positive_tags and not header_strings:
default_deployments: Final = [
deployment for deployment in candidate_deployments if "default" in _deployment_tags(deployment)
]
return default_deployments or candidate_deployments
match_any: Final[bool] = (
getattr(router, "tag_filtering_match_any", True)
@ -738,13 +762,20 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
)
matched_deployments: Final = [
deployment
for deployment in required_deployments
if is_valid_deployment_tag(_deployment_tags(deployment), positive_tags, match_any)
for deployment in candidate_deployments
if isinstance(deployment, Mapping)
and _match_deployment(
deployment=deployment,
request_tags=positive_tags,
header_strings=header_strings,
match_any=match_any,
)
is not None
]
if matched_deployments:
return matched_deployments
fallback_default_deployments: Final = [
deployment for deployment in required_deployments if "default" in _deployment_tags(deployment)
deployment for deployment in candidate_deployments if "default" in _deployment_tags(deployment)
]
return fallback_default_deployments
@ -786,9 +817,24 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
if model_id_deployment is not None and hasattr(model_id_deployment, "model_dump")
else model_id_deployment
)
specific_deployment_rows: list[object] | None = None
deployment_names: Final[object] = getattr(llm_router, "deployment_names", None)
specific_lookup: Final[object] = getattr(llm_router, "_get_deployment_by_litellm_model", None)
if (
model_id_deployment_row is None
and isinstance(deployment_names, Sequence)
and not isinstance(deployment_names, (str, bytes))
and model in deployment_names
and callable(specific_lookup)
):
specific_result: Final = specific_lookup(model=model)
if isinstance(specific_result, list):
specific_deployment_rows = specific_result
raw_listed_deployments: Final = (
[]
if model_id_deployment_row is not None
else specific_deployment_rows
if specific_deployment_rows is not None
else (
llm_router.get_model_list(
model_name=model,
@ -803,7 +849,12 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
isinstance(concrete_model_names, (list, tuple, set, frozenset)) and model in concrete_model_names
)
is_model_alias: Final = isinstance(model_group_aliases, Mapping) and model in model_group_aliases
if model_id_deployment_row is None and not is_concrete_model and not is_model_alias:
if (
model_id_deployment_row is None
and specific_deployment_rows is None
and not is_concrete_model
and not is_model_alias
):
pattern_router: Final[object | None] = getattr(llm_router, "pattern_router", None)
get_pattern_deployments: Final[object | None] = getattr(
pattern_router, "get_deployments_by_pattern", None
@ -924,6 +975,23 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
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_fallback_lookup: Final[object] = getattr(llm_router, "_get_first_default_fallback", None)
default_fallback_model: Final[object] = (
default_fallback_lookup() if callable(default_fallback_lookup) else None
)
if isinstance(default_fallback_model, str) and default_fallback_model != model:
fallback_deployments: Final = (
llm_router.get_model_list(
model_name=default_fallback_model,
team_id=resolved_team_id,
)
or []
)
if isinstance(fallback_deployments, list) and fallback_deployments:
fallback_request_data: Final = dict(request_data)
fallback_request_data["model"] = default_fallback_model
return BedrockGuardrail._router_allows_bedrock(fallback_request_data)
default_deployment = getattr(llm_router, "default_deployment", None)
if default_deployment is not None:
default_params: object = (

View file

@ -249,6 +249,91 @@ def test_bedrock_guardrail_matches_request_tag_pool(monkeypatch: pytest.MonkeyPa
)
def test_bedrock_guardrail_matches_regex_tag_pool(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"}, "model_info": {}},
{
"litellm_params": {
"custom_llm_provider": "bedrock",
"tag_regex": [r"^User-Agent: claude-code/"],
},
"model_info": {},
},
]
monkeypatch.setattr(proxy_server, "llm_router", router)
assert (
BedrockGuardrail._router_allows_bedrock(
{"model": "shared-alias", "metadata": {"user_agent": "claude-code/1.0"}}
)
is True
)
def test_bedrock_guardrail_honors_false_chain_tag_filtering_override(monkeypatch: pytest.MonkeyPatch):
from litellm.proxy import proxy_server
router = MagicMock()
router.enable_tag_filtering = True
deployments = [
{
"litellm_params": {"custom_llm_provider": "openai", "tags": ["slow"]},
"model_info": {"enable_tag_filtering": False},
},
{"litellm_params": {"custom_llm_provider": "bedrock", "tags": ["fast"]}, "model_info": {}},
]
router.get_model_list.return_value = deployments
router._get_all_deployments.return_value = deployments
monkeypatch.setattr(proxy_server, "llm_router", router)
assert (
BedrockGuardrail._router_allows_bedrock(
{"model": "shared-alias", "metadata": {"tags": ["fast"]}}
)
is False
)
def test_bedrock_guardrail_resolves_specific_deployment_name(monkeypatch: pytest.MonkeyPatch):
from litellm.proxy import proxy_server
router = MagicMock()
router.deployment_names = ["bedrock-deployment"]
router._get_deployment_by_litellm_model.return_value = [
{"litellm_params": {"custom_llm_provider": "bedrock"}, "model_info": {"id": "bedrock-deployment"}}
]
router.get_model_list.return_value = [
{"litellm_params": {"custom_llm_provider": "openai"}, "model_info": {}}
]
monkeypatch.setattr(proxy_server, "llm_router", router)
assert BedrockGuardrail._router_allows_bedrock({"model": "bedrock-deployment"}) is True
router._get_deployment_by_litellm_model.assert_called_once_with(model="bedrock-deployment")
router.get_model_list.assert_not_called()
def test_bedrock_guardrail_follows_default_fallback_group(monkeypatch: pytest.MonkeyPatch):
from litellm.proxy import proxy_server
router = MagicMock()
router._get_first_default_fallback.return_value = "bedrock-fallback"
router.default_deployment = None
def _get_model_list(model_name: str, team_id: str | None = None) -> list[dict[str, object]]:
if model_name == "bedrock-fallback":
return [{"litellm_params": {"custom_llm_provider": "bedrock"}, "model_info": {}}]
return []
router.get_model_list.side_effect = _get_model_list
monkeypatch.setattr(proxy_server, "llm_router", router)
assert BedrockGuardrail._router_allows_bedrock({"model": "unknown-model"}) is True
def test_bedrock_guardrail_explicit_non_bedrock_provider_wins_alias(monkeypatch: pytest.MonkeyPatch):
from litellm.proxy import proxy_server