mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix: resolve native Bedrock guardrail providers
This commit is contained in:
parent
99e4ba34c5
commit
31a529aee6
4 changed files with 34 additions and 13 deletions
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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