diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/__init__.py index c8ca68a8967..44b19f82218 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/__init__.py @@ -23,11 +23,11 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail" return _llm_shield_guardrail_callback -guardrail_initializer_registry: Final = { # mutable-ok: module-level registry, built once and never mutated +guardrail_initializer_registry: Final = { SupportedGuardrailIntegrations.LLM_SHIELD_PROXY.value: initialize_guardrail, } -guardrail_class_registry: Final = { # mutable-ok: module-level registry, built once and never mutated +guardrail_class_registry: Final = { SupportedGuardrailIntegrations.LLM_SHIELD_PROXY.value: LLMShieldProxyGuardrail, } diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py index 8c02816b844..d02be56fc52 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_shield_proxy/llm_shield_proxy.py @@ -7,6 +7,7 @@ import copy import functools +import itertools import json import os import re @@ -72,11 +73,9 @@ _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 @@ -92,14 +91,14 @@ _MAX_CONTENT_DEPTH: Final = 8 # generous; past it the request is refused, for the same reason as above. _MAX_JSON_DEPTH: Final = 64 -_Slot: TypeAlias = tuple[str, Callable[[str], None]] # mutable-ok: Callable's param list. +_Slot: TypeAlias = tuple[str, Callable[[str], None]] # One incremental rehydration step for a stream the caller has already bound to its # vault: (new text, carried window, final) -> (text safe to emit, window still held). -_StreamStep: TypeAlias = Callable[[str, str, bool], Awaitable[tuple[str, str]]] # mutable-ok: Callable's param list. +_StreamStep: TypeAlias = Callable[[str, str, bool], Awaitable[tuple[str, str]]] # A batch rehydration already bound to the request's vault. -_Rehydrate: TypeAlias = Callable[[Sequence[str]], Awaitable[Sequence[str]]] # mutable-ok: Callable's param list. +_Rehydrate: TypeAlias = Callable[[Sequence[str]], Awaitable[Sequence[str]]] # Anthropic /v1/messages delta types that carry restorable text, and the field holding # it. `thinking_delta` is left out on purpose: a thinking block is signed, and one @@ -172,17 +171,17 @@ _SCHEMA_MAP_KEYWORDS: Final = frozenset( # The accumulator the collectors below append into. It never escapes # _locate_request_texts, which freezes it into a tuple before returning. -_SlotSink: TypeAlias = list[_Slot] # mutable-ok: accumulator passed between collectors. +_SlotSink: TypeAlias = list[_Slot] # Sliding windows keyed by (choice index, tool-call index | None), threaded through one # stream. `None` is the content channel; an int is one tool call's accumulating # `arguments`. Content and each tool call are separate token streams, so each needs its # own window -- one shared window would splice one stream's held-back tail onto another. -_CarryWindows: TypeAlias = dict # mutable-ok: per-stream windows advanced in place. +_CarryWindows: TypeAlias = dict # A caller-owned list whose entries are rewritten in place, such as a Completions # `prompt` sent as an array of strings. -MutableSeq: TypeAlias = list # mutable-ok: the request payload's own list. +MutableSeq: TypeAlias = list def _collect(container: MutableRequest, key: str, slots: _SlotSink) -> None: @@ -280,7 +279,7 @@ def _collect_participant_name(message: MutableRequest, slots: _SlotSink) -> None def _collect_tool_arguments(message: MutableRequest, slots: _SlotSink) -> None: """Tool arguments carry the values a user asked the model to act on.""" for tool_call in message.get("tool_calls") or (): - function = tool_call.get("function") if isinstance(tool_call, dict) else None # rebind-ok: loop variable. + function = tool_call.get("function") if isinstance(tool_call, dict) else None if isinstance(function, dict): _collect(function, "arguments", slots) legacy: Final = message.get("function_call") @@ -552,7 +551,7 @@ def _continuation_delta(tool_index: int, text: str) -> list[dict[str, object]]: Clients concatenate tool-call fragments by index, so no id or name is needed. """ - return [{"index": tool_index, "function": {"arguments": text}}] # mutable-ok: delta.tool_calls is a list. + return [{"index": tool_index, "function": {"arguments": text}}] def _collect_response_item(item: object, slots: _SlotSink) -> None: @@ -638,9 +637,9 @@ class _AnthropicSSERestorer: self._step: Final = step self._carries: Final[dict[int, str]] = {} # mutable-ok: per-block windows advanced in place. self._delta_types: Final[dict[int, str]] = {} # mutable-ok: each block's delta type, for its flush. - self._pending = b"" # rebind-ok: the unfinished tail of the stream. - self._as_text = False # rebind-ok: set once if the stream arrives as str rather than bytes. - self._is_sse: bool | None = None # rebind-ok: undecided until the opening bytes settle it. + self._pending = b"" + self._as_text = False + self._is_sse: bool | None = None async def feed(self, chunk: bytes | str) -> tuple[bytes | str, ...]: """Restores every event this chunk completes; holds back an unfinished tail.""" @@ -669,9 +668,7 @@ class _AnthropicSSERestorer: # empty remainder after the last separator. parts: Final = _SSE_EVENT_BOUNDARY.split(buffered[:cut]) restored: Final = tuple( - [ # mutable-ok: an await needs a list comprehension; frozen at once. - await self._restore_event(parts[index]) + parts[index + 1] for index in range(0, len(parts) - 1, 2) - ] + [await self._restore_event(parts[index]) + parts[index + 1] for index in range(0, len(parts) - 1, 2)] ) return self._emit(b"".join(restored)) @@ -754,7 +751,7 @@ class _AnthropicSSERestorer: text, _ = await self._step("", carry, True) if not text: return b"" - event: Final[JsonBody] = { # mutable-ok: serialised on the next line. + event: Final[JsonBody] = { "type": "content_block_delta", "index": index, "delta": {"type": delta_type, field: text}, @@ -762,7 +759,7 @@ class _AnthropicSSERestorer: return f"event: content_block_delta\ndata: {json.dumps(event, ensure_ascii=False)}\n\n".encode() async def _flush_all(self) -> bytes: - flushed: Final = tuple([await self._flush(index) for index in tuple(self._carries)]) # mutable-ok: frozen. + flushed: Final = tuple([await self._flush(index) for index in tuple(self._carries)]) return b"".join(flushed) @@ -798,13 +795,13 @@ class _ResponsesStreamRestorer: if kind.endswith(".delta") and kind not in _RESPONSES_BINARY_DELTAS: await self._restore_delta(event, kind) return (event,) - slots: Final[_SlotSink] = [] # mutable-ok: accumulator, restored in one batch. + slots: Final[_SlotSink] = [] flushed: Final = await self._flush(_responses_stream_key(event, kind)) if kind.endswith(".done") else () if kind.endswith(".done"): _collect_event_text(event, slots) part: Final = _read_field(event, "part") if part is not None: - _collect_response_item({"content": [part]}, slots) # mutable-ok: a one-part view. + _collect_response_item({"content": [part]}, slots) _collect_response_item(_read_field(event, "item"), slots) elif kind in _RESPONSES_TERMINAL_EVENTS: for item in _read_list(_read_field(event, "response"), "output"): @@ -814,8 +811,8 @@ class _ResponsesStreamRestorer: async def finish(self) -> tuple[object, ...]: """Flushes every stream the provider never closed, e.g. a truncated reply.""" - flushed: Final = tuple([await self._flush(key) for key in tuple(self._carries)]) # mutable-ok: frozen. - return tuple(event for events in flushed for event in events) + flushed: Final = tuple([await self._flush(key) for key in tuple(self._carries)]) + return tuple(itertools.chain.from_iterable(flushed)) async def _restore_delta(self, event: object, kind: str) -> None: text: Final = _read_field(event, "delta") @@ -908,12 +905,12 @@ class LLMShieldProxyGuardrail(CustomGuardrail): @classmethod 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. + return [GuardrailEventHooks.pre_call, GuardrailEventHooks.post_call] # --- transport --------------------------------------------------------------- def _headers(self, session_id: str) -> JsonBody: - headers: Final[JsonBody] = { # mutable-ok: httpx requires a real dict. + headers: Final[JsonBody] = { "Content-Type": "application/json", "X-Session-ID": session_id, } @@ -951,12 +948,12 @@ class LLMShieldProxyGuardrail(CustomGuardrail): ) from exc async def _redact(self, texts: Sequence[str], session_id: str) -> Sequence[str]: - payload: Final[JsonBody] = {"texts": list(texts)} # mutable-ok: JSON body for httpx. + payload: Final[JsonBody] = {"texts": list(texts)} 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: Sequence[str], session_id: str) -> Sequence[str]: - payload: Final[JsonBody] = {"texts": list(texts)} # mutable-ok: JSON body for httpx. + payload: Final[JsonBody] = {"texts": list(texts)} body: Final = await self._call_shield(_REHYDRATE_PATH, session_id, payload) return self._same_length_or_raise(body.get("texts"), texts, "rehydrate") @@ -984,7 +981,7 @@ class LLMShieldProxyGuardrail(CustomGuardrail): # /v1/responses. The session id is a capability against the vault's rehydrate # endpoint, so handing it to the provider alongside the placeholders would let the # provider read back exactly what this guardrail exists to withhold. - metadata: Final = data.setdefault("litellm_metadata", {}) # mutable-ok: per-request store. + metadata: Final = data.setdefault("litellm_metadata", {}) if isinstance(metadata, dict): metadata[_SESSION_METADATA_KEY] = session_id return session_id @@ -1030,8 +1027,8 @@ class LLMShieldProxyGuardrail(CustomGuardrail): guardrail, and an agent that reads a file and quotes an address from it needs that address back. """ - slots: Final[_SlotSink] = [] # mutable-ok: accumulator, frozen on return. - privileged: Final[_SlotSink] = [] # mutable-ok: accumulator, frozen on return. + slots: Final[_SlotSink] = [] + privileged: Final[_SlotSink] = [] for message in data.get("messages") or (): if isinstance(message, dict): sink = privileged if message.get("role") in _PRIVILEGED_ROLES else slots @@ -1202,7 +1199,7 @@ class LLMShieldProxyGuardrail(CustomGuardrail): fields on the request side -- a function_call item holds `arguments`, a function_call_output holds `output` -- so the two directions stay symmetric. """ - slots: Final[_SlotSink] = [] # mutable-ok: accumulator, frozen on return. + slots: Final[_SlotSink] = [] for item in getattr(response, "output", None) or (): _collect_response_item(item, slots) return tuple(slots) @@ -1376,7 +1373,7 @@ class LLMShieldProxyGuardrail(CustomGuardrail): continuations.extend(_continuation_delta(tool_index, text)) if continuations: existing: Final = tuple(getattr(delta, "tool_calls", None) or ()) - delta.tool_calls = [*existing, *continuations] # mutable-ok: delta.tool_calls is a list. + delta.tool_calls = [*existing, *continuations] async def _flush_trailing( self, last_chunk: Any, carries: _CarryWindows, session_id: str @@ -1433,7 +1430,7 @@ class LLMShieldProxyGuardrail(CustomGuardrail): kept.index = index # The terminal signal, if there was one, already went out with the real chunk. kept.finish_reason = None - chunk.choices = [kept] # mutable-ok: the chunk model requires a list. + chunk.choices = [kept] return chunk async def _stream_step(self, text: str, carry: str, final: bool, session_id: str) -> tuple[str, str]: @@ -1441,8 +1438,7 @@ class LLMShieldProxyGuardrail(CustomGuardrail): body: Final = await self._call_shield( _REHYDRATE_STREAM_PATH, session_id, - # mutable-ok: JSON request body for httpx. - {"text": text, "carry": carry, "final": final}, # mutable-ok: JSON request body for httpx. + {"text": text, "carry": carry, "final": final}, ) emitted: Final = body.get("text") remaining: Final = body.get("carry") @@ -1500,7 +1496,7 @@ class LLMShieldProxyGuardrail(CustomGuardrail): write(replacement) # 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} # mutable-ok: TypedDict. + merged: Final[JsonBody] = {**inputs} if text_list: merged["texts"] = restored_values[: len(text_list)] if restored_calls: