mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix: scope Bedrock guardrail bearer keys
This commit is contained in:
parent
f005afa146
commit
99e4ba34c5
3 changed files with 111 additions and 4 deletions
|
|
@ -669,6 +669,24 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
# logic becomes shared across providers.
|
||||
|
||||
#### CALL HOOKS - proxy only ####
|
||||
@staticmethod
|
||||
def _get_bedrock_api_key(request_data: Mapping[str, object] | None) -> str | None:
|
||||
if not request_data:
|
||||
return None
|
||||
|
||||
top_level_provider: Final[object | None] = request_data.get("custom_llm_provider")
|
||||
request_provider: Final[str | None] = (
|
||||
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
|
||||
if custom_llm_provider not in ("bedrock", "bedrock_converse"):
|
||||
return None
|
||||
|
||||
api_key: Final[object | None] = request_data.get("api_key")
|
||||
return api_key if isinstance(api_key, str) else None
|
||||
|
||||
def _load_credentials(
|
||||
self,
|
||||
):
|
||||
|
|
@ -841,7 +859,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 = self._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(
|
||||
|
|
@ -851,8 +869,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
|
||||
|
|
@ -1828,7 +1844,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 = self._get_bedrock_api_key(request_data)
|
||||
|
||||
prepared_request: Final = self._prepare_request(
|
||||
credentials=credentials,
|
||||
|
|
|
|||
|
|
@ -1114,6 +1114,55 @@ async def test_make_apply_guardrail_request_skips_scan_without_credentials():
|
|||
mock_post.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("custom_llm_provider", "expected_auth_prefix"),
|
||||
[("nvidia_nim", "AWS4-HMAC-SHA256"), ("bedrock", "Bearer bedrock-key")],
|
||||
)
|
||||
async def test_make_apply_guardrail_request_scopes_api_key_to_bedrock_provider(
|
||||
custom_llm_provider, expected_auth_prefix, monkeypatch
|
||||
):
|
||||
monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False)
|
||||
guardrail = BedrockGuardrail(
|
||||
guardrailIdentifier="test-guardrail",
|
||||
guardrailVersion="DRAFT",
|
||||
)
|
||||
request_data = {
|
||||
"model": f"{custom_llm_provider}/test-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
|
||||
)
|
||||
|
||||
|
||||
@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,48 @@ 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")],
|
||||
)
|
||||
async def test_request_scopes_api_key_to_bedrock_provider(
|
||||
custom_llm_provider, expected_auth_prefix, monkeypatch
|
||||
):
|
||||
monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False)
|
||||
g = BedrockGuardrail(checks=CONTENT_FILTER_CHECKS)
|
||||
request_data = {
|
||||
"model": f"{custom_llm_provider}/test-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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue