mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix: harden Bedrock pre-call routing
This commit is contained in:
parent
250f30a211
commit
e673f4b48e
2 changed files with 179 additions and 29 deletions
|
|
@ -704,8 +704,13 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
if isinstance(chain_tag_filtering, bool)
|
||||
else getattr(router, "enable_tag_filtering", False)
|
||||
)
|
||||
router_settings_override: Final[object] = request_data.get("router_settings_override")
|
||||
trusted_request_tag_filtering: Final[bool] = (
|
||||
isinstance(router_settings_override, Mapping)
|
||||
and router_settings_override.get("enable_tag_filtering") is True
|
||||
)
|
||||
tag_filtering_enabled: Final = (
|
||||
request_data.get("enable_tag_filtering") is True or effective_tag_filtering is True
|
||||
trusted_request_tag_filtering or effective_tag_filtering is True
|
||||
)
|
||||
if not tag_filtering_enabled:
|
||||
return deployments
|
||||
|
|
@ -796,6 +801,21 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
return candidate_deployments or allowed_deployments
|
||||
return []
|
||||
|
||||
@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)
|
||||
router_settings_override: Final[object] = request_data.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
|
||||
for metadata_name in ("metadata", "litellm_metadata"):
|
||||
metadata: Final[object] = router_kwargs.get(metadata_name)
|
||||
if isinstance(metadata, Mapping) and "routing_decision" in metadata:
|
||||
router_kwargs[metadata_name] = {
|
||||
key: value for key, value in metadata.items() if key != "routing_decision"
|
||||
}
|
||||
return router_kwargs
|
||||
|
||||
@staticmethod
|
||||
def _router_deployment_field(deployment: object, field: str) -> object | None:
|
||||
model_info: Final[object | None] = (
|
||||
|
|
@ -808,6 +828,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
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):
|
||||
|
|
@ -830,15 +851,8 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
)
|
||||
try:
|
||||
resolved_team_id: Final = team_id if isinstance(team_id, str) else None
|
||||
router_kwargs: Final = dict(request_data)
|
||||
router_kwargs: Final = BedrockGuardrail._get_trusted_router_request_kwargs(request_data)
|
||||
effective_model: str = model
|
||||
for metadata in (request_data.get("litellm_metadata"), request_data.get("metadata")):
|
||||
if not isinstance(metadata, Mapping):
|
||||
continue
|
||||
routing_decision: Final[object] = metadata.get("routing_decision")
|
||||
if isinstance(routing_decision, Mapping) and isinstance(routing_decision.get("routed_model"), str):
|
||||
effective_model = routing_decision["routed_model"]
|
||||
break
|
||||
common_result: tuple[object, object] | None = None
|
||||
common_lookup: Final[object] = getattr(llm_router, "_common_checks_available_deployment", None)
|
||||
if callable(common_lookup):
|
||||
|
|
@ -1042,11 +1056,15 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
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,
|
||||
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(
|
||||
|
|
@ -1126,7 +1144,10 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
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 BedrockGuardrail._router_allows_bedrock(
|
||||
fallback_request_data,
|
||||
apply_tag_filtering=apply_tag_filtering,
|
||||
)
|
||||
|
||||
return False
|
||||
|
||||
|
|
@ -1175,29 +1196,49 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
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)
|
||||
routing_context: Final[object] = getattr(llm_router, "_get_routing_context", None)
|
||||
if callable(routing_context):
|
||||
try:
|
||||
context_result: Final = routing_context(
|
||||
model=request_data.get("model"), request_kwargs=router_request_kwargs
|
||||
)
|
||||
if isinstance(context_result, tuple) and context_result:
|
||||
routing_strategy = context_result[0]
|
||||
except Exception: # noqa: BLE001 # fall back to the router default
|
||||
pass
|
||||
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_lookup: Final[object] = getattr(llm_router, "async_get_healthy_deployments", None)
|
||||
if callable(async_lookup):
|
||||
model: Final[object | None] = request_data.get("model")
|
||||
if isinstance(model, str):
|
||||
effective_model = model
|
||||
for metadata in (request_data.get("litellm_metadata"), request_data.get("metadata")):
|
||||
if not isinstance(metadata, Mapping):
|
||||
continue
|
||||
routing_decision: Final[object] = metadata.get("routing_decision")
|
||||
if isinstance(routing_decision, Mapping) and isinstance(
|
||||
routing_decision.get("routed_model"), str
|
||||
):
|
||||
effective_model = routing_decision["routed_model"]
|
||||
break
|
||||
try:
|
||||
healthy_deployments: Final = await async_lookup(
|
||||
model=effective_model,
|
||||
request_kwargs=dict(request_data),
|
||||
messages=request_data.get("messages")
|
||||
if isinstance(request_data.get("messages"), list)
|
||||
request_kwargs=router_request_kwargs,
|
||||
messages=router_request_kwargs.get("messages")
|
||||
if isinstance(router_request_kwargs.get("messages"), list)
|
||||
else None,
|
||||
input=request_data.get("input")
|
||||
if isinstance(request_data.get("input"), (str, list))
|
||||
input=router_request_kwargs.get("input")
|
||||
if isinstance(router_request_kwargs.get("input"), (str, list))
|
||||
else None,
|
||||
specific_deployment=request_data.get("specific_deployment") is True,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -249,6 +249,115 @@ def test_bedrock_guardrail_matches_request_tag_pool(monkeypatch: pytest.MonkeyPa
|
|||
)
|
||||
|
||||
|
||||
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_uses_trusted_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"
|
||||
)
|
||||
|
||||
|
||||
@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()
|
||||
|
||||
|
||||
def test_bedrock_guardrail_matches_regex_tag_pool(monkeypatch: pytest.MonkeyPatch):
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue