diff --git a/litellm/proxy/guardrails/guardrail_hooks/aim/aim.py b/litellm/proxy/guardrails/guardrail_hooks/aim/aim.py index 3c03d70cef8..c95c41ee1d0 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/aim/aim.py +++ b/litellm/proxy/guardrails/guardrail_hooks/aim/aim.py @@ -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 {} diff --git a/tests/local_testing/test_aim_guardrails.py b/tests/local_testing/test_aim_guardrails.py index 3a9b6e9a3d1..b271875c2ec 100644 --- a/tests/local_testing/test_aim_guardrails.py +++ b/tests/local_testing/test_aim_guardrails.py @@ -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?"