mirror of
https://github.com/langchain-ai/langchain.git
synced 2026-10-05 17:35:28 +03:00
feat(langchain): redact streamed PII in flight on PIIMiddleware (#37616)
`PIIMiddleware` previously scrubbed detected PII only at the state level via its `after_model` / `before_model` hooks. Consumers reading the live stream — `astream_events(version="v3")` or `run.messages` / `run.tool_calls` / `run.values` — saw the raw model text, the raw tool-call args, the raw tool outputs, and the raw state snapshots until the run finished and the canonical conversation history was written. This change registers a stream transformer ahead of `MessagesTransformer` that redacts every wire surface of an agent run. The transformer holds a sliding lookback buffer (default 128 characters) per `(run_id, content-block index)` so PII patterns that straddle delta boundaries are caught before the safe prefix is released downstream. Anything older than the lookback is run through the configured detector and emitted; the trailing tail stays buffered until a later delta extends it past the cap or the block finishes. `_finalize_block` always re-runs detection over the full block snapshot so the finalized content lands fully redacted even when the in-flight buffer never released a tail (short responses, or PII arriving in the final delta). The `block` strategy is now supported on the streaming path via a buffering mode that withholds every delta until the block resolves — clean blocks release the full text at finalize, PII-bearing blocks zero the wire and let `after_model` / `apply_to_tool_results` raise `PIIDetectionError` on the original state message. Activation is gated on `apply_to_output=True`, matching the existing post-hoc semantics. The middleware's transformer factory is cloned by `StreamMux._make_child` into every subgraph scope, so attaching `PIIMiddleware` at the outer agent also redacts streamed deltas from sub-agents invoked inside tools. ## Tool-call and tools-channel coverage The transformer covers every wire surface of an agent run, not just AI message text: - **Streamed AI text deltas** (`content-block-delta` of type `text-delta`) — lookback machinery, redacted in place. - **Streamed tool-call args** (`content-block-delta` with `tool_call_chunk` / `server_tool_call_chunk` fields) — each delta carries the full cumulative args string; detection runs on the field directly and redacts in place. Verified empirically against `_compat_bridge.py` and the consumer-side `_merge_block_delta_into_store` snapshot-replace semantics. - **Finalized tool-call blocks** (`content-block-finish` with `tool_call` / `server_tool_call` / `invalid_tool_call`) — `args` dict walked recursively and each string leaf redacted. - **Tool execution events on the `tools` channel** — `tool-started.input`, `tool-output-delta`, `tool-finished.output`, `tool-error.message` all run through detection. String deltas use the same lookback machinery as text-deltas keyed by `tool_call_id`; structured payloads walk recursively. - **State snapshots on the `values` channel** — message lists are walked and each message's `.content` is redacted on a fresh copy. Graph state itself stays intact for the state-level enforcer (`apply_to_tool_results` via `before_model`) to act on independently. - **Legacy `(BaseMessage, metadata)` payloads** on the `messages` channel (Python 3.10 path, where `langgraph`'s `ASYNCIO_ACCEPTS_CONTEXT = sys.version_info >= (3, 11)` falls back to a code path that doesn't propagate the streaming callback into the chat model) — `.content` and `AIMessage.tool_calls[*].args` are scrubbed. For `block`, the event's `data` tuple is replaced with an empty-content copy so the original message stays in state for `after_model` to raise on. ## Worth a careful look - `_PIIStreamTransformer._mutate_text_delta` — lookback partition. Anything older than `lookback` characters is released after redaction; the tail stays buffered. Bulletproof against whitespace-permissive detectors (notably `credit_card`, whose regex matches across spaces). - `_PIIStreamTransformer._mutate_tool_call_chunk_delta` — direct in-place redaction of the cumulative args string. No buffer; the wire shape is cumulative-snapshot, the consumer-side merge is replace-not-append. - `_PIIStreamTransformer._mutate_legacy_payload` — the dual path: mutate-in-place for non-`block` (idempotent with `after_model`), replace-with-empty-copy for `block` (keeps original in graph state for `after_model` to raise on). - `_PIIStreamTransformer._redact_value` — the recursive walker. `BaseMessage` branch returns a fresh `.content`-redacted copy via `model_copy(update=...)` — never mutates in place — so tool-output payloads that wrap a `ToolMessage` and message lists in state snapshots flow through cleanly. - The new `transformers` attribute on `PIIMiddleware`: this is what makes `create_agent` pick the factory up. Multiple `PIIMiddleware` instances each register one transformer; ordering is preserved within the `before_builtins` lane. ## Compatibility Bumps `langgraph` to `>=1.2.1` for the `before_builtins` opt-in on `StreamTransformer`.
This commit is contained in:
1 parent
06e65072af
commit
d08245f70d
4 files changed
+2178
-8
No files matched your search
@@ -2,9 +2,11 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING, Any, Literal
|
||||
from functools import partial
|
||||
from typing import TYPE_CHECKING, Any, ClassVar, Literal
|
||||
|
||||
from langchain_core.messages import AIMessage, AnyMessage, HumanMessage, ToolMessage
|
||||
from langchain_core.messages import AIMessage, AnyMessage, BaseMessage, HumanMessage, ToolMessage
|
||||
from langgraph.stream import StreamTransformer
|
||||
from typing_extensions import override
|
||||
|
||||
from langchain.agents.middleware._redaction import (
|
||||
@@ -31,6 +33,460 @@ if TYPE_CHECKING:
|
||||
from collections.abc import Callable
|
||||
|
||||
from langgraph.runtime import Runtime
|
||||
from langgraph.stream._types import ProtocolEvent
|
||||
|
||||
|
||||
_DEFAULT_STREAM_LOOKBACK = 128
|
||||
"""Default trailing-buffer size for cross-delta PII detection.
|
||||
|
||||
The transformer always holds the last `lookback` characters in a per-content
|
||||
block buffer so that PII patterns straddling delta boundaries are detected
|
||||
before any text is released downstream. 128 comfortably covers the built-in
|
||||
detectors (the credit-card regex tops out at 19 characters; URLs and emails
|
||||
are typically well under 100) while bounding first-token latency.
|
||||
"""
|
||||
|
||||
|
||||
class _PIIStreamTransformer(StreamTransformer):
|
||||
"""Mutates `content-block-delta` text on `messages` events in flight.
|
||||
|
||||
Runs before built-in stream transformers so the redacted text is what
|
||||
every downstream consumer sees — both the main protocol event log and
|
||||
the `run.messages` projection that `MessagesTransformer` snapshots into.
|
||||
|
||||
Holds a sliding buffer of the most recent text per (run_id, content
|
||||
block index) so PII patterns that straddle delta boundaries are caught.
|
||||
Anything older than `lookback` characters is redacted with the resolved
|
||||
rule's strategy and emitted as the new delta text; the trailing tail
|
||||
stays in the buffer until a later delta extends it past the cap or the
|
||||
block's finish event flushes the snapshot.
|
||||
"""
|
||||
|
||||
before_builtins: ClassVar[bool] = True
|
||||
required_stream_modes: ClassVar[tuple[str, ...]] = ("messages", "tools", "values")
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
scope: tuple[str, ...] = (),
|
||||
*,
|
||||
rule: ResolvedRedactionRule,
|
||||
lookback: int = _DEFAULT_STREAM_LOOKBACK,
|
||||
) -> None:
|
||||
super().__init__(scope)
|
||||
self._rule = rule
|
||||
self._lookback = lookback
|
||||
# Text/reasoning deltas keyed by `(run_id, content_block_index)`.
|
||||
self._buffers: dict[tuple[str, int], str] = {}
|
||||
# Tool-output-delta buffers keyed by `tool_call_id`. Held in a
|
||||
# separate dict so `_drop_run` on the messages channel can't
|
||||
# sweep active tool-output state.
|
||||
self._tool_buffers: dict[str, str] = {}
|
||||
|
||||
def init(self) -> dict[str, Any]:
|
||||
# No projection — this transformer mutates events in place rather
|
||||
# than building a derived view.
|
||||
return {}
|
||||
|
||||
def process(self, event: ProtocolEvent) -> bool:
|
||||
method = event["method"]
|
||||
if method == "messages":
|
||||
return self._process_messages_event(event)
|
||||
if method == "tools":
|
||||
return self._process_tools_event(event)
|
||||
if method == "values":
|
||||
return self._process_values_event(event)
|
||||
return True
|
||||
|
||||
def _process_values_event(self, event: ProtocolEvent) -> bool:
|
||||
"""Redact the state snapshot on the `values` channel.
|
||||
|
||||
State snapshots emitted between nodes carry the full state dict,
|
||||
which typically includes the messages list. Walking the snapshot
|
||||
with `_redact_value` returns a fresh structure where every
|
||||
message has a redacted copy of its content — the original
|
||||
objects in graph state remain intact for the state-level
|
||||
enforcer (`apply_to_tool_results` via `before_model`) to act on
|
||||
independently when the agent loops back.
|
||||
"""
|
||||
data = event["params"].get("data")
|
||||
if data is None:
|
||||
return True
|
||||
event["params"]["data"] = self._redact_value(data)
|
||||
return True
|
||||
|
||||
def _process_messages_event(self, event: ProtocolEvent) -> bool:
|
||||
params = event["params"]
|
||||
data = params.get("data")
|
||||
if not isinstance(data, tuple) or len(data) != 2: # noqa: PLR2004
|
||||
return True
|
||||
payload, metadata = data
|
||||
|
||||
# Legacy `(BaseMessage, metadata)` shape: the langgraph→langchain
|
||||
# integration emits this when a model only implements `_generate`
|
||||
# (or when its `_astream` falls back), producing a single event
|
||||
# carrying the full message rather than streamed content-block
|
||||
# deltas. Swap in a redacted copy so the consumer sees scrubbed
|
||||
# text on the wire while the original stays intact in graph state
|
||||
# for `after_model` to act on independently. Under `block`,
|
||||
# `_redact_base_message` raises `PIIDetectionError` via
|
||||
# `apply_strategy` before we get here.
|
||||
if isinstance(payload, BaseMessage):
|
||||
redacted = self._redact_base_message(payload)
|
||||
if redacted is not payload:
|
||||
params["data"] = (redacted, metadata)
|
||||
return True
|
||||
|
||||
if not isinstance(payload, dict):
|
||||
return True
|
||||
kind = payload.get("event")
|
||||
run_id = str(metadata.get("run_id") or "") if metadata else ""
|
||||
|
||||
if kind == "content-block-delta":
|
||||
self._mutate_delta(payload, run_id)
|
||||
elif kind == "content-block-finish":
|
||||
self._finalize_block(payload, run_id)
|
||||
elif kind in {"message-finish", "error"}:
|
||||
self._drop_run(run_id)
|
||||
return True
|
||||
|
||||
def _process_tools_event(self, event: ProtocolEvent) -> bool:
|
||||
data = event["params"].get("data")
|
||||
if not isinstance(data, dict):
|
||||
return True
|
||||
kind = data.get("event")
|
||||
tool_call_id = data.get("tool_call_id")
|
||||
|
||||
if kind == "tool-started":
|
||||
# Tool inputs may be a dict (multi-arg tools), a string
|
||||
# (single-arg tools — `BaseTool._parse_input` passes the
|
||||
# raw string through), or a list (array-input tools).
|
||||
# `_redact_value` handles all three uniformly.
|
||||
if "input" in data:
|
||||
data["input"] = self._redact_value(data["input"])
|
||||
elif kind == "tool-output-delta":
|
||||
# Use the tool_call_id as buffer key when present; fall back
|
||||
# to a None-keyed slot for the rare malformed/custom emitter
|
||||
# case (the buffer becomes shared but at least redaction runs).
|
||||
self._mutate_tool_output_delta(
|
||||
data, tool_call_id if isinstance(tool_call_id, str) else ""
|
||||
)
|
||||
elif kind == "tool-finished":
|
||||
if "output" in data:
|
||||
data["output"] = self._redact_value(data["output"])
|
||||
if isinstance(tool_call_id, str):
|
||||
self._tool_buffers.pop(tool_call_id, None)
|
||||
elif kind == "tool-error":
|
||||
msg = data.get("message")
|
||||
if isinstance(msg, str) and msg:
|
||||
matches = self._rule.detector(msg)
|
||||
if matches:
|
||||
data["message"] = apply_strategy(msg, matches, self._rule.strategy)
|
||||
if isinstance(tool_call_id, str):
|
||||
self._tool_buffers.pop(tool_call_id, None)
|
||||
|
||||
return True
|
||||
|
||||
def _mutate_tool_output_delta(self, data: dict[str, Any], tool_call_id: str) -> None:
|
||||
"""Redact a `tool-output-delta` payload.
|
||||
|
||||
String deltas go through the same lookback machinery as
|
||||
text-deltas, keyed by `tool_call_id` in the disjoint
|
||||
`_tool_buffers` dict so `_drop_run` on the messages channel
|
||||
can't sweep active tool-output state.
|
||||
|
||||
Structured deltas (dict/list) walk recursively without
|
||||
buffering — they don't have a position-stable shape across
|
||||
deltas to buffer against.
|
||||
"""
|
||||
delta = data.get("delta")
|
||||
if isinstance(delta, str):
|
||||
held = self._tool_buffers.get(tool_call_id, "")
|
||||
combined = held + delta
|
||||
|
||||
matches = self._rule.detector(combined)
|
||||
if matches:
|
||||
# `apply_strategy` raises `PIIDetectionError` under
|
||||
# `strategy="block"`, failing the run immediately —
|
||||
# cleaner than withholding deltas until `after_model`
|
||||
# raises later.
|
||||
combined = apply_strategy(combined, matches, self._rule.strategy)
|
||||
|
||||
emit_end = max(0, len(combined) - self._lookback)
|
||||
self._tool_buffers[tool_call_id] = combined[emit_end:]
|
||||
data["delta"] = combined[:emit_end]
|
||||
elif isinstance(delta, (dict, list)):
|
||||
data["delta"] = self._redact_value(delta)
|
||||
|
||||
def _redact_tool_call_list(self, calls: list[Any] | None) -> tuple[list[Any], bool]:
|
||||
"""Walk a list of tool-call (or invalid-tool-call) dicts.
|
||||
|
||||
Returns `(new_list, changed)`. Each element's `args` is run
|
||||
through `_redact_value` regardless of its type — `tool_call.args`
|
||||
is a dict, `invalid_tool_call.args` is a raw JSON string, and
|
||||
`_redact_value` handles both shapes uniformly. If nothing
|
||||
changed, returns the input list and `changed=False`.
|
||||
"""
|
||||
if not calls:
|
||||
return calls or [], False
|
||||
new_calls: list[Any] = []
|
||||
changed = False
|
||||
for tc in calls:
|
||||
if isinstance(tc, dict) and "args" in tc and tc["args"] is not None:
|
||||
redacted = self._redact_value(tc["args"])
|
||||
if redacted != tc["args"]:
|
||||
new_tc = dict(tc)
|
||||
new_tc["args"] = redacted
|
||||
new_calls.append(new_tc)
|
||||
changed = True
|
||||
continue
|
||||
new_calls.append(tc)
|
||||
return new_calls, changed
|
||||
|
||||
def _redact_value(self, value: Any) -> Any:
|
||||
"""Recursively redact PII in string leaves of a nested structure.
|
||||
|
||||
Returns a new value where every `str` leaf that contains PII has
|
||||
been replaced (or emptied under `block`). Non-string leaves and
|
||||
the structure itself are preserved.
|
||||
|
||||
`BaseMessage` payloads (typically `ToolMessage` from
|
||||
`tool-finished.output`, or any message reached via the `values`
|
||||
channel) return a fresh copy with `.content` redacted plus
|
||||
`AIMessage.tool_calls[*].args` / `invalid_tool_calls[*].args`
|
||||
walked. The original object stays intact for state-level
|
||||
enforcers (`after_model`, `before_model` with
|
||||
`apply_to_tool_results`) to act on independently.
|
||||
|
||||
Scope mirrors the pre-streaming state-level surfaces:
|
||||
`.content` (string or list-of-content-blocks) and `tool_calls`
|
||||
args. Other message attributes (`additional_kwargs`,
|
||||
`response_metadata`, `ToolMessage.artifact`) are intentionally
|
||||
not walked here — they aren't scrubbed in graph state by the
|
||||
existing hooks, so scrubbing them on the wire would create
|
||||
a wire/state divergence.
|
||||
"""
|
||||
if isinstance(value, str):
|
||||
if not value:
|
||||
return value
|
||||
matches = self._rule.detector(value)
|
||||
if not matches:
|
||||
return value
|
||||
# `apply_strategy` raises `PIIDetectionError` under `block`
|
||||
# — the run fails immediately rather than buffering until a
|
||||
# state-level hook can raise.
|
||||
return apply_strategy(value, matches, self._rule.strategy)
|
||||
if isinstance(value, BaseMessage):
|
||||
return self._redact_base_message(value)
|
||||
if isinstance(value, dict):
|
||||
return {k: self._redact_value(v) for k, v in value.items()}
|
||||
if isinstance(value, list):
|
||||
return [self._redact_value(v) for v in value]
|
||||
if isinstance(value, tuple):
|
||||
return tuple(self._redact_value(v) for v in value)
|
||||
return value
|
||||
|
||||
def _redact_base_message(self, value: BaseMessage) -> BaseMessage:
|
||||
"""Return a fresh copy of `value` with PII-carrying surfaces redacted."""
|
||||
update: dict[str, Any] = {}
|
||||
|
||||
content = value.content
|
||||
if isinstance(content, str) and content:
|
||||
matches = self._rule.detector(content)
|
||||
if matches:
|
||||
update["content"] = apply_strategy(content, matches, self._rule.strategy)
|
||||
elif isinstance(content, list) and content:
|
||||
# Structured content-blocks shape:
|
||||
# `[{"type": "text", "text": "..."}, {"type": "tool_call", ...}, ...]`.
|
||||
redacted_content = self._redact_value(content)
|
||||
if redacted_content != content:
|
||||
update["content"] = redacted_content
|
||||
|
||||
# `AIMessage.tool_calls` and `.invalid_tool_calls` carry PII in
|
||||
# `args` independently of `.content`. `tool_call.args` is a
|
||||
# dict; `invalid_tool_call.args` is a raw JSON string —
|
||||
# `_redact_value` handles both shapes via the recursion.
|
||||
if isinstance(value, AIMessage):
|
||||
new_tc_list, tc_changed = self._redact_tool_call_list(value.tool_calls)
|
||||
if tc_changed:
|
||||
update["tool_calls"] = new_tc_list
|
||||
new_inv_list, inv_changed = self._redact_tool_call_list(value.invalid_tool_calls)
|
||||
if inv_changed:
|
||||
update["invalid_tool_calls"] = new_inv_list
|
||||
|
||||
if not update:
|
||||
return value
|
||||
return value.model_copy(update=update)
|
||||
|
||||
def _mutate_delta(self, payload: dict[str, Any], run_id: str) -> None:
|
||||
delta = payload.get("delta")
|
||||
if not isinstance(delta, dict):
|
||||
return
|
||||
delta_type = delta.get("type")
|
||||
if delta_type == "text-delta":
|
||||
self._mutate_string_field_delta(delta, payload, run_id, "text")
|
||||
return
|
||||
if delta_type == "reasoning-delta":
|
||||
# Reasoning content (chain-of-thought from extended-thinking
|
||||
# models) is a real PII surface — models echo back
|
||||
# user-supplied data or synthesize it from context. Run the
|
||||
# same lookback machinery as text-delta against the
|
||||
# `reasoning` field. Block indices are unique within a
|
||||
# message regardless of block type, so the buffer key
|
||||
# `(run_id, index)` naturally disjoint from text-delta keys.
|
||||
self._mutate_string_field_delta(delta, payload, run_id, "reasoning")
|
||||
return
|
||||
if delta_type == "block-delta":
|
||||
fields = delta.get("fields")
|
||||
if isinstance(fields, dict) and fields.get("type") in {
|
||||
"tool_call_chunk",
|
||||
"server_tool_call_chunk",
|
||||
}:
|
||||
self._mutate_tool_call_chunk_delta(fields)
|
||||
# Other delta types (`data-delta`, vendor block types) pass
|
||||
# through. The pre-streaming middleware scrubbed `.content` text
|
||||
# on state messages only; binary payloads and provider-specific
|
||||
# block shapes are out of scope for parity with that surface.
|
||||
|
||||
def _mutate_string_field_delta(
|
||||
self,
|
||||
delta: dict[str, Any],
|
||||
payload: dict[str, Any],
|
||||
run_id: str,
|
||||
field: str,
|
||||
) -> None:
|
||||
"""Apply the lookback-buffer redaction to a string field on a delta.
|
||||
|
||||
Shared by `text-delta` (`field="text"`) and `reasoning-delta`
|
||||
(`field="reasoning"`). Buffer is keyed by `(run_id, block_index)`;
|
||||
block indices are unique within a message so different block
|
||||
types share the same key space without collision.
|
||||
"""
|
||||
text = delta.get(field)
|
||||
if not isinstance(text, str) or not text:
|
||||
return
|
||||
index = payload.get("index")
|
||||
if not isinstance(index, int):
|
||||
return
|
||||
|
||||
key = (run_id, index)
|
||||
held = self._buffers.get(key, "")
|
||||
combined = held + text
|
||||
|
||||
# Run detection on the full accumulated buffer before splitting.
|
||||
# Detecting only on the about-to-emit prefix would miss matches
|
||||
# that straddle the lookback boundary — the detector's regex
|
||||
# needs a complete, boundary-anchored hit, so a truncated prefix
|
||||
# would fail to match and the partial PII would leak on the
|
||||
# wire. Under `strategy="block"`, `apply_strategy` raises
|
||||
# `PIIDetectionError` here, failing the run as soon as PII
|
||||
# arrives rather than buffering until `after_model`.
|
||||
matches = self._rule.detector(combined)
|
||||
if matches:
|
||||
combined = apply_strategy(combined, matches, self._rule.strategy)
|
||||
|
||||
emit_end = max(0, len(combined) - self._lookback)
|
||||
self._buffers[key] = combined[emit_end:]
|
||||
delta[field] = combined[:emit_end]
|
||||
|
||||
def _mutate_tool_call_chunk_delta(self, fields: dict[str, Any]) -> None:
|
||||
"""Redact cumulative tool-call args with lookback withholding.
|
||||
|
||||
Each `tool_call_chunk` `block-delta` event carries the full
|
||||
accumulated args string (verified against `_compat_bridge.py`
|
||||
— `delta_source = current` for these block types — and against
|
||||
the consumer-side `_merge_block_delta_into_store`, which
|
||||
replaces wholesale rather than appends).
|
||||
|
||||
Detection runs on the full cumulative args so any complete PII
|
||||
anywhere in the string is redacted before emission. Lookback
|
||||
withholding then trims the trailing the lookback window characters
|
||||
from what reaches the consumer — those characters might be the
|
||||
start of a partial PII match that completes in a future
|
||||
cumulative delta. The trimmed tail surfaces at `content-block-
|
||||
finish` where `_finalize_block` redacts the parsed args dict.
|
||||
|
||||
For args that fit within the lookback window (the typical case),
|
||||
this withholds the entire args string during streaming — the
|
||||
redacted args dict appears only at finalize. For args that
|
||||
exceed the lookback window, the safe prefix streams incrementally
|
||||
as the cumulative state grows. PII that appears more than
|
||||
the lookback window characters from the cumulative tail in a
|
||||
delta where it hasn't yet completed can still surface in the
|
||||
emit prefix — same residual exposure as PII longer than
|
||||
the lookback window on the text path. The `content-block-finish`
|
||||
snapshot redaction is the backstop.
|
||||
"""
|
||||
args = fields.get("args")
|
||||
if not isinstance(args, str) or not args:
|
||||
return
|
||||
|
||||
matches = self._rule.detector(args)
|
||||
if matches:
|
||||
# `apply_strategy` raises `PIIDetectionError` under
|
||||
# `strategy="block"` — the run fails the moment a complete
|
||||
# PII pattern surfaces in the cumulative args string.
|
||||
args = apply_strategy(args, matches, self._rule.strategy)
|
||||
|
||||
emit_end = max(0, len(args) - self._lookback)
|
||||
fields["args"] = args[:emit_end]
|
||||
|
||||
def _finalize_block(self, payload: dict[str, Any], run_id: str) -> None:
|
||||
index = payload.get("index")
|
||||
if not isinstance(index, int):
|
||||
return
|
||||
key = (run_id, index)
|
||||
# The finalized block carries the model's original concatenation
|
||||
# of deltas, not what we emitted on the wire. Re-run detection over
|
||||
# its full text so the snapshot matches the redacted stream.
|
||||
content = payload.get("content")
|
||||
if isinstance(content, dict):
|
||||
ctype = content.get("type")
|
||||
if ctype == "text":
|
||||
self._finalize_string_field(content, "text")
|
||||
elif ctype == "reasoning":
|
||||
self._finalize_string_field(content, "reasoning")
|
||||
elif (
|
||||
ctype in {"tool_call", "server_tool_call", "invalid_tool_call"}
|
||||
and "args" in content
|
||||
and content["args"] is not None
|
||||
):
|
||||
# `tool_call` / `server_tool_call` args are dicts;
|
||||
# `invalid_tool_call.args` is the raw unparsed JSON
|
||||
# string. `_redact_value` handles both shapes.
|
||||
content["args"] = self._redact_value(content["args"])
|
||||
self._buffers.pop(key, None)
|
||||
|
||||
def _finalize_string_field(self, content: dict[str, Any], field: str) -> None:
|
||||
"""Re-redact a string content-block field on `content-block-finish`.
|
||||
|
||||
Used for `text` and `reasoning` content blocks. Under
|
||||
`strategy="block"` `apply_strategy` raises `PIIDetectionError`,
|
||||
failing the run immediately.
|
||||
"""
|
||||
text = content.get(field)
|
||||
if not isinstance(text, str) or not text:
|
||||
return
|
||||
matches = self._rule.detector(text)
|
||||
if not matches:
|
||||
return
|
||||
content[field] = apply_strategy(text, matches, self._rule.strategy)
|
||||
|
||||
def _drop_run(self, run_id: str) -> None:
|
||||
# Release any buffered tails for this run_id — content-block-finish
|
||||
# should have already done so for normal completion, but message-finish
|
||||
# / error paths need an explicit sweep so abandoned blocks don't
|
||||
# accumulate in long-lived processes.
|
||||
stale = [key for key in self._buffers if key[0] == run_id]
|
||||
for key in stale:
|
||||
del self._buffers[key]
|
||||
|
||||
def finalize(self) -> None:
|
||||
self._buffers.clear()
|
||||
self._tool_buffers.clear()
|
||||
|
||||
def fail(self, err: BaseException) -> None: # noqa: ARG002
|
||||
self._buffers.clear()
|
||||
self._tool_buffers.clear()
|
||||
|
||||
|
||||
class PIIMiddleware(AgentMiddleware[AgentState[ResponseT], ContextT, ResponseT]):
|
||||
@@ -133,6 +589,32 @@ class PIIMiddleware(AgentMiddleware[AgentState[ResponseT], ContextT, ResponseT])
|
||||
* If `None`: Uses built-in detector for the `pii_type`
|
||||
apply_to_input: Whether to check user messages before model call.
|
||||
apply_to_output: Whether to check AI messages after model call.
|
||||
|
||||
When `True`, a stream transformer is also installed so
|
||||
that every wire surface of an agent run is redacted in
|
||||
flight:
|
||||
|
||||
* Streamed AI text deltas (`content-block-delta` of type
|
||||
`text-delta`)
|
||||
* Streamed tool-call arguments (`content-block-delta`
|
||||
with `tool_call_chunk` / `server_tool_call_chunk`
|
||||
fields, plus the finalized `tool_call` content block
|
||||
on `content-block-finish`)
|
||||
* Tool execution events on the `tools` channel
|
||||
(`tool-started.input`, `tool-output-delta`,
|
||||
`tool-finished.output`, `tool-error.message`)
|
||||
* State snapshots on the `values` channel — message
|
||||
lists are walked and each message's `.content` is
|
||||
redacted on a fresh copy (state itself stays intact
|
||||
for `before_model` / `after_model` to act on
|
||||
independently)
|
||||
|
||||
State-level redaction via `after_model` (and
|
||||
`before_model` with `apply_to_tool_results`) remains the
|
||||
canonical enforcer; the streaming transformer ensures
|
||||
consumers reading `astream_events(version="v3")` or
|
||||
`run.messages` / `run.tool_calls` / `run.values` never
|
||||
see PII on the wire.
|
||||
apply_to_tool_results: Whether to check tool result messages after tool execution.
|
||||
|
||||
Raises:
|
||||
@@ -153,6 +635,26 @@ class PIIMiddleware(AgentMiddleware[AgentState[ResponseT], ContextT, ResponseT])
|
||||
self.strategy = self._resolved_rule.strategy
|
||||
self.detector = self._resolved_rule.detector
|
||||
|
||||
# Stream transformer scrubs the streamed surface of the same
|
||||
# messages that the state-level hooks scrub in graph state.
|
||||
# Installed whenever any output-side scrubbing is enabled —
|
||||
# `apply_to_output` covers AI messages (text, tool-call args,
|
||||
# reasoning), `apply_to_tool_results` covers tool execution
|
||||
# (the `tools` channel + ToolMessage content on `values` and
|
||||
# `messages`). For `block` the transformer raises
|
||||
# `PIIDetectionError` directly from its event handler the
|
||||
# moment a complete PII pattern is detected, failing the run
|
||||
# via langgraph's `StreamMux.afail` path. The state-level
|
||||
# `after_model` / `before_model` hooks remain a backstop for
|
||||
# non-streaming consumers.
|
||||
if self.apply_to_output or self.apply_to_tool_results:
|
||||
self.transformers = (
|
||||
partial(
|
||||
_PIIStreamTransformer,
|
||||
rule=self._resolved_rule,
|
||||
),
|
||||
)
|
||||
|
||||
@property
|
||||
def name(self) -> str:
|
||||
"""Name of the middleware."""
|
||||
|
||||
@@ -25,7 +25,7 @@ version = "1.3.1"
|
||||
requires-python = ">=3.10.0,<4.0.0"
|
||||
dependencies = [
|
||||
"langchain-core>=1.4.0,<2.0.0",
|
||||
"langgraph>=1.2.0,<1.3.0",
|
||||
"langgraph>=1.2.1,<1.3.0",
|
||||
"pydantic>=2.7.4,<3.0.0",
|
||||
]
|
||||
|
||||
|
||||
File diff suppressed because it is too large.
Load diff
Generated
+4
-4
@@ -2025,7 +2025,7 @@ requires-dist = [
|
||||
{ name = "langchain-perplexity", marker = "extra == 'perplexity'" },
|
||||
{ name = "langchain-together", marker = "extra == 'together'" },
|
||||
{ name = "langchain-xai", marker = "extra == 'xai'" },
|
||||
{ name = "langgraph", specifier = ">=1.2.0,<1.3.0" },
|
||||
{ name = "langgraph", specifier = ">=1.2.1,<1.3.0" },
|
||||
{ name = "pydantic", specifier = ">=2.7.4,<3.0.0" },
|
||||
]
|
||||
provides-extras = ["community", "anthropic", "openai", "azure-ai", "google-vertexai", "google-genai", "fireworks", "ollama", "together", "mistralai", "huggingface", "groq", "aws", "baseten", "deepseek", "xai", "perplexity"]
|
||||
@@ -2603,7 +2603,7 @@ wheels = [
|
||||
|
||||
[[package]]
|
||||
name = "langgraph"
|
||||
version = "1.2.0"
|
||||
version = "1.2.1"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
dependencies = [
|
||||
{ name = "langchain-core" },
|
||||
@@ -2613,9 +2613,9 @@ dependencies = [
|
||||
{ name = "pydantic" },
|
||||
{ name = "xxhash" },
|
||||
]
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/58/61/d5d25e783035aa307d289b37e082258a6061c0fb4caa4a284f3bf1e87169/langgraph-1.2.0.tar.gz", hash = "sha256:4a9baaf62afc5d5f63144a50095140a34b9aa9b7cea695d25326d564775348e7", size = 690248, upload-time = "2026-05-12T03:46:39.164Z" }
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/1c/8e/34a57e338a319e3b32c1bd183c2a9a04f7f35d683d3f3d8f597f6eacbc4e/langgraph-1.2.1.tar.gz", hash = "sha256:28314f844678d9d307cbd63e7b48b0145bf17177d84b40ee2921061e07b6f966", size = 693750, upload-time = "2026-05-21T18:33:07.478Z" }
|
||||
wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/f6/e8/e3304ac0015c2bdb04ad9785e4ed65c788855ce7857ce6104dd2f5d322db/langgraph-1.2.0-py3-none-any.whl", hash = "sha256:03fd5895a8d4b70db1ff63ebc3bacead29dd20cd794a8b1a483e7ec9018f7a65", size = 234262, upload-time = "2026-05-12T03:46:37.971Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/73/8c/313912e26866893bd15be9b4ea3442dc86f69270b0ad01a4961d1eba7118/langgraph-1.2.1-py3-none-any.whl", hash = "sha256:5cc4020de8f1e2a048d773f6e9128646a2af8c68a8067ab9cab177a2fcc8d221", size = 235317, upload-time = "2026-05-21T18:33:05.687Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
||||
Reference in new issue
Block a user