From 69dcf5497858f480286f2307e7c8735472c43096 Mon Sep 17 00:00:00 2001 From: Kenan Yildirim Date: Mon, 6 Apr 2026 16:13:11 -0400 Subject: [PATCH] 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. --- .../crowdstrike_aidr/crowdstrike_aidr.py | 269 +++++++++--------- .../guardrail_hooks/test_crowdstrike_aidr.py | 132 ++++++--- 2 files changed, 231 insertions(+), 170 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/crowdstrike_aidr/crowdstrike_aidr.py b/litellm/proxy/guardrails/guardrail_hooks/crowdstrike_aidr/crowdstrike_aidr.py index 19c5d54213f..14d950ecdf4 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/crowdstrike_aidr/crowdstrike_aidr.py +++ b/litellm/proxy/guardrails/guardrail_hooks/crowdstrike_aidr/crowdstrike_aidr.py @@ -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: diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_crowdstrike_aidr.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_crowdstrike_aidr.py index fa8f001f485..c58c94cbbc7 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_crowdstrike_aidr.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_crowdstrike_aidr.py @@ -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: ", - }, + "role": "assistant", + "content": "Yes, here is an 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: " + # 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: "] @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: ", + }, + ] + }, + }, + }, + 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: " + assert result["structured_messages"] == inputs["structured_messages"]