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"]