fix: align Bedrock guardrail eligibility

This commit is contained in:
aiedwardyi 2026-08-26 11:09:15 +09:00
parent c58daa5d11
commit 250f30a211
No known key found for this signature in database

View file

@ -119,6 +119,7 @@ _BEDROCK_INVOKE_GUARDRAIL_CHECKS_PATH: Final = "/guardrail-checks/invoke"
# more text blocks is split across multiple messages so ALL content is scanned --
# never truncated (truncation would let a user hide content past the limit).
_BEDROCK_CHECKS_MAX_CONTENT_BLOCKS: Final = 10
_ROUTER_COOLDOWNS_UNSET: Final = object()
_BEDROCK_CHECKS_KNOWN_KEYS: Final = frozenset({"contentFilter", "promptAttack", "sensitiveInformation"})
# Keys in a sensitiveInformation result that pinpoint the PII location. They are
# stripped before the response is handed to standard logging / telemetry so the
@ -692,8 +693,9 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
router: object,
deployments: list[object],
request_data: Mapping[str, object],
model: str | None = None,
) -> list[object]:
model: Final[object] = request_data.get("model")
model: Final[object] = model or request_data.get("model")
chain_tag_filtering: Final[object] = (
_chain_tag_filtering_override(router, model, deployments) if isinstance(model, str) else None
)
@ -784,7 +786,15 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
fallback_default_deployments: Final = [
deployment for deployment in candidate_deployments if "default" in _deployment_tags(deployment)
]
return fallback_default_deployments
if fallback_default_deployments:
return fallback_default_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 []
@staticmethod
def _router_deployment_field(deployment: object, field: str) -> object | None:
@ -794,7 +804,11 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
return model_info.get(field) if isinstance(model_info, Mapping) else getattr(model_info, field, None)
@staticmethod
def _router_allows_bedrock(request_data: Mapping[str, object]) -> bool | None:
def _router_allows_bedrock(
request_data: Mapping[str, object],
*,
cooldown_deployments: Sequence[str] | None | object = _ROUTER_COOLDOWNS_UNSET,
) -> bool | None:
model: Final[object | None] = request_data.get("model")
if not isinstance(model, str):
return False
@ -817,11 +831,19 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
try:
resolved_team_id: Final = team_id if isinstance(team_id, str) else None
router_kwargs: Final = dict(request_data)
effective_model: str = model
for metadata in (request_data.get("litellm_metadata"), request_data.get("metadata")):
if not isinstance(metadata, Mapping):
continue
routing_decision: Final[object] = metadata.get("routing_decision")
if isinstance(routing_decision, Mapping) and isinstance(routing_decision.get("routed_model"), str):
effective_model = routing_decision["routed_model"]
break
common_result: tuple[object, object] | None = None
common_lookup: Final[object] = getattr(llm_router, "_common_checks_available_deployment", None)
if callable(common_lookup):
try:
raw_common_result: Final = common_lookup(model=model, request_kwargs=router_kwargs)
raw_common_result: Final = common_lookup(model=effective_model, request_kwargs=router_kwargs)
except Exception: # noqa: BLE001 # fall back for lightweight router test doubles
raw_common_result = None
if (
@ -830,6 +852,8 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
and isinstance(raw_common_result[1], (Mapping, list))
):
common_result = raw_common_result
if isinstance(raw_common_result[0], str):
effective_model = raw_common_result[0]
if common_result is not None:
raw_deployments: Final[object] = common_result[1]
@ -844,7 +868,9 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
router_matched: Final = bool(candidate_deployments)
else:
model_id_deployment: Final = (
llm_router.get_deployment(model_id=model) if llm_router.has_model_id(model) is True else None
llm_router.get_deployment(model_id=effective_model)
if llm_router.has_model_id(effective_model) is True
else None
)
model_id_deployment_row: Final = (
model_id_deployment.model_dump(exclude_none=True)
@ -857,14 +883,17 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
model_group_aliases: Final = getattr(llm_router, "model_group_alias", None)
concrete_model_names: Final[object | None] = getattr(llm_router, "model_names", None)
is_concrete_model: Final = (
isinstance(concrete_model_names, (list, tuple, set, frozenset)) and model in concrete_model_names
isinstance(concrete_model_names, (list, tuple, set, frozenset))
and effective_model in concrete_model_names
)
is_model_alias: Final = (
isinstance(model_group_aliases, Mapping) and effective_model in model_group_aliases
)
is_model_alias: Final = isinstance(model_group_aliases, Mapping) and model in model_group_aliases
raw_listed_deployments = (
[]
if model_id_deployment_row is not None
else (
llm_router.get_model_list(model_name=model, team_id=resolved_team_id) or []
llm_router.get_model_list(model_name=effective_model, team_id=resolved_team_id) or []
if is_concrete_model or is_model_alias
else []
)
@ -875,7 +904,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
pattern_router, "get_deployments_by_pattern", None
)
global_pattern_deployments: Final = (
get_pattern_deployments(model=model) if callable(get_pattern_deployments) else None
get_pattern_deployments(model=effective_model) if callable(get_pattern_deployments) else None
)
team_pattern_router: Final[object | None] = (
getattr(llm_router, "team_pattern_routers", {}).get(resolved_team_id)
@ -887,7 +916,9 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
team_pattern_router, "get_deployments_by_pattern", None
)
team_pattern_deployments: Final = (
get_team_pattern_deployments(model=model) if callable(get_team_pattern_deployments) else None
get_team_pattern_deployments(model=effective_model)
if callable(get_team_pattern_deployments)
else None
)
if isinstance(global_pattern_deployments, list) and global_pattern_deployments:
raw_listed_deployments = global_pattern_deployments
@ -900,15 +931,15 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
elif (
isinstance(deployment_names, Sequence)
and not isinstance(deployment_names, (str, bytes))
and model in deployment_names
and effective_model in deployment_names
and callable(specific_lookup)
):
specific_result: Final = specific_lookup(model=model)
specific_result: Final = specific_lookup(model=effective_model)
specific_deployment_rows = specific_result if isinstance(specific_result, list) else []
raw_listed_deployments = specific_deployment_rows
else:
raw_listed_deployments = (
llm_router.get_model_list(model_name=model, team_id=resolved_team_id) or []
llm_router.get_model_list(model_name=effective_model, team_id=resolved_team_id) or []
)
candidate_deployments = (
[model_id_deployment_row]
@ -931,7 +962,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
filter_deployments: Final = getattr(llm_router, "_filter_deployments_by_model_access_groups", None)
filtered_deployments: Final = (
filter_deployments(
model=model,
model=effective_model,
healthy_deployments=team_filtered_deployments,
request_kwargs=dict(request_data),
request_team_id=resolved_team_id,
@ -960,20 +991,43 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
if isinstance(health_filtered_deployments, list)
else access_filtered_deployments
)
pre_call_filter: Final[object] = getattr(llm_router, "_pre_call_checks", None)
request_messages: Final[object] = request_data.get("messages")
request_input: Final[object] = request_data.get("input")
if (
model_id_deployment_row is None
and getattr(llm_router, "enable_pre_call_checks", False) is True
and (isinstance(request_messages, list) or isinstance(request_input, (str, list)))
and callable(pre_call_filter)
):
pre_call_deployments: Final = pre_call_filter(
model=effective_model,
healthy_deployments=healthy_deployments,
messages=request_messages if isinstance(request_messages, list) else None,
input=request_input if isinstance(request_input, (str, list)) else None,
request_kwargs=router_kwargs,
)
if isinstance(pre_call_deployments, list):
healthy_deployments = pre_call_deployments
cooldown_cache: Final[object | None] = getattr(llm_router, "cooldown_cache", None)
cooldown_lookup: Final[object | None] = getattr(cooldown_cache, "get_active_cooldowns", None)
cooldown_deployments: Final = (
resolved_cooldown_deployments: Final = (
_get_cooldown_deployments(
litellm_router_instance=llm_router,
parent_otel_span=None,
)
if (common_result is None or model_id_deployment_row is None)
if cooldown_deployments is _ROUTER_COOLDOWNS_UNSET
and (common_result is None or model_id_deployment_row is None)
and callable(cooldown_lookup)
and callable(getattr(llm_router, "get_model_ids", None))
else cooldown_deployments
if cooldown_deployments is not _ROUTER_COOLDOWNS_UNSET
else []
)
cooldown_ids: Final[frozenset[str]] = frozenset(
deployment_id for deployment_id in cooldown_deployments if isinstance(deployment_id, str)
deployment_id
for deployment_id in (resolved_cooldown_deployments or [])
if isinstance(deployment_id, str)
)
if common_result is not None and model_id_deployment_row is not None:
deployments = candidate_deployments
@ -992,6 +1046,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
router=llm_router,
deployments=unblocked_deployments,
request_data=request_data,
model=effective_model,
)
if common_result is not None and model_id_deployment_row is None:
web_search_deployments: Final = filter_web_search_deployments(
@ -1007,14 +1062,14 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
)
if isinstance(plugin_deployments, list):
deployments = plugin_deployments
deployments = litellm.utils._get_order_filtered_deployments(
deployments,
target_order=router_kwargs.pop("_target_order", None),
)
deployments = litellm.utils._get_excluded_filtered_deployments(
deployments,
excluded_deployment_ids=router_kwargs.pop("_excluded_deployment_ids", None),
)
deployments = litellm.utils._get_order_filtered_deployments(
deployments,
target_order=router_kwargs.pop("_target_order", None),
)
except Exception: # noqa: BLE001 # optional router state must not break guardrail auth
return False
if not deployments:
@ -1024,7 +1079,11 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
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
provider = (
BedrockGuardrail._resolve_model_provider(effective_model)
if isinstance(effective_model, str)
else None
)
return provider in ("bedrock", "bedrock_converse") if isinstance(provider, str) else None
default_deployment = getattr(llm_router, "default_deployment", None)
@ -1056,7 +1115,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
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:
if isinstance(default_fallback_model, str) and default_fallback_model != effective_model:
fallback_deployments: Final = (
llm_router.get_model_list(
model_name=default_fallback_model,
@ -1097,6 +1156,107 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
providers.append(provider)
return bool(providers) and all(provider in ("bedrock", "bedrock_converse") for provider in providers)
@staticmethod
async def _async_get_bedrock_api_key(request_data: Mapping[str, object] | None) -> str | None:
if not request_data:
return None
api_key: Final[object | None] = request_data.get("api_key")
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
try:
from litellm.proxy.proxy_server import llm_router
except ImportError:
llm_router = None
if llm_router is not None:
async_lookup: Final[object] = getattr(llm_router, "async_get_healthy_deployments", None)
if callable(async_lookup):
model: Final[object | None] = request_data.get("model")
if isinstance(model, str):
effective_model = model
for metadata in (request_data.get("litellm_metadata"), request_data.get("metadata")):
if not isinstance(metadata, Mapping):
continue
routing_decision: Final[object] = metadata.get("routing_decision")
if isinstance(routing_decision, Mapping) and isinstance(
routing_decision.get("routed_model"), str
):
effective_model = routing_decision["routed_model"]
break
try:
healthy_deployments: Final = await async_lookup(
model=effective_model,
request_kwargs=dict(request_data),
messages=request_data.get("messages")
if isinstance(request_data.get("messages"), list)
else None,
input=request_data.get("input")
if isinstance(request_data.get("input"), (str, list))
else None,
specific_deployment=request_data.get("specific_deployment") is True,
)
deployments: Final[list[object]] = (
[healthy_deployments]
if isinstance(healthy_deployments, Mapping)
else healthy_deployments
if isinstance(healthy_deployments, list)
else []
)
if deployments:
providers: list[str] = []
for deployment in deployments:
params: object = (
deployment.get("litellm_params")
if isinstance(deployment, Mapping)
else getattr(deployment, "litellm_params", None)
)
provider: object = (
params.get("custom_llm_provider")
if isinstance(params, Mapping)
else getattr(params, "custom_llm_provider", None)
)
if not isinstance(provider, str):
deployment_model: object = (
params.get("model")
if isinstance(params, Mapping)
else getattr(params, "model", None)
)
provider = (
BedrockGuardrail._resolve_model_provider(deployment_model)
if isinstance(deployment_model, str)
else None
)
if not isinstance(provider, str):
return None
providers.append(provider)
return api_key if all(
provider in ("bedrock", "bedrock_converse") for provider in providers
) else None
except Exception: # noqa: BLE001 # fall back to the sync compatibility path
pass
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
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
)
return api_key if model_provider in ("bedrock", "bedrock_converse") else None
@staticmethod
def _get_bedrock_api_key(request_data: Mapping[str, object] | None) -> str | None:
if not request_data:
@ -1295,7 +1455,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
bedrock_request_data: Final[dict] = dict(
self.convert_to_bedrock_format(source=source, messages=messages, response=response)
)
api_key: Final = self._get_bedrock_api_key(request_data)
api_key: Final = await self._async_get_bedrock_api_key(request_data)
if request_data:
dynamic_request_body_params = self.get_guardrail_dynamic_request_body_params(request_data=request_data)
bedrock_request_data.update(
@ -2280,7 +2440,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
credentials, aws_region_name = self._load_credentials()
body: Final[dict[str, Any]] = {"messages": checks_messages, "checks": self.checks}
api_key: Final = self._get_bedrock_api_key(request_data)
api_key: Final = await self._async_get_bedrock_api_key(request_data)
prepared_request: Final = self._prepare_request(
credentials=credentials,