mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix: align Bedrock router selection
This commit is contained in:
parent
46d9a2a67d
commit
c58daa5d11
2 changed files with 245 additions and 121 deletions
|
|
@ -62,7 +62,7 @@ from litellm.router_strategy.tag_based_routing import (
|
|||
_split_tags,
|
||||
_strip_routing_prefix,
|
||||
)
|
||||
from litellm.router_utils.common_utils import filter_team_based_models
|
||||
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
|
||||
|
|
@ -695,16 +695,16 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
) -> list[object]:
|
||||
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
|
||||
_chain_tag_filtering_override(router, model, deployments) if isinstance(model, str) else None
|
||||
)
|
||||
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
|
||||
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
|
||||
|
||||
|
|
@ -749,7 +749,14 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
|
||||
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:
|
||||
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
|
||||
)
|
||||
if not has_positive_filter:
|
||||
default_deployments: Final = [
|
||||
deployment for deployment in candidate_deployments if "default" in _deployment_tags(deployment)
|
||||
]
|
||||
|
|
@ -809,94 +816,112 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
)
|
||||
try:
|
||||
resolved_team_id: Final = team_id if isinstance(team_id, str) else None
|
||||
model_id_deployment: Final = (
|
||||
llm_router.get_deployment(model_id=model) if llm_router.has_model_id(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)
|
||||
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,
|
||||
team_id=resolved_team_id,
|
||||
)
|
||||
or []
|
||||
router_kwargs: Final = dict(request_data)
|
||||
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)
|
||||
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 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
|
||||
)
|
||||
)
|
||||
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
|
||||
)
|
||||
is_model_alias: Final = isinstance(model_group_aliases, Mapping) and model in model_group_aliases
|
||||
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
|
||||
)
|
||||
global_pattern_deployments: Final = (
|
||||
get_pattern_deployments(model=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=model) if callable(get_team_pattern_deployments) else None
|
||||
)
|
||||
selected_listed_deployments: Final = (
|
||||
global_pattern_deployments
|
||||
if isinstance(global_pattern_deployments, list) and global_pattern_deployments
|
||||
else (
|
||||
team_pattern_deployments
|
||||
if isinstance(team_pattern_deployments, list) and team_pattern_deployments
|
||||
else raw_listed_deployments
|
||||
)
|
||||
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:
|
||||
selected_listed_deployments: Final = raw_listed_deployments
|
||||
candidate_deployments: Final[list[object]] = (
|
||||
[model_id_deployment_row]
|
||||
if model_id_deployment_row is not None
|
||||
else [deployment for deployment in selected_listed_deployments if isinstance(deployment, Mapping)]
|
||||
)
|
||||
router_matched: Final = bool(candidate_deployments)
|
||||
model_id_deployment: Final = (
|
||||
llm_router.get_deployment(model_id=model) if llm_router.has_model_id(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 model in concrete_model_names
|
||||
)
|
||||
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 []
|
||||
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=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=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 model in deployment_names
|
||||
and callable(specific_lookup)
|
||||
):
|
||||
specific_result: Final = specific_lookup(model=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 []
|
||||
)
|
||||
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)
|
||||
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=dict(request_data),
|
||||
request_kwargs=router_kwargs,
|
||||
)
|
||||
)
|
||||
team_filtered_deployments: Final[list[object]] = (
|
||||
|
|
@ -927,7 +952,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
healthy_deployments=access_filtered_deployments,
|
||||
parent_otel_span=None,
|
||||
)
|
||||
if callable(health_filter)
|
||||
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]] = (
|
||||
|
|
@ -942,27 +967,54 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
litellm_router_instance=llm_router,
|
||||
parent_otel_span=None,
|
||||
)
|
||||
if callable(cooldown_lookup) and callable(getattr(llm_router, "get_model_ids", None))
|
||||
if (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_ids: Final[frozenset[str]] = frozenset(
|
||||
deployment_id for deployment_id in cooldown_deployments if isinstance(deployment_id, str)
|
||||
)
|
||||
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: Final = BedrockGuardrail._filter_router_deployments_by_tags(
|
||||
router=llm_router,
|
||||
deployments=unblocked_deployments,
|
||||
request_data=request_data,
|
||||
)
|
||||
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,
|
||||
)
|
||||
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),
|
||||
)
|
||||
except Exception: # noqa: BLE001 # optional router state must not break guardrail auth
|
||||
return False
|
||||
if not deployments:
|
||||
|
|
@ -975,25 +1027,8 @@ 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:
|
||||
if isinstance(default_deployment, Mapping):
|
||||
default_params: object = (
|
||||
default_deployment.get("litellm_params")
|
||||
if isinstance(default_deployment, Mapping)
|
||||
|
|
@ -1016,6 +1051,24 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
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)
|
||||
|
||||
return False
|
||||
|
||||
providers: list[str] = []
|
||||
|
|
|
|||
|
|
@ -316,6 +316,77 @@ def test_bedrock_guardrail_resolves_specific_deployment_name(monkeypatch: pytest
|
|||
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()
|
||||
deployments = [
|
||||
{
|
||||
"litellm_params": {"custom_llm_provider": "openai", "order": 1},
|
||||
"model_info": {"id": "openai"},
|
||||
},
|
||||
{
|
||||
"litellm_params": {"custom_llm_provider": "bedrock", "order": 2},
|
||||
"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
|
||||
)
|
||||
|
||||
|
||||
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
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue