"""
Extract the text that guardrails scan from a chat completion request or response, and.
write guardrails' changes back.
- Input guardrails scan the **latest user message**: what the user just said. They do not
scan the LLMClient's own system prompt, which would make e.g. a prompt injection guardrail
trigger on the operator's instructions, nor the earlier turns of the conversation, which
were scanned when they were new.
- Output guardrails scan the reply's message content, and refusal.
A message's content may be a string, or a list of parts, of which the ``text`` parts are
scanned. Each segment has a path, e.g. ``messages[3].content`` or
``messages[3].content[0].text``, so that redactions can be written back to the right place.
"""
import copy
import re
from typing import Any
from smarter.apps.provider.services.text_completion.contracts import (
GuardrailStage,
TextSegment,
)
USER_ROLE = "user"
def content_segments(content: Any, path: str) -> list[TextSegment]:
"""Return the segments of a message's content: a string, or a list of parts."""
if isinstance(content, str):
return [TextSegment(path=path, text=content)] if content else []
segments = []
if isinstance(content, list):
for i, part in enumerate(content):
if isinstance(part, dict) and part.get("type") == "text" and isinstance(part.get("text"), str):
if part["text"]:
segments.append(TextSegment(path=f"{path}[{i}].text", text=part["text"]))
return segments
[docs]
def latest_user_message_index(messages: list[Any]) -> int | None:
"""Return the index of the latest message whose role is user, or ``None``."""
for i in range(len(messages) - 1, -1, -1):
message = messages[i]
if isinstance(message, dict) and message.get("role") == USER_ROLE:
return i
return None
def extract_pre_segments(payload: dict[str, Any]) -> list[TextSegment]:
"""Return the segments of the latest user message of a chat completion request."""
messages = payload.get("messages") or []
index = latest_user_message_index(messages)
if index is None:
return []
return content_segments(messages[index].get("content"), f"messages[{index}].content")
def extract_post_segments(payload: dict[str, Any]) -> list[TextSegment]:
"""Return the segments of each choice's message content, and refusal, of a chat completion response."""
segments: list[TextSegment] = []
for i, choice in enumerate(payload.get("choices") or []):
message = (choice or {}).get("message") or {}
segments.extend(content_segments(message.get("content"), f"choices[{i}].message.content"))
refusal = message.get("refusal")
if isinstance(refusal, str) and refusal:
segments.append(TextSegment(path=f"choices[{i}].message.refusal", text=refusal))
return segments
[docs]
def write_segment(payload: dict[str, Any], path: str, new_text: str) -> dict[str, Any]:
"""
Return a deep copy of the payload, with the text at ``path`` replaced.
If ``path`` does not resolve, the copy is returned unchanged.
"""
updated = copy.deepcopy(payload)
node, key = resolve_path(updated, path)
if node is not None and key is not None:
node[key] = new_text
return updated
[docs]
def read_segment(payload: dict[str, Any], path: str) -> str | None:
"""Return the text at ``path``, or ``None`` if it does not resolve to a string."""
node, key = resolve_path(payload, path)
if node is None:
return None
try:
value = node[key]
except (KeyError, IndexError, TypeError):
return None
return value if isinstance(value, str) else None
[docs]
def resolve_path(payload: Any, path: str) -> tuple[Any, Any]:
"""
Walk a ``messages[0].content``-style path to its final container.
:returns: ``(container, key)``, such that ``container[key]`` is the field that ``path``
names, or ``(None, None)`` if the path does not resolve.
"""
tokens = tokenize(path)
if not tokens:
return None, None
node: Any = payload
for token in tokens[:-1]:
try:
node = node[token]
except (KeyError, IndexError, TypeError):
return None, None
return node, tokens[-1]
def tokenize(path: str) -> list[Any]:
"""Split a dotted, bracketed path into its keys and indices, e.g. ``a[0].b`` into ``["a", 0, "b"]``."""
tokens: list[Any] = []
for part in path.split("."):
name, _, rest = part.partition("[")
if name:
tokens.append(name)
for index in re.findall(r"(\d+)\]", "[" + rest if rest else ""):
tokens.append(int(index))
return tokens
__all__ = ["extract_segments", "read_segment", "resolve_path", "write_segment", "latest_user_message_index"]