Merge pull request #13748 from hakasecurity/migrate-to-aim-new-firewall-api

Migrate to aim new firewall api
This commit is contained in:
Krish Dholakia 2025-08-19 22:32:22 -07:00 • committed by GitHub
commit 04feaf33c0
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 11 additions and 7 deletions

View file

@ -118,7 +118,7 @@ class AimGuardrail(CustomGuardrail):
litellm_call_id=call_id,
)
response = await self.async_handler.post(
f"{self.api_base}/detect/openai/v2",
f"{self.api_base}/fw/v1/analyze",
headers=headers,
json={"messages": data.get("messages", [])},
)
@ -183,14 +183,17 @@ class AimGuardrail(CustomGuardrail):
)
call_id = request_data.get("litellm_call_id")
response = await self.async_handler.post(
f"{self.api_base}/detect/output/v2",
f"{self.api_base}/fw/v1/analyze",
headers=self._build_aim_headers(
hook=hook,
key_alias=key_alias,
user_email=user_email,
litellm_call_id=call_id,
),
json={"output": output, "messages": request_data.get("messages", [])},
json={
"messages": request_data.get("messages", [])
+ [{"role": "assistant", "content": output}]
},
)
response.raise_for_status()
res = response.json()
@ -297,7 +300,7 @@ class AimGuardrail(CustomGuardrail):
)
call_id = request_data.get("litellm_call_id")
async with connect(
f"{self.ws_api_base}/detect/output/ws",
f"{self.ws_api_base}/fw/v1/analyze/stream",
additional_headers=self._build_aim_headers(
hook="output",
key_alias=user_api_key_dict.key_alias,

View file

@ -216,12 +216,13 @@ async def test_post_call__with_anonymized_entities__it_deanonymizes_output():
) as mock_post:
def mock_post_detect_side_effect(url, *args, **kwargs):
if url.endswith("/detect/openai/v2"):
request_body = kwargs.get("json", {})
if request_body["messages"][-1]["role"] == "user":
return response_with_detections
elif url.endswith("/detect/output/v2"):
elif request_body["messages"][-1]["role"] == "assistant":
return response_without_detections
else:
raise ValueError("Unexpected URL: {}".format(url))
raise ValueError("Unexpected request: {}".format(request_body))
mock_post.side_effect = mock_post_detect_side_effect