refactor(grayswan): keep the monitor conversation immutable

The structured-conversation change tripped the LIT001 and LIT002
ceilings plus two basedpyright rules. Annotate the conversation helpers
with Mapping and tuple views, build turns in one shot through a
ReadOnly MonitorTurn TypedDict, and drop the runtime str checks that
the declared texts type already guarantees. The payload serializes to
the same JSON as before, so the tests now assert the wire form the
API receives instead of the Python container types
This commit is contained in:
Youlian Simidjiyski 2026-08-21 13:19:03 -04:00
parent 6ea100003f
commit edac577789
2 changed files with 60 additions and 42 deletions

View file

@ -2,9 +2,11 @@
import os
import time
from collections.abc import Mapping, Sequence
from typing import TYPE_CHECKING, Any, Final, Literal, Optional
from fastapi import HTTPException
from typing_extensions import NotRequired, ReadOnly, TypedDict
from litellm._logging import verbose_proxy_logger
from litellm.integrations.custom_guardrail import (
@ -28,6 +30,14 @@ GRAYSWAN_BLOCK_ERROR_MSG: Final = "Blocked by Gray Swan Guardrail"
GRAYSWAN_CONVERSATION_CACHE_KEY: Final = "_grayswan_request_conversation"
class MonitorTurn(TypedDict):
"""One conversation turn in the monitor payload."""
role: ReadOnly[str]
content: ReadOnly[str]
tool_calls: NotRequired[ReadOnly[tuple[Mapping[str, Any], ...]]]
class GraySwanGuardrailMissingSecrets(Exception):
"""Raised when the Gray Swan API key is missing."""
@ -209,7 +219,7 @@ class GraySwanGuardrail(CustomGuardrail):
input_type,
)
dynamic_body: Final = self.get_guardrail_dynamic_request_body_params(request_data) or {}
dynamic_body: Final = self.get_guardrail_dynamic_request_body_params(request_data)
if dynamic_body:
verbose_proxy_logger.debug("Gray Swan Guardrail: dynamic extra_body=%s", safe_dumps(dynamic_body))
@ -527,9 +537,9 @@ class GraySwanGuardrail(CustomGuardrail):
def _build_monitor_input(
self,
inputs: GenericGuardrailAPIInputs,
request_data: dict,
request_data: dict, # mutable-ok: the shared per-request state dict every hook receives; the request scan caches on it
input_type: Literal["request", "response"],
) -> tuple[list[dict[str, Any]], list[dict[str, Any]] | None]:
) -> tuple[tuple[Mapping[str, Any], ...], tuple[Mapping[str, Any], ...] | None]:
"""Build the monitor conversation from the translation layer's scoped view.
Request scans send `structured_messages` and `tools` exactly as the unified
@ -547,58 +557,57 @@ class GraySwanGuardrail(CustomGuardrail):
return conversation, tools
response_turns: Final = self._build_response_turns(inputs)
if not response_turns:
return [], None
return (), None
cached: Final = self._cached_request_conversation(request_data)
if cached is None:
return self._texts_fallback(inputs, "assistant"), None
return [*cached[0], *response_turns], cached[1]
return (*cached[0], *response_turns), cached[1]
def _cache_request_conversation(
self,
request_data: dict,
conversation: list[dict[str, Any]],
tools: list[dict[str, Any]] | None,
request_data: dict, # mutable-ok: the shared per-request state dict every hook receives; caching on it is the point
conversation: tuple[Mapping[str, Any], ...],
tools: tuple[Mapping[str, Any], ...] | None,
) -> None:
metadata: Final = request_data.setdefault(
"metadata", {}
"metadata",
{}, # mutable-ok: request metadata is shared mutable state other hooks also write to
) # rebind-ok: response scans replay the request-time scoped conversation and request_data is the only object shared across hooks
if isinstance(metadata, dict):
metadata[GRAYSWAN_CONVERSATION_CACHE_KEY] = (conversation, tools)
def _cached_request_conversation(
self, request_data: dict
) -> tuple[list[dict[str, Any]], list[dict[str, Any]] | None] | None:
self, request_data: Mapping[str, Any]
) -> tuple[tuple[Mapping[str, Any], ...], tuple[Mapping[str, Any], ...] | None] | None:
metadata: Final = request_data.get("metadata")
cached: Final = metadata.get(GRAYSWAN_CONVERSATION_CACHE_KEY) if isinstance(metadata, dict) else None
if isinstance(cached, tuple) and len(cached) == 2:
return cached
return None
def _build_response_turns(self, inputs: GenericGuardrailAPIInputs) -> list[dict[str, Any]]:
texts: Final = [
text if isinstance(text, str) else str(text) for text in inputs.get("texts", []) if text is not None
]
def _build_response_turns(self, inputs: GenericGuardrailAPIInputs) -> tuple[Mapping[str, Any], ...]:
tool_calls: Final = self._sanitize_json_list(inputs.get("tool_calls"))
text_turns: Final = [{"role": "assistant", "content": text} for text in texts if text]
text_turns: Final = tuple(self._turn("assistant", text) for text in inputs.get("texts", ()) if text)
if not tool_calls:
return text_turns
base: Final = text_turns or [{"role": "assistant", "content": ""}]
return [*base[:-1], {**base[-1], "tool_calls": tool_calls}]
base: Final = text_turns or (self._turn("assistant", ""),)
final_turn: Final[MonitorTurn] = {**base[-1], "tool_calls": tool_calls}
return (*base[:-1], final_turn)
def _texts_fallback(self, inputs: GenericGuardrailAPIInputs, role: str) -> list[dict[str, Any]]:
return [
{"role": role, "content": text if isinstance(text, str) else str(text)}
for text in inputs.get("texts", [])
if text is not None
]
def _texts_fallback(self, inputs: GenericGuardrailAPIInputs, role: str) -> tuple[MonitorTurn, ...]:
return tuple(self._turn(role, text) for text in inputs.get("texts", ()))
def _sanitize_json_list(self, value: object) -> list[dict[str, Any]] | None:
def _turn(self, role: str, content: str) -> MonitorTurn:
turn: Final[MonitorTurn] = {"role": role, "content": content}
return turn
def _sanitize_json_list(self, value: object) -> tuple[Mapping[str, Any], ...] | None:
if not isinstance(value, list) or not value:
return None
sanitized: Final = safe_json_loads(safe_dumps(value), default=None)
if not isinstance(sanitized, list):
return None
items: Final = [item for item in sanitized if isinstance(item, dict)]
items: Final = tuple(item for item in sanitized if isinstance(item, dict))
if len(items) != len(sanitized):
verbose_proxy_logger.debug(
"Gray Swan Guardrail: dropped %d non-dict conversation item(s)",
@ -608,11 +617,11 @@ class GraySwanGuardrail(CustomGuardrail):
def _prepare_payload(
self,
messages: list[dict[str, Any]],
messages: Sequence[Mapping[str, Any]],
dynamic_body: dict,
request_data: dict,
logging_obj: Optional["LiteLLMLoggingObj"] = None,
tools: list[dict[str, Any]] | None = None,
tools: Sequence[Mapping[str, Any]] | None = None,
) -> dict[str, Any] | None:
payload: Final[dict[str, Any]] = {"messages": messages}
if tools:

View file

@ -1,3 +1,4 @@
import json
from typing import Optional
import pytest
@ -596,6 +597,11 @@ def test_ensure_litellm_metadata_noop_when_already_present() -> None:
assert data["litellm_metadata"] == {"existing": "value"}
def _wire(payload: dict) -> dict:
"""The JSON wire form of a monitor payload, as the Cygnal API receives it."""
return json.loads(json.dumps(payload))
@pytest.mark.asyncio
async def test_apply_guardrail_sends_structured_conversation_with_tools(
monkeypatch, grayswan_guardrail: GraySwanGuardrail
@ -632,8 +638,9 @@ async def test_apply_guardrail_sends_structured_conversation_with_tools(
input_type="request",
)
assert captured["payload"]["messages"] == structured_messages
assert captured["payload"]["tools"] == tools
wire = _wire(captured["payload"])
assert wire["messages"] == structured_messages
assert wire["tools"] == tools
@pytest.mark.asyncio
@ -654,7 +661,7 @@ async def test_apply_guardrail_scans_texts_not_raw_messages(monkeypatch, grayswa
input_type="request",
)
assert captured["payload"]["messages"] == [{"role": "user", "content": "scan exactly this"}]
assert _wire(captured["payload"])["messages"] == [{"role": "user", "content": "scan exactly this"}]
assert "tools" not in captured["payload"]
@ -691,8 +698,9 @@ async def test_apply_guardrail_response_appends_assistant_turn_with_tool_calls(
input_type="response",
)
assert captured["payload"]["messages"][:-1] == structured_messages
assert captured["payload"]["messages"][-1] == {
wire = _wire(captured["payload"])
assert wire["messages"][:-1] == structured_messages
assert wire["messages"][-1] == {
"role": "assistant",
"content": "sure, writing it now",
"tool_calls": response_tool_calls,
@ -725,9 +733,10 @@ async def test_apply_guardrail_scans_tool_call_only_response(
input_type="response",
)
assert captured["payload"]["messages"][0] == {"role": "user", "content": "go"}
assert captured["payload"]["messages"][-1]["tool_calls"] == response_tool_calls
assert captured["payload"]["messages"][-1]["content"] == ""
wire = _wire(captured["payload"])
assert wire["messages"][0] == {"role": "user", "content": "go"}
assert wire["messages"][-1]["tool_calls"] == response_tool_calls
assert wire["messages"][-1]["content"] == ""
@pytest.mark.asyncio
@ -748,7 +757,7 @@ async def test_apply_guardrail_wraps_texts_when_no_conversation_available(
input_type="request",
)
assert captured["payload"]["messages"] == [{"role": "user", "content": "a prompt for an image"}]
assert _wire(captured["payload"])["messages"] == [{"role": "user", "content": "a prompt for an image"}]
assert "tools" not in captured["payload"]
@ -770,7 +779,7 @@ async def test_apply_guardrail_response_without_request_scan_wraps_texts(
input_type="response",
)
assert captured["payload"]["messages"] == [{"role": "assistant", "content": "a model reply"}]
assert _wire(captured["payload"])["messages"] == [{"role": "assistant", "content": "a model reply"}]
assert "tools" not in captured["payload"]
@ -827,7 +836,7 @@ async def test_apply_guardrail_multiple_response_texts_get_separate_turns(
input_type="response",
)
assert captured["payload"]["messages"][-2:] == [
assert _wire(captured["payload"])["messages"][-2:] == [
{"role": "assistant", "content": "candidate one"},
{"role": "assistant", "content": "candidate two"},
]
@ -872,6 +881,6 @@ async def test_apply_guardrail_build_error_fails_open(monkeypatch, grayswan_guar
def test_sanitize_json_list_drops_non_dict_items(grayswan_guardrail: GraySwanGuardrail) -> None:
assert grayswan_guardrail._sanitize_json_list([{"role": "user", "content": "hi"}, "junk", 3]) == [
{"role": "user", "content": "hi"}
]
assert grayswan_guardrail._sanitize_json_list([{"role": "user", "content": "hi"}, "junk", 3]) == (
{"role": "user", "content": "hi"},
)