fix(responses): presidio PII masking for Azure WebSocket and streaming (#30003)

* fix(responses): Presidio PII masking for Azure WebSocket and streaming

Wire Presidio into native Responses WebSocket forwarding and fix streaming output unmasking so masked tokens are restored for HTTP and WS clients.

Co-authored-by: Cursor <cursoragent@cursor.com>

* Fix unused imports in responses handlers

* fix(responses): address Greptile review - Azure WebSocket model URL and PII logging

- Add model_in_websocket_url() to BaseResponsesAPIConfig (default True) so
  providers can opt out of ?model= being appended to WebSocket URLs.
- Override model_in_websocket_url() to return False for Azure, since Azure
  sends the model in the response.create body, not the URL query string.
- Use this flag in llm_http_handler to conditionally append ?model=.
- Pass masked message to _store_input() instead of the original PII-containing
  message so logging destinations do not receive unmasked PII.

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>

* fix(responses): mask nested response.create input format for Presidio PII

Handle the nested {"type":"response.create","response":{"input":[...]}}
format in _mask_response_create. Previously only the flat top-level input
was masked; the nested shape bypassed Presidio and forwarded raw PII
upstream. Now both shapes are normalized and masked before forwarding.

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>

* style: apply black formatting to llm_http_handler and streaming_iterator

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>

* style: suppress PLR0915 on async_responses_websocket

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>

* fix(responses): add apply_to_output masking on Responses API WebSocket path

Previously the WebSocket guardrail filter excluded callbacks with
apply_to_output=True, leaving model-generated PII unmasked before
returning to the client.

- Collect apply_to_output callbacks separately in llm_http_handler and
  pass them to ResponsesWebSocketStreaming as output_guardrail_callbacks.
- Add _mask_response_completed method that calls check_pii(output_parse_pii=False)
  on text blocks in response.completed events, masking model output PII.
- backend_to_client now chains unmask (pii_tokens) → mask (apply_to_output)
  before forwarding each event to the client.

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>

* fix(responses): unmask PII tokens in streaming delta events and warn on guardrail init failure

- Rename _unmask_response_completed -> _unmask_response_event and extend
  it to also unmask response.output_text.delta (and other delta types)
  so real-time streaming clients receive original values, not PII tokens.
- Split the broad except-and-swallow into ImportError (expected in SDK-only
  environments) vs Exception (unexpected — now logs a warning so operators
  know masking is disabled).

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>

* fix(responses): enforce authorized model on WebSocket frames and remove proxy import

Security: add _enforce_authorized_model to ResponsesWebSocketStreaming that
overwrites both flat and nested model fields in every response.create frame
with the connection-authorized model, preventing deployment-substitution
attacks where an authenticated user sends a different model name in the frame
body after connecting with an allowed model.

Layering: remove the _OPTIONAL_PresidioPIIMasking isinstance check and proxy
import from the SDK handler. Use duck-typed checks (callable check_pii +
get_presidio_settings_from_request_data) so any guardrail implementing the
interface works, not just Presidio.

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>

* fix(responses): add _unmask_pii_text to duck-typed contract and mask delta frames

- Add callable(_unmask_pii_text) check to the guardrail_callbacks filter so
  a custom guardrail missing that method cannot cause an AttributeError and
  silently kill the WebSocket session.
- Extend _mask_response_completed to also mask response.output_text.delta
  (and other delta types) for apply_to_output callbacks, so real-time
  streaming clients do not receive unredacted model-generated PII in deltas.

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>

* Fix Responses WebSocket guardrail edge cases

* fix(responses): log masked output and suppress deltas when apply_to_output active

- Move _store_event to after _mask_response_completed so logs receive the
  redacted form, not raw model output containing PII.
- Suppress delta event forwarding when output_guardrail_callbacks are
  present: per-fragment Presidio cannot catch PII that spans multiple
  chunks (e.g. "alice@" + "example.com"). Clients receive only the
  fully-masked response.completed, which Presidio scans on complete text.

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>

* fix(responses): mask and suppress response.output_item.done for apply_to_output

response.output_item.done carries completed item text in item.content[*].text
before response.completed arrives, allowing unmasked PII to reach the client.

- _unmask_response_event: unmask input-PII tokens in item.content[*].text
- _mask_response_completed: run check_pii on item.content[*].text for
  apply_to_output callbacks (same as response.completed handling)
- backend_to_client suppression: also skip response.output_item.done when
  output_guardrail_callbacks are active; client receives only the
  fully-masked response.completed

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>

* fix(types): cast response_obj to ResponsesAPIResponse to satisfy mypy

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>

* Revert "fix(types): cast response_obj to ResponsesAPIResponse to satisfy mypy"

This reverts commit d5969557628f9aff58948b9d37cc64d577f95a15.

* Revert "fix(responses): mask and suppress response.output_item.done for apply_to_output"

This reverts commit 219fd54ea3446f4399fde40c07ba0617e2834573.

* fix(types): accept dict responses in guardrail output write-back

Streaming response.completed events pass a dict response object, so widen
_apply_guardrail_responses_to_output to match its existing runtime handling.

Co-authored-by: Cursor <cursoragent@cursor.com>

* perf(responses): skip Presidio masking on suppressed WebSocket delta events

Delta events are dropped wholesale when apply_to_output masking is active,
so masking them first issued a wasted check_pii call per fragment. Move the
suppression check ahead of the unmask/mask passes; the event type is
invariant across both, so client-visible behavior is unchanged.

* test(responses): cover Responses WebSocket PII masking hooks

Add regression tests for the native Responses WebSocket guardrail path:
input masking and model enforcement in _mask_response_create, token
unmasking in _unmask_response_event, apply_to_output masking and delta
suppression in _mask_response_completed/backend_to_client, and the
get_websocket_url / model_in_websocket_url defaults for the base and
Azure configs. Raises diff coverage above the codecov patch target.

* fix(responses): suppress text-bearing done events under output PII masking

When apply_to_output masking is active on a native Responses WebSocket,
response.output_text.done, response.content_part.done, and
response.output_item.done carry the full model output before the masked
response.completed arrives, so an authenticated client could read
unmasked PII from those events. Suppress them alongside delta events; the
client receives only the fully-masked response.completed.

* refactor(responses): drop dead delta branch in WebSocket output masking

Delta events are suppressed in backend_to_client before _mask_response_completed
runs when output masking is active, so the method's delta-handling branch was
unreachable. Restrict it to response.completed and cover the Responses API
unmask path with a Pydantic ResponseCompletedEvent regression test.

* fix(presidio): flush buffered chat chunks on mixed unmask stream

_stream_pii_unmasking buffered ModelResponseStream chunks but returned
early once a /v1/responses event was seen, silently dropping the buffered
chat chunks. Flush them in order before switching to passthrough, mirroring
_stream_apply_output_masking, and cover it with a regression test.

* fix(responses): mask instructions and tool-call arguments in WebSocket PII path

Presidio masking on the native Responses WebSocket path left two gaps. On the
request side _mask_response_create only walked the input containers, so PII
placed in the instructions field of a response.create frame was forwarded
upstream and logged unmasked even with output_parse_pii enabled. Now both the
flat and nested instructions strings are masked alongside input.

On the response side _mask_response_completed only masked content text blocks,
so model-produced PII inside function-call arguments could reach the client when
apply_to_output was enabled, both via the standalone
response.function_call_arguments.done event and via the function_call output
items in response.completed. The done event is now suppressed under output
masking and completed function-call arguments are run through check_pii before
forwarding or logging.

* fix(responses): suppress reasoning_summary_text.done under output PII masking

* fix(responses): mask function_call_output.output in WebSocket PII path

response.create input items of type function_call_output carry
user-controlled text in output, not content, so the Presidio masking
pass forwarded that text upstream unmasked. Mask the output field
(string or list of text blocks) alongside content.

* fix(responses): mask reasoning summary PII in WebSocket output path

---------

Co-authored-by: Cursor <cursoragent@cursor.com>
Co-authored-by: Claude Sonnet 4.6 <noreply@anthropic.com>
Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com>
This commit is contained in:
Sameer Kankute 2026-06-12 19:57:03 +05:30 • committed by GitHub
parent 02bce7b393
commit ccc20b121f
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
9 changed files with 2058 additions and 88 deletions

View file

@ -185,6 +185,40 @@ class AzureOpenAIResponsesAPIConfig(OpenAIResponsesAPIConfig):
default_api_version=AZURE_DEFAULT_RESPONSES_API_VERSION,
)
def supports_native_websocket(self) -> bool:
return True
def get_websocket_url(
self,
api_base: Optional[str],
litellm_params: dict,
) -> str:
"""
Azure Responses WebSocket endpoint is at /openai/v1/responses with no
api-version query param. Auth is via Authorization header, model is sent
in the response.create body — not the URL.
"""
if api_base is None:
raise ValueError("api_base is required for Azure WebSocket")
parsed_url = httpx.URL(api_base)
path = parsed_url.path.rstrip("/")
# Strip existing /openai/responses path if the api_base already contains it
for suffix in ("/openai/v1/responses", "/openai/responses"):
if path.endswith(suffix):
path = path[: -len(suffix)]
break
scheme = "wss" if parsed_url.scheme == "https" else "ws"
return str(
parsed_url.copy_with(
scheme=scheme, path=f"{path}/openai/v1/responses", query=None
)
)
def model_in_websocket_url(self) -> bool:
# Azure sends the model in the response.create body, not the URL
return False
#########################################################
########## DELETE RESPONSE API TRANSFORMATION ##############
#########################################################

View file

@ -258,6 +258,31 @@ class BaseResponsesAPIConfig(ABC):
"""
return False
def get_websocket_url(
self,
api_base: Optional[str],
litellm_params: dict,
) -> str:
"""
Return the wss:// URL for the provider's native Responses WebSocket endpoint.
Defaults to converting the HTTP URL from get_complete_url. Providers whose
WebSocket path differs from their HTTP path (e.g. Azure uses
/openai/v1/responses without api-version) should override this.
"""
http_url = self.get_complete_url(
api_base=api_base, litellm_params=litellm_params
)
return http_url.replace("https://", "wss://").replace("http://", "ws://")
def model_in_websocket_url(self) -> bool:
"""
Return True if the model should be appended as a ?model= query param to
the WebSocket URL. Providers that identify the model via the request body
(e.g. Azure Responses API) should override this to return False.
"""
return True
#########################################################
########## CANCEL RESPONSE API TRANSFORMATION ##########
#########################################################

View file

@ -5671,7 +5671,7 @@ class BaseLLMHTTPHandler:
)
raise
async def async_responses_websocket(
async def async_responses_websocket( # noqa: PLR0915
self,
model: str,
websocket: Any,
@ -5724,7 +5724,11 @@ class BaseLLMHTTPHandler:
import websockets
from websockets.asyncio.client import ClientConnection
litellm_params = GenericLiteLLMParams()
litellm_params = GenericLiteLLMParams(
api_base=api_base,
api_key=api_key,
**kwargs,
)
headers = responses_api_provider_config.validate_environment(
headers={},
model=model,
@ -5733,21 +5737,21 @@ class BaseLLMHTTPHandler:
if api_key:
headers["Authorization"] = f"Bearer {api_key}"
http_url = responses_api_provider_config.get_complete_url(
ws_url = responses_api_provider_config.get_websocket_url(
api_base=api_base,
litellm_params={},
litellm_params=dict(litellm_params),
)
ws_url = http_url.replace("https://", "wss://").replace("http://", "ws://")
# OpenAI's WebSocket responses endpoint requires ?model= in the URL,
# matching the Realtime API convention (wss://.../v1/realtime?model=...).
# Use urllib.parse so existing query params (e.g. api-version) are preserved.
_parsed = urlparse(ws_url)
_qs = parse_qs(_parsed.query)
if "model" not in _qs:
_qs["model"] = [model]
ws_url = urlunparse(
_parsed._replace(query=urlencode({k: v[0] for k, v in _qs.items()}))
)
# Some providers (e.g. OpenAI) require ?model= in the WebSocket URL.
# Providers that send the model in the request body (e.g. Azure) set
# model_in_websocket_url() to False to suppress this append.
if responses_api_provider_config.model_in_websocket_url():
_parsed = urlparse(ws_url)
_qs = parse_qs(_parsed.query)
if "model" not in _qs:
_qs["model"] = [model]
ws_url = urlunparse(
_parsed._replace(query=urlencode({k: v[0] for k, v in _qs.items()}))
)
try:
ssl_context = get_shared_realtime_ssl_context()
@ -5775,6 +5779,41 @@ class BaseLLMHTTPHandler:
_request_data: Dict[str, Any] = {}
if litellm_metadata:
_request_data["litellm_metadata"] = litellm_metadata
_ws_guardrail_callbacks: list = []
_ws_output_guardrail_callbacks: list = []
try:
import litellm as _litellm
# Use duck-typing so any guardrail that exposes the PII
# masking interface works, not just _OPTIONAL_PresidioPIIMasking.
# This avoids a layering violation (SDK importing from proxy).
_ws_guardrail_callbacks = [
cb
for cb in _litellm.callbacks
if callable(getattr(cb, "check_pii", None))
and callable(
getattr(cb, "get_presidio_settings_from_request_data", None)
)
and callable(getattr(cb, "_unmask_pii_text", None))
and getattr(cb, "output_parse_pii", False)
]
_ws_output_guardrail_callbacks = [
cb
for cb in _litellm.callbacks
if callable(getattr(cb, "check_pii", None))
and callable(
getattr(cb, "get_presidio_settings_from_request_data", None)
)
and getattr(cb, "apply_to_output", False)
]
except Exception as _guardrail_exc:
verbose_logger.warning(
"Responses WebSocket: failed to collect guardrail "
"callbacks — PII masking will be skipped. Error: %s",
_guardrail_exc,
)
streaming = ResponsesWebSocketStreaming(
websocket=websocket,
backend_ws=cast(ClientConnection, backend_ws),
@ -5782,6 +5821,9 @@ class BaseLLMHTTPHandler:
user_api_key_dict=user_api_key_dict,
request_data=_request_data,
first_message=first_message,
guardrail_callbacks=_ws_guardrail_callbacks,
output_guardrail_callbacks=_ws_output_guardrail_callbacks,
authorized_model=model,
)
await streaming.bidirectional_forward()

View file

@ -35,7 +35,6 @@ from pydantic import BaseModel
from litellm._logging import verbose_proxy_logger
from litellm.completion_extras.litellm_responses_transformation.transformation import (
LiteLLMResponsesTransformationHandler,
OpenAiResponsesToChatCompletionStreamIterator,
)
from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation
@ -479,90 +478,137 @@ class OpenAIResponsesHandler(BaseTranslation):
) -> List[Any]:
"""
Process output streaming response by applying guardrails to text content.
Mirrors the Chat Completions handler pattern: extract text from the final
chunk, apply the guardrail, then write the result back in-place so the
caller sees the modified content (e.g. PII tokens replaced).
For ``response.completed`` events (the normal end-of-stream signal) we
use the same per-item extraction + task-mapping approach as
``process_output_response`` so that unmasking / blocking works correctly
for every output item.
"""
if not responses_so_far:
return responses_so_far
final_chunk = responses_so_far[-1]
# Accept both plain dicts and Pydantic models (BaseLiteLLMOpenAIResponseObject
# exposes a .get() shim, so all the .get() calls below work for both).
if not (isinstance(final_chunk, dict) or hasattr(final_chunk, "get")):
return responses_so_far
# ------------------------------------------------------------------ #
# Case 1: response.completed — full response is available in the #
# final chunk; iterate output items, apply guardrail, write back. #
# ------------------------------------------------------------------ #
if final_chunk.get("type") == "response.completed":
response_obj = final_chunk.get("response") or {}
if not hasattr(response_obj, "get"):
return responses_so_far
outputs: List[Any] = response_obj.get("output") or []
texts_to_check: List[str] = []
tool_calls_to_check: List[ChatCompletionToolCallChunk] = []
task_mappings: List[Tuple[int, int]] = []
for output_idx, output_item in enumerate(outputs):
self._extract_output_text_and_images(
output_item=output_item,
output_idx=output_idx,
texts_to_check=texts_to_check,
images_to_check=[],
task_mappings=task_mappings,
tool_calls_to_check=tool_calls_to_check,
)
if texts_to_check or tool_calls_to_check:
if request_data is None:
request_data = {}
if "response" not in request_data:
request_data["response"] = response_obj
if "litellm_metadata" not in request_data:
user_metadata = self.transform_user_api_key_dict_to_metadata(
user_api_key_dict
)
if user_metadata:
request_data["litellm_metadata"] = user_metadata
inputs = GenericGuardrailAPIInputs(texts=texts_to_check)
if tool_calls_to_check:
inputs["tool_calls"] = cast(
List[ChatCompletionToolCallChunk], tool_calls_to_check
)
response_model = response_obj.get("model")
if response_model:
inputs["model"] = response_model
guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
inputs=inputs,
request_data=request_data,
input_type="response",
logging_obj=litellm_logging_obj,
)
guardrailed_texts = guardrailed_inputs.get("texts", [])
# Write guardrailed texts back into the output items in-place.
# final_chunk is a reference into responses_so_far so this
# mutates the list that the caller holds.
await self._apply_guardrail_responses_to_output(
response=response_obj,
responses=guardrailed_texts,
task_mappings=task_mappings,
)
return responses_so_far
# ------------------------------------------------------------------ #
# Case 2: response.output_item.done — extract tool calls only. #
# ------------------------------------------------------------------ #
if final_chunk.get("type") == "response.output_item.done":
# convert openai response to model response
model_response_stream = OpenAiResponsesToChatCompletionStreamIterator.translate_responses_chunk_to_openai_stream(
final_chunk
)
tool_calls = model_response_stream.choices[0].delta.tool_calls
if tool_calls:
inputs = GenericGuardrailAPIInputs()
inputs["tool_calls"] = cast(
List[ChatCompletionToolCallChunk], tool_calls
)
# Include model information if available
if (
hasattr(model_response_stream, "model")
and model_response_stream.model
):
inputs["model"] = model_response_stream.model
_guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
await guardrail_to_apply.apply_guardrail(
inputs=inputs,
request_data=request_data if request_data is not None else {},
input_type="response",
logging_obj=litellm_logging_obj,
)
return responses_so_far
elif final_chunk.get("type") == "response.completed":
# convert openai response to model response
outputs = final_chunk.get("response", {}).get("output", [])
return responses_so_far
model_response_choices = LiteLLMResponsesTransformationHandler._convert_response_output_to_choices(
output_items=outputs,
handle_raw_dict_callback=None,
)
if model_response_choices:
tool_calls = model_response_choices[0].message.tool_calls
text = model_response_choices[0].message.content
guardrail_inputs = GenericGuardrailAPIInputs()
if text:
guardrail_inputs["texts"] = [text]
if tool_calls:
guardrail_inputs["tool_calls"] = cast(
List[ChatCompletionToolCallChunk], tool_calls
)
# Include model information from the response if available
response_model = final_chunk.get("response", {}).get("model")
if response_model:
guardrail_inputs["model"] = response_model
if tool_calls or text:
_guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
inputs=guardrail_inputs,
request_data=request_data if request_data is not None else {},
input_type="response",
logging_obj=litellm_logging_obj,
)
return responses_so_far
else:
verbose_proxy_logger.debug(
"Skipping output guardrail - model response has no choices"
)
# model_response_stream = OpenAiResponsesToChatCompletionStreamIterator.translate_responses_chunk_to_openai_stream(final_chunk)
# tool_calls = model_response_stream.choices[0].tool_calls
# convert openai response to model response
# ------------------------------------------------------------------ #
# Fallback: apply guardrail to the accumulated text string. #
# No structured write-back is possible here; guardrails that only #
# need to block/flag (not rewrite) still work correctly. #
# ------------------------------------------------------------------ #
string_so_far = self.get_streaming_string_so_far(responses_so_far)
inputs = GenericGuardrailAPIInputs(texts=[string_so_far])
# Try to get model from the final chunk if available
if isinstance(final_chunk, dict):
if string_so_far:
fallback_inputs = GenericGuardrailAPIInputs(texts=[string_so_far])
response_model = (
final_chunk.get("response", {}).get("model")
if isinstance(final_chunk.get("response"), dict)
else None
)
if response_model:
inputs["model"] = response_model
_guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
inputs=inputs,
request_data=request_data if request_data is not None else {},
input_type="response",
logging_obj=litellm_logging_obj,
)
fallback_inputs["model"] = response_model
await guardrail_to_apply.apply_guardrail(
inputs=fallback_inputs,
request_data=request_data if request_data is not None else {},
input_type="response",
logging_obj=litellm_logging_obj,
)
return responses_so_far
def _check_streaming_has_ended(self, responses_so_far: List[Any]) -> bool:
@ -721,7 +767,7 @@ class OpenAIResponsesHandler(BaseTranslation):
async def _apply_guardrail_responses_to_output(
self,
response: "ResponsesAPIResponse",
response: Union["ResponsesAPIResponse", Dict[Any, Any]],
responses: List[str],
task_mappings: List[Tuple[int, int]],
) -> None:

View file

@ -1194,9 +1194,10 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
return
if not all_chunks:
verbose_proxy_logger.warning(
"Presidio apply_to_output: streaming response contained only "
"bytes chunks (Anthropic native SSE). Output PII masking was "
"skipped for this response."
"Presidio apply_to_output: streaming response contained no "
"ModelResponseStream chunks (e.g. raw SSE bytes or an empty "
"upstream stream). Output PII masking was skipped for this "
"response."
)
return
@ -1258,6 +1259,37 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
return "\n".join(result_lines).encode("utf-8")
def _unmask_responses_api_completed_chunk(
self, chunk: Any, pii_tokens: Dict[str, str]
) -> None:
"""
Unmask PII tokens in-place for a ``response.completed`` Responses API event.
The chunk carries a ``response`` attribute (ResponsesAPIResponse) whose
``output`` list holds message items. Each item has a ``content`` list of
blocks; text blocks expose a ``.text`` string attribute. We walk the tree
and replace every PII token with its original value.
"""
response_obj = getattr(chunk, "response", None)
if response_obj is None:
return
output = getattr(response_obj, "output", None) or []
for output_item in output:
content = getattr(output_item, "content", None) or []
for content_block in content:
if isinstance(content_block, dict):
if isinstance(content_block.get("text"), str):
content_block["text"] = self._unmask_pii_text(
content_block["text"], pii_tokens
)
elif hasattr(content_block, "text") and isinstance(
content_block.text, str
):
content_block.text = self._unmask_pii_text(
content_block.text, pii_tokens
)
async def _stream_pii_unmasking(
self,
response: Any,
@ -1274,16 +1306,36 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
pii_tokens: Dict[str, str] = metadata.get("pii_tokens", {})
remaining_chunks: List[ModelResponseStream] = []
saw_non_chat_chunk = False
try:
async for chunk in response:
if isinstance(chunk, ModelResponseStream):
remaining_chunks.append(chunk)
if saw_non_chat_chunk:
yield chunk
else:
remaining_chunks.append(chunk)
elif isinstance(chunk, bytes):
if pii_tokens:
yield self._unmask_sse_bytes_chunk(chunk, pii_tokens) # type: ignore[misc]
else:
yield chunk # type: ignore[misc]
continue
else:
# /v1/responses events: unmask response.completed text in-place.
# A mixed stream can't be reassembled, so flush buffered chat
# chunks in order before passthrough instead of dropping them.
if remaining_chunks and not saw_non_chat_chunk:
for buffered_chunk in remaining_chunks:
yield buffered_chunk
remaining_chunks = []
chunk_type = getattr(chunk, "type", None)
if chunk_type == "response.completed" and pii_tokens:
self._unmask_responses_api_completed_chunk(chunk, pii_tokens)
saw_non_chat_chunk = True
yield chunk
if saw_non_chat_chunk:
return
if not remaining_chunks:
return

View file

@ -1231,6 +1231,10 @@ RESPONSES_WS_LOGGED_EVENT_TYPES = [
"error",
]
RESPONSES_WS_MASKABLE_TEXT_BLOCK_TYPES = frozenset(
{"input_text", "output_text", "text"}
)
class ResponsesWebSocketStreaming:
"""
@ -1253,6 +1257,9 @@ class ResponsesWebSocketStreaming:
user_api_key_dict: Optional[Any] = None,
request_data: Optional[Dict] = None,
first_message: Optional[str] = None,
guardrail_callbacks: Optional[List[Any]] = None,
output_guardrail_callbacks: Optional[List[Any]] = None,
authorized_model: Optional[str] = None,
):
self.websocket = websocket
self.backend_ws = backend_ws
@ -1262,6 +1269,11 @@ class ResponsesWebSocketStreaming:
self.messages: list[Dict] = []
self.input_messages: list[Dict[str, str]] = []
self.first_message = first_message
self.guardrail_callbacks: List[Any] = guardrail_callbacks or []
self.output_guardrail_callbacks: List[Any] = output_guardrail_callbacks or []
# Model name authorized at connection time; enforced on every
# response.create frame to prevent deployment-substitution attacks.
self.authorized_model: Optional[str] = authorized_model
def _should_store_event(self, event_obj: dict) -> bool:
return event_obj.get("type") in RESPONSES_WS_LOGGED_EVENT_TYPES
@ -1352,8 +1364,33 @@ class ResponsesWebSocketStreaming:
else:
response_str = raw_response
self._store_event(response_str)
await self.websocket.send_text(response_str)
# When apply_to_output masking is active, suppress delta events
# and the text-bearing "done" events. Per-fragment Presidio
# cannot reliably catch PII spanning multiple delta chunks (e.g.
# "alice@" + "example.com"), and the done events carry the full
# output text that response.completed already delivers in
# fully-masked form; forwarding them would leak unmasked PII
# before response.completed arrives. The client receives only the
# masked response.completed.
if self.output_guardrail_callbacks:
try:
_evt_type = json.loads(response_str).get("type")
except (json.JSONDecodeError, TypeError):
_evt_type = None
if (
_evt_type in self._DELTA_EVENT_TYPES
or _evt_type in self._OUTPUT_DONE_EVENT_TYPES
):
continue
unmasked_str = self._unmask_response_event(response_str)
output_masked_str = await self._mask_response_completed(unmasked_str)
# Log the output-masked form so PII redacted by apply_to_output
# guardrails does not appear in success logs.
self._store_event(output_masked_str)
await self.websocket.send_text(output_masked_str)
except websockets.exceptions.ConnectionClosed as e: # type: ignore
verbose_logger.debug("Responses WS backend connection closed: %s", e)
@ -1362,20 +1399,316 @@ class ResponsesWebSocketStreaming:
finally:
await self._log_messages()
def _enforce_authorized_model(self, msg_obj: dict) -> bool:
"""
Overwrite any ``model`` field in a ``response.create`` frame with the
connection-authorized model to prevent deployment-substitution attacks.
Handles both shapes:
flat: ``{"type": "response.create", "model": "...", ...}``
nested: ``{"type": "response.create", "response": {"model": "...", ...}}``
Returns True if the object was modified.
"""
if not self.authorized_model:
return False
modified = False
nested = msg_obj.get("response")
if isinstance(nested, dict):
if nested.get("model") != self.authorized_model:
nested["model"] = self.authorized_model
modified = True
if "model" in msg_obj and msg_obj["model"] != self.authorized_model:
msg_obj["model"] = self.authorized_model
modified = True
elif msg_obj.get("model") != self.authorized_model:
msg_obj["model"] = self.authorized_model
modified = True
return modified
async def _mask_response_create(self, message: str) -> str:
"""
Enforce the authorized model and apply Presidio PII masking to a
``response.create`` message before it is forwarded to the upstream
provider.
- Overwrites any ``model`` field with the connection-authorized model
to prevent deployment-substitution attacks (always applied).
- Walks the ``input`` and ``instructions`` fields, calls ``check_pii``
on every text block, and stores the resulting ``pii_tokens`` map in
``self.request_data["metadata"]`` for later unmasking.
Non-``response.create`` messages are returned unchanged.
"""
try:
msg_obj = json.loads(message)
except (json.JSONDecodeError, TypeError):
return message
if msg_obj.get("type") != "response.create":
return message
# Always enforce the authorized model, even when PII masking is off.
model_modified = self._enforce_authorized_model(msg_obj)
if not self.guardrail_callbacks:
return json.dumps(msg_obj) if model_modified else message
if "metadata" not in self.request_data:
self.request_data["metadata"] = {}
modified = model_modified
for cb in self.guardrail_callbacks:
presidio_config = cb.get_presidio_settings_from_request_data(
self.request_data
)
# response.create carries client text in two shapes:
# flat: {"type": "response.create", "input": ..., "instructions": ...}
# nested: {"type": "response.create", "response": {"input": ..., "instructions": ...}}
# Mask "input" and "instructions" in both shapes so PII is never
# forwarded unmasked regardless of where the client places it.
nested_response = (
msg_obj.get("response")
if isinstance(msg_obj.get("response"), dict)
else None
)
text_containers: list[tuple[dict, str]] = []
for container in (msg_obj, nested_response):
if container is None:
continue
if "input" in container:
text_containers.append((container, "input"))
if isinstance(container.get("instructions"), str):
text_containers.append((container, "instructions"))
for container, key in text_containers:
field_value = container[key]
if isinstance(field_value, str):
container[key] = await cb.check_pii(
text=field_value,
output_parse_pii=True,
presidio_config=presidio_config,
request_data=self.request_data,
)
modified = True
elif isinstance(field_value, list):
for item in field_value:
if not isinstance(item, dict):
continue
for item_field in ("content", "output"):
value = item.get(item_field)
if isinstance(value, str):
item[item_field] = await cb.check_pii(
text=value,
output_parse_pii=True,
presidio_config=presidio_config,
request_data=self.request_data,
)
modified = True
elif isinstance(value, list):
for block in value:
if (
isinstance(block, dict)
and block.get("type")
in RESPONSES_WS_MASKABLE_TEXT_BLOCK_TYPES
and isinstance(block.get("text"), str)
):
block["text"] = await cb.check_pii(
text=block["text"],
output_parse_pii=True,
presidio_config=presidio_config,
request_data=self.request_data,
)
modified = True
return json.dumps(msg_obj) if modified else message
# Delta event types whose ``delta`` field may contain PII tokens.
_DELTA_EVENT_TYPES = frozenset(
{
"response.output_text.delta",
"response.reasoning_summary_text.delta",
"response.refusal.delta",
"response.function_call_arguments.delta",
}
)
# Terminal events that carry the full output text or tool-call arguments
# already delivered by ``response.completed``. Suppressed when output masking
# is active so the unmasked copy never reaches the client before the masked
# completed event.
_OUTPUT_DONE_EVENT_TYPES = frozenset(
{
"response.output_text.done",
"response.content_part.done",
"response.output_item.done",
"response.function_call_arguments.done",
"response.reasoning_summary_text.done",
"response.reasoning_summary_part.done",
}
)
def _unmask_response_event(self, response_str: str) -> str:
"""
Apply Presidio PII unmasking to backend events before forwarding to
the client.
Handles two shapes:
- ``response.completed``: walks ``response.output[*].content[*].text``
- streaming delta events (``response.output_text.delta``, etc.):
replaces tokens in the ``delta`` field
Uses the ``pii_tokens`` map stored during ``_mask_response_create`` to
replace every token (e.g. ``<EMAIL_ADDRESS_1>``) with the original
value. Events with no stored tokens are returned unchanged.
"""
if not self.guardrail_callbacks:
return response_str
pii_tokens: Dict[str, str] = (self.request_data.get("metadata") or {}).get(
"pii_tokens", {}
)
if not pii_tokens:
return response_str
try:
evt_obj = json.loads(response_str)
except (json.JSONDecodeError, TypeError):
return response_str
cb = self.guardrail_callbacks[0]
event_type = evt_obj.get("type")
if event_type == "response.completed":
modified = False
response_obj = evt_obj.get("response") or {}
if not isinstance(response_obj, dict):
return response_str
for output_item in response_obj.get("output") or []:
if not isinstance(output_item, dict):
continue
content = output_item.get("content") or []
if not isinstance(content, list):
continue
for content_block in content:
if not isinstance(content_block, dict):
continue
text = content_block.get("text")
if isinstance(text, str):
unmasked = cb._unmask_pii_text(text, pii_tokens)
if unmasked != text:
content_block["text"] = unmasked
modified = True
return json.dumps(evt_obj) if modified else response_str
if event_type in self._DELTA_EVENT_TYPES:
delta = evt_obj.get("delta")
if isinstance(delta, str):
unmasked = cb._unmask_pii_text(delta, pii_tokens)
if unmasked != delta:
evt_obj["delta"] = unmasked
return json.dumps(evt_obj)
return response_str
async def _mask_response_completed(self, response_str: str) -> str:
"""
Apply Presidio output masking (apply_to_output=True) to the
``response.completed`` event before it is forwarded to the client.
Walks ``response.output[*].content[*].text`` and masks every text block,
as well as ``response.output[*].arguments`` on function-call items and
``response.output[*].summary[*].text`` on reasoning items. Delta and
``*.done`` events are suppressed upstream in ``backend_to_client`` when
output masking is active, so only the authoritative full-output view
reaches this method; events of other types are returned unchanged.
"""
if not self.output_guardrail_callbacks:
return response_str
try:
evt_obj = json.loads(response_str)
except (json.JSONDecodeError, TypeError):
return response_str
if evt_obj.get("type") != "response.completed":
return response_str
modified = False
for cb in self.output_guardrail_callbacks:
presidio_config = cb.get_presidio_settings_from_request_data(
self.request_data
)
response_obj = evt_obj.get("response") or {}
if not isinstance(response_obj, dict):
continue
for output_item in response_obj.get("output") or []:
if not isinstance(output_item, dict):
continue
arguments = output_item.get("arguments")
if isinstance(arguments, str):
masked_args = await cb.check_pii(
text=arguments,
output_parse_pii=False,
presidio_config=presidio_config,
request_data=self.request_data,
)
if masked_args != arguments:
output_item["arguments"] = masked_args
modified = True
summary = output_item.get("summary") or []
if isinstance(summary, list):
for summary_block in summary:
if not isinstance(summary_block, dict):
continue
summary_text = summary_block.get("text")
if isinstance(summary_text, str):
masked_summary = await cb.check_pii(
text=summary_text,
output_parse_pii=False,
presidio_config=presidio_config,
request_data=self.request_data,
)
if masked_summary != summary_text:
summary_block["text"] = masked_summary
modified = True
content = output_item.get("content") or []
if not isinstance(content, list):
continue
for content_block in content:
if not isinstance(content_block, dict):
continue
text = content_block.get("text")
if isinstance(text, str):
masked = await cb.check_pii(
text=text,
output_parse_pii=False,
presidio_config=presidio_config,
request_data=self.request_data,
)
if masked != text:
content_block["text"] = masked
modified = True
return json.dumps(evt_obj) if modified else response_str
async def client_to_backend(self) -> None:
"""Forward response.create events from client to backend."""
try:
if self.first_message is not None:
self._store_input(self.first_message)
self._store_event(self.first_message)
await self.backend_ws.send(self.first_message) # type: ignore[union-attr]
masked_first = await self._mask_response_create(self.first_message)
self._store_input(masked_first)
self._store_event(masked_first)
await self.backend_ws.send(masked_first) # type: ignore[union-attr]
while True:
message = await self.websocket.receive_text()
self._store_input(message)
self._store_event(message)
await self.backend_ws.send(message) # type: ignore[union-attr]
masked = await self._mask_response_create(message)
self._store_input(masked)
self._store_event(masked)
await self.backend_ws.send(masked) # type: ignore[union-attr]
except Exception as e:
verbose_logger.debug("Responses WS client_to_backend ended: %s", e)

View file

@ -900,6 +900,20 @@ class TestOpenAIResponsesHandlerStreamingOutputProcessing:
# Should return the responses unchanged
assert result == responses_so_far
@pytest.mark.asyncio
async def test_process_output_streaming_response_null_response(self):
handler = OpenAIResponsesHandler()
guardrail = MockPassThroughGuardrail(guardrail_name="test")
responses_so_far = [{"type": "response.completed", "response": None}]
result = await handler.process_output_streaming_response(
responses_so_far=responses_so_far,
guardrail_to_apply=guardrail,
litellm_logging_obj=None,
)
assert result == responses_so_far
@pytest.mark.asyncio
async def test_process_output_streaming_response_unrecognized_output_type(self):
"""Test that streaming response with unrecognized output types doesn't raise IndexError
@ -996,6 +1010,105 @@ class TestOpenAIResponsesHandlerStreamingOutputProcessing:
# Should return the responses
assert result == responses_so_far
@pytest.mark.asyncio
async def test_process_output_streaming_response_writes_back_guardrailed_text(self):
"""Guardrailed text must be written back into the response.completed chunk in-place."""
class RewriteGuardrail(CustomGuardrail):
"""Replaces '<TOKEN_1>' with 'john@example.com' to simulate PII unmasking."""
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,
request_data: dict,
input_type: Literal["request", "response"],
logging_obj: Optional[Any] = None,
) -> GenericGuardrailAPIInputs:
texts = inputs.get("texts", [])
inputs["texts"] = [
t.replace("<TOKEN_1>", "john@example.com") for t in texts
]
return inputs
handler = OpenAIResponsesHandler()
guardrail = RewriteGuardrail(guardrail_name="test-rewrite")
responses_so_far = [
{"type": "response.output_text.delta", "delta": "send to "},
{"type": "response.output_text.delta", "delta": "<TOKEN_1>"},
{
"type": "response.completed",
"response": {
"id": "resp_123",
"model": "gpt-4o",
"output": [
{
"type": "message",
"id": "msg_123",
"status": "completed",
"role": "assistant",
"content": [
{"type": "output_text", "text": "send to <TOKEN_1>"},
],
}
],
"status": "completed",
},
},
]
result = await handler.process_output_streaming_response(
responses_so_far=responses_so_far,
guardrail_to_apply=guardrail,
litellm_logging_obj=None,
)
completed_chunk = next(
c
for c in result
if isinstance(c, dict) and c.get("type") == "response.completed"
)
output_text = completed_chunk["response"]["output"][0]["content"][0]["text"]
assert (
output_text == "send to john@example.com"
), f"Expected PII token to be unmasked in response.completed output, got: {output_text!r}"
@pytest.mark.asyncio
async def test_process_output_streaming_response_pass_through_unchanged(self):
"""A pass-through guardrail must not modify the output text."""
handler = OpenAIResponsesHandler()
guardrail = MockPassThroughGuardrail(guardrail_name="pass-through")
original_text = "No PII here, just normal text."
responses_so_far = [
{
"type": "response.completed",
"response": {
"id": "resp_456",
"model": "gpt-4o",
"output": [
{
"type": "message",
"id": "msg_456",
"status": "completed",
"role": "assistant",
"content": [{"type": "output_text", "text": original_text}],
}
],
"status": "completed",
},
}
]
result = await handler.process_output_streaming_response(
responses_so_far=responses_so_far,
guardrail_to_apply=guardrail,
litellm_logging_obj=None,
)
output_text = result[-1]["response"]["output"][0]["content"][0]["text"]
assert output_text == original_text
class TestGetStructuredMessages:
"""Test the get_structured_messages method for Responses API handler."""

View file

@ -2329,6 +2329,164 @@ async def test_apply_to_output_streaming_bytes_only_logs_warning():
assert "Output PII masking was skipped" in warning_msg
@pytest.mark.asyncio
async def test_output_parse_pii_streaming_responses_events_passthrough(
mock_user_api_key,
):
"""
Regression test: when output_parse_pii=True and pii_tokens exist, /v1/responses
streaming events must pass through instead of being dropped.
"""
guardrail = _OPTIONAL_PresidioPIIMasking(
mock_testing=True,
output_parse_pii=True,
)
response_events = [
{"type": "response.created", "response": {"id": "resp_1"}},
{"type": "response.output_text.delta", "delta": "Hello"},
{
"type": "response.completed",
"response": {"id": "resp_1", "status": "completed"},
},
]
async def mock_stream():
for event in response_events:
yield event
collected = []
async for chunk in guardrail.async_post_call_streaming_iterator_hook(
user_api_key_dict=mock_user_api_key,
response=mock_stream(),
request_data={
"metadata": {
"pii_tokens": {"<EMAIL_ADDRESS_1>": "john@example.com"},
}
},
):
collected.append(chunk)
assert collected == response_events
@pytest.mark.asyncio
async def test_output_parse_pii_streaming_responses_completed_event_unmasked(
mock_user_api_key,
):
"""
When output_parse_pii=True, a /v1/responses ``response.completed`` event
(a Pydantic ResponseCompletedEvent, as produced in production) must have its
output text unmasked in-place before being forwarded to the client.
"""
from litellm.types.llms.openai import (
ResponseCompletedEvent,
ResponsesAPIResponse,
ResponsesAPIStreamEvents,
)
from litellm.types.responses.main import GenericResponseOutputItem, OutputText
guardrail = _OPTIONAL_PresidioPIIMasking(
mock_testing=True,
output_parse_pii=True,
)
completed_event = ResponseCompletedEvent(
type=ResponsesAPIStreamEvents.RESPONSE_COMPLETED,
response=ResponsesAPIResponse(
id="resp_1",
created_at=1,
output=[
GenericResponseOutputItem(
type="message",
id="msg_1",
status="completed",
role="assistant",
content=[
OutputText(
type="output_text",
text="Reach me at <EMAIL_ADDRESS_1> today.",
annotations=[],
)
],
)
],
parallel_tool_calls=False,
tool_choice="auto",
tools=[],
),
)
async def mock_stream():
yield completed_event
collected = []
async for chunk in guardrail.async_post_call_streaming_iterator_hook(
user_api_key_dict=mock_user_api_key,
response=mock_stream(),
request_data={
"metadata": {
"pii_tokens": {"<EMAIL_ADDRESS_1>": "john@example.com"},
}
},
):
collected.append(chunk)
assert collected == [completed_event]
assert (
collected[0].response.output[0].content[0].text
== "Reach me at john@example.com today."
)
@pytest.mark.asyncio
async def test_output_parse_pii_streaming_mixed_chunks_flushes_buffered(
mock_user_api_key,
):
"""
Regression test: when output_parse_pii=True and a stream mixes buffered
ModelResponseStream chunks with a /v1/responses event, the buffered chat
chunks must still be forwarded (in order) instead of being dropped at the
saw_non_chat_chunk early return.
"""
guardrail = _OPTIONAL_PresidioPIIMasking(
mock_testing=True,
output_parse_pii=True,
)
class FakeResponsesEvent:
def __init__(self, event_type: str):
self.type = event_type
model_chunk = ModelResponseStream(
id="chatcmpl-mixed-unmask-1",
choices=[],
created=1,
model="gpt-4",
object="chat.completion.chunk",
system_fingerprint=None,
)
response_completed = FakeResponsesEvent("response.completed")
async def mock_stream():
yield model_chunk
yield response_completed
collected = []
async for chunk in guardrail.async_post_call_streaming_iterator_hook(
user_api_key_dict=mock_user_api_key,
response=mock_stream(),
request_data={
"metadata": {
"pii_tokens": {"<EMAIL_ADDRESS_1>": "john@example.com"},
}
},
):
collected.append(chunk)
assert collected == [model_chunk, response_completed]
@pytest.mark.asyncio
async def test_anonymize_text_uses_correct_positions_no_parse_pii():
"""