mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix(guardrails): send only new messages since last assistant turn to CrowdStrike AIDR (#31974)
Previously, every guardrail request forwarded the full conversation history to CrowdStrike AIDR. In a multi-turn conversation this means every prior message gets re-scanned on every new call, even though those messages were already evaluated in earlier turns. CrowdStrike AIDR internally has a conversation boundary optimization in place for just this scenario (ref. <https://aidr-docs.crowdstrike.com/docs/aidr/apis#messages-array-optional---array-of-message-objects-containing-a-conversation-segment-with-the-ai-system>). However, it is nevertheless wasteful to send so much data to the API when only a subset of it will be processed. It also risks hitting the documented 1 MiB request size limit. So now we filter down to system messages plus either the messages after the last assistant turn, or the last assistant message itself when that is what is being guarded. We also preserve the original, full message history within the guardrail in order to stitch back any transformations. Co-authored-by: Kenan Yildirim <kenan@kenany.me>
This commit is contained in:
parent
36a7860c10
commit
4428c1b681
2 changed files with 925 additions and 137 deletions
|
|
@ -1,17 +1,29 @@
|
|||
from collections.abc import Mapping, Sequence
|
||||
import json
|
||||
import os
|
||||
from typing import TYPE_CHECKING, Annotated, Literal, Optional, Type, Union, cast
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
from typing_extensions import Any, override
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Annotated,
|
||||
Literal,
|
||||
NamedTuple,
|
||||
Optional,
|
||||
Union,
|
||||
cast,
|
||||
)
|
||||
|
||||
from fastapi import HTTPException
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
from typing_extensions import Any, override
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.integrations.custom_guardrail import (
|
||||
CustomGuardrail,
|
||||
log_guardrail_information,
|
||||
)
|
||||
from litellm.llms.base_llm.guardrail_translation.utils import (
|
||||
effective_skip_system_message_for_guardrail,
|
||||
effective_skip_tool_message_for_guardrail,
|
||||
)
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
get_async_httpx_client,
|
||||
httpxSpecialProvider,
|
||||
|
|
@ -19,7 +31,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.llms.openai import AllMessageValues, OpenAIChatCompletionToolParam
|
||||
from litellm.types.utils import GenericGuardrailAPIInputs
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -64,6 +76,37 @@ class _GuardInput(BaseModel):
|
|||
tools: Optional[Sequence[OpenAIChatCompletionToolParam]] = None
|
||||
|
||||
|
||||
class _GuardChatCompletionsResult(BaseModel):
|
||||
guard_output: Optional[_GuardInput] = None
|
||||
"""Updated structured prompt."""
|
||||
blocked: Optional[bool] = None
|
||||
"""Whether or not the prompt triggered a block detection."""
|
||||
transformed: Optional[bool] = None
|
||||
"""Whether or not the original input was transformed."""
|
||||
detectors: Optional[dict[str, Any]] = None
|
||||
"""Result of the policy analyzing and input prompt."""
|
||||
|
||||
|
||||
class _GuardChatCompletionsResponse(BaseModel):
|
||||
result: Optional[_GuardChatCompletionsResult] = None
|
||||
|
||||
|
||||
class _FilteredMessages(NamedTuple):
|
||||
"""Subset of a conversation selected for guardrail analysis."""
|
||||
|
||||
messages: list[AllMessageValues]
|
||||
"""Messages subset."""
|
||||
indices: tuple[int, ...]
|
||||
"""Positions of the subset's messages in the original list."""
|
||||
|
||||
|
||||
class _GuardInputWithIndices(NamedTuple):
|
||||
guard_input: _GuardInput
|
||||
"""Guard API payload."""
|
||||
sent_indices: tuple[int, ...]
|
||||
"""Positions of the guard input's messages in the original list."""
|
||||
|
||||
|
||||
def _normalize_content(raw: object) -> str | list[_ContentPart] | None:
|
||||
if raw is None:
|
||||
return None
|
||||
|
|
@ -99,7 +142,16 @@ def _extract_text_from_content(content: object) -> str:
|
|||
return ""
|
||||
|
||||
|
||||
def _merge_metadata_bags(request_data: Mapping[str, Any]) -> Optional[dict[str, Any]]:
|
||||
def _extract_text_from_message(message: _Message) -> str:
|
||||
content = message.content
|
||||
if isinstance(content, str):
|
||||
return content
|
||||
if content is None:
|
||||
return ""
|
||||
return "\n".join(part.text for part in content if isinstance(part, _TextContentPart))
|
||||
|
||||
|
||||
def _merge_metadata_bags(request_data: Mapping[str, Any]) -> dict[str, Any] | None:
|
||||
merged: dict[str, Any] = {}
|
||||
present = False
|
||||
for bag in (request_data.get("metadata"), request_data.get("litellm_metadata")):
|
||||
|
|
@ -109,6 +161,75 @@ def _merge_metadata_bags(request_data: Mapping[str, Any]) -> Optional[dict[str,
|
|||
return merged if present else None
|
||||
|
||||
|
||||
def _messages_since_last_assistant(
|
||||
messages: list[AllMessageValues],
|
||||
) -> _FilteredMessages:
|
||||
if not messages:
|
||||
return _FilteredMessages([], ())
|
||||
|
||||
if messages[-1]["role"] == "assistant":
|
||||
indices = tuple(i for i, m in enumerate(messages) if m["role"] == "system") + (len(messages) - 1,)
|
||||
return _FilteredMessages([messages[i] for i in indices], indices)
|
||||
|
||||
last_assistant_idx = -1
|
||||
for i in range(len(messages) - 1, -1, -1):
|
||||
if messages[i]["role"] == "assistant":
|
||||
last_assistant_idx = i
|
||||
break
|
||||
|
||||
system_indices = tuple(i for i in range(last_assistant_idx + 1) if messages[i]["role"] == "system")
|
||||
tail_indices = tuple(range(last_assistant_idx + 1, len(messages)))
|
||||
indices = system_indices + tail_indices
|
||||
return _FilteredMessages([messages[i] for i in indices], indices)
|
||||
|
||||
|
||||
def _merge_request_transforms(
|
||||
guard_output: _GuardInput,
|
||||
structured_messages: list[AllMessageValues] | None,
|
||||
texts: list[str],
|
||||
sent_indices: tuple[int, ...],
|
||||
) -> list[str]:
|
||||
returned_texts = [_extract_text_from_message(msg) for msg in guard_output.messages]
|
||||
original_texts = (
|
||||
[_extract_text_from_content(m.get("content")) for m in structured_messages] if structured_messages else texts
|
||||
)
|
||||
replacements = {
|
||||
idx: returned_texts[pos]
|
||||
for pos, idx in enumerate(sent_indices)
|
||||
if pos < len(returned_texts) and idx < len(original_texts)
|
||||
}
|
||||
return [replacements.get(idx, original) for idx, original in enumerate(original_texts)]
|
||||
|
||||
|
||||
def _apply_message_redaction(original: AllMessageValues, redacted: _Message) -> AllMessageValues:
|
||||
content = original.get("content")
|
||||
if isinstance(content, str):
|
||||
return cast(AllMessageValues, {**original, "content": _extract_text_from_message(redacted)})
|
||||
if isinstance(content, list) and _extract_text_from_content(content):
|
||||
redacted_content = redacted.content
|
||||
new_content = (
|
||||
[part.model_dump() for part in redacted_content] if isinstance(redacted_content, list) else redacted_content
|
||||
)
|
||||
return cast(AllMessageValues, {**original, "content": new_content})
|
||||
return original
|
||||
|
||||
|
||||
def _redacted_messages(
|
||||
processed_messages: list[AllMessageValues],
|
||||
guard_output: _GuardInput,
|
||||
sent_indices: tuple[int, ...],
|
||||
full_messages: list[AllMessageValues],
|
||||
) -> list[AllMessageValues] | None:
|
||||
redactions = {
|
||||
id(processed_messages[idx]): _apply_message_redaction(processed_messages[idx], guard_output.messages[pos])
|
||||
for pos, idx in enumerate(sent_indices)
|
||||
if pos < len(guard_output.messages) and idx < len(processed_messages)
|
||||
}
|
||||
if not redactions.keys() <= {id(message) for message in full_messages}:
|
||||
return None
|
||||
return [redactions.get(id(message), message) for message in full_messages]
|
||||
|
||||
|
||||
class CrowdStrikeAIDRHandler(CustomGuardrail):
|
||||
"""
|
||||
CrowdStrike AIDR AI Guardrail handler to interact with the CrowdStrike AIDR
|
||||
|
|
@ -118,17 +239,17 @@ class CrowdStrikeAIDRHandler(CustomGuardrail):
|
|||
def __init__(
|
||||
self,
|
||||
guardrail_name: str,
|
||||
api_key: Optional[str] = None,
|
||||
api_base: Optional[str] = None,
|
||||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
**kwargs,
|
||||
):
|
||||
) -> None:
|
||||
"""
|
||||
Initializes the CrowdStrikeAIDRHandler.
|
||||
|
||||
Args:
|
||||
guardrail_name (str): The name of the guardrail instance.
|
||||
api_key (Optional[str]): The CrowdStrike AIDR API key. Reads from CS_AIDR_TOKEN env var if None.
|
||||
api_base (Optional[str]): The CrowdStrike AIDR API base URL. Reads from CS_AIDR_BASE_URL env var if None.
|
||||
api_key (str | None): The CrowdStrike AIDR API key. Reads from CS_AIDR_TOKEN env var if None.
|
||||
api_base (str | None): The CrowdStrike AIDR API base URL. Reads from CS_AIDR_BASE_URL env var if None.
|
||||
**kwargs: Additional arguments passed to the CustomGuardrail base class.
|
||||
"""
|
||||
self.async_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.GuardrailCallback)
|
||||
|
|
@ -151,7 +272,9 @@ class CrowdStrikeAIDRHandler(CustomGuardrail):
|
|||
f"Initialized CrowdStrike AIDR Guardrail: name={guardrail_name}, api_base={self.api_base}"
|
||||
)
|
||||
|
||||
async def _call_crowdstrike_aidr_guard(self, payload: dict[str, Any], hook_name: str) -> dict[str, Any]:
|
||||
async def _call_crowdstrike_aidr_guard(
|
||||
self, payload: dict[str, Any], hook_name: str
|
||||
) -> _GuardChatCompletionsResult:
|
||||
"""
|
||||
Makes the API call to the CrowdStrike AIDR AI Guard endpoint.
|
||||
The function itself will raise an error if a response should be blocked,
|
||||
|
|
@ -167,7 +290,7 @@ class CrowdStrikeAIDRHandler(CustomGuardrail):
|
|||
Exception: For other API call failures.
|
||||
|
||||
Returns:
|
||||
dict: The API response body
|
||||
The parsed `result` body of the API response.
|
||||
"""
|
||||
endpoint = f"{self.api_base}/v1/guard_chat_completions"
|
||||
|
||||
|
|
@ -181,11 +304,12 @@ class CrowdStrikeAIDRHandler(CustomGuardrail):
|
|||
)
|
||||
|
||||
response = await self.async_handler.post(url=endpoint, json=payload, headers=headers)
|
||||
assert response is not None
|
||||
response.raise_for_status()
|
||||
|
||||
result: dict[str, Any] = response.json()
|
||||
result = _GuardChatCompletionsResponse.model_validate(response.json()).result or _GuardChatCompletionsResult()
|
||||
|
||||
if result.get("result", {}).get("blocked"):
|
||||
if result.blocked:
|
||||
verbose_proxy_logger.warning(
|
||||
f"CrowdStrike AIDR Guardrail ({hook_name}): Request blocked. Response: {result}"
|
||||
)
|
||||
|
|
@ -197,25 +321,28 @@ class CrowdStrikeAIDRHandler(CustomGuardrail):
|
|||
},
|
||||
)
|
||||
verbose_proxy_logger.debug(
|
||||
f"CrowdStrike AIDR Guardrail ({hook_name}): Request passed. Response: {result.get('result', {}).get('detectors')}"
|
||||
f"CrowdStrike AIDR Guardrail ({hook_name}): Request passed. Response: {result.detectors}"
|
||||
)
|
||||
|
||||
return result
|
||||
|
||||
def _build_guard_input_for_request(self, inputs: GenericGuardrailAPIInputs) -> Optional[_GuardInput]:
|
||||
def _build_guard_input_for_request(self, inputs: GenericGuardrailAPIInputs) -> _GuardInputWithIndices | None:
|
||||
guard_input = _GuardInput(messages=[], tools=[])
|
||||
structured_messages = inputs.get("structured_messages")
|
||||
texts = inputs.get("texts", [])
|
||||
tools = inputs.get("tools")
|
||||
|
||||
if structured_messages:
|
||||
for message in structured_messages:
|
||||
filtered = _messages_since_last_assistant(structured_messages)
|
||||
for message in filtered.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))
|
||||
indices = filtered.indices
|
||||
elif texts:
|
||||
guard_input.messages = [_Message(role="user", content=text) for text in texts]
|
||||
indices = tuple(range(len(texts)))
|
||||
else:
|
||||
verbose_proxy_logger.warning("CrowdStrike AIDR Guardrail: No messages or texts provided for input request")
|
||||
return None
|
||||
|
|
@ -223,37 +350,36 @@ class CrowdStrikeAIDRHandler(CustomGuardrail):
|
|||
if tools:
|
||||
guard_input.tools = tools
|
||||
|
||||
return guard_input
|
||||
return _GuardInputWithIndices(guard_input, indices)
|
||||
|
||||
def _build_guard_input_for_response(
|
||||
self, inputs: GenericGuardrailAPIInputs, request_data: Mapping[str, Any]
|
||||
) -> Optional[_GuardInput]:
|
||||
def _build_guard_input_for_response(self, inputs: GenericGuardrailAPIInputs) -> _GuardInput:
|
||||
output_texts: list[str] = inputs.get("texts", [])
|
||||
if len(output_texts) == 0:
|
||||
verbose_proxy_logger.warning("CrowdStrike AIDR Guardrail: No text in output response.")
|
||||
return None
|
||||
|
||||
input_messages = request_data.get("messages", [])
|
||||
|
||||
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]
|
||||
messages=[_Message(role="assistant", content=text) for text in output_texts],
|
||||
tools=inputs.get("tools", []),
|
||||
)
|
||||
|
||||
def _extract_transformed_texts(
|
||||
def _extract_transformed_texts(self, guard_output: _GuardInput, num_assistant_messages: int) -> list[str]:
|
||||
tail = guard_output.messages[-num_assistant_messages:] if num_assistant_messages > 0 else []
|
||||
return [_extract_text_from_message(msg) for msg in tail]
|
||||
|
||||
def _writeback_messages(
|
||||
self,
|
||||
guard_output: Mapping[str, Any],
|
||||
num_assistant_messages: int,
|
||||
) -> list[str]:
|
||||
transformed_messages = guard_output.get("messages", [])
|
||||
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]
|
||||
structured_messages: list[AllMessageValues],
|
||||
guard_output: _GuardInput,
|
||||
sent_indices: tuple[int, ...],
|
||||
request_data: dict,
|
||||
) -> list[AllMessageValues] | None:
|
||||
if effective_skip_system_message_for_guardrail(self) or effective_skip_tool_message_for_guardrail(self):
|
||||
request_messages = request_data.get("messages")
|
||||
full_messages = (
|
||||
cast("list[AllMessageValues]", request_messages)
|
||||
if isinstance(request_messages, list)
|
||||
else structured_messages
|
||||
)
|
||||
else:
|
||||
full_messages = structured_messages
|
||||
return _redacted_messages(structured_messages, guard_output, sent_indices, full_messages)
|
||||
|
||||
@log_guardrail_information
|
||||
@override
|
||||
|
|
@ -273,15 +399,18 @@ class CrowdStrikeAIDRHandler(CustomGuardrail):
|
|||
tool_calls = inputs.get("tool_calls")
|
||||
|
||||
# Build guard_input based on input_type
|
||||
sent_indices: tuple[int, ...] = ()
|
||||
if input_type == "request":
|
||||
guard_input = self._build_guard_input_for_request(inputs)
|
||||
if guard_input is None:
|
||||
request_result = self._build_guard_input_for_request(inputs)
|
||||
if request_result is None:
|
||||
return inputs
|
||||
guard_input = request_result.guard_input
|
||||
sent_indices = request_result.sent_indices
|
||||
event_type = "input"
|
||||
hook_name = "apply_guardrail (request)"
|
||||
else:
|
||||
guard_input = self._build_guard_input_for_response(inputs, request_data)
|
||||
if guard_input is None:
|
||||
guard_input = self._build_guard_input_for_response(inputs)
|
||||
if len(guard_input.messages) == 0:
|
||||
return inputs
|
||||
event_type = "output"
|
||||
hook_name = "apply_guardrail (response)"
|
||||
|
|
@ -307,29 +436,20 @@ class CrowdStrikeAIDRHandler(CustomGuardrail):
|
|||
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)
|
||||
result = await self._call_crowdstrike_aidr_guard(ai_guard_payload, hook_name)
|
||||
|
||||
if "body" in request_data or "messages" in request_data:
|
||||
add_guardrail_to_applied_guardrails_header(request_data=request_data, guardrail_name=self.guardrail_name)
|
||||
|
||||
result = ai_guard_response.get("result", {})
|
||||
if not result.get("transformed"):
|
||||
if not result.transformed or result.guard_output is None:
|
||||
return inputs
|
||||
|
||||
guard_output = result.get("guard_output", {})
|
||||
guard_output = result.guard_output
|
||||
|
||||
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
|
||||
]
|
||||
transformed_texts = _merge_request_transforms(guard_output, structured_messages, texts, sent_indices)
|
||||
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)
|
||||
transformed_texts = self._extract_transformed_texts(guard_output, len(texts))
|
||||
|
||||
result_inputs: GenericGuardrailAPIInputs = {"texts": transformed_texts}
|
||||
if tools:
|
||||
|
|
@ -337,13 +457,18 @@ class CrowdStrikeAIDRHandler(CustomGuardrail):
|
|||
if tool_calls:
|
||||
result_inputs["tool_calls"] = tool_calls
|
||||
if structured_messages:
|
||||
result_inputs["structured_messages"] = structured_messages
|
||||
rebuilt = (
|
||||
self._writeback_messages(structured_messages, guard_output, sent_indices, request_data)
|
||||
if input_type == "request"
|
||||
else None
|
||||
)
|
||||
result_inputs["structured_messages"] = rebuilt if rebuilt is not None else structured_messages
|
||||
|
||||
return result_inputs
|
||||
|
||||
@override
|
||||
@staticmethod
|
||||
def get_config_model() -> Optional[Type["GuardrailConfigModel"]]:
|
||||
def get_config_model() -> type["GuardrailConfigModel"] | None:
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.crowdstrike_aidr import (
|
||||
CrowdStrikeAIDRGuardrailConfigModel,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -93,9 +93,7 @@ async def test_apply_guardrail_request_blocked(
|
|||
],
|
||||
}
|
||||
request_data = {"messages": inputs["structured_messages"]}
|
||||
guardrail_endpoint = (
|
||||
f"{crowdstrike_aidr_guardrail.api_base}/v1/guard_chat_completions"
|
||||
)
|
||||
guardrail_endpoint = f"{crowdstrike_aidr_guardrail.api_base}/v1/guard_chat_completions"
|
||||
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
|
||||
|
|
@ -108,9 +106,7 @@ async def test_apply_guardrail_request_blocked(
|
|||
),
|
||||
),
|
||||
) as mock_method:
|
||||
with pytest.raises(
|
||||
HTTPException, match="Violated CrowdStrike AIDR guardrail policy"
|
||||
):
|
||||
with pytest.raises(HTTPException, match="Violated CrowdStrike AIDR guardrail policy"):
|
||||
await crowdstrike_aidr_guardrail.apply_guardrail(
|
||||
inputs=inputs,
|
||||
request_data=request_data,
|
||||
|
|
@ -121,10 +117,7 @@ async def test_apply_guardrail_request_blocked(
|
|||
called_kwargs = mock_method.call_args.kwargs
|
||||
assert called_kwargs["json"]["event_type"] == "input"
|
||||
# Should include messages
|
||||
assert (
|
||||
called_kwargs["json"]["guard_input"]["messages"]
|
||||
== inputs["structured_messages"]
|
||||
)
|
||||
assert called_kwargs["json"]["guard_input"]["messages"] == inputs["structured_messages"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -141,9 +134,7 @@ async def test_apply_guardrail_request_transformed(
|
|||
],
|
||||
}
|
||||
request_data = {"messages": inputs["structured_messages"]}
|
||||
guardrail_endpoint = (
|
||||
f"{crowdstrike_aidr_guardrail.api_base}/v1/guard_chat_completions"
|
||||
)
|
||||
guardrail_endpoint = f"{crowdstrike_aidr_guardrail.api_base}/v1/guard_chat_completions"
|
||||
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
|
||||
|
|
@ -179,10 +170,7 @@ async def test_apply_guardrail_request_transformed(
|
|||
called_kwargs = mock_method.call_args.kwargs
|
||||
assert called_kwargs["json"]["event_type"] == "input"
|
||||
# Should include messages
|
||||
assert (
|
||||
called_kwargs["json"]["guard_input"]["messages"]
|
||||
== inputs["structured_messages"]
|
||||
)
|
||||
assert called_kwargs["json"]["guard_input"]["messages"] == inputs["structured_messages"]
|
||||
# Verify the transformed output
|
||||
assert result["texts"][0] == "Here is an SSN for one my employees: <US_SSN>"
|
||||
|
||||
|
|
@ -196,9 +184,7 @@ async def test_apply_guardrail_request_ok(
|
|||
"structured_messages": [{"role": "user", "content": "Hello, how are you?"}],
|
||||
}
|
||||
request_data = {"messages": inputs["structured_messages"]}
|
||||
guardrail_endpoint = (
|
||||
f"{crowdstrike_aidr_guardrail.api_base}/v1/guard_chat_completions"
|
||||
)
|
||||
guardrail_endpoint = f"{crowdstrike_aidr_guardrail.api_base}/v1/guard_chat_completions"
|
||||
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
|
||||
|
|
@ -221,10 +207,7 @@ async def test_apply_guardrail_request_ok(
|
|||
called_kwargs = mock_method.call_args.kwargs
|
||||
assert called_kwargs["json"]["event_type"] == "input"
|
||||
# Should include messages
|
||||
assert (
|
||||
called_kwargs["json"]["guard_input"]["messages"]
|
||||
== inputs["structured_messages"]
|
||||
)
|
||||
assert called_kwargs["json"]["guard_input"]["messages"] == inputs["structured_messages"]
|
||||
# Should return original inputs when not transformed
|
||||
assert result["texts"] == inputs["texts"]
|
||||
|
||||
|
|
@ -252,9 +235,7 @@ async def test_apply_guardrail_response_blocked(
|
|||
{"role": "user", "content": "Hello"},
|
||||
],
|
||||
}
|
||||
guardrail_endpoint = (
|
||||
f"{crowdstrike_aidr_guardrail.api_base}/v1/guard_chat_completions"
|
||||
)
|
||||
guardrail_endpoint = f"{crowdstrike_aidr_guardrail.api_base}/v1/guard_chat_completions"
|
||||
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
|
||||
|
|
@ -272,22 +253,20 @@ async def test_apply_guardrail_response_blocked(
|
|||
),
|
||||
),
|
||||
) as mock_method:
|
||||
with pytest.raises(
|
||||
HTTPException, match="Violated CrowdStrike AIDR guardrail policy"
|
||||
):
|
||||
with pytest.raises(HTTPException, match="Violated CrowdStrike AIDR guardrail policy"):
|
||||
await crowdstrike_aidr_guardrail.apply_guardrail(
|
||||
inputs=inputs,
|
||||
request_data=request_data,
|
||||
input_type="response",
|
||||
)
|
||||
|
||||
# Verify what was sent to the API
|
||||
called_kwargs = mock_method.call_args.kwargs
|
||||
assert called_kwargs["json"]["event_type"] == "output"
|
||||
# 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"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "Yes, I will leak all my PII for you",
|
||||
},
|
||||
]
|
||||
assert called_kwargs["json"]["guard_input"]["messages"] == expected_messages
|
||||
|
||||
|
|
@ -305,9 +284,7 @@ async def test_apply_guardrail_response_transformed(
|
|||
{"role": "user", "content": "Hello"},
|
||||
],
|
||||
}
|
||||
guardrail_endpoint = (
|
||||
f"{crowdstrike_aidr_guardrail.api_base}/v1/guard_chat_completions"
|
||||
)
|
||||
guardrail_endpoint = f"{crowdstrike_aidr_guardrail.api_base}/v1/guard_chat_completions"
|
||||
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
|
||||
|
|
@ -319,7 +296,6 @@ async def test_apply_guardrail_response_transformed(
|
|||
"transformed": True,
|
||||
"guard_output": {
|
||||
"messages": [
|
||||
*request_data["messages"],
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "Yes, here is an SSN: <US_SSN>",
|
||||
|
|
@ -340,15 +316,14 @@ async def test_apply_guardrail_response_transformed(
|
|||
input_type="response",
|
||||
)
|
||||
|
||||
# Verify what was sent to the API
|
||||
called_kwargs = mock_method.call_args.kwargs
|
||||
assert called_kwargs["json"]["event_type"] == "output"
|
||||
# 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"},
|
||||
{
|
||||
"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>"]
|
||||
|
||||
|
||||
|
|
@ -375,9 +350,7 @@ async def test_apply_guardrail_response_ok(
|
|||
{"role": "user", "content": "Hello"},
|
||||
],
|
||||
}
|
||||
guardrail_endpoint = (
|
||||
f"{crowdstrike_aidr_guardrail.api_base}/v1/guard_chat_completions"
|
||||
)
|
||||
guardrail_endpoint = f"{crowdstrike_aidr_guardrail.api_base}/v1/guard_chat_completions"
|
||||
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
|
||||
|
|
@ -401,13 +374,13 @@ async def test_apply_guardrail_response_ok(
|
|||
input_type="response",
|
||||
)
|
||||
|
||||
# Verify what was sent to the API
|
||||
called_kwargs = mock_method.call_args.kwargs
|
||||
assert called_kwargs["json"]["event_type"] == "output"
|
||||
# Should include history + assistant in messages
|
||||
expected_messages = [
|
||||
*request_data["messages"],
|
||||
{"role": "assistant", "content": "Hello! How can I help you today?"},
|
||||
{
|
||||
"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
|
||||
|
|
@ -431,9 +404,7 @@ async def test_apply_guardrail_sends_user_id_model_and_extra_info(
|
|||
"user_api_key_user_email": "alice@example.com",
|
||||
},
|
||||
}
|
||||
guardrail_endpoint = (
|
||||
f"{crowdstrike_aidr_guardrail.api_base}/v1/guard_chat_completions"
|
||||
)
|
||||
guardrail_endpoint = f"{crowdstrike_aidr_guardrail.api_base}/v1/guard_chat_completions"
|
||||
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
|
||||
|
|
@ -472,9 +443,7 @@ async def test_apply_guardrail_empty_extra_info_when_no_email(
|
|||
"user_api_key_user_email": None,
|
||||
},
|
||||
}
|
||||
guardrail_endpoint = (
|
||||
f"{crowdstrike_aidr_guardrail.api_base}/v1/guard_chat_completions"
|
||||
)
|
||||
guardrail_endpoint = f"{crowdstrike_aidr_guardrail.api_base}/v1/guard_chat_completions"
|
||||
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
|
||||
|
|
@ -505,9 +474,7 @@ async def test_apply_guardrail_no_metadata_skips_user_fields(
|
|||
"structured_messages": [{"role": "user", "content": "Hello"}],
|
||||
}
|
||||
request_data = {"messages": inputs["structured_messages"]}
|
||||
guardrail_endpoint = (
|
||||
f"{crowdstrike_aidr_guardrail.api_base}/v1/guard_chat_completions"
|
||||
)
|
||||
guardrail_endpoint = f"{crowdstrike_aidr_guardrail.api_base}/v1/guard_chat_completions"
|
||||
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
|
||||
|
|
@ -533,12 +500,41 @@ async def test_apply_guardrail_no_metadata_skips_user_fields(
|
|||
@pytest.mark.parametrize(
|
||||
"litellm_metadata, metadata",
|
||||
[
|
||||
(None, {"user_api_key_user_id": "uid-abc", "user_api_key_user_email": "alice@example.com"}),
|
||||
({"trace_id": "t1"}, {"user_api_key_user_id": "uid-abc", "user_api_key_user_email": "alice@example.com"}),
|
||||
(["unexpected"], {"user_api_key_user_id": "uid-abc", "user_api_key_user_email": "alice@example.com"}),
|
||||
({"user_api_key_user_id": "uid-abc", "user_api_key_user_email": "alice@example.com"}, {"trace_id": "t1"}),
|
||||
(
|
||||
None,
|
||||
{
|
||||
"user_api_key_user_id": "uid-abc",
|
||||
"user_api_key_user_email": "alice@example.com",
|
||||
},
|
||||
),
|
||||
(
|
||||
{"trace_id": "t1"},
|
||||
{
|
||||
"user_api_key_user_id": "uid-abc",
|
||||
"user_api_key_user_email": "alice@example.com",
|
||||
},
|
||||
),
|
||||
(
|
||||
["unexpected"],
|
||||
{
|
||||
"user_api_key_user_id": "uid-abc",
|
||||
"user_api_key_user_email": "alice@example.com",
|
||||
},
|
||||
),
|
||||
(
|
||||
{
|
||||
"user_api_key_user_id": "uid-abc",
|
||||
"user_api_key_user_email": "alice@example.com",
|
||||
},
|
||||
{"trace_id": "t1"},
|
||||
),
|
||||
],
|
||||
ids=[
|
||||
"identity_in_metadata_llm_none",
|
||||
"identity_in_metadata_llm_user_dict",
|
||||
"identity_in_metadata_llm_non_mapping",
|
||||
"identity_in_litellm_metadata",
|
||||
],
|
||||
ids=["identity_in_metadata_llm_none", "identity_in_metadata_llm_user_dict", "identity_in_metadata_llm_non_mapping", "identity_in_litellm_metadata"],
|
||||
)
|
||||
async def test_apply_guardrail_reads_identity_from_either_metadata_bag(
|
||||
crowdstrike_aidr_guardrail: CrowdStrikeAIDRHandler,
|
||||
|
|
@ -556,9 +552,7 @@ async def test_apply_guardrail_reads_identity_from_either_metadata_bag(
|
|||
"litellm_metadata": litellm_metadata,
|
||||
"metadata": metadata,
|
||||
}
|
||||
guardrail_endpoint = (
|
||||
f"{crowdstrike_aidr_guardrail.api_base}/v1/guard_chat_completions"
|
||||
)
|
||||
guardrail_endpoint = f"{crowdstrike_aidr_guardrail.api_base}/v1/guard_chat_completions"
|
||||
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
|
||||
|
|
@ -593,17 +587,13 @@ async def test_apply_guardrail_request_skipped_messages_stay_aligned(
|
|||
{"role": "user", "content": "Hello, help me with my task"},
|
||||
{
|
||||
"role": "tool",
|
||||
"content": [
|
||||
{"type": "tool_result", "tool_use_id": "t1", "content": "ok"}
|
||||
],
|
||||
"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"
|
||||
)
|
||||
guardrail_endpoint = f"{crowdstrike_aidr_guardrail.api_base}/v1/guard_chat_completions"
|
||||
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
|
||||
|
|
@ -644,4 +634,677 @@ async def test_apply_guardrail_request_skipped_messages_stay_aligned(
|
|||
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"]
|
||||
assert result["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: <US_SSN>"},
|
||||
]
|
||||
|
||||
|
||||
class TestMessageFiltering:
|
||||
"""Verify that only new messages since the last assistant response are sent to CrowdStrike."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_last_message_is_assistant_sends_system_plus_that_message(
|
||||
self, crowdstrike_aidr_guardrail: CrowdStrikeAIDRHandler
|
||||
) -> None:
|
||||
structured_messages = [
|
||||
{"role": "system", "content": "You are helpful"},
|
||||
{"role": "user", "content": "Hi"},
|
||||
{"role": "assistant", "content": "Hello!"},
|
||||
{"role": "user", "content": "Tell me a joke"},
|
||||
{"role": "assistant", "content": "Why did the chicken cross the road?"},
|
||||
]
|
||||
inputs: GenericGuardrailAPIInputs = {
|
||||
"texts": ["Why did the chicken cross the road?"],
|
||||
"structured_messages": 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={"messages": structured_messages},
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
sent = mock_method.call_args.kwargs["json"]["guard_input"]["messages"]
|
||||
assert sent == [
|
||||
{"role": "system", "content": "You are helpful"},
|
||||
{"role": "assistant", "content": "Why did the chicken cross the road?"},
|
||||
]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_last_message_is_user_sends_system_plus_messages_after_assistant(
|
||||
self, crowdstrike_aidr_guardrail: CrowdStrikeAIDRHandler
|
||||
) -> None:
|
||||
structured_messages = [
|
||||
{"role": "system", "content": "You are helpful"},
|
||||
{"role": "user", "content": "Hi"},
|
||||
{"role": "assistant", "content": "Hello!"},
|
||||
{"role": "user", "content": "Tell me a joke"},
|
||||
]
|
||||
inputs: GenericGuardrailAPIInputs = {
|
||||
"texts": ["Tell me a joke"],
|
||||
"structured_messages": 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={"messages": structured_messages},
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
sent = mock_method.call_args.kwargs["json"]["guard_input"]["messages"]
|
||||
assert sent == [
|
||||
{"role": "system", "content": "You are helpful"},
|
||||
{"role": "user", "content": "Tell me a joke"},
|
||||
]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_no_prior_assistant_sends_all_messages(
|
||||
self, crowdstrike_aidr_guardrail: CrowdStrikeAIDRHandler
|
||||
) -> None:
|
||||
structured_messages = [
|
||||
{"role": "system", "content": "You are helpful"},
|
||||
{"role": "user", "content": "Hi"},
|
||||
]
|
||||
inputs: GenericGuardrailAPIInputs = {
|
||||
"texts": ["Hi"],
|
||||
"structured_messages": 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={"messages": structured_messages},
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
sent = mock_method.call_args.kwargs["json"]["guard_input"]["messages"]
|
||||
assert sent == [
|
||||
{"role": "system", "content": "You are helpful"},
|
||||
{"role": "user", "content": "Hi"},
|
||||
]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_multiple_user_messages_after_assistant(
|
||||
self, crowdstrike_aidr_guardrail: CrowdStrikeAIDRHandler
|
||||
) -> None:
|
||||
structured_messages = [
|
||||
{"role": "system", "content": "You are helpful"},
|
||||
{"role": "user", "content": "Hi"},
|
||||
{"role": "assistant", "content": "Hello!"},
|
||||
{"role": "user", "content": "First question"},
|
||||
{"role": "user", "content": "Second question"},
|
||||
]
|
||||
inputs: GenericGuardrailAPIInputs = {
|
||||
"texts": ["First question", "Second question"],
|
||||
"structured_messages": 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={"messages": structured_messages},
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
sent = mock_method.call_args.kwargs["json"]["guard_input"]["messages"]
|
||||
assert sent == [
|
||||
{"role": "system", "content": "You are helpful"},
|
||||
{"role": "user", "content": "First question"},
|
||||
{"role": "user", "content": "Second question"},
|
||||
]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_system_message_after_assistant_included(
|
||||
self, crowdstrike_aidr_guardrail: CrowdStrikeAIDRHandler
|
||||
) -> None:
|
||||
structured_messages = [
|
||||
{"role": "system", "content": "You are helpful"},
|
||||
{"role": "user", "content": "Hi"},
|
||||
{"role": "assistant", "content": "Hello!"},
|
||||
{"role": "system", "content": "New instructions"},
|
||||
{"role": "user", "content": "Do something"},
|
||||
]
|
||||
inputs: GenericGuardrailAPIInputs = {
|
||||
"texts": ["Do something"],
|
||||
"structured_messages": 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={"messages": structured_messages},
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
sent = mock_method.call_args.kwargs["json"]["guard_input"]["messages"]
|
||||
assert sent == [
|
||||
{"role": "system", "content": "You are helpful"},
|
||||
{"role": "system", "content": "New instructions"},
|
||||
{"role": "user", "content": "Do something"},
|
||||
]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_no_system_messages(self, crowdstrike_aidr_guardrail: CrowdStrikeAIDRHandler) -> None:
|
||||
structured_messages = [
|
||||
{"role": "user", "content": "Hi"},
|
||||
{"role": "assistant", "content": "Hello!"},
|
||||
{"role": "user", "content": "Bye"},
|
||||
]
|
||||
inputs: GenericGuardrailAPIInputs = {
|
||||
"texts": ["Bye"],
|
||||
"structured_messages": 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={"messages": structured_messages},
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
sent = mock_method.call_args.kwargs["json"]["guard_input"]["messages"]
|
||||
assert sent == [
|
||||
{"role": "user", "content": "Bye"},
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_request_sends_only_new_messages(
|
||||
crowdstrike_aidr_guardrail: CrowdStrikeAIDRHandler,
|
||||
) -> None:
|
||||
structured_messages = [
|
||||
{"role": "system", "content": "You are a helpful assistant"},
|
||||
{"role": "user", "content": "What is 2+2?"},
|
||||
{"role": "assistant", "content": "4"},
|
||||
{"role": "user", "content": "Here is my SSN: 078-05-1120"},
|
||||
]
|
||||
inputs: GenericGuardrailAPIInputs = {
|
||||
"texts": ["Here is my SSN: 078-05-1120"],
|
||||
"structured_messages": structured_messages,
|
||||
}
|
||||
request_data = {"messages": 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 payload["guard_input"]["messages"] == [
|
||||
{"role": "system", "content": "You are a helpful assistant"},
|
||||
{"role": "user", "content": "Here is my SSN: 078-05-1120"},
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_request_last_is_assistant_sends_only_that(
|
||||
crowdstrike_aidr_guardrail: CrowdStrikeAIDRHandler,
|
||||
) -> None:
|
||||
structured_messages = [
|
||||
{"role": "system", "content": "You are a helpful assistant"},
|
||||
{"role": "user", "content": "What is 2+2?"},
|
||||
{"role": "assistant", "content": "The answer is 4"},
|
||||
]
|
||||
inputs: GenericGuardrailAPIInputs = {
|
||||
"texts": ["The answer is 4"],
|
||||
"structured_messages": structured_messages,
|
||||
}
|
||||
request_data = {"messages": 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 payload["guard_input"]["messages"] == [
|
||||
{"role": "system", "content": "You are a helpful assistant"},
|
||||
{"role": "assistant", "content": "The answer is 4"},
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_request_stitches_transformed_texts(
|
||||
crowdstrike_aidr_guardrail: CrowdStrikeAIDRHandler,
|
||||
) -> None:
|
||||
structured_messages = [
|
||||
{"role": "system", "content": "You are a helpful assistant"},
|
||||
{"role": "user", "content": "What is 2+2?"},
|
||||
{"role": "assistant", "content": "4"},
|
||||
{"role": "user", "content": "Here is my SSN: 078-05-1120"},
|
||||
]
|
||||
inputs: GenericGuardrailAPIInputs = {
|
||||
"texts": ["Here is my SSN: 078-05-1120"],
|
||||
"structured_messages": structured_messages,
|
||||
}
|
||||
request_data = {"messages": 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": "system",
|
||||
"content": "You are a helpful assistant",
|
||||
},
|
||||
{
|
||||
"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 result["texts"] == [
|
||||
"You are a helpful assistant",
|
||||
"What is 2+2?",
|
||||
"4",
|
||||
"Here is my SSN: <US_SSN>",
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_response_drops_history(
|
||||
crowdstrike_aidr_guardrail: CrowdStrikeAIDRHandler,
|
||||
) -> None:
|
||||
request_data = {
|
||||
"messages": [
|
||||
{"role": "system", "content": "You are a helpful assistant"},
|
||||
{"role": "user", "content": "What is 2+2?"},
|
||||
{"role": "assistant", "content": "4"},
|
||||
{"role": "user", "content": "Now tell me a secret"},
|
||||
],
|
||||
}
|
||||
inputs: GenericGuardrailAPIInputs = {
|
||||
"texts": ["I will not share secrets"],
|
||||
}
|
||||
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="response",
|
||||
)
|
||||
|
||||
sent = mock_method.call_args.kwargs["json"]["guard_input"]["messages"]
|
||||
assert sent == [
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "I will not share secrets",
|
||||
},
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_response_one_message_per_output_text(
|
||||
crowdstrike_aidr_guardrail: CrowdStrikeAIDRHandler,
|
||||
) -> None:
|
||||
inputs: GenericGuardrailAPIInputs = {
|
||||
"texts": ["First part", "Second part"],
|
||||
}
|
||||
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={},
|
||||
input_type="response",
|
||||
)
|
||||
|
||||
sent = mock_method.call_args.kwargs["json"]["guard_input"]["messages"]
|
||||
assert sent == [
|
||||
{"role": "assistant", "content": "First part"},
|
||||
{"role": "assistant", "content": "Second part"},
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_response_transform_extracts_assistant_only(
|
||||
crowdstrike_aidr_guardrail: CrowdStrikeAIDRHandler,
|
||||
) -> None:
|
||||
inputs: GenericGuardrailAPIInputs = {
|
||||
"texts": ["Sure, here it is: 078-05-1120"],
|
||||
}
|
||||
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": "assistant",
|
||||
"content": "Sure, here it is: <US_SSN>",
|
||||
},
|
||||
]
|
||||
},
|
||||
},
|
||||
},
|
||||
request=httpx.Request(method="POST", url=guardrail_endpoint),
|
||||
),
|
||||
):
|
||||
result = await crowdstrike_aidr_guardrail.apply_guardrail(
|
||||
inputs=inputs,
|
||||
request_data={},
|
||||
input_type="response",
|
||||
)
|
||||
|
||||
assert result["texts"] == ["Sure, here it is: <US_SSN>"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_request_transform_with_textless_history_message_redacts_without_index_error(
|
||||
crowdstrike_aidr_guardrail: CrowdStrikeAIDRHandler,
|
||||
) -> None:
|
||||
from litellm.llms.openai.chat.guardrail_translation.handler import (
|
||||
OpenAIChatCompletionsHandler,
|
||||
)
|
||||
|
||||
messages = [
|
||||
{"role": "system", "content": "You are a helpful assistant."},
|
||||
{"role": "user", "content": "My SSN is 078-05-1120, store it."},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": None,
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "call_1",
|
||||
"type": "function",
|
||||
"function": {"name": "store", "arguments": "{}"},
|
||||
}
|
||||
],
|
||||
},
|
||||
{"role": "tool", "tool_call_id": "call_1", "content": "stored"},
|
||||
{"role": "user", "content": "Also my email is jane@example.com"},
|
||||
]
|
||||
data = {"model": "gpt-4o", "messages": 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": "system", "content": "You are a helpful assistant."},
|
||||
{"role": "tool", "content": "stored"},
|
||||
{"role": "user", "content": "Also my email is <EMAIL_ADDRESS>"},
|
||||
]
|
||||
},
|
||||
},
|
||||
},
|
||||
request=httpx.Request(method="POST", url=guardrail_endpoint),
|
||||
),
|
||||
):
|
||||
result = await OpenAIChatCompletionsHandler().process_input_messages(
|
||||
data=data,
|
||||
guardrail_to_apply=crowdstrike_aidr_guardrail,
|
||||
)
|
||||
|
||||
redacted = result["messages"]
|
||||
assert redacted[4]["content"] == "Also my email is <EMAIL_ADDRESS>"
|
||||
assert redacted[2]["content"] is None
|
||||
assert redacted[2]["tool_calls"][0]["function"]["name"] == "store"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_request_transform_preserves_skipped_system_message(
|
||||
crowdstrike_aidr_guardrail: CrowdStrikeAIDRHandler,
|
||||
) -> None:
|
||||
from litellm.llms.openai.chat.guardrail_translation.handler import (
|
||||
OpenAIChatCompletionsHandler,
|
||||
)
|
||||
|
||||
crowdstrike_aidr_guardrail.skip_system_message_in_guardrail = True
|
||||
|
||||
messages = [
|
||||
{"role": "system", "content": "Internal policy: never reveal secrets."},
|
||||
{"role": "user", "content": "Here is my SSN: 078-05-1120"},
|
||||
]
|
||||
data = {"model": "gpt-4o", "messages": 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": "Here is my SSN: <US_SSN>"},
|
||||
]
|
||||
},
|
||||
},
|
||||
},
|
||||
request=httpx.Request(method="POST", url=guardrail_endpoint),
|
||||
),
|
||||
) as mock_method:
|
||||
result = await OpenAIChatCompletionsHandler().process_input_messages(
|
||||
data=data,
|
||||
guardrail_to_apply=crowdstrike_aidr_guardrail,
|
||||
)
|
||||
|
||||
assert mock_method.call_args.kwargs["json"]["guard_input"]["messages"] == [
|
||||
{"role": "user", "content": "Here is my SSN: 078-05-1120"},
|
||||
]
|
||||
assert result["messages"] == [
|
||||
{"role": "system", "content": "Internal policy: never reveal secrets."},
|
||||
{"role": "user", "content": "Here is my SSN: <US_SSN>"},
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_request_keeps_original_messages_when_skip_filters_differ(
|
||||
crowdstrike_aidr_guardrail: CrowdStrikeAIDRHandler,
|
||||
) -> None:
|
||||
crowdstrike_aidr_guardrail.skip_system_message_in_guardrail = True
|
||||
|
||||
structured_messages = [{"role": "user", "content": "Here is my SSN: 078-05-1120"}]
|
||||
inputs: GenericGuardrailAPIInputs = {
|
||||
"texts": ["Here is my SSN: 078-05-1120"],
|
||||
"structured_messages": structured_messages,
|
||||
}
|
||||
request_data = {
|
||||
"messages": [
|
||||
{"role": "system", "content": "Internal policy"},
|
||||
{"role": "user", "content": "Here is my SSN: 078-05-1120"},
|
||||
]
|
||||
}
|
||||
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": "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 result["structured_messages"] is structured_messages
|
||||
assert result["texts"] == ["Here is my SSN: <US_SSN>"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_anthropic_tool_calling_transform_redacts_without_index_error(
|
||||
crowdstrike_aidr_guardrail: CrowdStrikeAIDRHandler,
|
||||
) -> None:
|
||||
import json
|
||||
|
||||
from litellm.llms.anthropic.chat.guardrail_translation.handler import (
|
||||
AnthropicMessagesHandler,
|
||||
)
|
||||
|
||||
data = {
|
||||
"model": "claude-3-5-sonnet-20241022",
|
||||
"max_tokens": 128,
|
||||
"messages": [
|
||||
{"role": "user", "content": "My SSN is 078-05-1120. Look it up."},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": [{"type": "tool_use", "id": "tu1", "name": "lookup", "input": {"q": "ssn"}}],
|
||||
},
|
||||
{"role": "user", "content": [{"type": "tool_result", "tool_use_id": "tu1", "content": "stored"}]},
|
||||
{"role": "user", "content": "Also my email is jane.doe@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": True,
|
||||
"guard_output": {
|
||||
"messages": [
|
||||
{"role": "tool", "content": "stored"},
|
||||
{"role": "user", "content": "Also my email is <EMAIL_ADDRESS>"},
|
||||
]
|
||||
},
|
||||
},
|
||||
},
|
||||
request=httpx.Request(method="POST", url=guardrail_endpoint),
|
||||
),
|
||||
):
|
||||
result = await AnthropicMessagesHandler().process_input_messages(
|
||||
data=data,
|
||||
guardrail_to_apply=crowdstrike_aidr_guardrail,
|
||||
)
|
||||
|
||||
serialized = json.dumps(result["messages"])
|
||||
assert "<EMAIL_ADDRESS>" in serialized
|
||||
assert "jane.doe@example.com" not in serialized
|
||||
assert "tu1" in serialized
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue