mirror of
https://github.com/BerriAI/litellm.git
synced 2026-08-28 05:25:59 +00:00
Merge 695fa583d4 into 3300fc3a96
This commit is contained in:
commit
78e5dfa06e
5 changed files with 1731 additions and 10 deletions
|
|
@ -30,7 +30,10 @@ from litellm.caching import DualCache
|
|||
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 redact_nested_match_and_regex_keys
|
||||
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
|
||||
from litellm.llms.anthropic.chat.guardrail_translation.handler import AnthropicMessagesHandler
|
||||
from litellm.llms.base_llm.guardrail_translation.utils import (
|
||||
|
|
@ -51,6 +54,18 @@ from litellm.proxy.guardrails.anthropic_sse import (
|
|||
is_raw_sse_stream,
|
||||
model_response_text,
|
||||
)
|
||||
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
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.guardrails import BedrockChecksConfigModel, GuardrailEventHooks
|
||||
from litellm.types.llms.openai import AllMessageValues, ChatCompletionUserMessage
|
||||
|
|
@ -106,6 +121,19 @@ _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()
|
||||
|
||||
|
||||
class _RouterCandidates(NamedTuple):
|
||||
"""What the router would consider for a request, before the eligibility filters."""
|
||||
|
||||
effective_model: str
|
||||
common_result: tuple[object, object] | None
|
||||
model_id_deployment_row: object | None
|
||||
candidate_deployments: Sequence[object]
|
||||
router_matched: bool
|
||||
|
||||
|
||||
_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
|
||||
|
|
@ -665,10 +693,803 @@ 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: # noqa: BLE001 # provider resolution has a safe prefix fallback
|
||||
return model.partition("/")[0]
|
||||
|
||||
@staticmethod
|
||||
def _filter_router_deployments_by_tags(
|
||||
router: object,
|
||||
deployments: list[object],
|
||||
request_data: Mapping[str, object],
|
||||
model: str | None = None,
|
||||
) -> list[object]:
|
||||
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
|
||||
)
|
||||
router_settings_override: Final[object] = request_data.get("router_settings_override")
|
||||
override_tag_filtering: Final[object] = (
|
||||
router_settings_override.get("enable_tag_filtering")
|
||||
if isinstance(router_settings_override, Mapping)
|
||||
else None
|
||||
)
|
||||
effective_tag_filtering: Final = (
|
||||
True
|
||||
if override_tag_filtering is True
|
||||
else chain_tag_filtering
|
||||
if isinstance(chain_tag_filtering, bool)
|
||||
else getattr(router, "enable_tag_filtering", False)
|
||||
)
|
||||
tag_filtering_enabled: Final = effective_tag_filtering is True
|
||||
if not tag_filtering_enabled:
|
||||
return deployments
|
||||
|
||||
def _deployment_tags(deployment: object) -> tuple[str, ...]:
|
||||
params: Final[object | None] = (
|
||||
deployment.get("litellm_params")
|
||||
if isinstance(deployment, Mapping)
|
||||
else getattr(deployment, "litellm_params", None)
|
||||
)
|
||||
tags: Final[object | None] = (
|
||||
params.get("tags") if isinstance(params, Mapping) else getattr(params, "tags", None)
|
||||
)
|
||||
return (
|
||||
tuple(tag for tag in tags if isinstance(tag, str))
|
||||
if isinstance(tags, Sequence) and not isinstance(tags, str)
|
||||
else ()
|
||||
)
|
||||
|
||||
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
|
||||
|
||||
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, 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))
|
||||
]
|
||||
candidate_deployments: Final = [
|
||||
deployment for deployment in allowed_deployments if required_set.issubset(_deployment_tags(deployment))
|
||||
]
|
||||
|
||||
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 []
|
||||
has_regex_deployments: Final = any(
|
||||
isinstance(deployment, Mapping) and bool((deployment.get("litellm_params") or {}).get("tag_regex"))
|
||||
for deployment in candidate_deployments
|
||||
)
|
||||
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 or _fail_open_deployments()
|
||||
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)
|
||||
if isinstance(getattr(router, "tag_filtering_match_any", True), bool)
|
||||
else True
|
||||
)
|
||||
matched_deployments: Final = [
|
||||
deployment
|
||||
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
|
||||
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]:
|
||||
router_kwargs: Final = dict(request_data)
|
||||
router_kwargs.pop("enable_tag_filtering", None)
|
||||
for metadata_name in ("metadata", "litellm_metadata"):
|
||||
metadata: Final[object] = router_kwargs.get(metadata_name)
|
||||
if isinstance(metadata, Mapping):
|
||||
router_kwargs[metadata_name] = {
|
||||
key: value for key, value in metadata.items() if key != "routing_decision"
|
||||
}
|
||||
router_settings_override: Final[object] = router_kwargs.get("router_settings_override")
|
||||
if (
|
||||
isinstance(router_settings_override, Mapping)
|
||||
and router_settings_override.get("enable_tag_filtering") is True
|
||||
):
|
||||
router_kwargs["enable_tag_filtering"] = True
|
||||
return router_kwargs
|
||||
|
||||
@staticmethod
|
||||
def _router_deployment_field(deployment: object, field: str) -> object | None:
|
||||
model_info: Final[object | None] = (
|
||||
deployment.get("model_info") if isinstance(deployment, Mapping) else getattr(deployment, "model_info", None)
|
||||
)
|
||||
return model_info.get(field) if isinstance(model_info, Mapping) else getattr(model_info, field, None)
|
||||
|
||||
@staticmethod
|
||||
def _router_deployment_provider(deployment: object) -> str | None:
|
||||
params: Final[object] = (
|
||||
deployment.get("litellm_params")
|
||||
if isinstance(deployment, Mapping)
|
||||
else getattr(deployment, "litellm_params", None)
|
||||
)
|
||||
provider: Final[object] = (
|
||||
params.get("custom_llm_provider")
|
||||
if isinstance(params, Mapping)
|
||||
else getattr(params, "custom_llm_provider", None)
|
||||
)
|
||||
if isinstance(provider, str):
|
||||
return provider
|
||||
deployment_model: Final[object] = (
|
||||
params.get("model") if isinstance(params, Mapping) else getattr(params, "model", None)
|
||||
)
|
||||
return BedrockGuardrail._resolve_model_provider(deployment_model) if isinstance(deployment_model, str) else None
|
||||
|
||||
@staticmethod
|
||||
def _router_deployments_for_provider_check(deployments: Sequence[object]) -> list[object]:
|
||||
if not deployments:
|
||||
return []
|
||||
|
||||
def _weight(deployment: object, weight_by: str) -> object | None:
|
||||
params: Final[object] = (
|
||||
deployment.get("litellm_params")
|
||||
if isinstance(deployment, Mapping)
|
||||
else getattr(deployment, "litellm_params", None)
|
||||
)
|
||||
return params.get(weight_by) if isinstance(params, Mapping) else getattr(params, weight_by, None)
|
||||
|
||||
for weight_by in ("weight", "rpm", "tpm"):
|
||||
first_weight: Final[object | None] = _weight(deployments[0], weight_by)
|
||||
if first_weight is None:
|
||||
continue
|
||||
try:
|
||||
weights: Final[list[object]] = [
|
||||
0 if (value := _weight(deployment, weight_by)) is None else value for deployment in deployments
|
||||
]
|
||||
if sum(weights) > 0:
|
||||
return [deployment for deployment, weight in zip(deployments, weights) if weight > 0]
|
||||
except (TypeError, ValueError):
|
||||
return list(deployments)
|
||||
return list(deployments)
|
||||
|
||||
@staticmethod
|
||||
def _router_candidate_deployments(
|
||||
llm_router: object,
|
||||
request_data: Mapping[str, object],
|
||||
router_kwargs: dict[str, object], # mutable-ok: the router pops its own routing keys off these kwargs
|
||||
model: str,
|
||||
resolved_team_id: str | None,
|
||||
) -> _RouterCandidates:
|
||||
"""Deployments the router would consider, before any of the eligibility filters."""
|
||||
effective_model: str = model
|
||||
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=effective_model,
|
||||
request_kwargs=router_kwargs,
|
||||
specific_deployment=request_data.get("specific_deployment") is True,
|
||||
)
|
||||
except Exception: # noqa: BLE001 # fall back for lightweight router test doubles
|
||||
raw_common_result = None
|
||||
if (
|
||||
isinstance(raw_common_result, tuple)
|
||||
and len(raw_common_result) == 2
|
||||
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]
|
||||
model_id_deployment_row: Final[object | None] = (
|
||||
raw_deployments if isinstance(raw_deployments, Mapping) else None
|
||||
)
|
||||
candidate_deployments: Final[list[object]] = (
|
||||
[raw_deployments]
|
||||
if isinstance(raw_deployments, Mapping)
|
||||
else [deployment for deployment in raw_deployments if isinstance(deployment, Mapping)]
|
||||
)
|
||||
router_matched: Final = bool(candidate_deployments)
|
||||
else:
|
||||
model_id_deployment: Final = (
|
||||
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)
|
||||
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)
|
||||
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 effective_model in concrete_model_names
|
||||
)
|
||||
is_model_alias: Final = isinstance(model_group_aliases, Mapping) and effective_model in model_group_aliases
|
||||
raw_listed_deployments = (
|
||||
[]
|
||||
if model_id_deployment_row is not None
|
||||
else (
|
||||
llm_router.get_model_list(model_name=effective_model, team_id=resolved_team_id) or []
|
||||
if is_concrete_model or is_model_alias
|
||||
else []
|
||||
)
|
||||
)
|
||||
if model_id_deployment_row 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
|
||||
)
|
||||
global_pattern_deployments: Final = (
|
||||
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)
|
||||
if resolved_team_id is not None
|
||||
and isinstance(getattr(llm_router, "team_pattern_routers", None), Mapping)
|
||||
else None
|
||||
)
|
||||
get_team_pattern_deployments: Final[object | None] = getattr(
|
||||
team_pattern_router, "get_deployments_by_pattern", None
|
||||
)
|
||||
team_pattern_deployments: Final = (
|
||||
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
|
||||
elif isinstance(team_pattern_deployments, list) and team_pattern_deployments:
|
||||
raw_listed_deployments = team_pattern_deployments
|
||||
else:
|
||||
default_deployment = getattr(llm_router, "default_deployment", None)
|
||||
if isinstance(default_deployment, Mapping):
|
||||
raw_listed_deployments = [default_deployment]
|
||||
elif (
|
||||
isinstance(deployment_names, Sequence)
|
||||
and not isinstance(deployment_names, (str, bytes))
|
||||
and effective_model in deployment_names
|
||||
and callable(specific_lookup)
|
||||
):
|
||||
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=effective_model, team_id=resolved_team_id) or []
|
||||
)
|
||||
candidate_deployments = (
|
||||
[model_id_deployment_row]
|
||||
if model_id_deployment_row is not None
|
||||
else [deployment for deployment in raw_listed_deployments if isinstance(deployment, Mapping)]
|
||||
)
|
||||
router_matched = bool(candidate_deployments)
|
||||
return _RouterCandidates(
|
||||
effective_model,
|
||||
common_result,
|
||||
model_id_deployment_row,
|
||||
candidate_deployments,
|
||||
router_matched,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _filter_router_deployments(
|
||||
llm_router: object,
|
||||
request_data: Mapping[str, object],
|
||||
router_kwargs: dict[str, object], # mutable-ok: the router pops its own routing keys off these kwargs
|
||||
*,
|
||||
effective_model: str,
|
||||
resolved_team_id: str | None,
|
||||
common_result: tuple[object, object] | None,
|
||||
model_id_deployment_row: object | None,
|
||||
candidate_deployments: Sequence[object],
|
||||
cooldown_deployments: Sequence[str] | None | object,
|
||||
apply_tag_filtering: bool,
|
||||
) -> Sequence[object]:
|
||||
"""The router's own eligibility chain, in the router's order.
|
||||
|
||||
Order filtering runs before the weighted-failover exclusion, matching Router
|
||||
so a guardrail verdict cannot disagree with the deployment actually picked.
|
||||
"""
|
||||
team_filtered_result: Final = (
|
||||
candidate_deployments
|
||||
if model_id_deployment_row is not None
|
||||
else filter_team_based_models(
|
||||
healthy_deployments=candidate_deployments,
|
||||
request_kwargs=router_kwargs,
|
||||
)
|
||||
)
|
||||
team_filtered_deployments: Final[list[object]] = (
|
||||
team_filtered_result if isinstance(team_filtered_result, list) else candidate_deployments
|
||||
)
|
||||
|
||||
filter_deployments: Final = getattr(llm_router, "_filter_deployments_by_model_access_groups", None)
|
||||
filtered_deployments: Final = (
|
||||
filter_deployments(
|
||||
model=effective_model,
|
||||
healthy_deployments=team_filtered_deployments,
|
||||
request_kwargs=dict(request_data),
|
||||
request_team_id=resolved_team_id,
|
||||
)
|
||||
if callable(filter_deployments)
|
||||
and isinstance(team_filtered_deployments, list)
|
||||
and model_id_deployment_row is None
|
||||
else None
|
||||
)
|
||||
access_filtered_deployments: Final[list[object]] = (
|
||||
filtered_deployments if isinstance(filtered_deployments, list) else team_filtered_deployments
|
||||
)
|
||||
health_filter: Final[object | None] = getattr(llm_router, "_filter_health_check_unhealthy_deployments", None)
|
||||
health_filtered_deployments: Final = (
|
||||
health_filter(
|
||||
healthy_deployments=access_filtered_deployments,
|
||||
parent_otel_span=None,
|
||||
)
|
||||
if (common_result is None or model_id_deployment_row is None) and callable(health_filter)
|
||||
else access_filtered_deployments
|
||||
)
|
||||
healthy_deployments: Final[list[object]] = (
|
||||
health_filtered_deployments
|
||||
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)
|
||||
resolved_cooldown_deployments: Final = (
|
||||
_get_cooldown_deployments(
|
||||
litellm_router_instance=llm_router,
|
||||
parent_otel_span=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 (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
|
||||
else:
|
||||
cooldown_filtered_deployments: Final = [
|
||||
deployment
|
||||
for deployment in healthy_deployments
|
||||
if BedrockGuardrail._router_deployment_field(deployment, "id") not in cooldown_ids
|
||||
]
|
||||
unblocked_deployments: Final = [
|
||||
deployment
|
||||
for deployment in cooldown_filtered_deployments
|
||||
if BedrockGuardrail._router_deployment_field(deployment, "blocked") is not True
|
||||
]
|
||||
deployments = (
|
||||
BedrockGuardrail._filter_router_deployments_by_tags(
|
||||
router=llm_router,
|
||||
deployments=unblocked_deployments,
|
||||
request_data=request_data,
|
||||
model=effective_model,
|
||||
)
|
||||
if apply_tag_filtering
|
||||
else unblocked_deployments
|
||||
)
|
||||
if common_result is not None and model_id_deployment_row is None:
|
||||
web_search_deployments: Final = filter_web_search_deployments(
|
||||
healthy_deployments=deployments,
|
||||
request_kwargs=router_kwargs,
|
||||
)
|
||||
deployments = web_search_deployments if isinstance(web_search_deployments, list) else deployments
|
||||
plugin_filter: Final[object] = getattr(llm_router, "_filter_by_routing_plugin_candidates", None)
|
||||
if callable(plugin_filter):
|
||||
plugin_deployments: Final = plugin_filter(
|
||||
healthy_deployments=deployments,
|
||||
request_kwargs=router_kwargs,
|
||||
)
|
||||
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),
|
||||
)
|
||||
return deployments
|
||||
|
||||
@staticmethod
|
||||
def _router_verdict_without_deployments(
|
||||
llm_router: object,
|
||||
request_data: Mapping[str, object],
|
||||
*,
|
||||
effective_model: str,
|
||||
resolved_team_id: str | None,
|
||||
router_matched: bool,
|
||||
apply_tag_filtering: bool,
|
||||
) -> bool | None:
|
||||
"""Verdict when the filters left nothing: pass-through, default deployment, or fallback."""
|
||||
if router_matched:
|
||||
return False
|
||||
router_settings: Final[object] = getattr(llm_router, "router_general_settings", None)
|
||||
if getattr(router_settings, "pass_through_all_models", False) is True:
|
||||
requested_provider: Final[object] = request_data.get("custom_llm_provider")
|
||||
passthrough_provider: Final[object] = (
|
||||
requested_provider
|
||||
if isinstance(requested_provider, str)
|
||||
else BedrockGuardrail._resolve_model_provider(effective_model)
|
||||
)
|
||||
return (
|
||||
passthrough_provider in ("bedrock", "bedrock_converse")
|
||||
if isinstance(passthrough_provider, str)
|
||||
else None
|
||||
)
|
||||
|
||||
default_deployment: Final[object] = getattr(llm_router, "default_deployment", None)
|
||||
if isinstance(default_deployment, Mapping):
|
||||
default_params: Final[object] = default_deployment.get("litellm_params")
|
||||
configured_provider: Final[object] = (
|
||||
default_params.get("custom_llm_provider")
|
||||
if isinstance(default_params, Mapping)
|
||||
else getattr(default_params, "custom_llm_provider", None)
|
||||
)
|
||||
default_model: Final[object] = (
|
||||
default_params.get("model")
|
||||
if isinstance(default_params, Mapping)
|
||||
else getattr(default_params, "model", None)
|
||||
)
|
||||
default_provider: Final[object] = (
|
||||
configured_provider
|
||||
if isinstance(configured_provider, str)
|
||||
else BedrockGuardrail._resolve_model_provider(default_model)
|
||||
if isinstance(default_model, str)
|
||||
else None
|
||||
)
|
||||
return default_provider in ("bedrock", "bedrock_converse") if isinstance(default_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 != effective_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,
|
||||
apply_tag_filtering=apply_tag_filtering,
|
||||
)
|
||||
|
||||
return False
|
||||
|
||||
@staticmethod
|
||||
def _router_allows_bedrock(
|
||||
request_data: Mapping[str, object],
|
||||
*,
|
||||
cooldown_deployments: Sequence[str] | None | object = _ROUTER_COOLDOWNS_UNSET,
|
||||
apply_tag_filtering: bool = True,
|
||||
) -> bool | None:
|
||||
model: Final[object | None] = request_data.get("model")
|
||||
if not isinstance(model, str):
|
||||
return False
|
||||
|
||||
selected_deployment: Final[object | None] = request_data.get("deployment")
|
||||
if selected_deployment is not None and not isinstance(selected_deployment, Mapping):
|
||||
selected_provider: Final[str | None] = BedrockGuardrail._router_deployment_provider(selected_deployment)
|
||||
return selected_provider in ("bedrock", "bedrock_converse") if selected_provider is not None else False
|
||||
|
||||
try:
|
||||
from litellm.proxy.proxy_server import llm_router
|
||||
except ImportError:
|
||||
return None
|
||||
if llm_router is None:
|
||||
return None
|
||||
|
||||
team_id: Final[object | None] = next(
|
||||
(
|
||||
metadata.get("user_api_key_team_id")
|
||||
for metadata in (request_data.get("litellm_metadata"), request_data.get("metadata"))
|
||||
if isinstance(metadata, Mapping) and isinstance(metadata.get("user_api_key_team_id"), str)
|
||||
),
|
||||
None,
|
||||
)
|
||||
try:
|
||||
resolved_team_id: Final = team_id if isinstance(team_id, str) else None
|
||||
router_kwargs: Final = BedrockGuardrail._get_trusted_router_request_kwargs(request_data)
|
||||
candidates: Final = BedrockGuardrail._router_candidate_deployments(
|
||||
llm_router, request_data, router_kwargs, model, resolved_team_id
|
||||
)
|
||||
deployments: Final = BedrockGuardrail._filter_router_deployments(
|
||||
llm_router,
|
||||
request_data,
|
||||
router_kwargs,
|
||||
effective_model=candidates.effective_model,
|
||||
resolved_team_id=resolved_team_id,
|
||||
common_result=candidates.common_result,
|
||||
model_id_deployment_row=candidates.model_id_deployment_row,
|
||||
candidate_deployments=candidates.candidate_deployments,
|
||||
cooldown_deployments=cooldown_deployments,
|
||||
apply_tag_filtering=apply_tag_filtering,
|
||||
)
|
||||
except Exception: # noqa: BLE001 # optional router state must not break guardrail auth
|
||||
return False
|
||||
if not deployments:
|
||||
return BedrockGuardrail._router_verdict_without_deployments(
|
||||
llm_router,
|
||||
request_data,
|
||||
effective_model=candidates.effective_model,
|
||||
resolved_team_id=resolved_team_id,
|
||||
router_matched=candidates.router_matched,
|
||||
apply_tag_filtering=apply_tag_filtering,
|
||||
)
|
||||
|
||||
providers: list[str] = []
|
||||
for deployment in deployments:
|
||||
provider: Final[str | None] = BedrockGuardrail._router_deployment_provider(deployment)
|
||||
if provider is None:
|
||||
return False
|
||||
providers.append(provider)
|
||||
return bool(providers) and all(provider in ("bedrock", "bedrock_converse") for provider in providers)
|
||||
|
||||
@staticmethod
|
||||
async def _async_router_bedrock_verdict(
|
||||
llm_router: object,
|
||||
request_data: Mapping[str, object],
|
||||
router_request_kwargs: dict[str, object], # mutable-ok: the pre-routing hook writes resolved params back here
|
||||
routing_strategy: object | None,
|
||||
) -> bool | None:
|
||||
"""Whether the router's async healthy-deployment set is all-Bedrock.
|
||||
|
||||
None means no verdict, so the caller falls through to the sync path.
|
||||
"""
|
||||
async_lookup: Final[object] = getattr(llm_router, "async_get_healthy_deployments", None)
|
||||
if not callable(async_lookup):
|
||||
return None
|
||||
model: Final[object | None] = request_data.get("model")
|
||||
if not isinstance(model, str):
|
||||
return None
|
||||
|
||||
effective_model = model
|
||||
effective_messages: object | None = (
|
||||
router_request_kwargs.get("messages") if isinstance(router_request_kwargs.get("messages"), list) else None
|
||||
)
|
||||
effective_input: object | None = (
|
||||
router_request_kwargs.get("input") if isinstance(router_request_kwargs.get("input"), (str, list)) else None
|
||||
)
|
||||
try:
|
||||
pre_routing_lookup: Final[object] = getattr(llm_router, "async_pre_routing_hook", None)
|
||||
if callable(pre_routing_lookup):
|
||||
pre_routing_result = pre_routing_lookup(
|
||||
model=model,
|
||||
request_kwargs=router_request_kwargs,
|
||||
messages=effective_messages,
|
||||
input=effective_input,
|
||||
specific_deployment=request_data.get("specific_deployment") is True,
|
||||
)
|
||||
if asyncio.iscoroutine(pre_routing_result):
|
||||
pre_routing_result = await pre_routing_result
|
||||
routed_model: Final[object] = getattr(pre_routing_result, "model", None)
|
||||
if isinstance(routed_model, str):
|
||||
effective_model = routed_model
|
||||
routed_messages: Final[object] = getattr(pre_routing_result, "messages", None)
|
||||
effective_messages = routed_messages if isinstance(routed_messages, list) else None
|
||||
routed_params: Final[object] = getattr(pre_routing_result, "litellm_params", None)
|
||||
if isinstance(routed_params, Mapping):
|
||||
router_request_kwargs.update(routed_params)
|
||||
healthy_deployments: Final = await async_lookup(
|
||||
model=effective_model,
|
||||
request_kwargs=router_request_kwargs,
|
||||
messages=effective_messages,
|
||||
input=effective_input,
|
||||
specific_deployment=request_data.get("specific_deployment") is True,
|
||||
)
|
||||
except Exception as exc: # noqa: BLE001 # fall back to the sync compatibility path
|
||||
verbose_proxy_logger.debug("Bedrock guardrail: async router lookup failed, using the sync path: %s", exc)
|
||||
return None
|
||||
|
||||
deployments: list[object] = (
|
||||
[healthy_deployments]
|
||||
if isinstance(healthy_deployments, Mapping)
|
||||
else healthy_deployments
|
||||
if isinstance(healthy_deployments, list)
|
||||
else []
|
||||
)
|
||||
if routing_strategy == "simple-shuffle":
|
||||
deployments = BedrockGuardrail._router_deployments_for_provider_check(deployments)
|
||||
if not deployments:
|
||||
return None
|
||||
|
||||
providers: list[str] = []
|
||||
for deployment in deployments:
|
||||
provider: Final[str | None] = BedrockGuardrail._router_deployment_provider(deployment)
|
||||
if provider is None:
|
||||
return False
|
||||
providers.append(provider)
|
||||
return 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:
|
||||
router_request_kwargs: Final = BedrockGuardrail._get_trusted_router_request_kwargs(request_data)
|
||||
routing_strategy: object | None = getattr(llm_router, "routing_strategy", None)
|
||||
if hasattr(routing_strategy, "value"):
|
||||
routing_strategy = routing_strategy.value
|
||||
if isinstance(routing_strategy, str) and routing_strategy not in {
|
||||
"usage-based-routing-v2",
|
||||
"simple-shuffle",
|
||||
"cost-based-routing",
|
||||
"latency-based-routing",
|
||||
"least-busy",
|
||||
}:
|
||||
router_allows_bedrock: Final = BedrockGuardrail._router_allows_bedrock(
|
||||
request_data,
|
||||
cooldown_deployments=[],
|
||||
apply_tag_filtering=False,
|
||||
)
|
||||
if router_allows_bedrock is not None:
|
||||
return api_key if router_allows_bedrock else None
|
||||
|
||||
async_verdict: Final = await BedrockGuardrail._async_router_bedrock_verdict(
|
||||
llm_router,
|
||||
request_data,
|
||||
router_request_kwargs,
|
||||
routing_strategy,
|
||||
)
|
||||
if async_verdict is not None:
|
||||
return api_key if async_verdict else None
|
||||
|
||||
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:
|
||||
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
|
||||
|
||||
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
|
||||
|
||||
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
|
||||
|
||||
def _load_credentials(
|
||||
self,
|
||||
):
|
||||
|
|
@ -843,7 +1664,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
bedrock_request_data: Final[dict] = dict(
|
||||
self.convert_to_bedrock_format(source=source, messages=messages, response=response)
|
||||
)
|
||||
api_key: str | None = None
|
||||
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(
|
||||
|
|
@ -853,8 +1674,6 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
if key not in _BEDROCK_DYNAMIC_BODY_DENYLIST
|
||||
}
|
||||
)
|
||||
if request_data.get("api_key") is not None:
|
||||
api_key = request_data["api_key"]
|
||||
|
||||
event_type: Final = (
|
||||
logging_event_type
|
||||
|
|
@ -1830,7 +2649,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[str | None] = request_data.get("api_key") if request_data else None
|
||||
api_key: Final = await self._async_get_bedrock_api_key(request_data)
|
||||
|
||||
prepared_request: Final = self._prepare_request(
|
||||
credentials=credentials,
|
||||
|
|
|
|||
|
|
@ -10,7 +10,6 @@ from litellm.proxy._types import UserAPIKeyAuth
|
|||
from litellm.caching import DualCache
|
||||
from unittest.mock import MagicMock, AsyncMock, patch
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bedrock_guardrails_pii_masking():
|
||||
# Create proper mock objects
|
||||
|
|
|
|||
|
|
@ -30,6 +30,767 @@ 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_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_rejects_blocked_model_with_pass_through(monkeypatch: pytest.MonkeyPatch):
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
router = MagicMock()
|
||||
router.get_model_list.return_value = [
|
||||
{"litellm_params": {"custom_llm_provider": "openai"}, "model_info": {"blocked": True}}
|
||||
]
|
||||
router.router_general_settings.pass_through_all_models = True
|
||||
router.default_deployment = None
|
||||
monkeypatch.setattr(proxy_server, "llm_router", router)
|
||||
|
||||
assert (
|
||||
BedrockGuardrail._get_bedrock_api_key(
|
||||
{
|
||||
"model": "blocked-alias",
|
||||
"custom_llm_provider": "bedrock",
|
||||
"api_key": "bedrock-key",
|
||||
}
|
||||
)
|
||||
is None
|
||||
)
|
||||
|
||||
|
||||
def test_bedrock_guardrail_rejects_access_filtered_model_with_pass_through(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
):
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
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 = []
|
||||
router.router_general_settings.pass_through_all_models = True
|
||||
router.default_deployment = None
|
||||
monkeypatch.setattr(proxy_server, "llm_router", router)
|
||||
|
||||
assert (
|
||||
BedrockGuardrail._get_bedrock_api_key(
|
||||
{
|
||||
"model": "scoped-alias",
|
||||
"custom_llm_provider": "bedrock",
|
||||
"api_key": "bedrock-key",
|
||||
}
|
||||
)
|
||||
is None
|
||||
)
|
||||
|
||||
|
||||
def test_bedrock_guardrail_resolves_model_id_before_wildcards(monkeypatch: pytest.MonkeyPatch):
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
router = MagicMock()
|
||||
router.has_model_id.return_value = True
|
||||
deployment = MagicMock()
|
||||
deployment.model_dump.return_value = {
|
||||
"litellm_params": {"custom_llm_provider": "bedrock"},
|
||||
"model_info": {"id": "deployment-id"},
|
||||
}
|
||||
router.get_deployment.return_value = 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": "deployment-id"}) is True
|
||||
router.get_model_list.assert_not_called()
|
||||
|
||||
|
||||
def test_bedrock_guardrail_ignores_cooling_non_bedrock_deployments(monkeypatch: pytest.MonkeyPatch):
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
router = MagicMock()
|
||||
router.get_model_list.return_value = [
|
||||
{
|
||||
"litellm_params": {"custom_llm_provider": "openai"},
|
||||
"model_info": {"id": "openai-deployment"},
|
||||
},
|
||||
{
|
||||
"litellm_params": {"custom_llm_provider": "bedrock"},
|
||||
"model_info": {"id": "bedrock-deployment"},
|
||||
},
|
||||
]
|
||||
router.get_model_ids.return_value = ["openai-deployment", "bedrock-deployment"]
|
||||
router.cooldown_cache.get_active_cooldowns.return_value = [("openai-deployment", 1.0)]
|
||||
monkeypatch.setattr(proxy_server, "llm_router", router)
|
||||
|
||||
assert BedrockGuardrail._router_allows_bedrock({"model": "shared-alias"}) is True
|
||||
|
||||
|
||||
def test_bedrock_guardrail_matches_global_wildcard_precedence(monkeypatch: pytest.MonkeyPatch):
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
router = MagicMock()
|
||||
router.model_names = []
|
||||
router.model_group_alias = {}
|
||||
router.get_model_list.return_value = [
|
||||
{"litellm_params": {"custom_llm_provider": "bedrock"}, "model_info": {}}
|
||||
]
|
||||
router.pattern_router.get_deployments_by_pattern.return_value = [
|
||||
{"litellm_params": {"custom_llm_provider": "bedrock"}, "model_info": {}}
|
||||
]
|
||||
team_pattern_router = MagicMock()
|
||||
team_pattern_router.get_deployments_by_pattern.return_value = [
|
||||
{"litellm_params": {"custom_llm_provider": "openai"}, "model_info": {}}
|
||||
]
|
||||
router.team_pattern_routers = {"team-id": team_pattern_router}
|
||||
monkeypatch.setattr(proxy_server, "llm_router", router)
|
||||
|
||||
assert (
|
||||
BedrockGuardrail._router_allows_bedrock(
|
||||
{"model": "provider/model", "litellm_metadata": {"user_api_key_team_id": "team-id"}}
|
||||
)
|
||||
is True
|
||||
)
|
||||
|
||||
|
||||
def test_bedrock_guardrail_matches_request_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", "tags": ["slow"]}, "model_info": {}},
|
||||
{"litellm_params": {"custom_llm_provider": "bedrock", "tags": ["fast"]}, "model_info": {}},
|
||||
]
|
||||
monkeypatch.setattr(proxy_server, "llm_router", router)
|
||||
|
||||
assert (
|
||||
BedrockGuardrail._router_allows_bedrock(
|
||||
{"model": "shared-alias", "metadata": {"tags": ["fast"]}}
|
||||
)
|
||||
is True
|
||||
)
|
||||
|
||||
|
||||
def test_bedrock_guardrail_ignores_client_routing_decision(monkeypatch: pytest.MonkeyPatch):
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
router = MagicMock()
|
||||
router.get_model_list.return_value = [
|
||||
{"litellm_params": {"custom_llm_provider": "openai"}, "model_info": {}},
|
||||
]
|
||||
monkeypatch.setattr(proxy_server, "llm_router", router)
|
||||
|
||||
assert (
|
||||
BedrockGuardrail._get_bedrock_api_key(
|
||||
{
|
||||
"model": "shared-alias",
|
||||
"api_key": "bedrock-key",
|
||||
"metadata": {"routing_decision": {"routed_model": "bedrock-alias"}},
|
||||
}
|
||||
)
|
||||
is None
|
||||
)
|
||||
|
||||
|
||||
def test_bedrock_guardrail_ignores_client_tag_filtering_override(monkeypatch: pytest.MonkeyPatch):
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
router = MagicMock()
|
||||
router.enable_tag_filtering = False
|
||||
router.get_model_list.return_value = [
|
||||
{"litellm_params": {"custom_llm_provider": "openai", "tags": ["slow"]}, "model_info": {}},
|
||||
{"litellm_params": {"custom_llm_provider": "bedrock", "tags": ["fast"]}, "model_info": {}},
|
||||
]
|
||||
monkeypatch.setattr(proxy_server, "llm_router", router)
|
||||
|
||||
assert (
|
||||
BedrockGuardrail._get_bedrock_api_key(
|
||||
{
|
||||
"model": "shared-alias",
|
||||
"api_key": "bedrock-key",
|
||||
"enable_tag_filtering": True,
|
||||
"metadata": {"tags": ["fast"]},
|
||||
}
|
||||
)
|
||||
is None
|
||||
)
|
||||
|
||||
|
||||
def test_bedrock_guardrail_honors_router_settings_tag_filtering_override(monkeypatch: pytest.MonkeyPatch):
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
router = MagicMock()
|
||||
router.enable_tag_filtering = False
|
||||
router.get_model_list.return_value = [
|
||||
{
|
||||
"litellm_params": {"custom_llm_provider": "openai", "tags": ["slow"]},
|
||||
"model_info": {},
|
||||
},
|
||||
{
|
||||
"litellm_params": {"custom_llm_provider": "bedrock", "tags": ["fast"]},
|
||||
"model_info": {},
|
||||
},
|
||||
]
|
||||
monkeypatch.setattr(proxy_server, "llm_router", router)
|
||||
|
||||
assert (
|
||||
BedrockGuardrail._get_bedrock_api_key(
|
||||
{
|
||||
"model": "shared-alias",
|
||||
"api_key": "bedrock-key",
|
||||
"router_settings_override": {"enable_tag_filtering": True},
|
||||
"metadata": {"tags": ["fast"]},
|
||||
}
|
||||
)
|
||||
== "bedrock-key"
|
||||
)
|
||||
|
||||
|
||||
def test_bedrock_guardrail_keeps_all_required_tag_matches(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": ["required", "default"]},
|
||||
"model_info": {},
|
||||
},
|
||||
{
|
||||
"litellm_params": {"custom_llm_provider": "openai", "tags": ["required"]},
|
||||
"model_info": {},
|
||||
},
|
||||
]
|
||||
monkeypatch.setattr(proxy_server, "llm_router", router)
|
||||
|
||||
assert (
|
||||
BedrockGuardrail._router_allows_bedrock(
|
||||
{"model": "shared-alias", "metadata": {"tags": ["&required"]}}
|
||||
)
|
||||
is False
|
||||
)
|
||||
|
||||
|
||||
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
|
||||
|
||||
router = MagicMock()
|
||||
router.routing_strategy = "usage-based-routing"
|
||||
router.model_names = ["shared-alias"]
|
||||
router.model_group_alias = {}
|
||||
router.has_model_id.return_value = False
|
||||
router.get_model_list.return_value = [
|
||||
{"litellm_params": {"custom_llm_provider": "openai", "tags": ["slow"]}, "model_info": {}},
|
||||
{"litellm_params": {"custom_llm_provider": "bedrock", "tags": ["fast"]}, "model_info": {}},
|
||||
]
|
||||
router._filter_health_check_unhealthy_deployments.side_effect = lambda healthy_deployments, **_: healthy_deployments
|
||||
router._filter_deployments_by_model_access_groups.side_effect = (
|
||||
lambda **kwargs: kwargs["healthy_deployments"]
|
||||
)
|
||||
router.get_model_ids.return_value = []
|
||||
router.cooldown_cache.get_active_cooldowns.return_value = []
|
||||
router.pattern_router = None
|
||||
router.default_deployment = None
|
||||
router.router_general_settings.pass_through_all_models = False
|
||||
router.async_get_healthy_deployments = AsyncMock(
|
||||
return_value=[{"litellm_params": {"custom_llm_provider": "bedrock"}, "model_info": {}}]
|
||||
)
|
||||
monkeypatch.setattr(proxy_server, "llm_router", router)
|
||||
|
||||
assert (
|
||||
await BedrockGuardrail._async_get_bedrock_api_key(
|
||||
{
|
||||
"model": "shared-alias",
|
||||
"api_key": "bedrock-key",
|
||||
"metadata": {"tags": ["fast"]},
|
||||
}
|
||||
)
|
||||
is None
|
||||
)
|
||||
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()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bedrock_guardrail_async_uses_pre_routed_model(monkeypatch: pytest.MonkeyPatch):
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
router = MagicMock()
|
||||
router.routing_strategy = "usage-based-routing-v2"
|
||||
router.async_pre_routing_hook = AsyncMock(
|
||||
return_value=MagicMock(model="bedrock-model", messages=[{"role": "user", "content": "hi"}])
|
||||
)
|
||||
router.async_get_healthy_deployments = AsyncMock(
|
||||
return_value=[{"litellm_params": {"custom_llm_provider": "bedrock"}, "model_info": {}}]
|
||||
)
|
||||
monkeypatch.setattr(proxy_server, "llm_router", router)
|
||||
|
||||
assert (
|
||||
await BedrockGuardrail._async_get_bedrock_api_key(
|
||||
{
|
||||
"model": "router-alias",
|
||||
"api_key": "bedrock-key",
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
}
|
||||
)
|
||||
== "bedrock-key"
|
||||
)
|
||||
assert router.async_get_healthy_deployments.await_args.kwargs["model"] == "bedrock-model"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bedrock_guardrail_async_honors_nested_tag_override(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)
|
||||
|
||||
assert (
|
||||
await BedrockGuardrail._async_get_bedrock_api_key(
|
||||
{
|
||||
"model": "router-alias",
|
||||
"api_key": "bedrock-key",
|
||||
"router_settings_override": {"enable_tag_filtering": True},
|
||||
}
|
||||
)
|
||||
== "bedrock-key"
|
||||
)
|
||||
assert router.async_get_healthy_deployments.await_args.kwargs["request_kwargs"]["enable_tag_filtering"] is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bedrock_guardrail_async_ignores_zero_weight_provider(monkeypatch: pytest.MonkeyPatch):
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
router = MagicMock()
|
||||
router.routing_strategy = "simple-shuffle"
|
||||
router.async_get_healthy_deployments = AsyncMock(
|
||||
return_value=[
|
||||
{"litellm_params": {"custom_llm_provider": "openai", "weight": 0}, "model_info": {}},
|
||||
{"litellm_params": {"custom_llm_provider": "bedrock", "weight": 1}, "model_info": {}},
|
||||
]
|
||||
)
|
||||
monkeypatch.setattr(proxy_server, "llm_router", router)
|
||||
|
||||
assert (
|
||||
await BedrockGuardrail._async_get_bedrock_api_key(
|
||||
{"model": "router-alias", "api_key": "bedrock-key"}
|
||||
)
|
||||
== "bedrock-key"
|
||||
)
|
||||
|
||||
|
||||
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_ignores_user_agent_without_regex_route(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"}, "model_info": {}}
|
||||
]
|
||||
monkeypatch.setattr(proxy_server, "llm_router", router)
|
||||
|
||||
assert BedrockGuardrail._router_allows_bedrock(
|
||||
{"model": "shared-alias", "metadata": {"user_agent": "client/1.0"}}
|
||||
) is True
|
||||
|
||||
|
||||
def test_bedrock_guardrail_applies_router_post_filters(monkeypatch: pytest.MonkeyPatch):
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
router = MagicMock()
|
||||
# Equal order: the min-order filter runs before exclusion, so an openai row ordered
|
||||
# ahead of bedrock would decide the verdict on its own and never exercise exclusion.
|
||||
deployments = [
|
||||
{
|
||||
"litellm_params": {"custom_llm_provider": "openai", "order": 1},
|
||||
"model_info": {"id": "openai"},
|
||||
},
|
||||
{
|
||||
"litellm_params": {"custom_llm_provider": "bedrock", "order": 1},
|
||||
"model_info": {"id": "bedrock"},
|
||||
},
|
||||
]
|
||||
router._common_checks_available_deployment.return_value = ("shared-alias", deployments)
|
||||
router._filter_health_check_unhealthy_deployments.return_value = deployments
|
||||
router.routing_plugins = []
|
||||
router.default_deployment = None
|
||||
monkeypatch.setattr(proxy_server, "llm_router", router)
|
||||
|
||||
assert BedrockGuardrail._router_allows_bedrock({"model": "shared-alias"}) is False
|
||||
assert (
|
||||
BedrockGuardrail._router_allows_bedrock(
|
||||
{"model": "shared-alias", "_excluded_deployment_ids": ["openai"]}
|
||||
)
|
||||
is True
|
||||
)
|
||||
assert (
|
||||
BedrockGuardrail._router_allows_bedrock({"model": "shared-alias", "_target_order": 2}) is False
|
||||
)
|
||||
|
||||
|
||||
def test_bedrock_guardrail_applies_web_search_filter(monkeypatch: pytest.MonkeyPatch):
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
router = MagicMock()
|
||||
deployments = [
|
||||
{
|
||||
"litellm_params": {"custom_llm_provider": "openai"},
|
||||
"model_info": {"id": "openai", "supports_web_search": True},
|
||||
},
|
||||
{
|
||||
"litellm_params": {"custom_llm_provider": "bedrock"},
|
||||
"model_info": {"id": "bedrock", "supports_web_search": False},
|
||||
},
|
||||
]
|
||||
router._common_checks_available_deployment.return_value = ("shared-alias", deployments)
|
||||
router._filter_health_check_unhealthy_deployments.return_value = deployments
|
||||
router.routing_plugins = []
|
||||
monkeypatch.setattr(proxy_server, "llm_router", router)
|
||||
|
||||
assert (
|
||||
BedrockGuardrail._router_allows_bedrock(
|
||||
{"model": "shared-alias", "tools": [{"type": "web_search"}]}
|
||||
)
|
||||
is False
|
||||
)
|
||||
|
||||
|
||||
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
|
||||
|
||||
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 = [
|
||||
{"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"}
|
||||
# filter_team_based_models drops by model_info.id, so a row without one takes every
|
||||
# other id-less row down with it.
|
||||
router.get_model_list.return_value = [
|
||||
{
|
||||
"litellm_params": {"custom_llm_provider": "openai"},
|
||||
"model_info": {"id": "openai", "team_id": "other-team"},
|
||||
},
|
||||
{
|
||||
"litellm_params": {"custom_llm_provider": "bedrock"},
|
||||
"model_info": {"id": "bedrock", "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"""
|
||||
|
|
@ -1114,6 +1875,99 @@ async def test_make_apply_guardrail_request_skips_scan_without_credentials():
|
|||
mock_post.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("model", "expected_auth_prefix"),
|
||||
[
|
||||
("nvidia_nim/test-model", "AWS4-HMAC-SHA256"),
|
||||
("bedrock/test-model", "Bearer bedrock-key"),
|
||||
("amazon.nova-lite-v1:0", "Bearer bedrock-key"),
|
||||
],
|
||||
)
|
||||
async def test_make_apply_guardrail_request_scopes_api_key_to_bedrock_provider(
|
||||
model: str, expected_auth_prefix: str, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False)
|
||||
guardrail = BedrockGuardrail(
|
||||
guardrailIdentifier="test-guardrail",
|
||||
guardrailVersion="DRAFT",
|
||||
)
|
||||
request_data = {
|
||||
"model": model,
|
||||
"api_key": "bedrock-key",
|
||||
}
|
||||
mock_credentials = MagicMock()
|
||||
mock_credentials.access_key = "test-access-key"
|
||||
mock_credentials.secret_key = "test-secret-key"
|
||||
mock_credentials.token = None
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.json.return_value = {"action": "NONE", "assessments": []}
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
guardrail,
|
||||
"_load_credentials",
|
||||
return_value=(mock_credentials, "us-east-1"),
|
||||
),
|
||||
patch.object(
|
||||
guardrail.async_handler,
|
||||
"post",
|
||||
new=AsyncMock(return_value=mock_response),
|
||||
) as mock_post,
|
||||
):
|
||||
result = await guardrail.make_bedrock_api_request(
|
||||
source="INPUT",
|
||||
messages=[{"role": "user", "content": "hello"}],
|
||||
request_data=request_data,
|
||||
)
|
||||
|
||||
assert result["action"] == "NONE"
|
||||
assert mock_post.await_args.kwargs["headers"]["Authorization"].startswith(
|
||||
expected_auth_prefix
|
||||
)
|
||||
|
||||
|
||||
def test_bedrock_api_key_rejects_caller_provider_spoofing(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
router = MagicMock()
|
||||
router.get_model_list.return_value = [
|
||||
{"litellm_params": {"model": "gpt-4o", "custom_llm_provider": "openai"}}
|
||||
]
|
||||
monkeypatch.setattr(proxy_server, "llm_router", router)
|
||||
|
||||
assert (
|
||||
BedrockGuardrail._get_bedrock_api_key(
|
||||
{
|
||||
"model": "shared-alias",
|
||||
"custom_llm_provider": "bedrock",
|
||||
"api_key": "bedrock-key",
|
||||
}
|
||||
)
|
||||
is None
|
||||
)
|
||||
|
||||
|
||||
def test_bedrock_api_key_accepts_alias_with_only_bedrock_deployments(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
router = MagicMock()
|
||||
router.get_model_list.return_value = [
|
||||
{"litellm_params": {"model": "amazon.nova-lite-v1:0", "custom_llm_provider": "bedrock"}}
|
||||
]
|
||||
monkeypatch.setattr(proxy_server, "llm_router", router)
|
||||
|
||||
assert (
|
||||
BedrockGuardrail._get_bedrock_api_key(
|
||||
{"model": "bedrock-alias", "api_key": "bedrock-key"}
|
||||
)
|
||||
== "bedrock-key"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bedrock_apply_guardrail_response_uses_OUTPUT_source():
|
||||
"""input_type='response' must call Bedrock with source=OUTPUT and assistant content.
|
||||
|
|
|
|||
|
|
@ -395,6 +395,52 @@ async def test_request_uses_checks_path_and_body():
|
|||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("model", "expected_auth_prefix"),
|
||||
[
|
||||
("nvidia_nim/test-model", "AWS4-HMAC-SHA256"),
|
||||
("bedrock/test-model", "Bearer bedrock-key"),
|
||||
("amazon.nova-lite-v1:0", "Bearer bedrock-key"),
|
||||
],
|
||||
)
|
||||
async def test_request_scopes_api_key_to_bedrock_provider(
|
||||
model: str, expected_auth_prefix: str, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False)
|
||||
g = BedrockGuardrail(checks=CONTENT_FILTER_CHECKS)
|
||||
request_data = {
|
||||
"model": model,
|
||||
"api_key": "bedrock-key",
|
||||
}
|
||||
mock_credentials = MagicMock()
|
||||
mock_credentials.access_key = "test-access-key"
|
||||
mock_credentials.secret_key = "test-secret-key"
|
||||
mock_credentials.token = None
|
||||
mock_post = AsyncMock(
|
||||
return_value=_mock_http_response(200, {"results": {}})
|
||||
)
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
g,
|
||||
"_load_credentials",
|
||||
return_value=(mock_credentials, "us-east-1"),
|
||||
),
|
||||
patch.object(g.async_handler, "post", new=mock_post),
|
||||
):
|
||||
result = await g.make_bedrock_api_request(
|
||||
source="INPUT",
|
||||
messages=[{"role": "user", "content": "hello"}],
|
||||
request_data=request_data,
|
||||
)
|
||||
|
||||
assert result == BedrockGuardrailResponse()
|
||||
assert mock_post.await_args.kwargs["headers"]["Authorization"].startswith(
|
||||
expected_auth_prefix
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_empty_messages_passes_without_api_call():
|
||||
g = BedrockGuardrail(checks=CONTENT_FILTER_CHECKS)
|
||||
|
|
|
|||
|
|
@ -889,7 +889,10 @@ async def test_bedrock_guardrail_make_api_request_passes_api_key():
|
|||
mock_response.status_code = 200
|
||||
mock_response.json.return_value = {"action": "NONE", "outputs": []}
|
||||
|
||||
test_request_data = {"api_key": "test-api-key-789"}
|
||||
test_request_data = {
|
||||
"model": "bedrock/test-model",
|
||||
"api_key": "test-api-key-789",
|
||||
}
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue