Merge pull request #14438 from hakasecurity/change-aim-headers

rename aim headers + tests
This commit is contained in:
Krish Dholakia 2025-09-19 23:32:50 -07:00 • committed by GitHub
commit 1f9afcb349
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 14 additions and 5 deletions

View file

@ -80,7 +80,6 @@ class AimGuardrail(CustomGuardrail):
],
) -> Union[Exception, str, dict, None]:
verbose_proxy_logger.debug("Inside AIM Pre-Call Hook")
return await self.call_aim_guardrail(
data, hook="pre_call", key_alias=user_api_key_dict.key_alias
)
@ -246,13 +245,13 @@ class AimGuardrail(CustomGuardrail):
"x-aim-litellm-version": litellm_version,
}
# Used by Aim to track together single call input and output
| ({"x-aim-litellm-call-id": litellm_call_id} if litellm_call_id else {})
| ({"x-aim-call-id": litellm_call_id} if litellm_call_id else {})
# Used by Aim to track guardrails violations by user.
| ({"x-aim-user-email": user_email} if user_email else {})
| (
{
# Used by Aim apply only the guardrails that are associated with the key alias.
"x-aim-litellm-key-alias": key_alias,
"x-aim-gateway-key-alias": key_alias,
}
if key_alias
else {}

View file

@ -209,6 +209,7 @@ async def test_post_call__with_anonymized_entities__it_deanonymizes_output():
"messages": [
{"role": "user", "content": "Hi my name id Brian"},
],
"litellm_call_id": "test-call-id",
}
with patch(
@ -217,6 +218,13 @@ async def test_post_call__with_anonymized_entities__it_deanonymizes_output():
def mock_post_detect_side_effect(url, *args, **kwargs):
request_body = kwargs.get("json", {})
request_headers = kwargs.get("headers", {})
assert (
request_headers["x-aim-call-id"] == "test-call-id"
), "Wrong header: x-aim-call-id"
assert (
request_headers["x-aim-gateway-key-alias"] == "test-key"
), "Wrong header: x-aim-gateway-key-alias"
if request_body["messages"][-1]["role"] == "user":
return response_with_detections
elif request_body["messages"][-1]["role"] == "assistant":
@ -229,7 +237,7 @@ async def test_post_call__with_anonymized_entities__it_deanonymizes_output():
data = await aim_guardrail.async_pre_call_hook(
data=data,
cache=DualCache(),
user_api_key_dict=UserAPIKeyAuth(),
user_api_key_dict=UserAPIKeyAuth(key_alias="test-key"),
call_type="completion",
)
assert data["messages"][0]["content"] == "Hi my name is [NAME_1]"
@ -249,7 +257,9 @@ async def test_post_call__with_anonymized_entities__it_deanonymizes_output():
)
result = await aim_guardrail.async_post_call_success_hook(
data=data, response=llm_response(), user_api_key_dict=UserAPIKeyAuth()
data=data,
response=llm_response(),
user_api_key_dict=UserAPIKeyAuth(key_alias="test-key"),
)
assert result["choices"][0]["message"]["content"] == "Hello Brian! How are you?"