diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py b/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py index 125bea5590e..fe8b70c08b8 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py @@ -7,7 +7,7 @@ import os import uuid -from collections.abc import AsyncGenerator +from collections.abc import AsyncGenerator, Callable, Mapping, Sequence from typing import ( TYPE_CHECKING, Any, # noqa: TID251 # **kwargs forwards verbatim to CustomGuardrail.__init__ @@ -15,6 +15,7 @@ from typing import ( Final, Literal, Optional, + TypeAlias, ) import httpx @@ -53,6 +54,19 @@ _SESSION_METADATA_KEY: Final = "llm_shield_session_id" _DEFAULT_TIMEOUT_SECONDS: Final = 10.0 +# The proxy's own request dict. Mutable by design: a pre-call guardrail rewrites +# the caller's payload in place, which is the entire point of the hook. +# mutable-ok: the shape is fixed by CustomLogger's hook signatures. +MutableRequest: TypeAlias = dict + +# A JSON body on its way to httpx, which requires a real dict rather than a view. +# mutable-ok: handed straight to the HTTP client. +JsonBody: TypeAlias = dict + +# One redactable span: the text as it stands, and the write that puts the +# replacement back where it came from. +_Slot: TypeAlias = tuple[str, Callable[[str], None]] # mutable-ok: Callable's param list. + class LLMShieldGuardrail(CustomGuardrail): """Redacts PII before it leaves the proxy and restores it in the response. @@ -78,7 +92,7 @@ class LLMShieldGuardrail(CustomGuardrail): guardrail_name: str = GUARDRAIL_NAME, api_base: str | None = None, api_key: str | None = None, - **kwargs: Any, + **kwargs: Any, # noqa: LIT008 # kwargs-ok: forwarded verbatim to CustomGuardrail.__init__ ) -> None: self.async_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.GuardrailCallback) self.api_base: Final = (api_base or os.environ.get("LLM_SHIELD_API_BASE") or _DEFAULT_API_BASE).rstrip("/") @@ -86,18 +100,21 @@ class LLMShieldGuardrail(CustomGuardrail): super().__init__(guardrail_name=guardrail_name, **kwargs) @classmethod - def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: - return [GuardrailEventHooks.pre_call, GuardrailEventHooks.post_call] + def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: # mutable-ok: parent's signature. + return [GuardrailEventHooks.pre_call, GuardrailEventHooks.post_call] # mutable-ok: parent's signature. # --- transport --------------------------------------------------------------- - def _headers(self, session_id: str) -> dict: - headers = {"Content-Type": "application/json", "X-Session-ID": session_id} + def _headers(self, session_id: str) -> JsonBody: + headers: Final[JsonBody] = { # mutable-ok: httpx requires a real dict. + "Content-Type": "application/json", + "X-Session-ID": session_id, + } if self.api_key: headers["Authorization"] = f"Bearer {self.api_key}" return headers - async def _call_shield(self, path: str, session_id: str, payload: dict) -> dict: + async def _call_shield(self, path: str, session_id: str, payload: JsonBody) -> Mapping[str, object]: """Posts to LLM Shield, failing closed on any transport or status error. A redaction guardrail that fails open sends the very data it exists to @@ -105,7 +122,7 @@ class LLMShieldGuardrail(CustomGuardrail): blocks the request instead of passing it through. """ try: - response = await self.async_handler.post( + response: Final = await self.async_handler.post( f"{self.api_base}{path}", headers=self._headers(session_id), json=payload, @@ -126,69 +143,89 @@ class LLMShieldGuardrail(CustomGuardrail): message="LLM Shield is unreachable; blocking the request.", ) from exc - async def _redact(self, texts: list, session_id: str) -> list: - body = await self._call_shield(_REDACT_PATH, session_id, {"texts": texts}) + async def _redact(self, texts: Sequence[str], session_id: str) -> Sequence[str]: + payload: Final[JsonBody] = {"texts": list(texts)} # mutable-ok: JSON body for httpx. + body: Final = await self._call_shield(_REDACT_PATH, session_id, payload) return self._same_length_or_raise(body.get("texts"), texts, "redact") - async def _rehydrate(self, texts: list, session_id: str) -> list: - body = await self._call_shield(_REHYDRATE_PATH, session_id, {"texts": texts}) + async def _rehydrate(self, texts: Sequence[str], session_id: str) -> Sequence[str]: + payload: Final[JsonBody] = {"texts": list(texts)} # mutable-ok: JSON body for httpx. + body: Final = await self._call_shield(_REHYDRATE_PATH, session_id, payload) return self._same_length_or_raise(body.get("texts"), texts, "rehydrate") - def _same_length_or_raise(self, returned: Any, sent: list, operation: str) -> list: + def _same_length_or_raise(self, returned: object, sent: Sequence[str], operation: str) -> Sequence[str]: """Guards the positional mapping the callers rely on to write results back.""" if not isinstance(returned, list) or len(returned) != len(sent): raise GuardrailRaisedException( guardrail_name=self.guardrail_name, message=f"LLM Shield {operation} returned an unexpected payload; blocking the request.", ) - return returned + return tuple(returned) # --- session ------------------------------------------------------------------ - def _session_id(self, data: dict) -> str: + def _session_id(self, data: MutableRequest) -> str: """Returns a session id stable across this request's hooks.""" - metadata = data.setdefault("metadata", {}) + metadata: Final = data.setdefault("metadata", {}) # mutable-ok: per-request store. if not isinstance(metadata, dict): return f"litellm-{uuid.uuid4().hex}" - existing = metadata.get(_SESSION_METADATA_KEY) + existing: Final = metadata.get(_SESSION_METADATA_KEY) if isinstance(existing, str) and existing: return existing - session_id = get_session_id_from_request_data(data) or f"litellm-{uuid.uuid4().hex}" + session_id: Final = get_session_id_from_request_data(data) or f"litellm-{uuid.uuid4().hex}" metadata[_SESSION_METADATA_KEY] = session_id return session_id - # --- message traversal -------------------------------------------------------- + # --- request traversal -------------------------------------------------------- @staticmethod - def _locate_texts(messages: list) -> list: - """Finds every text span in a message list. + def _locate_request_texts(data: MutableRequest) -> Sequence[_Slot]: + """Finds every redactable span in an outbound request. - Returns ``(message_index, part_index_or_None, text)``. The list form is the - multimodal shape, where only ``text`` parts carry redactable content. + Returns ``(text, write)`` pairs. Any shape missed here reaches the provider + in the clear, so this walks all of the request shapes that carry caller text: + + - chat ``messages``, both string and multimodal list ``content`` + - tool call ``arguments``, which routinely carry the values a user asked + the model to look up + - the Responses API ``input``, as a bare string or a list of items """ - located = [] - for message_index, message in enumerate(messages): - if not isinstance(message, dict): - continue - content = message.get("content") - if isinstance(content, str) and content: - located.append((message_index, None, content)) - elif isinstance(content, list): - for part_index, part in enumerate(content): - if not isinstance(part, dict) or part.get("type") != "text": - continue - text = part.get("text") - if isinstance(text, str) and text: - located.append((message_index, part_index, text)) - return located + slots: Final[list[_Slot]] = [] # mutable-ok: accumulator, frozen on return. - @staticmethod - def _write_back(messages: list, located: list, replacements: list) -> None: - for (message_index, part_index, _), replacement in zip(located, replacements): - if part_index is None: - messages[message_index]["content"] = replacement - else: - messages[message_index]["content"][part_index]["text"] = replacement + def add(container: MutableRequest, key: str, value: object) -> None: + if isinstance(value, str) and value: + slots.append((value, lambda new, c=container, k=key: c.__setitem__(k, new))) + + def add_content(container: MutableRequest) -> None: + """Adds `content`, which is either a string or a list of typed parts.""" + content: Final = container.get("content") + if isinstance(content, str): + add(container, "content", content) + return + for part in content if isinstance(content, list) else (): + if isinstance(part, dict): + add(part, "text", part.get("text")) + + def add_tool_calls(message: MutableRequest) -> None: + for tool_call in message.get("tool_calls") or (): + function = tool_call.get("function") if isinstance(tool_call, dict) else None + if isinstance(function, dict): + add(function, "arguments", function.get("arguments")) + + for message in data.get("messages") or (): + if isinstance(message, dict): + add_content(message) + add_tool_calls(message) + + request_input: Final = data.get("input") + if isinstance(request_input, str): + add(data, "input", request_input) + else: + for item in request_input if isinstance(request_input, list) else (): + if isinstance(item, dict): + add_content(item) + + return tuple(slots) # --- hooks -------------------------------------------------------------------- @@ -197,29 +234,26 @@ class LLMShieldGuardrail(CustomGuardrail): self, user_api_key_dict: UserAPIKeyAuth, cache: "DualCache", - data: dict, + data: MutableRequest, call_type: str, - ) -> dict | None: - """Replaces PII in the outbound messages with vault placeholders.""" + ) -> MutableRequest | None: + """Replaces PII anywhere in the outbound request with vault placeholders.""" if self.should_run_guardrail(data=data, event_type=GuardrailEventHooks.pre_call) is not True: return data - messages = data.get("messages") - if not isinstance(messages, list): + slots: Final = self._locate_request_texts(data) + if not slots: return data - located = self._locate_texts(messages) - if not located: - return data - - redacted = await self._redact([text for _, _, text in located], self._session_id(data)) - self._write_back(messages, located, redacted) + redacted: Final = await self._redact(tuple(text for text, _ in slots), self._session_id(data)) + for (_, write), replacement in zip(slots, redacted): + write(replacement) return data @log_guardrail_information async def async_post_call_success_hook( self, - data: dict, + data: MutableRequest, user_api_key_dict: UserAPIKeyAuth, response: Any, ) -> Any: @@ -230,27 +264,31 @@ class LLMShieldGuardrail(CustomGuardrail): if self._is_anthropic_message_response(response): return await self._restore_anthropic_response(response, data) - choices = getattr(response, "choices", None) + text_blocks: Final = self._responses_api_text_blocks(response) + if text_blocks: + return await self._restore_responses_api_response(response, text_blocks, data) + + choices: Final = getattr(response, "choices", None) if not choices: return response - pending = [] - for choice in choices: - message = getattr(choice, "message", None) - content = getattr(message, "content", None) - if isinstance(content, str) and content: - pending.append((message, content)) - + pending: Final = tuple( + (choice.message, choice.message.content) + for choice in choices + if getattr(choice, "message", None) is not None + and isinstance(getattr(choice.message, "content", None), str) + and choice.message.content + ) if not pending: return response - restored = await self._rehydrate([text for _, text in pending], self._session_id(data)) + restored: Final = await self._rehydrate(tuple(text for _, text in pending), self._session_id(data)) for (message, _), replacement in zip(pending, restored): message.content = replacement return response @staticmethod - def _is_anthropic_message_response(response: Any) -> bool: + def _is_anthropic_message_response(response: object) -> bool: """Anthropic's native /v1/messages reply arrives as a plain dict.""" return ( isinstance(response, dict) @@ -258,30 +296,68 @@ class LLMShieldGuardrail(CustomGuardrail): and isinstance(response.get("content"), list) ) - async def _restore_anthropic_response(self, response: dict, data: dict) -> dict: + async def _restore_anthropic_response(self, response: MutableRequest, data: MutableRequest) -> MutableRequest: """Restores text blocks in an Anthropic native message reply. This shape has no `choices`, so without its own branch the reply would go back to the caller still carrying placeholders. """ - blocks = [ + blocks: Final = tuple( block for block in response["content"] if isinstance(block, dict) and block.get("type") == "text" and isinstance(block.get("text"), str) - ] + ) if not blocks: return response - restored = await self._rehydrate([block["text"] for block in blocks], self._session_id(data)) + restored: Final = await self._rehydrate(tuple(block["text"] for block in blocks), self._session_id(data)) for block, replacement in zip(blocks, restored): block["text"] = replacement return response + @staticmethod + def _responses_api_text_blocks(response: object) -> Sequence[object]: + """Text blocks in a Responses API reply. + + That shape carries `output` items rather than `choices`, so it needs its own + walk; without one the reply goes back to the caller still holding + placeholders even though the request was redacted correctly. Blocks come + through as dicts or as objects depending on how far the reply has been + deserialised, so both are handled. + """ + blocks: Final[list[object]] = [] # mutable-ok: accumulator, frozen on return. + for item in getattr(response, "output", None) or (): + for block in getattr(item, "content", None) or (): + if isinstance(block, dict): + if isinstance(block.get("text"), str) and block["text"]: + blocks.append(block) + elif isinstance(getattr(block, "text", None), str) and block.text: + blocks.append(block) + return tuple(blocks) + + @staticmethod + def _block_text(block: object) -> str: + return block["text"] if isinstance(block, dict) else block.text + + async def _restore_responses_api_response( + self, response: Any, blocks: Sequence[object], data: MutableRequest + ) -> Any: + """Puts the original values back into a Responses API reply.""" + restored: Final = await self._rehydrate( + tuple(self._block_text(block) for block in blocks), self._session_id(data) + ) + for block, replacement in zip(blocks, restored): + if isinstance(block, dict): + block["text"] = replacement + else: + block.text = replacement + return response + async def async_post_call_streaming_iterator_hook( self, user_api_key_dict: UserAPIKeyAuth, response: Any, - request_data: dict, + request_data: MutableRequest, ) -> AsyncGenerator[Any, None]: """Restores original values incrementally, without buffering the stream. @@ -295,9 +371,9 @@ class LLMShieldGuardrail(CustomGuardrail): yield chunk return - session_id = self._session_id(request_data) - carry = "" - last_chunk = None + session_id: Final = self._session_id(request_data) + carry = "" # rebind-ok: the sliding window advances with every delta. + last_chunk = None # rebind-ok: tracks the most recent chunk for the final flush. async for chunk in response: last_chunk = chunk @@ -309,53 +385,54 @@ class LLMShieldGuardrail(CustomGuardrail): # Nothing to restore in this chunk, but a final chunk still has to # flush whatever the window is holding. if is_final and carry: - body = await self._stream_step("", carry, True, session_id) - carry = body["carry"] - if body["text"] and delta is not None: - delta.content = body["text"] + emitted, carry = await self._stream_step("", carry, True, session_id) + if emitted and delta is not None: + delta.content = emitted yield chunk continue - body = await self._stream_step(text, carry, is_final, session_id) - carry = body["carry"] - delta.content = body["text"] + emitted, carry = await self._stream_step(text, carry, is_final, session_id) + delta.content = emitted yield chunk # A stream that ended without a finish_reason can still leave text held back. if carry and last_chunk is not None: - body = await self._stream_step("", carry, True, session_id) - if body["text"]: - trailing = last_chunk.model_copy(deep=True) - trailing_delta = self._stream_delta(trailing) + flushed: Final = await self._stream_step("", carry, True, session_id) + trailing_text, carry = flushed # rebind-ok: window advances. + if trailing_text: + trailing: Final = last_chunk.model_copy(deep=True) + trailing_delta: Final = self._stream_delta(trailing) if trailing_delta is not None: - trailing_delta.content = body["text"] + trailing_delta.content = trailing_text yield trailing - async def _stream_step(self, text: str, carry: str, final: bool, session_id: str) -> dict: - body = await self._call_shield( + async def _stream_step(self, text: str, carry: str, final: bool, session_id: str) -> tuple[str, str]: + """Returns ``(text safe to emit now, window still being held)``.""" + body: Final = await self._call_shield( _REHYDRATE_STREAM_PATH, session_id, - {"text": text, "carry": carry, "final": final}, + # mutable-ok: JSON request body for httpx. + {"text": text, "carry": carry, "final": final}, # mutable-ok: JSON request body for httpx. ) - emitted = body.get("text") - remaining = body.get("carry") + emitted: Final = body.get("text") + remaining: Final = body.get("carry") if not isinstance(emitted, str) or not isinstance(remaining, str): raise GuardrailRaisedException( guardrail_name=self.guardrail_name, message="LLM Shield stream rehydration returned an unexpected payload.", ) - return {"text": emitted, "carry": remaining} + return emitted, remaining @staticmethod - def _stream_delta(chunk: Any) -> Any: - choices = getattr(chunk, "choices", None) + def _stream_delta(chunk: object) -> Any: + choices: Final = getattr(chunk, "choices", None) if not choices: return None return getattr(choices[0], "delta", None) @staticmethod - def _is_final_chunk(chunk: Any) -> bool: - choices = getattr(chunk, "choices", None) + def _is_final_chunk(chunk: object) -> bool: + choices: Final = getattr(chunk, "choices", None) if not choices: return False return bool(getattr(choices[0], "finish_reason", None)) @@ -366,17 +443,21 @@ class LLMShieldGuardrail(CustomGuardrail): async def apply_guardrail( self, inputs: GenericGuardrailAPIInputs, - request_data: dict, + request_data: MutableRequest, input_type: Literal["request", "response"], logging_obj: Optional["LiteLLMLoggingObj"] = None, ) -> GenericGuardrailAPIInputs: - texts = inputs.get("texts") + texts: Final = inputs.get("texts") if not texts: return inputs - session_id = self._session_id(request_data) - if input_type == "request": - inputs["texts"] = await self._redact(list(texts), session_id) - else: - inputs["texts"] = await self._rehydrate(list(texts), session_id) - return inputs + session_id: Final = self._session_id(request_data) + replaced: Final = ( + await self._redact(tuple(texts), session_id) + if input_type == "request" + else await self._rehydrate(tuple(texts), session_id) + ) + # Return a new mapping rather than rewriting the caller's, so this stays a + # pure transform of the inputs it was handed. + merged: Final[JsonBody] = {**inputs, "texts": list(replaced)} # mutable-ok: TypedDict. + return merged diff --git a/ruff-strict.toml b/ruff-strict.toml index ae092bdde7d..f9d026c6011 100644 --- a/ruff-strict.toml +++ b/ruff-strict.toml @@ -30,6 +30,10 @@ external = [ # grows over time; typing it concretely (`object`) broke that forwarding call outright — # basedpyright turned every named param into a reportArgumentType error. Any is correct here. "litellm/proxy/guardrails/guardrail_hooks/alice/alice.py" = ["ANN401"] +# Same reason: `**kwargs` forwards verbatim to CustomGuardrail.__init__, and the lifecycle +# hook signatures inherit `Any` for `response` from CustomLogger, so narrowing them here +# would break the override rather than describe it. +"litellm/proxy/guardrails/guardrail_hooks/llm_shield/llm_shield.py" = ["ANN401"] [lint.mccabe] max-complexity = 15 diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield.py index c2d50bf5fa7..7a07c9761a8 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_llm_shield.py @@ -1,3 +1,4 @@ +from types import SimpleNamespace from unittest.mock import AsyncMock import pytest @@ -133,10 +134,16 @@ class TestRedaction: assert data["messages"][0]["content"][1]["image_url"]["url"] == "http://x/y.png" @pytest.mark.asyncio - async def test_request_without_messages_is_untouched(self): + async def test_request_without_text_is_untouched(self): + """No text to redact means no call to LLM Shield. + + This deliberately uses a request with no caller text at all. An earlier + version used a Responses-API `input`, which asserted the very bypass that + let `input` reach the provider unredacted. + """ guardrail = _guardrail() mock = _mock_post(guardrail) - data = {"input": "no messages here"} + data = {"model": "gpt-4o", "temperature": 0.2} await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion") @@ -156,6 +163,91 @@ class TestRedaction: assert len(sessions) == 1 +class TestRequestCoverage: + """Every request shape that carries caller text must be redacted. + + A shape missed here is not a cosmetic gap: the guardrail reports as enabled + while the raw value goes to the provider. + """ + + @pytest.mark.asyncio + async def test_responses_api_string_input_is_redacted(self): + """Measured against a live provider: `input` reached the model unredacted.""" + guardrail = _guardrail() + mock = _mock_post(guardrail, {"texts": ["Email [EMAIL_1] the invoice"]}) + + data = {"input": "Email jane.doe@example.com the invoice"} + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="aresponses") + + assert mock.call_args_list[0].kwargs["json"]["texts"] == ["Email jane.doe@example.com the invoice"] + assert data["input"] == "Email [EMAIL_1] the invoice" + + @pytest.mark.asyncio + async def test_responses_api_list_input_is_redacted(self): + guardrail = _guardrail() + _mock_post(guardrail, {"texts": ["[EMAIL_1]", "[PHONE_1]"]}) + + data = { + "input": [ + {"role": "user", "content": "jane.doe@example.com"}, + {"role": "user", "content": [{"type": "input_text", "text": "555-0100"}]}, + ] + } + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="aresponses") + + assert data["input"][0]["content"] == "[EMAIL_1]" + assert data["input"][1]["content"][0]["text"] == "[PHONE_1]" + + @pytest.mark.asyncio + async def test_tool_call_arguments_are_redacted(self): + """Tool arguments carry the values the user asked the model to act on.""" + guardrail = _guardrail() + _mock_post(guardrail, {"texts": ['{"email": "[EMAIL_1]"}']}) + + data = { + "messages": [ + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_1", + "type": "function", + "function": {"name": "send", "arguments": '{"email": "jane.doe@example.com"}'}, + } + ], + } + ] + } + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion") + + assert data["messages"][0]["tool_calls"][0]["function"]["arguments"] == '{"email": "[EMAIL_1]"}' + + @pytest.mark.asyncio + async def test_every_shape_in_one_request_is_redacted(self): + guardrail = _guardrail() + mock = _mock_post(guardrail, {"texts": ["a", "b", "c", "d"]}) + + data = { + "messages": [ + {"role": "user", "content": "one"}, + {"role": "user", "content": [{"type": "text", "text": "two"}]}, + { + "role": "assistant", + "tool_calls": [{"function": {"name": "f", "arguments": "three"}}], + }, + ], + "input": "four", + } + await guardrail.async_pre_call_hook(user_api_key_dict=None, cache=None, data=data, call_type="completion") + + assert mock.call_args_list[0].kwargs["json"]["texts"] == ["one", "two", "three", "four"] + assert data["messages"][0]["content"] == "a" + assert data["messages"][1]["content"][0]["text"] == "b" + assert data["messages"][2]["tool_calls"][0]["function"]["arguments"] == "c" + assert data["input"] == "d" + + class TestRestoration: @pytest.mark.asyncio async def test_openai_shape_is_restored(self): @@ -169,6 +261,36 @@ class TestRestoration: assert result.choices[0].message.content == "a@b.com" + @pytest.mark.asyncio + async def test_responses_api_shape_is_restored(self): + """The Responses API reply carries output items, not choices. + + Measured against a live provider: once the request side was fixed the reply + came back still holding the placeholder, because this shape has no choices + to walk. + """ + guardrail = _guardrail(event_hook="post_call") + _mock_post(guardrail, {"texts": ["a@b.com"]}) + + response = SimpleNamespace(output=[SimpleNamespace(content=[{"type": "output_text", "text": "[EMAIL_1]"}])]) + result = await guardrail.async_post_call_success_hook( + data={"messages": []}, user_api_key_dict=None, response=response + ) + + assert result.output[0].content[0]["text"] == "a@b.com" + + @pytest.mark.asyncio + async def test_responses_api_object_blocks_are_restored(self): + """Blocks arrive as objects too, depending on how far the reply is parsed.""" + guardrail = _guardrail(event_hook="post_call") + _mock_post(guardrail, {"texts": ["a@b.com"]}) + + block = SimpleNamespace(text="[EMAIL_1]") + response = SimpleNamespace(output=[SimpleNamespace(content=[block])]) + await guardrail.async_post_call_success_hook(data={"messages": []}, user_api_key_dict=None, response=response) + + assert block.text == "a@b.com" + @pytest.mark.asyncio async def test_anthropic_message_shape_is_restored(self): """The /v1/messages reply is a plain dict with no choices. diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/add_guardrail_form.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/add_guardrail_form.tsx index 29df7c8bf3d..a02cae097a7 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/add_guardrail_form.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/add_guardrail_form.tsx @@ -73,7 +73,9 @@ interface GuardrailPreset { provider: string; categoryName?: string; guardrailNameSuggestion: string; - mode: string; + // A guardrail that both rewrites the request and repairs the response needs two + // modes seeded, not one; the form already normalises either shape. + mode: string | string[]; defaultOn: boolean; } diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_configs.ts b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_configs.ts index b445cc9c5ad..3579457bcd5 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_configs.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_configs.ts @@ -2,7 +2,9 @@ export interface GuardrailPreset { provider: string; categoryName?: string; guardrailNameSuggestion: string; - mode: string; + // A guardrail that both rewrites the request and repairs the response needs two + // modes seeded, not one; the form already normalises either shape. + mode: string | string[]; defaultOn: boolean; } @@ -321,7 +323,9 @@ export const GUARDRAIL_PRESETS: Record = { llm_shield: { provider: "LLM Shield", guardrailNameSuggestion: "LLM Shield", - mode: "pre_call", + // Both halves are required. With only pre_call the request is redacted and the + // placeholders are handed straight back to the caller. + mode: ["pre_call", "post_call"], defaultOn: false, }, };