mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-12 23:01:41 +00:00
feat(guardrails): capture user and model metadata in CrowdStrike AIDR
This commit is contained in:
parent
db7f25d22c
commit
6fc715c5bd
2 changed files with 137 additions and 3 deletions
|
|
@ -1,4 +1,5 @@
|
|||
import os
|
||||
from collections.abc import Mapping
|
||||
from typing import TYPE_CHECKING, Literal, Optional, Type
|
||||
from typing_extensions import Any, override
|
||||
|
||||
|
|
@ -310,11 +311,27 @@ class CrowdStrikeAIDRHandler(CustomGuardrail):
|
|||
event_type = "output"
|
||||
hook_name = "apply_guardrail (response)"
|
||||
|
||||
ai_guard_payload = {
|
||||
ai_guard_payload: dict[str, Any] = {
|
||||
"guard_input": guard_input,
|
||||
"event_type": event_type,
|
||||
}
|
||||
|
||||
model = inputs.get("model")
|
||||
if model:
|
||||
ai_guard_payload["model"] = model
|
||||
|
||||
metadata = request_data.get("litellm_metadata", request_data.get("metadata"))
|
||||
if isinstance(metadata, Mapping):
|
||||
user_id = metadata.get("user_api_key_user_id")
|
||||
if user_id:
|
||||
ai_guard_payload["user_id"] = user_id
|
||||
|
||||
extra_info: dict[str, str] = {}
|
||||
user_email = metadata.get("user_api_key_user_email")
|
||||
if user_email:
|
||||
extra_info["user_name"] = user_email
|
||||
ai_guard_payload["extra_info"] = extra_info
|
||||
|
||||
ai_guard_response = await self._call_crowdstrike_aidr_guard(
|
||||
ai_guard_payload, hook_name
|
||||
)
|
||||
|
|
|
|||
|
|
@ -41,7 +41,8 @@ def test_crowdstrike_aidr_guardrail_config() -> None:
|
|||
)
|
||||
|
||||
|
||||
def test_crowdstrike_aidr_guardrail_config_no_api_key() -> None:
|
||||
def test_crowdstrike_aidr_guardrail_config_no_api_key(monkeypatch) -> None:
|
||||
monkeypatch.delenv("CS_AIDR_TOKEN", raising=False)
|
||||
with pytest.raises(CrowdStrikeAIDRGuardrailMissingSecrets):
|
||||
init_guardrails_v2(
|
||||
all_guardrails=[
|
||||
|
|
@ -59,7 +60,8 @@ def test_crowdstrike_aidr_guardrail_config_no_api_key() -> None:
|
|||
)
|
||||
|
||||
|
||||
def test_crowdstrike_aidr_guardrail_config_no_api_base() -> None:
|
||||
def test_crowdstrike_aidr_guardrail_config_no_api_base(monkeypatch) -> None:
|
||||
monkeypatch.delenv("CS_AIDR_BASE_URL", raising=False)
|
||||
with pytest.raises(CrowdStrikeAIDRGuardrailMissingSecrets):
|
||||
init_guardrails_v2(
|
||||
all_guardrails=[
|
||||
|
|
@ -428,3 +430,118 @@ async def test_apply_guardrail_response_ok(
|
|||
)
|
||||
# Should return original inputs when not transformed
|
||||
assert result["texts"] == inputs["texts"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_sends_user_id_model_and_extra_info(
|
||||
crowdstrike_aidr_guardrail: CrowdStrikeAIDRHandler,
|
||||
) -> None:
|
||||
inputs: GenericGuardrailAPIInputs = {
|
||||
"texts": ["Hello"],
|
||||
"structured_messages": [{"role": "user", "content": "Hello"}],
|
||||
"model": "gpt-4o",
|
||||
}
|
||||
request_data = {
|
||||
"messages": inputs["structured_messages"],
|
||||
"model": "gpt-4o",
|
||||
"litellm_metadata": {
|
||||
"user_api_key_user_id": "uid-abc",
|
||||
"user_api_key_user_email": "alice@example.com",
|
||||
},
|
||||
}
|
||||
guardrail_endpoint = (
|
||||
f"{crowdstrike_aidr_guardrail.api_base}/v1/guard_chat_completions"
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
|
||||
return_value=httpx.Response(
|
||||
status_code=200,
|
||||
json={"result": {"blocked": False, "transformed": False}},
|
||||
request=httpx.Request(method="POST", url=guardrail_endpoint),
|
||||
),
|
||||
) as mock_method:
|
||||
await crowdstrike_aidr_guardrail.apply_guardrail(
|
||||
inputs=inputs,
|
||||
request_data=request_data,
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
payload = mock_method.call_args.kwargs["json"]
|
||||
assert payload["user_id"] == "uid-abc"
|
||||
assert payload["model"] == "gpt-4o"
|
||||
assert payload["extra_info"] == {"user_name": "alice@example.com"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_empty_extra_info_when_no_email(
|
||||
crowdstrike_aidr_guardrail: CrowdStrikeAIDRHandler,
|
||||
) -> None:
|
||||
inputs: GenericGuardrailAPIInputs = {
|
||||
"texts": ["Hello"],
|
||||
"structured_messages": [{"role": "user", "content": "Hello"}],
|
||||
"model": "gemini-flash",
|
||||
}
|
||||
request_data = {
|
||||
"messages": inputs["structured_messages"],
|
||||
"model": "gemini-flash",
|
||||
"litellm_metadata": {
|
||||
"user_api_key_user_id": "uid-no-email",
|
||||
"user_api_key_user_email": None,
|
||||
},
|
||||
}
|
||||
guardrail_endpoint = (
|
||||
f"{crowdstrike_aidr_guardrail.api_base}/v1/guard_chat_completions"
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
|
||||
return_value=httpx.Response(
|
||||
status_code=200,
|
||||
json={"result": {"blocked": False, "transformed": False}},
|
||||
request=httpx.Request(method="POST", url=guardrail_endpoint),
|
||||
),
|
||||
) as mock_method:
|
||||
await crowdstrike_aidr_guardrail.apply_guardrail(
|
||||
inputs=inputs,
|
||||
request_data=request_data,
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
payload = mock_method.call_args.kwargs["json"]
|
||||
assert payload["user_id"] == "uid-no-email"
|
||||
assert payload["model"] == "gemini-flash"
|
||||
assert payload["extra_info"] == {}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_no_metadata_skips_user_fields(
|
||||
crowdstrike_aidr_guardrail: CrowdStrikeAIDRHandler,
|
||||
) -> None:
|
||||
inputs: GenericGuardrailAPIInputs = {
|
||||
"texts": ["Hello"],
|
||||
"structured_messages": [{"role": "user", "content": "Hello"}],
|
||||
}
|
||||
request_data = {"messages": inputs["structured_messages"]}
|
||||
guardrail_endpoint = (
|
||||
f"{crowdstrike_aidr_guardrail.api_base}/v1/guard_chat_completions"
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
|
||||
return_value=httpx.Response(
|
||||
status_code=200,
|
||||
json={"result": {"blocked": False, "transformed": False}},
|
||||
request=httpx.Request(method="POST", url=guardrail_endpoint),
|
||||
),
|
||||
) as mock_method:
|
||||
await crowdstrike_aidr_guardrail.apply_guardrail(
|
||||
inputs=inputs,
|
||||
request_data=request_data,
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
payload = mock_method.call_args.kwargs["json"]
|
||||
assert "user_id" not in payload
|
||||
assert "model" not in payload
|
||||
assert "extra_info" not in payload
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue