mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-13 23:11:40 +00:00
fix(guardrails): improve CrowdStrike AIDR input handling
Added stricter data models to ensure that everything is converted to the format that the CrowdStrike AIDR API expects. Also greatly simplified how LLM responses are handled while fixing streaming responses at the same time.
This commit is contained in:
parent
b3377b2d17
commit
69dcf54978
2 changed files with 231 additions and 170 deletions
|
|
@ -1,5 +1,8 @@
|
|||
from collections.abc import Mapping, Sequence
|
||||
import json
|
||||
import os
|
||||
from typing import TYPE_CHECKING, Literal, Optional, Type
|
||||
from typing import TYPE_CHECKING, Annotated, Literal, Optional, Type, Union, cast
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
from typing_extensions import Any, override
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
|
@ -16,6 +19,7 @@ from litellm.llms.custom_httpx.http_handler import (
|
|||
from litellm.proxy.common_utils.callback_utils import (
|
||||
add_guardrail_to_applied_guardrails_header,
|
||||
)
|
||||
from litellm.types.llms.openai import OpenAIChatCompletionToolParam
|
||||
from litellm.types.utils import GenericGuardrailAPIInputs
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -29,6 +33,78 @@ class CrowdStrikeAIDRGuardrailMissingSecrets(Exception):
|
|||
pass
|
||||
|
||||
|
||||
class _TextContentPart(BaseModel):
|
||||
model_config = ConfigDict(extra="forbid")
|
||||
|
||||
type: Literal["text"] = "text"
|
||||
text: str
|
||||
|
||||
|
||||
class _ImageUrl(BaseModel):
|
||||
url: str
|
||||
|
||||
|
||||
class _ImageUrlContentPart(BaseModel):
|
||||
model_config = ConfigDict(extra="forbid")
|
||||
|
||||
type: Literal["image_url"] = "image_url"
|
||||
image_url: _ImageUrl
|
||||
|
||||
|
||||
_ContentPart = Annotated[
|
||||
Union[_TextContentPart, _ImageUrlContentPart], Field(discriminator="type")
|
||||
]
|
||||
|
||||
|
||||
class _Message(BaseModel):
|
||||
role: str
|
||||
content: Optional[Union[str, list[_ContentPart]]] = None
|
||||
|
||||
|
||||
class _GuardInput(BaseModel):
|
||||
messages: list[_Message]
|
||||
tools: Optional[Sequence[OpenAIChatCompletionToolParam]] = None
|
||||
|
||||
|
||||
def _normalize_content(raw: object) -> str | list[_ContentPart] | None:
|
||||
if raw is None:
|
||||
return None
|
||||
if isinstance(raw, str):
|
||||
return raw
|
||||
if not isinstance(raw, list):
|
||||
return json.dumps(raw)
|
||||
parts: list[_ContentPart] = []
|
||||
for block in raw:
|
||||
if not isinstance(block, dict):
|
||||
parts.append(_TextContentPart(text=json.dumps(block)))
|
||||
continue
|
||||
|
||||
t = block.get("type")
|
||||
if t == "text" and isinstance(block.get("text"), str):
|
||||
parts.append(_TextContentPart(text=cast(str, block["text"])))
|
||||
elif t == "image_url":
|
||||
iu = block.get("image_url")
|
||||
url = iu if isinstance(iu, str) else str((iu or {}).get("url", ""))
|
||||
parts.append(_ImageUrlContentPart(image_url=_ImageUrl(url=url)))
|
||||
|
||||
# Any other types are not recognized by the CrowdStrike AIDR API.
|
||||
|
||||
return parts
|
||||
|
||||
|
||||
def _extract_text_from_content(content: object) -> str:
|
||||
if isinstance(content, str):
|
||||
return content
|
||||
if isinstance(content, list):
|
||||
parts = [
|
||||
item.get("text", "")
|
||||
for item in content
|
||||
if isinstance(item, dict) and item.get("type") == "text"
|
||||
]
|
||||
return "\n".join(parts)
|
||||
return ""
|
||||
|
||||
|
||||
class CrowdStrikeAIDRHandler(CustomGuardrail):
|
||||
"""
|
||||
CrowdStrike AIDR AI Guardrail handler to interact with the CrowdStrike AIDR
|
||||
|
|
@ -130,17 +206,23 @@ class CrowdStrikeAIDRHandler(CustomGuardrail):
|
|||
|
||||
def _build_guard_input_for_request(
|
||||
self, inputs: GenericGuardrailAPIInputs
|
||||
) -> Optional[dict[str, Any]]:
|
||||
guard_input: dict[str, Any] = {}
|
||||
) -> Optional[_GuardInput]:
|
||||
guard_input = _GuardInput(messages=[], tools=[])
|
||||
structured_messages = inputs.get("structured_messages")
|
||||
texts = inputs.get("texts", [])
|
||||
tools = inputs.get("tools")
|
||||
|
||||
if structured_messages:
|
||||
guard_input["messages"] = structured_messages
|
||||
for message in structured_messages:
|
||||
content = _normalize_content(message.get("content"))
|
||||
if content is None or len(content) == 0:
|
||||
content = ""
|
||||
guard_input.messages.append(
|
||||
_Message(role=message["role"], content=content)
|
||||
)
|
||||
elif texts:
|
||||
guard_input["messages"] = [
|
||||
{"role": "user", "content": text} for text in texts
|
||||
guard_input.messages = [
|
||||
_Message(role="user", content=text) for text in texts
|
||||
]
|
||||
else:
|
||||
verbose_proxy_logger.warning(
|
||||
|
|
@ -149,131 +231,53 @@ class CrowdStrikeAIDRHandler(CustomGuardrail):
|
|||
return None
|
||||
|
||||
if tools:
|
||||
guard_input["tools"] = tools
|
||||
guard_input.tools = tools
|
||||
|
||||
return guard_input
|
||||
|
||||
def _build_guard_input_for_response(
|
||||
self,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
request_data: dict,
|
||||
logging_obj: Optional["LiteLLMLoggingObj"],
|
||||
) -> Optional[dict[str, Any]]:
|
||||
guard_input: dict[str, Any] = {}
|
||||
response = request_data.get("response")
|
||||
if not response:
|
||||
self, inputs: GenericGuardrailAPIInputs, request_data: Mapping[str, Any]
|
||||
) -> Optional[_GuardInput]:
|
||||
output_texts: list[str] = inputs.get("texts", [])
|
||||
if len(output_texts) == 0:
|
||||
verbose_proxy_logger.warning(
|
||||
"CrowdStrike AIDR Guardrail: No response object in request_data for output response"
|
||||
"CrowdStrike AIDR Guardrail: No text in output response."
|
||||
)
|
||||
return None
|
||||
|
||||
# Extract choices from the response
|
||||
if hasattr(response, "choices") and response.choices:
|
||||
guard_input["choices"] = []
|
||||
for choice in response.choices:
|
||||
choice_dict = {}
|
||||
if hasattr(choice, "message"):
|
||||
message = choice.message
|
||||
choice_dict["message"] = {
|
||||
"role": getattr(message, "role", "assistant"),
|
||||
"content": getattr(message, "content", ""),
|
||||
}
|
||||
guard_input["choices"].append(choice_dict)
|
||||
input_messages = request_data.get("messages", [])
|
||||
|
||||
input_messages = None
|
||||
if "body" in request_data:
|
||||
input_messages = request_data["body"].get("messages")
|
||||
if not input_messages:
|
||||
input_messages = request_data.get("messages")
|
||||
if not input_messages and logging_obj:
|
||||
try:
|
||||
if hasattr(logging_obj, "model_call_details"):
|
||||
model_call_details = logging_obj.model_call_details
|
||||
if isinstance(model_call_details, dict):
|
||||
input_messages = model_call_details.get("messages")
|
||||
except Exception:
|
||||
pass
|
||||
return _GuardInput(
|
||||
messages=[
|
||||
_Message(role=role, content=content)
|
||||
for (role, content) in (
|
||||
(message["role"], _normalize_content(message.get("content")))
|
||||
for message in input_messages
|
||||
)
|
||||
if content is not None and len(content) > 0
|
||||
]
|
||||
+ [_Message(role="assistant", content=text) for text in output_texts]
|
||||
)
|
||||
|
||||
guard_input["messages"] = input_messages if input_messages else []
|
||||
|
||||
if tools := inputs.get("tools"):
|
||||
guard_input["tools"] = tools
|
||||
elif tools := request_data.get("body", {}).get("tools"):
|
||||
guard_input["tools"] = tools
|
||||
|
||||
return guard_input
|
||||
|
||||
def _extract_transformed_texts_from_messages(
|
||||
def _extract_transformed_texts(
|
||||
self,
|
||||
guard_output: dict[str, Any],
|
||||
structured_messages: Optional[list],
|
||||
texts: list[str],
|
||||
guard_output: Mapping[str, Any],
|
||||
num_assistant_messages: int,
|
||||
) -> list[str]:
|
||||
transformed_texts: list[str] = []
|
||||
transformed_messages = guard_output.get("messages", [])
|
||||
|
||||
if structured_messages and len(transformed_messages) == len(
|
||||
structured_messages
|
||||
):
|
||||
for msg in transformed_messages:
|
||||
if isinstance(msg, dict):
|
||||
content = msg.get("content")
|
||||
if isinstance(content, str):
|
||||
transformed_texts.append(content)
|
||||
elif isinstance(content, list):
|
||||
text_found = False
|
||||
for item in content:
|
||||
if isinstance(item, dict) and item.get("type") == "text":
|
||||
transformed_texts.append(item.get("text", ""))
|
||||
text_found = True
|
||||
break
|
||||
if not text_found:
|
||||
transformed_texts.append("")
|
||||
else:
|
||||
for msg in transformed_messages:
|
||||
if isinstance(msg, dict):
|
||||
content = msg.get("content")
|
||||
if isinstance(content, str):
|
||||
transformed_texts.append(content)
|
||||
elif isinstance(content, list):
|
||||
for item in content:
|
||||
if isinstance(item, dict) and item.get("type") == "text":
|
||||
transformed_texts.append(item.get("text", ""))
|
||||
break
|
||||
|
||||
while len(transformed_texts) < len(texts):
|
||||
transformed_texts.append(texts[len(transformed_texts)])
|
||||
return transformed_texts[: len(texts)]
|
||||
|
||||
def _extract_transformed_texts_from_choices(
|
||||
self, guard_output: dict[str, Any], texts: list[str]
|
||||
) -> list[str]:
|
||||
transformed_texts: list[str] = []
|
||||
transformed_choices = guard_output.get("choices", [])
|
||||
|
||||
for choice in transformed_choices:
|
||||
if isinstance(choice, dict):
|
||||
message = choice.get("message", {})
|
||||
content = message.get("content")
|
||||
if isinstance(content, str):
|
||||
transformed_texts.append(content)
|
||||
elif isinstance(content, list):
|
||||
text_found = False
|
||||
for item in content:
|
||||
if isinstance(item, dict) and item.get("type") == "text":
|
||||
transformed_texts.append(item.get("text", ""))
|
||||
text_found = True
|
||||
break
|
||||
if not text_found:
|
||||
transformed_texts.append("")
|
||||
else:
|
||||
transformed_texts.append("")
|
||||
else:
|
||||
transformed_texts.append("")
|
||||
|
||||
while len(transformed_texts) < len(texts):
|
||||
transformed_texts.append(texts[len(transformed_texts)])
|
||||
return transformed_texts[: len(texts)]
|
||||
tail = (
|
||||
transformed_messages[-num_assistant_messages:]
|
||||
if num_assistant_messages > 0
|
||||
else []
|
||||
)
|
||||
return [
|
||||
(
|
||||
_extract_text_from_content(msg.get("content"))
|
||||
if isinstance(msg, dict)
|
||||
else ""
|
||||
)
|
||||
for msg in tail
|
||||
]
|
||||
|
||||
@log_guardrail_information
|
||||
@override
|
||||
|
|
@ -302,16 +306,14 @@ class CrowdStrikeAIDRHandler(CustomGuardrail):
|
|||
event_type = "input"
|
||||
hook_name = "apply_guardrail (request)"
|
||||
else:
|
||||
guard_input = self._build_guard_input_for_response(
|
||||
inputs, request_data, logging_obj
|
||||
)
|
||||
guard_input = self._build_guard_input_for_response(inputs, request_data)
|
||||
if guard_input is None:
|
||||
return inputs
|
||||
event_type = "output"
|
||||
hook_name = "apply_guardrail (response)"
|
||||
|
||||
ai_guard_payload = {
|
||||
"guard_input": guard_input,
|
||||
"guard_input": guard_input.model_dump(mode="json"),
|
||||
"event_type": event_type,
|
||||
}
|
||||
|
||||
|
|
@ -326,18 +328,27 @@ class CrowdStrikeAIDRHandler(CustomGuardrail):
|
|||
|
||||
result = ai_guard_response.get("result", {})
|
||||
if not result.get("transformed"):
|
||||
# Not transformed, return original inputs.
|
||||
return inputs
|
||||
|
||||
guard_output = result.get("guard_output", {})
|
||||
|
||||
transformed_texts = (
|
||||
self._extract_transformed_texts_from_messages(
|
||||
guard_output, structured_messages, texts
|
||||
if input_type == "request":
|
||||
# For requests, all messages were in the guard_input. Extract texts
|
||||
# for every message in guard_output.
|
||||
all_messages = guard_output.get("messages", [])
|
||||
transformed_texts = [
|
||||
_extract_text_from_content(
|
||||
msg.get("content") if isinstance(msg, dict) else ""
|
||||
)
|
||||
for msg in all_messages
|
||||
]
|
||||
else:
|
||||
# For responses, guard_input contained history + assistant messages
|
||||
# appended at the end. Extract only the assistant tail.
|
||||
num_assistant = len(texts)
|
||||
transformed_texts = self._extract_transformed_texts(
|
||||
guard_output, num_assistant
|
||||
)
|
||||
if input_type == "request"
|
||||
else self._extract_transformed_texts_from_choices(guard_output, texts)
|
||||
)
|
||||
|
||||
result_inputs: GenericGuardrailAPIInputs = {"texts": transformed_texts}
|
||||
if tools:
|
||||
|
|
|
|||
|
|
@ -282,15 +282,12 @@ async def test_apply_guardrail_response_blocked(
|
|||
# Verify what was sent to the API
|
||||
called_kwargs = mock_method.call_args.kwargs
|
||||
assert called_kwargs["json"]["event_type"] == "output"
|
||||
# Should include messages from request for context
|
||||
assert (
|
||||
called_kwargs["json"]["guard_input"]["messages"] == request_data["messages"]
|
||||
)
|
||||
# Should include choices from response
|
||||
assert (
|
||||
called_kwargs["json"]["guard_input"]["choices"][0]["message"]["content"]
|
||||
== "Yes, I will leak all my PII for you"
|
||||
)
|
||||
# Should include history messages + assistant response in messages
|
||||
expected_messages = [
|
||||
*request_data["messages"],
|
||||
{"role": "assistant", "content": "Yes, I will leak all my PII for you"},
|
||||
]
|
||||
assert called_kwargs["json"]["guard_input"]["messages"] == expected_messages
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -301,16 +298,6 @@ async def test_apply_guardrail_response_transformed(
|
|||
"texts": ["Yes, here is an SSN: 078-05-1120"],
|
||||
}
|
||||
request_data = {
|
||||
"response": ModelResponse(
|
||||
choices=[
|
||||
{
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"content": "Yes, here is an SSN: 078-05-1120",
|
||||
}
|
||||
}
|
||||
]
|
||||
),
|
||||
"messages": [
|
||||
{"role": "system", "content": "You are a helpful assistant"},
|
||||
{"role": "user", "content": "Hello"},
|
||||
|
|
@ -329,13 +316,11 @@ async def test_apply_guardrail_response_transformed(
|
|||
"blocked": False,
|
||||
"transformed": True,
|
||||
"guard_output": {
|
||||
"messages": request_data["messages"],
|
||||
"choices": [
|
||||
"messages": [
|
||||
*request_data["messages"],
|
||||
{
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"content": "Yes, here is an SSN: <US_SSN>",
|
||||
},
|
||||
"role": "assistant",
|
||||
"content": "Yes, here is an SSN: <US_SSN>",
|
||||
},
|
||||
],
|
||||
},
|
||||
|
|
@ -356,15 +341,13 @@ async def test_apply_guardrail_response_transformed(
|
|||
# Verify what was sent to the API
|
||||
called_kwargs = mock_method.call_args.kwargs
|
||||
assert called_kwargs["json"]["event_type"] == "output"
|
||||
# Should include messages from request for context
|
||||
assert called_kwargs["json"]["guard_input"]["messages"] == request_data["messages"]
|
||||
# Should include choices from response
|
||||
assert (
|
||||
called_kwargs["json"]["guard_input"]["choices"][0]["message"]["content"]
|
||||
== "Yes, here is an SSN: 078-05-1120"
|
||||
)
|
||||
# Verify the transformed output
|
||||
assert result["texts"][0] == "Yes, here is an SSN: <US_SSN>"
|
||||
# Should include history + assistant in messages
|
||||
assert called_kwargs["json"]["guard_input"]["messages"] == [
|
||||
*request_data["messages"],
|
||||
{"role": "assistant", "content": "Yes, here is an SSN: 078-05-1120"},
|
||||
]
|
||||
# Verify the transformed output extracts only the assistant message
|
||||
assert result["texts"] == ["Yes, here is an SSN: <US_SSN>"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -419,12 +402,79 @@ async def test_apply_guardrail_response_ok(
|
|||
# Verify what was sent to the API
|
||||
called_kwargs = mock_method.call_args.kwargs
|
||||
assert called_kwargs["json"]["event_type"] == "output"
|
||||
# Should include messages from request for context
|
||||
assert called_kwargs["json"]["guard_input"]["messages"] == request_data["messages"]
|
||||
# Should include choices from response
|
||||
assert (
|
||||
called_kwargs["json"]["guard_input"]["choices"][0]["message"]["content"]
|
||||
== "Hello! How can I help you today?"
|
||||
)
|
||||
# Should include history + assistant in messages
|
||||
expected_messages = [
|
||||
*request_data["messages"],
|
||||
{"role": "assistant", "content": "Hello! How can I help you today?"},
|
||||
]
|
||||
assert called_kwargs["json"]["guard_input"]["messages"] == expected_messages
|
||||
# Should return original inputs when not transformed
|
||||
assert result["texts"] == inputs["texts"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_request_skipped_messages_stay_aligned(
|
||||
crowdstrike_aidr_guardrail: CrowdStrikeAIDRHandler,
|
||||
) -> None:
|
||||
inputs: GenericGuardrailAPIInputs = {
|
||||
"texts": [
|
||||
"Hello, help me with my task",
|
||||
"",
|
||||
"Here is my SSN: 078-05-1120",
|
||||
],
|
||||
"structured_messages": [
|
||||
{"role": "user", "content": "Hello, help me with my task"},
|
||||
{
|
||||
"role": "tool",
|
||||
"content": [
|
||||
{"type": "tool_result", "tool_use_id": "t1", "content": "ok"}
|
||||
],
|
||||
},
|
||||
{"role": "user", "content": "Here is my SSN: 078-05-1120"},
|
||||
],
|
||||
}
|
||||
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": True,
|
||||
"guard_output": {
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": "Hello, help me with my task",
|
||||
},
|
||||
{
|
||||
"role": "tool",
|
||||
"content": "",
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": "Here is my SSN: <US_SSN>",
|
||||
},
|
||||
]
|
||||
},
|
||||
},
|
||||
},
|
||||
request=httpx.Request(method="POST", url=guardrail_endpoint),
|
||||
),
|
||||
):
|
||||
result = await crowdstrike_aidr_guardrail.apply_guardrail(
|
||||
inputs=inputs,
|
||||
request_data=request_data,
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
assert len(result["texts"]) == len(inputs["structured_messages"])
|
||||
assert result["texts"][0] == "Hello, help me with my task"
|
||||
assert result["texts"][1] == ""
|
||||
assert result["texts"][2] == "Here is my SSN: <US_SSN>"
|
||||
assert result["structured_messages"] == inputs["structured_messages"]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue