diff --git a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py index 06772f0ad07..8ade4a011f7 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py +++ b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py @@ -669,6 +669,14 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): # 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: + return model.partition("/")[0] + @staticmethod def _get_bedrock_api_key(request_data: Mapping[str, object] | None) -> str | None: if not request_data: @@ -679,8 +687,10 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): top_level_provider if isinstance(top_level_provider, str) else None ) model: Final[object | None] = request_data.get("model") - model_provider: Final = model.partition("/")[0] if isinstance(model, str) else None - custom_llm_provider: Final = request_provider or model_provider + model_provider: Final[str | None] = ( + BedrockGuardrail._resolve_model_provider(model) if isinstance(model, str) else None + ) + custom_llm_provider: Final[str | None] = request_provider or model_provider if custom_llm_provider not in ("bedrock", "bedrock_converse"): return None diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py index c111072a06a..a1dbbaa5953 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py @@ -1116,19 +1116,23 @@ async def test_make_apply_guardrail_request_skips_scan_without_credentials(): @pytest.mark.asyncio @pytest.mark.parametrize( - ("custom_llm_provider", "expected_auth_prefix"), - [("nvidia_nim", "AWS4-HMAC-SHA256"), ("bedrock", "Bearer bedrock-key")], + ("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( - custom_llm_provider, expected_auth_prefix, monkeypatch -): + 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": f"{custom_llm_provider}/test-model", + "model": model, "api_key": "bedrock-key", } mock_credentials = MagicMock() diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_invoke_guardrail_checks.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_invoke_guardrail_checks.py index 634e2254187..a50e4bcffa0 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_invoke_guardrail_checks.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_invoke_guardrail_checks.py @@ -397,16 +397,20 @@ async def test_request_uses_checks_path_and_body(): @pytest.mark.asyncio @pytest.mark.parametrize( - ("custom_llm_provider", "expected_auth_prefix"), - [("nvidia_nim", "AWS4-HMAC-SHA256"), ("bedrock", "Bearer bedrock-key")], + ("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( - custom_llm_provider, expected_auth_prefix, monkeypatch -): + 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": f"{custom_llm_provider}/test-model", + "model": model, "api_key": "bedrock-key", } mock_credentials = MagicMock() diff --git a/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py b/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py index 45f5afef1bc..ab7a8f8acac 100644 --- a/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py +++ b/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py @@ -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(