From c37d47f89fc957d38b9525bb20fb49adfe4456ff Mon Sep 17 00:00:00 2001 From: Will Killian Date: Tue, 7 Jul 2026 10:09:48 -0400 Subject: [PATCH 1/3] feat: expose event sanitizers in Python Signed-off-by: Will Killian --- crates/python/src/py_api/mod.rs | 122 +++++++++++ crates/python/src/py_callable.rs | 69 +++++- crates/python/src/py_plugin.rs | 111 +++++++++- .../python/tests/coverage/coverage_tests.rs | 12 ++ .../coverage/py_callable_coverage_tests.rs | 54 +++++ .../coverage/py_plugin_coverage_tests.rs | 29 ++- python/nemo_relay/__init__.py | 14 +- python/nemo_relay/__init__.pyi | 10 +- python/nemo_relay/_native.pyi | 34 ++- python/nemo_relay/guardrails.py | 56 +++++ python/nemo_relay/pii_redaction.py | 2 + python/nemo_relay/pii_redaction.pyi | 1 + python/nemo_relay/plugin.py | 15 ++ python/nemo_relay/plugin.pyi | 8 + python/nemo_relay/scope_local.py | 60 ++++++ .../plugin/src/nemo_relay_plugin/__init__.py | 6 + python/plugin/src/nemo_relay_plugin/_api.py | 65 +++++- .../plugin/test_public_api_docstrings.py | 1 + python/tests/plugin/test_worker_sdk.py | 56 +++++ python/tests/test_event_sanitizers.py | 198 ++++++++++++++++++ python/tests/test_pii_redaction_plugin.py | 4 + 21 files changed, 907 insertions(+), 20 deletions(-) create mode 100644 python/tests/test_event_sanitizers.py diff --git a/crates/python/src/py_api/mod.rs b/crates/python/src/py_api/mod.rs index 2603e92f3..d87370638 100644 --- a/crates/python/src/py_api/mod.rs +++ b/crates/python/src/py_api/mod.rs @@ -875,6 +875,44 @@ fn llm_stream_call_execute<'py>( // Guardrail registrations (macro-generated) // --------------------------------------------------------------------------- +macro_rules! py_event_guardrail_api { + ($register_name:ident, $deregister_name:ident, $core_register:path, $core_deregister:path) => { + #[pyfunction] + fn $register_name(name: &str, priority: i32, guardrail: Py) -> PyResult<()> { + $core_register( + name, + priority, + py_callable::wrap_py_event_sanitize_fn(guardrail), + ) + .map_err(to_py_err) + } + + #[pyfunction] + fn $deregister_name(name: &str) -> PyResult { + $core_deregister(name).map_err(to_py_err) + } + }; +} + +py_event_guardrail_api!( + register_mark_sanitize_guardrail, + deregister_mark_sanitize_guardrail, + core_registry_api::register_mark_sanitize_guardrail, + core_registry_api::deregister_mark_sanitize_guardrail +); +py_event_guardrail_api!( + register_scope_sanitize_start_guardrail, + deregister_scope_sanitize_start_guardrail, + core_registry_api::register_scope_sanitize_start_guardrail, + core_registry_api::deregister_scope_sanitize_start_guardrail +); +py_event_guardrail_api!( + register_scope_sanitize_end_guardrail, + deregister_scope_sanitize_end_guardrail, + core_registry_api::register_scope_sanitize_end_guardrail, + core_registry_api::deregister_scope_sanitize_end_guardrail +); + /// Macro that generates a register/deregister pair for tool guardrails /// whose callback signature is `(tool_name: str, json: Any) -> Any`. macro_rules! py_guardrail_tool_api { @@ -1272,6 +1310,52 @@ fn parse_uuid(scope_uuid: &str) -> PyResult { .map_err(|e| PyErr::new::(format!("invalid UUID: {e}"))) } +macro_rules! py_scope_event_guardrail_api { + ($register_name:ident, $deregister_name:ident, $core_register:path, $core_deregister:path) => { + #[pyfunction] + fn $register_name( + scope_uuid: &str, + name: &str, + priority: i32, + guardrail: Py, + ) -> PyResult<()> { + let uuid = parse_uuid(scope_uuid)?; + $core_register( + &uuid, + name, + priority, + py_callable::wrap_py_event_sanitize_fn(guardrail), + ) + .map_err(to_py_err) + } + + #[pyfunction] + fn $deregister_name(scope_uuid: &str, name: &str) -> PyResult { + let uuid = parse_uuid(scope_uuid)?; + $core_deregister(&uuid, name).map_err(to_py_err) + } + }; +} + +py_scope_event_guardrail_api!( + scope_register_mark_sanitize_guardrail, + scope_deregister_mark_sanitize_guardrail, + core_registry_api::scope_register_mark_sanitize_guardrail, + core_registry_api::scope_deregister_mark_sanitize_guardrail +); +py_scope_event_guardrail_api!( + scope_register_scope_sanitize_start_guardrail, + scope_deregister_scope_sanitize_start_guardrail, + core_registry_api::scope_register_scope_sanitize_start_guardrail, + core_registry_api::scope_deregister_scope_sanitize_start_guardrail +); +py_scope_event_guardrail_api!( + scope_register_scope_sanitize_end_guardrail, + scope_deregister_scope_sanitize_end_guardrail, + core_registry_api::scope_register_scope_sanitize_end_guardrail, + core_registry_api::scope_deregister_scope_sanitize_end_guardrail +); + /// Macro that generates a scope-local register/deregister pair for guardrails /// whose callback signature is `(tool_name: str, json: Any) -> Any`. macro_rules! py_scope_local_guardrail_tool_api { @@ -1651,6 +1735,23 @@ pub fn register(m: &Bound<'_, PyModule>) -> PyResult<()> { m )?)?; + // Mark and scope event guardrails + m.add_function(wrap_pyfunction!(register_mark_sanitize_guardrail, m)?)?; + m.add_function(wrap_pyfunction!(deregister_mark_sanitize_guardrail, m)?)?; + m.add_function(wrap_pyfunction!( + register_scope_sanitize_start_guardrail, + m + )?)?; + m.add_function(wrap_pyfunction!( + deregister_scope_sanitize_start_guardrail, + m + )?)?; + m.add_function(wrap_pyfunction!(register_scope_sanitize_end_guardrail, m)?)?; + m.add_function(wrap_pyfunction!( + deregister_scope_sanitize_end_guardrail, + m + )?)?; + // Tool intercepts m.add_function(wrap_pyfunction!(register_tool_request_intercept, m)?)?; m.add_function(wrap_pyfunction!(deregister_tool_request_intercept, m)?)?; @@ -1727,6 +1828,27 @@ pub fn register(m: &Bound<'_, PyModule>) -> PyResult<()> { scope_deregister_tool_conditional_execution_guardrail, m )?)?; + m.add_function(wrap_pyfunction!(scope_register_mark_sanitize_guardrail, m)?)?; + m.add_function(wrap_pyfunction!( + scope_deregister_mark_sanitize_guardrail, + m + )?)?; + m.add_function(wrap_pyfunction!( + scope_register_scope_sanitize_start_guardrail, + m + )?)?; + m.add_function(wrap_pyfunction!( + scope_deregister_scope_sanitize_start_guardrail, + m + )?)?; + m.add_function(wrap_pyfunction!( + scope_register_scope_sanitize_end_guardrail, + m + )?)?; + m.add_function(wrap_pyfunction!( + scope_deregister_scope_sanitize_end_guardrail, + m + )?)?; // Scope-local tool intercepts m.add_function(wrap_pyfunction!(scope_register_tool_request_intercept, m)?)?; diff --git a/crates/python/src/py_callable.rs b/crates/python/src/py_callable.rs index 8d59fe72a..25b1707a2 100644 --- a/crates/python/src/py_callable.rs +++ b/crates/python/src/py_callable.rs @@ -25,16 +25,16 @@ use std::pin::Pin; use std::sync::Arc; use nemo_relay::api::runtime::{ - EventSubscriberFn, LlmConditionalFn, LlmExecutionNextFn, LlmRequestInterceptFn, - LlmSanitizeRequestFn, LlmSanitizeResponseFn, LlmStreamExecutionNextFn, ToolConditionalFn, - ToolExecutionNextFn, ToolInterceptFn, ToolSanitizeFn, + EventSanitizeFn, EventSubscriberFn, LlmConditionalFn, LlmExecutionNextFn, + LlmRequestInterceptFn, LlmSanitizeRequestFn, LlmSanitizeResponseFn, LlmStreamExecutionNextFn, + ToolConditionalFn, ToolExecutionNextFn, ToolInterceptFn, ToolSanitizeFn, }; use nemo_relay::error::{FlowError, Result as FlowResult}; use pyo3::prelude::*; use serde_json::Value as Json; use tokio_stream::Stream; -use nemo_relay::api::event::Event; +use nemo_relay::api::event::{Event, EventSanitizeFields}; use nemo_relay::api::llm::{LlmRequest, LlmRequestInterceptOutcome}; use nemo_relay::api::tool::ToolExecutionInterceptOutcome; use nemo_relay::codec::request::AnnotatedLlmRequest as AnnotatedLLMRequest; @@ -854,6 +854,67 @@ pub fn wrap_py_event_subscriber(py_fn: Py) -> EventSubscriberFn { }) } +/// Wrap a Python callable ``(Event, EventSanitizeFields) -> EventSanitizeFields``. +pub fn wrap_py_event_sanitize_fn(py_fn: Py) -> EventSanitizeFn { + Arc::new(move |event: &Event, fields: EventSanitizeFields| { + Python::attach(|py| { + let py_event = match event { + Event::Scope(inner) => Py::new( + py, + crate::py_types::PyScopeEvent { + inner: inner.clone(), + }, + ) + .map(|value| value.into_any()), + Event::Mark(inner) => Py::new( + py, + crate::py_types::PyMarkEvent { + inner: inner.clone(), + }, + ) + .map(|value| value.into_any()), + }; + let py_event = match py_event { + Ok(value) => value, + Err(error) => { + eprintln!("nemo_relay: failed to convert event sanitizer context: {error}"); + return fields.clone(); + } + }; + let fields_json = match serde_json::to_value(&fields) { + Ok(value) => value, + Err(error) => { + eprintln!("nemo_relay: failed to serialize event sanitizer fields: {error}"); + return fields.clone(); + } + }; + let py_fields = match json_to_py(py, &fields_json) { + Ok(value) => value, + Err(error) => { + eprintln!("nemo_relay: failed to convert event sanitizer fields: {error}"); + return fields.clone(); + } + }; + let result = match py_fn.call1(py, (py_event, py_fields)) { + Ok(value) => value, + Err(error) => { + eprintln!("nemo_relay: Python event sanitizer callable failed: {error}"); + return fields.clone(); + } + }; + py_to_json(result.bind(py)) + .ok() + .and_then(|value| serde_json::from_value(value).ok()) + .unwrap_or_else(|| { + eprintln!( + "nemo_relay: event sanitizer must return data, category_profile, and metadata fields" + ); + fields.clone() + }) + }) + }) +} + // --------------------------------------------------------------------------- // LLM Codec wrapper // --------------------------------------------------------------------------- diff --git a/crates/python/src/py_plugin.rs b/crates/python/src/py_plugin.rs index 09209b03a..c3e98664e 100644 --- a/crates/python/src/py_plugin.rs +++ b/crates/python/src/py_plugin.rs @@ -16,16 +16,21 @@ use nemo_relay::api::registry::{ deregister_llm_conditional_execution_guardrail, deregister_llm_execution_intercept, deregister_llm_request_intercept, deregister_llm_sanitize_request_guardrail, deregister_llm_sanitize_response_guardrail, deregister_llm_stream_execution_intercept, - deregister_tool_conditional_execution_guardrail, deregister_tool_execution_intercept, - deregister_tool_request_intercept, deregister_tool_sanitize_request_guardrail, - deregister_tool_sanitize_response_guardrail, register_llm_conditional_execution_guardrail, - register_llm_execution_intercept, register_llm_request_intercept, - register_llm_sanitize_request_guardrail, register_llm_sanitize_response_guardrail, - register_llm_stream_execution_intercept, register_tool_conditional_execution_guardrail, + deregister_mark_sanitize_guardrail, deregister_scope_sanitize_end_guardrail, + deregister_scope_sanitize_start_guardrail, deregister_tool_conditional_execution_guardrail, + deregister_tool_execution_intercept, deregister_tool_request_intercept, + deregister_tool_sanitize_request_guardrail, deregister_tool_sanitize_response_guardrail, + register_llm_conditional_execution_guardrail, register_llm_execution_intercept, + register_llm_request_intercept, register_llm_sanitize_request_guardrail, + register_llm_sanitize_response_guardrail, register_llm_stream_execution_intercept, + register_mark_sanitize_guardrail, register_scope_sanitize_end_guardrail, + register_scope_sanitize_start_guardrail, register_tool_conditional_execution_guardrail, register_tool_execution_intercept, register_tool_request_intercept, register_tool_sanitize_request_guardrail, register_tool_sanitize_response_guardrail, }; +use nemo_relay::api::runtime::EventSanitizeFn; use nemo_relay::api::subscriber::{deregister_subscriber, register_subscriber}; +use nemo_relay::error::Result as FlowResult; use nemo_relay::plugin::{ ConfigDiagnostic, DiagnosticLevel, Plugin, PluginConfig, PluginError, PluginRegistration, PluginRegistrationContext, active_plugin_report, clear_plugin_configuration, deregister_plugin, @@ -35,11 +40,11 @@ use nemo_relay::plugin::{ use crate::convert::{json_to_py, py_to_json}; use crate::py_callable::{ - wrap_py_event_subscriber, wrap_py_llm_conditional_fn, wrap_py_llm_exec_intercept_fn, - wrap_py_llm_request_intercept_fn, wrap_py_llm_sanitize_request_fn, - wrap_py_llm_sanitize_response_fn, wrap_py_llm_stream_exec_intercept_fn, - wrap_py_tool_conditional_fn, wrap_py_tool_exec_intercept_fn, wrap_py_tool_fn, - wrap_py_tool_request_intercept_fn, + wrap_py_event_sanitize_fn, wrap_py_event_subscriber, wrap_py_llm_conditional_fn, + wrap_py_llm_exec_intercept_fn, wrap_py_llm_request_intercept_fn, + wrap_py_llm_sanitize_request_fn, wrap_py_llm_sanitize_response_fn, + wrap_py_llm_stream_exec_intercept_fn, wrap_py_tool_conditional_fn, + wrap_py_tool_exec_intercept_fn, wrap_py_tool_fn, wrap_py_tool_request_intercept_fn, }; #[cfg(test)] @@ -206,10 +211,94 @@ impl PyPluginContext { fn qualify_name(&self, name: &str) -> String { format!("{}{}", self.namespace_prefix, name) } + + fn register_event_sanitizer( + &self, + name: &str, + priority: i32, + callback: Py, + register: fn(&str, i32, EventSanitizeFn) -> FlowResult<()>, + deregister: fn(&str) -> FlowResult, + label: &'static str, + ) -> PyResult<()> { + let qualified_name = self.qualify_name(name); + register( + &qualified_name, + priority, + wrap_py_event_sanitize_fn(callback), + ) + .map_err(to_py_err)?; + + let name_owned = qualified_name; + let mut guard = self.registrations.lock().map_err(|e| { + pyo3::exceptions::PyRuntimeError::new_err(format!("plugin context lock poisoned: {e}")) + })?; + guard.push(PluginRegistration::new( + "plugin", + name_owned.clone(), + Box::new(move || { + deregister(&name_owned).map(|_| ()).map_err(|e| { + PluginError::RegistrationFailed(format!("{label} deregistration failed: {e}")) + }) + }), + )); + Ok(()) + } } #[pymethods] impl PyPluginContext { + #[pyo3(signature = (name: "str", priority: "int", callback: "object") -> "None", text_signature = "(name: str, priority: int, callback: object) -> None")] + fn register_mark_sanitize_guardrail( + &self, + name: &str, + priority: i32, + callback: Py, + ) -> PyResult<()> { + self.register_event_sanitizer( + name, + priority, + callback, + register_mark_sanitize_guardrail, + deregister_mark_sanitize_guardrail, + "mark sanitize guardrail", + ) + } + + #[pyo3(signature = (name: "str", priority: "int", callback: "object") -> "None", text_signature = "(name: str, priority: int, callback: object) -> None")] + fn register_scope_sanitize_start_guardrail( + &self, + name: &str, + priority: i32, + callback: Py, + ) -> PyResult<()> { + self.register_event_sanitizer( + name, + priority, + callback, + register_scope_sanitize_start_guardrail, + deregister_scope_sanitize_start_guardrail, + "scope start sanitize guardrail", + ) + } + + #[pyo3(signature = (name: "str", priority: "int", callback: "object") -> "None", text_signature = "(name: str, priority: int, callback: object) -> None")] + fn register_scope_sanitize_end_guardrail( + &self, + name: &str, + priority: i32, + callback: Py, + ) -> PyResult<()> { + self.register_event_sanitizer( + name, + priority, + callback, + register_scope_sanitize_end_guardrail, + deregister_scope_sanitize_end_guardrail, + "scope end sanitize guardrail", + ) + } + #[pyo3( signature = (name: "str", callback: "object") -> "None", text_signature = "(name: str, callback: object) -> None" diff --git a/crates/python/tests/coverage/coverage_tests.rs b/crates/python/tests/coverage/coverage_tests.rs index 49e98ebf1..16093c2d5 100644 --- a/crates/python/tests/coverage/coverage_tests.rs +++ b/crates/python/tests/coverage/coverage_tests.rs @@ -248,6 +248,12 @@ fn test_register_exposes_all_native_api_functions() { "llm_call_end", "llm_call_execute", "llm_stream_call_execute", + "register_mark_sanitize_guardrail", + "deregister_mark_sanitize_guardrail", + "register_scope_sanitize_start_guardrail", + "deregister_scope_sanitize_start_guardrail", + "register_scope_sanitize_end_guardrail", + "deregister_scope_sanitize_end_guardrail", "register_tool_sanitize_request_guardrail", "deregister_tool_sanitize_request_guardrail", "register_tool_sanitize_response_guardrail", @@ -278,6 +284,12 @@ fn test_register_exposes_all_native_api_functions() { "scope_deregister_tool_sanitize_response_guardrail", "scope_register_tool_conditional_execution_guardrail", "scope_deregister_tool_conditional_execution_guardrail", + "scope_register_mark_sanitize_guardrail", + "scope_deregister_mark_sanitize_guardrail", + "scope_register_scope_sanitize_start_guardrail", + "scope_deregister_scope_sanitize_start_guardrail", + "scope_register_scope_sanitize_end_guardrail", + "scope_deregister_scope_sanitize_end_guardrail", "scope_register_tool_request_intercept", "scope_deregister_tool_request_intercept", "scope_register_tool_execution_intercept", diff --git a/crates/python/tests/coverage/py_callable_coverage_tests.rs b/crates/python/tests/coverage/py_callable_coverage_tests.rs index 10f3908a8..2ea3fb12a 100644 --- a/crates/python/tests/coverage/py_callable_coverage_tests.rs +++ b/crates/python/tests/coverage/py_callable_coverage_tests.rs @@ -595,3 +595,57 @@ async def collect_stream(awaitable): }); }); } + +#[test] +fn event_sanitize_wrapper_covers_conversion_success_and_fail_open_paths() { + use nemo_relay::api::event::{BaseEvent, MarkEvent}; + + let _python = crate::test_support::init_python_test(); + Python::attach(|py| { + let module = load_module( + py, + r#" +def sanitize(event, fields): + assert event.kind == "mark" + fields["data"] = {"safe": event.name} + fields["metadata"] = None + return fields + +def raises(event, fields): + raise RuntimeError("sanitize boom") + +def invalid(event, fields): + return "not fields" +"#, + ); + let event = Event::Mark(MarkEvent::new( + BaseEvent::builder().name("checkpoint").build(), + None, + None, + )); + let fields = EventSanitizeFields { + data: Some(json!({"secret": true})), + category_profile: None, + metadata: Some(json!({"secret": true})), + }; + + let sanitized = wrap_py_event_sanitize_fn(module.getattr("sanitize").unwrap().unbind())( + &event, + fields.clone(), + ); + assert_eq!(sanitized.data, Some(json!({"safe": "checkpoint"}))); + assert_eq!(sanitized.metadata, None); + + let raised = wrap_py_event_sanitize_fn(module.getattr("raises").unwrap().unbind())( + &event, + fields.clone(), + ); + assert_eq!(raised, fields); + + let invalid = wrap_py_event_sanitize_fn(module.getattr("invalid").unwrap().unbind())( + &event, + fields.clone(), + ); + assert_eq!(invalid, fields); + }); +} diff --git a/crates/python/tests/coverage/py_plugin_coverage_tests.rs b/crates/python/tests/coverage/py_plugin_coverage_tests.rs index 33b9c822e..03c8557d8 100644 --- a/crates/python/tests/coverage/py_plugin_coverage_tests.rs +++ b/crates/python/tests/coverage/py_plugin_coverage_tests.rs @@ -120,6 +120,9 @@ fn plugin_context_registers_all_runtime_hooks_and_drains_registrations() { def subscriber(event): return None +def event_sanitize(event, fields): + return fields + def tool_fn(name, value): return value @@ -163,6 +166,27 @@ async def tool_execution_intercept(name, value, next): helpers.getattr("subscriber").unwrap().unbind(), ) .unwrap(); + context + .register_mark_sanitize_guardrail( + "mark_sanitize", + 1, + helpers.getattr("event_sanitize").unwrap().unbind(), + ) + .unwrap(); + context + .register_scope_sanitize_start_guardrail( + "scope_start_sanitize", + 1, + helpers.getattr("event_sanitize").unwrap().unbind(), + ) + .unwrap(); + context + .register_scope_sanitize_end_guardrail( + "scope_end_sanitize", + 1, + helpers.getattr("event_sanitize").unwrap().unbind(), + ) + .unwrap(); context .register_tool_sanitize_request_guardrail( "tool_sanitize_request", @@ -250,7 +274,7 @@ async def tool_execution_intercept(name, value, next): .unwrap(); let registrations = context.drain_registrations().unwrap(); - assert_eq!(registrations.len(), 12); + assert_eq!(registrations.len(), 15); assert!( registrations .iter() @@ -258,6 +282,9 @@ async def tool_execution_intercept(name, value, next): ); assert!(deregister_subscriber("demo.subscriber").unwrap()); + assert!(deregister_mark_sanitize_guardrail("demo.mark_sanitize").unwrap()); + assert!(deregister_scope_sanitize_start_guardrail("demo.scope_start_sanitize").unwrap()); + assert!(deregister_scope_sanitize_end_guardrail("demo.scope_end_sanitize").unwrap()); assert!(deregister_tool_sanitize_request_guardrail("demo.tool_sanitize_request").unwrap()); assert!( deregister_tool_sanitize_response_guardrail("demo.tool_sanitize_response").unwrap() diff --git a/python/nemo_relay/__init__.py b/python/nemo_relay/__init__.py index bd18a47d9..afdb79170 100644 --- a/python/nemo_relay/__init__.py +++ b/python/nemo_relay/__init__.py @@ -79,7 +79,7 @@ async def main(): import contextvars import typing from collections.abc import Callable as AbcCallable -from typing import AsyncIterator, Awaitable, Callable, Literal, Optional, TypeAlias +from typing import AsyncIterator, Awaitable, Callable, Literal, Optional, TypeAlias, TypedDict # Native bitflag classes exported at the top level for user code. # Native LLM request and normalized codec view types. @@ -138,11 +138,21 @@ async def main(): #: configuration dataclasses. UnsupportedBehavior: TypeAlias = Literal["ignore", "warn", "error"] + +class EventSanitizeFields(TypedDict): + """Observability fields returned by mark and scope event sanitizers.""" + + data: Json | None + category_profile: JsonObject | None + metadata: Json | None + + #: Guardrail callback that sanitizes emitted tool request or response payloads. #: Arguments are the tool name and JSON payload. The return value is the JSON #: payload recorded on the emitted event. Exceptions propagate through the #: lifecycle call that invoked the guardrail. ToolSanitizeGuardrail: TypeAlias = Callable[[str, Json], Json] +EventSanitizeGuardrail: TypeAlias = Callable[["Event", EventSanitizeFields], EventSanitizeFields] #: Guardrail callback that can block tool execution by returning a rejection #: message. Returning ``None`` allows execution to continue. ToolConditionalExecutionGuardrail: TypeAlias = Callable[[str, Json], Optional[str]] @@ -481,6 +491,8 @@ def worker() -> None: "JsonObject", "Json", "UnsupportedBehavior", + "EventSanitizeFields", + "EventSanitizeGuardrail", "ToolSanitizeGuardrail", "ToolConditionalExecutionGuardrail", "LlmSanitizeRequestGuardrail", diff --git a/python/nemo_relay/__init__.pyi b/python/nemo_relay/__init__.pyi index 61400d129..4a95aea2a 100644 --- a/python/nemo_relay/__init__.pyi +++ b/python/nemo_relay/__init__.pyi @@ -23,7 +23,7 @@ from __future__ import annotations import contextvars from collections.abc import AsyncIterator, Awaitable, Callable -from typing import Literal, Optional, TypeAlias +from typing import Literal, Optional, TypeAlias, TypedDict from nemo_relay import adaptive as adaptive from nemo_relay import codecs as codecs @@ -142,7 +142,15 @@ Json: TypeAlias = JsonValue UnsupportedBehavior: TypeAlias = Literal["ignore", "warn", "error"] """Policy used by config helpers when unknown fields or values are encountered.""" +class EventSanitizeFields(TypedDict): + """Observability fields returned by mark and scope event sanitizers.""" + + data: Json | None + category_profile: JsonObject | None + metadata: Json | None + ToolSanitizeGuardrail: TypeAlias = Callable[[str, Json], Json] +EventSanitizeGuardrail: TypeAlias = Callable[[Event, EventSanitizeFields], EventSanitizeFields] """Guardrail callback that sanitizes emitted tool request or response payloads. Arguments: diff --git a/python/nemo_relay/_native.pyi b/python/nemo_relay/_native.pyi index 81a211cff..3b7f65aca 100644 --- a/python/nemo_relay/_native.pyi +++ b/python/nemo_relay/_native.pyi @@ -24,16 +24,23 @@ from __future__ import annotations from collections.abc import AsyncIterator, Awaitable, Callable, Generator, Mapping, Sequence from datetime import datetime -from typing import ClassVar, Literal, Optional, TypeAlias +from typing import ClassVar, Literal, Optional, TypeAlias, TypedDict _JsonPrimitive: TypeAlias = str | int | float | bool | None _JsonValue: TypeAlias = _JsonPrimitive | list["_JsonValue"] | dict[str, "_JsonValue"] _JsonObject: TypeAlias = dict[str, _JsonValue] _Json: TypeAlias = _JsonValue + +class _EventSanitizeFields(TypedDict): + data: _Json | None + category_profile: _JsonObject | None + metadata: _Json | None + _ToolSanitizeGuardrail: TypeAlias = Callable[[str, _Json], _Json] _ToolConditionalExecutionGuardrail: TypeAlias = Callable[[str, _Json], Optional[str]] _LlmSanitizeRequestGuardrail: TypeAlias = Callable[["LLMRequest"], "LLMRequest"] _LlmSanitizeResponseGuardrail: TypeAlias = Callable[[_JsonObject], _JsonObject] +_EventSanitizeGuardrail: TypeAlias = Callable[[ScopeEvent | MarkEvent, _EventSanitizeFields], _EventSanitizeFields] _LlmConditionalExecutionGuardrail: TypeAlias = Callable[["LLMRequest"], Optional[str]] _ToolRequestIntercept: TypeAlias = Callable[[str, _Json], _Json] _ToolExecutionIntercept: TypeAlias = Callable[ @@ -1180,7 +1187,32 @@ class PluginContext: Python plugin protocols expose the public shape. The native class exists for runtime registration callbacks. """ + def register_mark_sanitize_guardrail(self, name: str, priority: int, callback: _EventSanitizeGuardrail) -> None: ... + def register_scope_sanitize_start_guardrail( + self, name: str, priority: int, callback: _EventSanitizeGuardrail + ) -> None: ... + def register_scope_sanitize_end_guardrail( + self, name: str, priority: int, callback: _EventSanitizeGuardrail + ) -> None: ... +def register_mark_sanitize_guardrail(name: str, priority: int, guardrail: _EventSanitizeGuardrail) -> None: ... +def deregister_mark_sanitize_guardrail(name: str) -> bool: ... +def register_scope_sanitize_start_guardrail(name: str, priority: int, guardrail: _EventSanitizeGuardrail) -> None: ... +def deregister_scope_sanitize_start_guardrail(name: str) -> bool: ... +def register_scope_sanitize_end_guardrail(name: str, priority: int, guardrail: _EventSanitizeGuardrail) -> None: ... +def deregister_scope_sanitize_end_guardrail(name: str) -> bool: ... +def scope_register_mark_sanitize_guardrail( + scope_uuid: str, name: str, priority: int, guardrail: _EventSanitizeGuardrail +) -> None: ... +def scope_deregister_mark_sanitize_guardrail(scope_uuid: str, name: str) -> bool: ... +def scope_register_scope_sanitize_start_guardrail( + scope_uuid: str, name: str, priority: int, guardrail: _EventSanitizeGuardrail +) -> None: ... +def scope_deregister_scope_sanitize_start_guardrail(scope_uuid: str, name: str) -> bool: ... +def scope_register_scope_sanitize_end_guardrail( + scope_uuid: str, name: str, priority: int, guardrail: _EventSanitizeGuardrail +) -> None: ... +def scope_deregister_scope_sanitize_end_guardrail(scope_uuid: str, name: str) -> bool: ... def create_scope_stack() -> ScopeStack: """Create a fresh native scope stack. diff --git a/python/nemo_relay/guardrails.py b/python/nemo_relay/guardrails.py index 478fcdde1..92bc8d9b5 100644 --- a/python/nemo_relay/guardrails.py +++ b/python/nemo_relay/guardrails.py @@ -21,6 +21,7 @@ def redact(tool_name, args): """ from nemo_relay import ( + EventSanitizeGuardrail, LlmConditionalExecutionGuardrail, LlmSanitizeRequestGuardrail, LlmSanitizeResponseGuardrail, @@ -36,6 +37,13 @@ def redact(tool_name, args): from nemo_relay._native import ( deregister_llm_sanitize_response_guardrail as _native_deregister_llm_sanitize_response, ) +from nemo_relay._native import deregister_mark_sanitize_guardrail as _native_deregister_mark_sanitize +from nemo_relay._native import ( + deregister_scope_sanitize_end_guardrail as _native_deregister_scope_sanitize_end, +) +from nemo_relay._native import ( + deregister_scope_sanitize_start_guardrail as _native_deregister_scope_sanitize_start, +) from nemo_relay._native import ( deregister_tool_conditional_execution_guardrail as _native_deregister_tool_conditional_execution, ) @@ -54,6 +62,13 @@ def redact(tool_name, args): from nemo_relay._native import ( register_llm_sanitize_response_guardrail as _native_register_llm_sanitize_response, ) +from nemo_relay._native import register_mark_sanitize_guardrail as _native_register_mark_sanitize +from nemo_relay._native import ( + register_scope_sanitize_end_guardrail as _native_register_scope_sanitize_end, +) +from nemo_relay._native import ( + register_scope_sanitize_start_guardrail as _native_register_scope_sanitize_start, +) from nemo_relay._native import ( register_tool_conditional_execution_guardrail as _native_register_tool_conditional_execution, ) @@ -64,6 +79,41 @@ def redact(tool_name, args): register_tool_sanitize_response_guardrail as _native_register_tool_sanitize_response, ) +# --------------------------------------------------------------------------- +# Mark and scope event guardrails +# --------------------------------------------------------------------------- + + +def register_mark_sanitize(name: str, priority: int, guardrail: EventSanitizeGuardrail) -> None: + """Register a sanitizer for mark event observability fields.""" + return _native_register_mark_sanitize(name, priority, guardrail) + + +def deregister_mark_sanitize(name: str) -> bool: + """Remove a global mark event sanitizer by name.""" + return _native_deregister_mark_sanitize(name) + + +def register_scope_sanitize_start(name: str, priority: int, guardrail: EventSanitizeGuardrail) -> None: + """Register a sanitizer for every scope start event category.""" + return _native_register_scope_sanitize_start(name, priority, guardrail) + + +def deregister_scope_sanitize_start(name: str) -> bool: + """Remove a global scope-start event sanitizer by name.""" + return _native_deregister_scope_sanitize_start(name) + + +def register_scope_sanitize_end(name: str, priority: int, guardrail: EventSanitizeGuardrail) -> None: + """Register a sanitizer for every scope end event category.""" + return _native_register_scope_sanitize_end(name, priority, guardrail) + + +def deregister_scope_sanitize_end(name: str) -> bool: + """Remove a global scope-end event sanitizer by name.""" + return _native_deregister_scope_sanitize_end(name) + + # --------------------------------------------------------------------------- # Tool guardrails # --------------------------------------------------------------------------- @@ -312,6 +362,12 @@ def deregister_llm_conditional_execution(name: str) -> bool: __all__ = [ + "register_mark_sanitize", + "deregister_mark_sanitize", + "register_scope_sanitize_start", + "deregister_scope_sanitize_start", + "register_scope_sanitize_end", + "deregister_scope_sanitize_end", "register_tool_sanitize_request", "deregister_tool_sanitize_request", "register_tool_sanitize_response", diff --git a/python/nemo_relay/pii_redaction.py b/python/nemo_relay/pii_redaction.py index 3f463f517..53c2b77f8 100644 --- a/python/nemo_relay/pii_redaction.py +++ b/python/nemo_relay/pii_redaction.py @@ -134,6 +134,7 @@ class PiiRedactionConfig: output: bool = True tool_input: bool = True tool_output: bool = True + mark: bool = True priority: int = 100 codec: Literal["openai_chat", "openai_responses", "anthropic_messages"] | str | None = None builtin: BuiltinConfig | None = None @@ -150,6 +151,7 @@ def to_dict(self) -> JsonObject: "output": self.output, "tool_input": self.tool_input, "tool_output": self.tool_output, + "mark": self.mark, "priority": self.priority, "codec": self.codec, "builtin": self.builtin, diff --git a/python/nemo_relay/pii_redaction.pyi b/python/nemo_relay/pii_redaction.pyi index de1032701..244f6a3ef 100644 --- a/python/nemo_relay/pii_redaction.pyi +++ b/python/nemo_relay/pii_redaction.pyi @@ -56,6 +56,7 @@ class PiiRedactionConfig: output: bool = ... tool_input: bool = ... tool_output: bool = ... + mark: bool = ... priority: int = ... codec: Literal["openai_chat", "openai_responses", "anthropic_messages"] | str | None = ... builtin: BuiltinConfig | None = ... diff --git a/python/nemo_relay/plugin.py b/python/nemo_relay/plugin.py index 7e7b2f0f3..9dbb14f2e 100644 --- a/python/nemo_relay/plugin.py +++ b/python/nemo_relay/plugin.py @@ -15,6 +15,7 @@ from typing import TYPE_CHECKING, AsyncIterator, Callable, Literal, Protocol, TypedDict, cast from nemo_relay import ( + EventSanitizeGuardrail, Json, JsonObject, LlmConditionalExecutionGuardrail, @@ -82,6 +83,20 @@ def register_subscriber(self, name: str, callback: Callable[[Event], None]) -> N """Register an infallible event subscriber for this component.""" ... + def register_mark_sanitize_guardrail(self, name: str, priority: int, callback: EventSanitizeGuardrail) -> None: + """Register a mark event sanitizer for this component.""" + ... + + def register_scope_sanitize_start_guardrail( + self, name: str, priority: int, callback: EventSanitizeGuardrail + ) -> None: + """Register a scope-start event sanitizer for this component.""" + ... + + def register_scope_sanitize_end_guardrail(self, name: str, priority: int, callback: EventSanitizeGuardrail) -> None: + """Register a scope-end event sanitizer for this component.""" + ... + def register_tool_sanitize_request_guardrail( self, name: str, priority: int, callback: ToolSanitizeGuardrail ) -> None: diff --git a/python/nemo_relay/plugin.pyi b/python/nemo_relay/plugin.pyi index 9e286830d..3c0e5fc87 100644 --- a/python/nemo_relay/plugin.pyi +++ b/python/nemo_relay/plugin.pyi @@ -6,6 +6,7 @@ from typing import AsyncContextManager, Literal, Protocol, TypedDict from nemo_relay import ( Event, + EventSanitizeGuardrail, JsonObject, LlmConditionalExecutionGuardrail, LlmExecutionIntercept, @@ -35,6 +36,13 @@ class ConfigReport(TypedDict): class PluginContext(Protocol): def register_subscriber(self, name: str, callback: Callable[[Event], None]) -> None: ... + def register_mark_sanitize_guardrail(self, name: str, priority: int, callback: EventSanitizeGuardrail) -> None: ... + def register_scope_sanitize_start_guardrail( + self, name: str, priority: int, callback: EventSanitizeGuardrail + ) -> None: ... + def register_scope_sanitize_end_guardrail( + self, name: str, priority: int, callback: EventSanitizeGuardrail + ) -> None: ... def register_tool_sanitize_request_guardrail( self, name: str, priority: int, callback: ToolSanitizeGuardrail ) -> None: ... diff --git a/python/nemo_relay/scope_local.py b/python/nemo_relay/scope_local.py index 7c1191c5f..08975721a 100644 --- a/python/nemo_relay/scope_local.py +++ b/python/nemo_relay/scope_local.py @@ -37,6 +37,15 @@ def redact(tool_name, args): from nemo_relay._native import ( scope_deregister_llm_stream_execution_intercept as _deregister_llm_stream_execution, ) +from nemo_relay._native import ( + scope_deregister_mark_sanitize_guardrail as _deregister_mark_sanitize, +) +from nemo_relay._native import ( + scope_deregister_scope_sanitize_end_guardrail as _deregister_scope_sanitize_end, +) +from nemo_relay._native import ( + scope_deregister_scope_sanitize_start_guardrail as _deregister_scope_sanitize_start, +) from nemo_relay._native import ( scope_deregister_subscriber as _deregister_subscriber, ) @@ -73,6 +82,15 @@ def redact(tool_name, args): from nemo_relay._native import ( scope_register_llm_stream_execution_intercept as _register_llm_stream_execution, ) +from nemo_relay._native import ( + scope_register_mark_sanitize_guardrail as _register_mark_sanitize, +) +from nemo_relay._native import ( + scope_register_scope_sanitize_end_guardrail as _register_scope_sanitize_end, +) +from nemo_relay._native import ( + scope_register_scope_sanitize_start_guardrail as _register_scope_sanitize_start, +) from nemo_relay._native import ( scope_register_subscriber as _register_subscriber, ) @@ -92,6 +110,41 @@ def redact(tool_name, args): scope_register_tool_sanitize_response_guardrail as _register_tool_sanitize_response, ) +# --------------------------------------------------------------------------- +# Mark and scope event guardrails (scope-local) +# --------------------------------------------------------------------------- + + +def register_mark_sanitize(scope_handle, name, priority, guardrail): + """Register a scope-local mark event sanitizer.""" + return _register_mark_sanitize(scope_handle.uuid, name, priority, guardrail) + + +def deregister_mark_sanitize(scope_handle, name): + """Remove a scope-local mark event sanitizer.""" + return _deregister_mark_sanitize(scope_handle.uuid, name) + + +def register_scope_sanitize_start(scope_handle, name, priority, guardrail): + """Register a scope-local sanitizer for scope start events.""" + return _register_scope_sanitize_start(scope_handle.uuid, name, priority, guardrail) + + +def deregister_scope_sanitize_start(scope_handle, name): + """Remove a scope-local scope-start event sanitizer.""" + return _deregister_scope_sanitize_start(scope_handle.uuid, name) + + +def register_scope_sanitize_end(scope_handle, name, priority, guardrail): + """Register a scope-local sanitizer for scope end events.""" + return _register_scope_sanitize_end(scope_handle.uuid, name, priority, guardrail) + + +def deregister_scope_sanitize_end(scope_handle, name): + """Remove a scope-local scope-end event sanitizer.""" + return _deregister_scope_sanitize_end(scope_handle.uuid, name) + + # --------------------------------------------------------------------------- # Tool guardrails (scope-local) # --------------------------------------------------------------------------- @@ -607,6 +660,13 @@ def deregister_subscriber(scope_handle, name): __all__ = [ + # Mark and scope event guardrails + "register_mark_sanitize", + "deregister_mark_sanitize", + "register_scope_sanitize_start", + "deregister_scope_sanitize_start", + "register_scope_sanitize_end", + "deregister_scope_sanitize_end", # Tool guardrails "register_tool_sanitize_request", "deregister_tool_sanitize_request", diff --git a/python/plugin/src/nemo_relay_plugin/__init__.py b/python/plugin/src/nemo_relay_plugin/__init__.py index 7d11e0876..2de65f21d 100644 --- a/python/plugin/src/nemo_relay_plugin/__init__.py +++ b/python/plugin/src/nemo_relay_plugin/__init__.py @@ -18,6 +18,7 @@ Public data types: Json: Any JSON-serializable Python value. Event: A Relay event represented as a JSON object. + EventSanitizeFields: Mutable event observability fields. LlmRequest: A Relay LLM request represented as a JSON object. AnnotatedLlmRequest: An annotated Relay LLM request represented as a JSON object. @@ -31,6 +32,7 @@ Public callback aliases: SubscriberCallback: Event subscriber callback. + EventSanitizeCallback: Mark or scope event sanitizer callback. ToolSanitizeCallback: Tool request or response sanitizer callback. ToolConditionalCallback: Tool execution guardrail callback. ToolRequestCallback: Tool request intercept callback. @@ -59,6 +61,8 @@ ConfigDiagnostic, DiagnosticLevel, Event, + EventSanitizeCallback, + EventSanitizeFields, Json, LlmConditionalCallback, LlmExecutionCallback, @@ -91,6 +95,8 @@ "ConfigDiagnostic", "DiagnosticLevel", "Event", + "EventSanitizeCallback", + "EventSanitizeFields", "Json", "LlmConditionalCallback", "LlmExecutionCallback", diff --git a/python/plugin/src/nemo_relay_plugin/_api.py b/python/plugin/src/nemo_relay_plugin/_api.py index 615e8d868..a5471c37f 100644 --- a/python/plugin/src/nemo_relay_plugin/_api.py +++ b/python/plugin/src/nemo_relay_plugin/_api.py @@ -69,7 +69,7 @@ from enum import Enum from importlib import metadata from pathlib import Path -from typing import Any, Protocol, TypeAlias +from typing import Any, Protocol, TypeAlias, TypedDict from urllib.parse import urlsplit grpc: Any = importlib.import_module("grpc") @@ -80,6 +80,16 @@ Json: TypeAlias = Any #: A Relay event represented as a JSON object. Event: TypeAlias = dict[str, Any] + + +class EventSanitizeFields(TypedDict): + """Observability fields returned by event sanitizer callbacks.""" + + data: Json | None + category_profile: dict[str, Any] | None + metadata: Json | None + + #: A Relay LLM request represented as a JSON object. LlmRequest: TypeAlias = dict[str, Any] #: An annotated Relay LLM request represented as a JSON object. @@ -393,6 +403,10 @@ def register(self, ctx: PluginContext, config: Json) -> None | Awaitable[None]: SubscriberCallback: TypeAlias = Callable[[Event], None | Awaitable[None]] +EventSanitizeCallback: TypeAlias = Callable[ + [Event, EventSanitizeFields], + EventSanitizeFields | Awaitable[EventSanitizeFields], +] ToolSanitizeCallback: TypeAlias = Callable[[str, Json], Json | Awaitable[Json]] ToolConditionalCallback: TypeAlias = Callable[[str, Json], str | None | Awaitable[str | None]] ToolRequestCallback: TypeAlias = Callable[[str, Json], Json | Awaitable[Json]] @@ -418,6 +432,7 @@ def register(self, ctx: PluginContext, config: Json) -> None | Awaitable[None]: class _Handlers: registrations: list[Any] subscribers: dict[str, SubscriberCallback] + event_sanitizers: dict[str, EventSanitizeCallback] tool_sanitize_requests: dict[str, ToolSanitizeCallback] tool_sanitize_responses: dict[str, ToolSanitizeCallback] tool_conditionals: dict[str, ToolConditionalCallback] @@ -435,6 +450,7 @@ def empty(cls) -> _Handlers: return cls( registrations=[], subscribers={}, + event_sanitizers={}, tool_sanitize_requests={}, tool_sanitize_responses={}, tool_conditionals={}, @@ -499,6 +515,34 @@ def register_subscriber(self, name: str, callback: SubscriberCallback) -> None: self._push_registration(name, pb.SUBSCRIBER, 0, False) self._handlers.subscribers[name] = callback + def _register_event_sanitizer( + self, + name: str, + callback: EventSanitizeCallback, + surface: int, + priority: int, + ) -> None: + self._push_registration(name, surface, priority, False) + self._handlers.event_sanitizers[name] = callback + + def register_mark_sanitize_guardrail( + self, name: str, callback: EventSanitizeCallback, *, priority: int = 0 + ) -> None: + """Register a sanitizer for mark event observability fields.""" + self._register_event_sanitizer(name, callback, pb.MARK_SANITIZE_GUARDRAIL, priority) + + def register_scope_sanitize_start_guardrail( + self, name: str, callback: EventSanitizeCallback, *, priority: int = 0 + ) -> None: + """Register a sanitizer for scope start event observability fields.""" + self._register_event_sanitizer(name, callback, pb.SCOPE_SANITIZE_START_GUARDRAIL, priority) + + def register_scope_sanitize_end_guardrail( + self, name: str, callback: EventSanitizeCallback, *, priority: int = 0 + ) -> None: + """Register a sanitizer for scope end event observability fields.""" + self._register_event_sanitizer(name, callback, pb.SCOPE_SANITIZE_END_GUARDRAIL, priority) + def register_tool_sanitize_request_guardrail( self, name: str, @@ -1422,6 +1466,22 @@ async def _invoke_result(self, request: Any) -> Any: event = _decode_required_envelope(request.event, "event", EVENT_SCHEMA) await _maybe_await(self._handler(self._handlers.subscribers, request.registration_name)(event)) return pb.InvokeResponse(empty=pb.EmptyResult()) + if request.surface in { + pb.MARK_SANITIZE_GUARDRAIL, + pb.SCOPE_SANITIZE_START_GUARDRAIL, + pb.SCOPE_SANITIZE_END_GUARDRAIL, + }: + event = _decode_required_envelope(request.event, "event", EVENT_SCHEMA) + fields: EventSanitizeFields = { + "data": event.get("data"), + "category_profile": event.get("category_profile"), + "metadata": event.get("metadata"), + } + return _json_response( + await _maybe_await( + self._handler(self._handlers.event_sanitizers, request.registration_name)(event, fields) + ) + ) if request.surface == pb.TOOL_SANITIZE_REQUEST_GUARDRAIL: return _json_response( await _maybe_await( @@ -1584,6 +1644,9 @@ def _plugin_id(plugin: _SupportsWorkerPlugin) -> str: def _all_surfaces() -> list[int]: return [ pb.SUBSCRIBER, + pb.MARK_SANITIZE_GUARDRAIL, + pb.SCOPE_SANITIZE_START_GUARDRAIL, + pb.SCOPE_SANITIZE_END_GUARDRAIL, pb.TOOL_SANITIZE_REQUEST_GUARDRAIL, pb.TOOL_SANITIZE_RESPONSE_GUARDRAIL, pb.TOOL_CONDITIONAL_EXECUTION_GUARDRAIL, diff --git a/python/tests/plugin/test_public_api_docstrings.py b/python/tests/plugin/test_public_api_docstrings.py index c37a141ba..31abf7427 100644 --- a/python/tests/plugin/test_public_api_docstrings.py +++ b/python/tests/plugin/test_public_api_docstrings.py @@ -21,6 +21,7 @@ _TYPE_ALIASES = { "AnnotatedLlmRequest", "Event", + "EventSanitizeCallback", "Json", "LlmRequest", "SubscriberCallback", diff --git a/python/tests/plugin/test_worker_sdk.py b/python/tests/plugin/test_worker_sdk.py index 3f33fbb27..da37a5662 100644 --- a/python/tests/plugin/test_worker_sdk.py +++ b/python/tests/plugin/test_worker_sdk.py @@ -184,6 +184,9 @@ def register(self, ctx: PluginContext, config: Json) -> None: async def subscriber(event: Json) -> None: await ctx.runtime.emit_mark("tests.subscriber", event) + async def event_sanitize(event: Json, fields: Json) -> Json: + return {**fields, "data": {"sanitized": event["name"]}, "metadata": None} + def tool_sanitize(name: str, value: Json) -> Json: return _tag(value, f"sanitize_{name}") @@ -229,6 +232,9 @@ async def llm_stream_execution(name: str, request: Json, next_call: Any) -> Asyn yield _tag(chunk, "llm_stream_execution") ctx.register_subscriber("subscriber", subscriber) + ctx.register_mark_sanitize_guardrail("mark_sanitize", event_sanitize, priority=1) + ctx.register_scope_sanitize_start_guardrail("scope_start_sanitize", event_sanitize, priority=2) + ctx.register_scope_sanitize_end_guardrail("scope_end_sanitize", event_sanitize, priority=3) ctx.register_tool_sanitize_request_guardrail("tool_sanitize", tool_sanitize, priority=1) ctx.register_tool_sanitize_response_guardrail("tool_sanitize", tool_sanitize, priority=2) ctx.register_tool_conditional_execution_guardrail("tool_conditional", tool_block, priority=3) @@ -270,6 +276,9 @@ def test_generated_proto_matches_worker_contract(): assert pb.SUBSCRIBER == 1 assert pb.TOOL_SANITIZE_REQUEST_GUARDRAIL == 10 assert pb.LLM_STREAM_EXECUTION_INTERCEPT == 25 + assert pb.MARK_SANITIZE_GUARDRAIL == 30 + assert pb.SCOPE_SANITIZE_START_GUARDRAIL == 31 + assert pb.SCOPE_SANITIZE_END_GUARDRAIL == 32 assert pb.CUSTOM == 10 @@ -308,6 +317,9 @@ async def test_health_handshake_validate_register_and_all_surfaces(service: _Wor ] assert registrations == [ ("subscriber", pb.SUBSCRIBER, 0, False), + ("mark_sanitize", pb.MARK_SANITIZE_GUARDRAIL, 1, False), + ("scope_start_sanitize", pb.SCOPE_SANITIZE_START_GUARDRAIL, 2, False), + ("scope_end_sanitize", pb.SCOPE_SANITIZE_END_GUARDRAIL, 3, False), ("tool_sanitize", pb.TOOL_SANITIZE_REQUEST_GUARDRAIL, 1, False), ("tool_sanitize", pb.TOOL_SANITIZE_RESPONSE_GUARDRAIL, 2, False), ("tool_conditional", pb.TOOL_CONDITIONAL_EXECUTION_GUARDRAIL, 3, False), @@ -590,6 +602,47 @@ def callback(tool_name: str, value: Json) -> Json: assert ("shared", pb.TOOL_SANITIZE_RESPONSE_GUARDRAIL) in registrations +@pytest.mark.parametrize( + "surface", + [ + pb.MARK_SANITIZE_GUARDRAIL, + pb.SCOPE_SANITIZE_START_GUARDRAIL, + pb.SCOPE_SANITIZE_END_GUARDRAIL, + ], +) +async def test_event_sanitizer_surfaces_receive_context_and_return_all_fields( + service: _WorkerService, surface: int +) -> None: + await _register(service) + registration = { + pb.MARK_SANITIZE_GUARDRAIL: "mark_sanitize", + pb.SCOPE_SANITIZE_START_GUARDRAIL: "scope_start_sanitize", + pb.SCOPE_SANITIZE_END_GUARDRAIL: "scope_end_sanitize", + }[surface] + response = await service.Invoke( + _invoke_request( + registration, + surface, + event=_json_envelope( + EVENT_SCHEMA, + { + "kind": "mark" if surface == pb.MARK_SANITIZE_GUARDRAIL else "scope", + "name": "worker-event", + "data": {"secret": True}, + "category_profile": {"subtype": "test"}, + "metadata": {"secret": True}, + }, + ), + ), + AbortContext(), + ) + assert _envelope_value(response.json.value) == { + "data": {"sanitized": "worker-event"}, + "category_profile": {"subtype": "test"}, + "metadata": None, + } + + async def test_validate_accepts_missing_config_and_dict_diagnostics(): class DictDiagnosticPlugin(WorkerPlugin): plugin_id = "tests.dict_diagnostic" @@ -2323,6 +2376,9 @@ def _llm_stream_next(runtime: PluginRuntime, request: Json) -> AsyncIterator[Jso def _all_expected_surfaces() -> list[int]: return [ pb.SUBSCRIBER, + pb.MARK_SANITIZE_GUARDRAIL, + pb.SCOPE_SANITIZE_START_GUARDRAIL, + pb.SCOPE_SANITIZE_END_GUARDRAIL, pb.TOOL_SANITIZE_REQUEST_GUARDRAIL, pb.TOOL_SANITIZE_RESPONSE_GUARDRAIL, pb.TOOL_CONDITIONAL_EXECUTION_GUARDRAIL, diff --git a/python/tests/test_event_sanitizers.py b/python/tests/test_event_sanitizers.py new file mode 100644 index 000000000..2949a48b5 --- /dev/null +++ b/python/tests/test_event_sanitizers.py @@ -0,0 +1,198 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from typing import cast + +import pytest + +import nemo_relay +from nemo_relay import EventSanitizeFields, guardrails, plugin, scope, scope_local, subscribers + + +def _capture_events(): + events = [] + name = "test-event-sanitizer-capture" + subscribers.register(name, events.append) + return name, events + + +def test_global_mark_sanitizers_order_convert_fields_and_remove_values() -> None: + capture_name, events = _capture_events() + calls: list[tuple[str, object]] = [] + + def first(event: nemo_relay.Event, fields: EventSanitizeFields) -> EventSanitizeFields: + calls.append((event.name, fields["data"])) + return { + "data": {"stage": "first"}, + "category_profile": fields["category_profile"], + "metadata": None, + } + + def second(event: nemo_relay.Event, fields: EventSanitizeFields) -> EventSanitizeFields: + calls.append((event.kind, fields["data"])) + return { + "data": {"stage": "second"}, + "category_profile": fields["category_profile"], + "metadata": fields["metadata"], + } + + guardrails.register_mark_sanitize("python-mark-second", 20, second) + guardrails.register_mark_sanitize("python-mark-first", 10, first) + try: + scope.event("checkpoint", data={"secret": "raw"}, metadata={"secret": "raw"}) + subscribers.flush() + finally: + guardrails.deregister_mark_sanitize("python-mark-first") + guardrails.deregister_mark_sanitize("python-mark-second") + subscribers.deregister(capture_name) + + mark = events[-1] + assert mark.data == {"stage": "second"} + assert mark.metadata is None + assert calls == [("checkpoint", {"secret": "raw"}), ("mark", {"stage": "first"})] + + +def test_invalid_mark_sanitizer_result_fails_open() -> None: + capture_name, events = _capture_events() + guardrails.register_mark_sanitize( + "python-mark-invalid", + 0, + cast(nemo_relay.EventSanitizeGuardrail, lambda _event, _fields: "invalid"), + ) + try: + scope.event("checkpoint", data={"kept": True}) + subscribers.flush() + finally: + guardrails.deregister_mark_sanitize("python-mark-invalid") + subscribers.deregister(capture_name) + + assert events[-1].data == {"kept": True} + + +def test_scope_start_and_end_sanitizers_cover_category_profile() -> None: + capture_name, events = _capture_events() + + def sanitize(_event: nemo_relay.Event, fields: EventSanitizeFields) -> EventSanitizeFields: + profile = dict(fields["category_profile"] or {}) + profile["subtype"] = "sanitized" + return {"data": None, "category_profile": profile, "metadata": {"safe": True}} + + guardrails.register_scope_sanitize_start("python-scope-start", 0, sanitize) + guardrails.register_scope_sanitize_end("python-scope-end", 0, sanitize) + try: + handle = scope.push( + "generic", + nemo_relay.ScopeType.Custom, + data={"secret": "start"}, + metadata={"secret": "start"}, + input={"secret": "input"}, + ) + scope.pop(handle, output={"secret": "output"}, metadata={"secret": "end"}) + subscribers.flush() + finally: + guardrails.deregister_scope_sanitize_start("python-scope-start") + guardrails.deregister_scope_sanitize_end("python-scope-end") + subscribers.deregister(capture_name) + + lifecycle = [event for event in events if event.name == "generic"] + assert len(lifecycle) == 2 + assert all(event.data is None for event in lifecycle) + assert all(event.metadata == {"safe": True} for event in lifecycle) + assert all(event.category_profile["subtype"] == "sanitized" for event in lifecycle) + + +def test_scope_local_event_sanitizers_are_inherited_and_cleaned_up() -> None: + capture_name, events = _capture_events() + + def sanitize(_event: nemo_relay.Event, fields: EventSanitizeFields) -> EventSanitizeFields: + return { + "data": {"scope_local": True}, + "category_profile": fields["category_profile"], + "metadata": fields["metadata"], + } + + owner = scope.push("owner", nemo_relay.ScopeType.Agent) + scope_local.register_mark_sanitize(owner, "python-local-mark", 0, sanitize) + scope.event("inside", data={"raw": True}) + child = scope.push("child", nemo_relay.ScopeType.Function) + scope.event("inherited", data={"raw": True}) + scope.pop(child) + scope.pop(owner) + scope.event("outside", data={"raw": True}) + subscribers.flush() + subscribers.deregister(capture_name) + + marks = {event.name: event for event in events if event.kind == "mark"} + assert marks["inside"].data == {"scope_local": True} + assert marks["inherited"].data == {"scope_local": True} + assert marks["outside"].data == {"raw": True} + + +async def test_in_process_plugin_event_sanitizers_are_removed_on_clear() -> None: + class EventPlugin: + def validate(self, _config): + return None + + def register(self, _config, context): + def sanitize(_event: nemo_relay.Event, fields: EventSanitizeFields) -> EventSanitizeFields: + return { + "data": {"plugin": True}, + "category_profile": fields["category_profile"], + "metadata": fields["metadata"], + } + + context.register_mark_sanitize_guardrail("mark", 0, sanitize) + + kind = "python.test_event_sanitizer" + capture_name, events = _capture_events() + plugin.register(kind, cast(plugin.Plugin, EventPlugin())) + try: + await plugin.initialize(plugin.PluginConfig(components=[plugin.ComponentSpec(kind=kind)])) + scope.event("configured", data={"raw": True}) + subscribers.flush() + plugin.clear() + scope.event("cleared", data={"raw": True}) + subscribers.flush() + finally: + plugin.clear() + plugin.deregister(kind) + subscribers.deregister(capture_name) + + marks = {event.name: event for event in events if event.kind == "mark"} + assert marks["configured"].data == {"plugin": True} + assert marks["cleared"].data == {"raw": True} + + +async def test_in_process_plugin_rolls_back_event_sanitizer_when_registration_fails() -> None: + class FailingPlugin: + def validate(self, _config): + return None + + def register(self, _config, context): + def sanitize(_event: nemo_relay.Event, fields: EventSanitizeFields) -> EventSanitizeFields: + return { + "data": {"leaked": True}, + "category_profile": fields["category_profile"], + "metadata": fields["metadata"], + } + + context.register_mark_sanitize_guardrail("mark", 0, sanitize) + raise RuntimeError("registration failed") + + kind = "python.test_event_sanitizer_rollback" + plugin.register(kind, cast(plugin.Plugin, FailingPlugin())) + try: + with pytest.raises(RuntimeError, match="registration failed"): + await plugin.initialize(plugin.PluginConfig(components=[plugin.ComponentSpec(kind=kind)])) + capture_name, events = _capture_events() + try: + scope.event("after-failure", data={"raw": True}) + subscribers.flush() + finally: + subscribers.deregister(capture_name) + assert events[-1].data == {"raw": True} + finally: + plugin.clear() + plugin.deregister(kind) diff --git a/python/tests/test_pii_redaction_plugin.py b/python/tests/test_pii_redaction_plugin.py index 4b3943dd8..cbed8c4c2 100644 --- a/python/tests/test_pii_redaction_plugin.py +++ b/python/tests/test_pii_redaction_plugin.py @@ -37,6 +37,10 @@ def test_defaults_and_component_wrapper(self): assert isinstance(wrapped_config, dict) assert wrapped_config["version"] == 1 assert wrapped_config["mode"] == "builtin" + assert wrapped_config["mark"] is True + + opted_out = PiiRedactionConfig(mark=False).to_dict() + assert opted_out["mark"] is False def test_validation_rejects_bad_values(self): report = validate_config( From f224e4bd8ed3e16f4c141f2bbab165a6666fb577 Mon Sep 17 00:00:00 2001 From: Will Killian Date: Thu, 9 Jul 2026 10:24:46 -0400 Subject: [PATCH 2/3] fix: address event sanitizer review feedback Signed-off-by: Will Killian --- crates/python/src/py_plugin.rs | 478 +++++++------------- python/plugin/src/nemo_relay_plugin/_api.py | 89 +++- python/tests/plugin/test_worker_sdk.py | 31 +- python/tests/test_event_sanitizers.py | 61 ++- 4 files changed, 279 insertions(+), 380 deletions(-) diff --git a/crates/python/src/py_plugin.rs b/crates/python/src/py_plugin.rs index c3e98664e..3178471dc 100644 --- a/crates/python/src/py_plugin.rs +++ b/crates/python/src/py_plugin.rs @@ -28,7 +28,6 @@ use nemo_relay::api::registry::{ register_tool_execution_intercept, register_tool_request_intercept, register_tool_sanitize_request_guardrail, register_tool_sanitize_response_guardrail, }; -use nemo_relay::api::runtime::EventSanitizeFn; use nemo_relay::api::subscriber::{deregister_subscriber, register_subscriber}; use nemo_relay::error::Result as FlowResult; use nemo_relay::plugin::{ @@ -212,22 +211,15 @@ impl PyPluginContext { format!("{}{}", self.namespace_prefix, name) } - fn register_event_sanitizer( + fn register_callback( &self, name: &str, - priority: i32, - callback: Py, - register: fn(&str, i32, EventSanitizeFn) -> FlowResult<()>, + register: impl FnOnce(&str) -> FlowResult<()>, deregister: fn(&str) -> FlowResult, label: &'static str, ) -> PyResult<()> { let qualified_name = self.qualify_name(name); - register( - &qualified_name, - priority, - wrap_py_event_sanitize_fn(callback), - ) - .map_err(to_py_err)?; + register(&qualified_name).map_err(to_py_err)?; let name_owned = qualified_name; let mut guard = self.registrations.lock().map_err(|e| { @@ -255,11 +247,15 @@ impl PyPluginContext { priority: i32, callback: Py, ) -> PyResult<()> { - self.register_event_sanitizer( + self.register_callback( name, - priority, - callback, - register_mark_sanitize_guardrail, + |qualified_name| { + register_mark_sanitize_guardrail( + qualified_name, + priority, + wrap_py_event_sanitize_fn(callback), + ) + }, deregister_mark_sanitize_guardrail, "mark sanitize guardrail", ) @@ -272,11 +268,15 @@ impl PyPluginContext { priority: i32, callback: Py, ) -> PyResult<()> { - self.register_event_sanitizer( + self.register_callback( name, - priority, - callback, - register_scope_sanitize_start_guardrail, + |qualified_name| { + register_scope_sanitize_start_guardrail( + qualified_name, + priority, + wrap_py_event_sanitize_fn(callback), + ) + }, deregister_scope_sanitize_start_guardrail, "scope start sanitize guardrail", ) @@ -289,11 +289,15 @@ impl PyPluginContext { priority: i32, callback: Py, ) -> PyResult<()> { - self.register_event_sanitizer( + self.register_callback( name, - priority, - callback, - register_scope_sanitize_end_guardrail, + |qualified_name| { + register_scope_sanitize_end_guardrail( + qualified_name, + priority, + wrap_py_event_sanitize_fn(callback), + ) + }, deregister_scope_sanitize_end_guardrail, "scope end sanitize guardrail", ) @@ -304,26 +308,14 @@ impl PyPluginContext { text_signature = "(name: str, callback: object) -> None" )] fn register_subscriber(&self, name: &str, callback: Py) -> PyResult<()> { - let qualified_name = self.qualify_name(name); - register_subscriber(&qualified_name, wrap_py_event_subscriber(callback)) - .map_err(to_py_err)?; - - let name_owned = qualified_name; - let mut guard = self.registrations.lock().map_err(|e| { - pyo3::exceptions::PyRuntimeError::new_err(format!("plugin context lock poisoned: {e}")) - })?; - guard.push(PluginRegistration::new( - "plugin", - name_owned.clone(), - Box::new(move || { - deregister_subscriber(&name_owned).map(|_| ()).map_err(|e| { - PluginError::RegistrationFailed(format!( - "subscriber deregistration failed: {e}" - )) - }) - }), - )); - Ok(()) + self.register_callback( + name, + |qualified_name| { + register_subscriber(qualified_name, wrap_py_event_subscriber(callback)) + }, + deregister_subscriber, + "subscriber", + ) } #[pyo3(signature = (name: "str", priority: "int", callback: "object") -> "None", text_signature = "(name: str, priority: int, callback: object) -> None")] @@ -333,32 +325,18 @@ impl PyPluginContext { priority: i32, callback: Py, ) -> PyResult<()> { - let qualified_name = self.qualify_name(name); - register_tool_sanitize_request_guardrail( - &qualified_name, - priority, - wrap_py_tool_fn(callback), + self.register_callback( + name, + |qualified_name| { + register_tool_sanitize_request_guardrail( + qualified_name, + priority, + wrap_py_tool_fn(callback), + ) + }, + deregister_tool_sanitize_request_guardrail, + "tool sanitize request guardrail", ) - .map_err(to_py_err)?; - - let name_owned = qualified_name; - let mut guard = self.registrations.lock().map_err(|e| { - pyo3::exceptions::PyRuntimeError::new_err(format!("plugin context lock poisoned: {e}")) - })?; - guard.push(PluginRegistration::new( - "plugin", - name_owned.clone(), - Box::new(move || { - deregister_tool_sanitize_request_guardrail(&name_owned) - .map(|_| ()) - .map_err(|e| { - PluginError::RegistrationFailed(format!( - "tool sanitize request guardrail deregistration failed: {e}" - )) - }) - }), - )); - Ok(()) } #[pyo3(signature = (name: "str", priority: "int", callback: "object") -> "None", text_signature = "(name: str, priority: int, callback: object) -> None")] @@ -368,32 +346,18 @@ impl PyPluginContext { priority: i32, callback: Py, ) -> PyResult<()> { - let qualified_name = self.qualify_name(name); - register_tool_sanitize_response_guardrail( - &qualified_name, - priority, - wrap_py_tool_fn(callback), + self.register_callback( + name, + |qualified_name| { + register_tool_sanitize_response_guardrail( + qualified_name, + priority, + wrap_py_tool_fn(callback), + ) + }, + deregister_tool_sanitize_response_guardrail, + "tool sanitize response guardrail", ) - .map_err(to_py_err)?; - - let name_owned = qualified_name; - let mut guard = self.registrations.lock().map_err(|e| { - pyo3::exceptions::PyRuntimeError::new_err(format!("plugin context lock poisoned: {e}")) - })?; - guard.push(PluginRegistration::new( - "plugin", - name_owned.clone(), - Box::new(move || { - deregister_tool_sanitize_response_guardrail(&name_owned) - .map(|_| ()) - .map_err(|e| { - PluginError::RegistrationFailed(format!( - "tool sanitize response guardrail deregistration failed: {e}" - )) - }) - }), - )); - Ok(()) } #[pyo3(signature = (name: "str", priority: "int", callback: "object") -> "None", text_signature = "(name: str, priority: int, callback: object) -> None")] @@ -403,32 +367,18 @@ impl PyPluginContext { priority: i32, callback: Py, ) -> PyResult<()> { - let qualified_name = self.qualify_name(name); - register_tool_conditional_execution_guardrail( - &qualified_name, - priority, - wrap_py_tool_conditional_fn(callback), + self.register_callback( + name, + |qualified_name| { + register_tool_conditional_execution_guardrail( + qualified_name, + priority, + wrap_py_tool_conditional_fn(callback), + ) + }, + deregister_tool_conditional_execution_guardrail, + "tool conditional execution guardrail", ) - .map_err(to_py_err)?; - - let name_owned = qualified_name; - let mut guard = self.registrations.lock().map_err(|e| { - pyo3::exceptions::PyRuntimeError::new_err(format!("plugin context lock poisoned: {e}")) - })?; - guard.push(PluginRegistration::new( - "plugin", - name_owned.clone(), - Box::new(move || { - deregister_tool_conditional_execution_guardrail(&name_owned) - .map(|_| ()) - .map_err(|e| { - PluginError::RegistrationFailed(format!( - "tool conditional execution guardrail deregistration failed: {e}" - )) - }) - }), - )); - Ok(()) } #[pyo3(signature = (name: "str", priority: "int", callback: "object") -> "None", text_signature = "(name: str, priority: int, callback: object) -> None")] @@ -438,32 +388,18 @@ impl PyPluginContext { priority: i32, callback: Py, ) -> PyResult<()> { - let qualified_name = self.qualify_name(name); - register_llm_sanitize_request_guardrail( - &qualified_name, - priority, - wrap_py_llm_sanitize_request_fn(callback), + self.register_callback( + name, + |qualified_name| { + register_llm_sanitize_request_guardrail( + qualified_name, + priority, + wrap_py_llm_sanitize_request_fn(callback), + ) + }, + deregister_llm_sanitize_request_guardrail, + "llm sanitize request guardrail", ) - .map_err(to_py_err)?; - - let name_owned = qualified_name; - let mut guard = self.registrations.lock().map_err(|e| { - pyo3::exceptions::PyRuntimeError::new_err(format!("plugin context lock poisoned: {e}")) - })?; - guard.push(PluginRegistration::new( - "plugin", - name_owned.clone(), - Box::new(move || { - deregister_llm_sanitize_request_guardrail(&name_owned) - .map(|_| ()) - .map_err(|e| { - PluginError::RegistrationFailed(format!( - "llm sanitize request guardrail deregistration failed: {e}" - )) - }) - }), - )); - Ok(()) } #[pyo3(signature = (name: "str", priority: "int", callback: "object") -> "None", text_signature = "(name: str, priority: int, callback: object) -> None")] @@ -473,32 +409,18 @@ impl PyPluginContext { priority: i32, callback: Py, ) -> PyResult<()> { - let qualified_name = self.qualify_name(name); - register_llm_sanitize_response_guardrail( - &qualified_name, - priority, - wrap_py_llm_sanitize_response_fn(callback), + self.register_callback( + name, + |qualified_name| { + register_llm_sanitize_response_guardrail( + qualified_name, + priority, + wrap_py_llm_sanitize_response_fn(callback), + ) + }, + deregister_llm_sanitize_response_guardrail, + "llm sanitize response guardrail", ) - .map_err(to_py_err)?; - - let name_owned = qualified_name; - let mut guard = self.registrations.lock().map_err(|e| { - pyo3::exceptions::PyRuntimeError::new_err(format!("plugin context lock poisoned: {e}")) - })?; - guard.push(PluginRegistration::new( - "plugin", - name_owned.clone(), - Box::new(move || { - deregister_llm_sanitize_response_guardrail(&name_owned) - .map(|_| ()) - .map_err(|e| { - PluginError::RegistrationFailed(format!( - "llm sanitize response guardrail deregistration failed: {e}" - )) - }) - }), - )); - Ok(()) } #[pyo3(signature = (name: "str", priority: "int", callback: "object") -> "None", text_signature = "(name: str, priority: int, callback: object) -> None")] @@ -508,32 +430,18 @@ impl PyPluginContext { priority: i32, callback: Py, ) -> PyResult<()> { - let qualified_name = self.qualify_name(name); - register_llm_conditional_execution_guardrail( - &qualified_name, - priority, - wrap_py_llm_conditional_fn(callback), + self.register_callback( + name, + |qualified_name| { + register_llm_conditional_execution_guardrail( + qualified_name, + priority, + wrap_py_llm_conditional_fn(callback), + ) + }, + deregister_llm_conditional_execution_guardrail, + "llm conditional execution guardrail", ) - .map_err(to_py_err)?; - - let name_owned = qualified_name; - let mut guard = self.registrations.lock().map_err(|e| { - pyo3::exceptions::PyRuntimeError::new_err(format!("plugin context lock poisoned: {e}")) - })?; - guard.push(PluginRegistration::new( - "plugin", - name_owned.clone(), - Box::new(move || { - deregister_llm_conditional_execution_guardrail(&name_owned) - .map(|_| ()) - .map_err(|e| { - PluginError::RegistrationFailed(format!( - "llm conditional execution guardrail deregistration failed: {e}" - )) - }) - }), - )); - Ok(()) } #[pyo3(signature = ( @@ -549,33 +457,19 @@ impl PyPluginContext { break_chain: bool, callback: Py, ) -> PyResult<()> { - let qualified_name = self.qualify_name(name); - register_llm_request_intercept( - &qualified_name, - priority, - break_chain, - wrap_py_llm_request_intercept_fn(callback), + self.register_callback( + name, + |qualified_name| { + register_llm_request_intercept( + qualified_name, + priority, + break_chain, + wrap_py_llm_request_intercept_fn(callback), + ) + }, + deregister_llm_request_intercept, + "llm request intercept", ) - .map_err(to_py_err)?; - - let name_owned = qualified_name; - let mut guard = self.registrations.lock().map_err(|e| { - pyo3::exceptions::PyRuntimeError::new_err(format!("plugin context lock poisoned: {e}")) - })?; - guard.push(PluginRegistration::new( - "plugin", - name_owned.clone(), - Box::new(move || { - deregister_llm_request_intercept(&name_owned) - .map(|_| ()) - .map_err(|e| { - PluginError::RegistrationFailed(format!( - "llm request intercept deregistration failed: {e}" - )) - }) - }), - )); - Ok(()) } #[pyo3(signature = (name: "str", priority: "int", callback: "object") -> "None", text_signature = "(name: str, priority: int, callback: object) -> None")] @@ -585,32 +479,18 @@ impl PyPluginContext { priority: i32, callback: Py, ) -> PyResult<()> { - let qualified_name = self.qualify_name(name); - register_llm_execution_intercept( - &qualified_name, - priority, - wrap_py_llm_exec_intercept_fn(callback), + self.register_callback( + name, + |qualified_name| { + register_llm_execution_intercept( + qualified_name, + priority, + wrap_py_llm_exec_intercept_fn(callback), + ) + }, + deregister_llm_execution_intercept, + "llm execution intercept", ) - .map_err(to_py_err)?; - - let name_owned = qualified_name; - let mut guard = self.registrations.lock().map_err(|e| { - pyo3::exceptions::PyRuntimeError::new_err(format!("plugin context lock poisoned: {e}")) - })?; - guard.push(PluginRegistration::new( - "plugin", - name_owned.clone(), - Box::new(move || { - deregister_llm_execution_intercept(&name_owned) - .map(|_| ()) - .map_err(|e| { - PluginError::RegistrationFailed(format!( - "llm execution intercept deregistration failed: {e}" - )) - }) - }), - )); - Ok(()) } #[pyo3(signature = (name: "str", priority: "int", callback: "object") -> "None", text_signature = "(name: str, priority: int, callback: object) -> None")] @@ -620,32 +500,18 @@ impl PyPluginContext { priority: i32, callback: Py, ) -> PyResult<()> { - let qualified_name = self.qualify_name(name); - register_llm_stream_execution_intercept( - &qualified_name, - priority, - wrap_py_llm_stream_exec_intercept_fn(callback), + self.register_callback( + name, + |qualified_name| { + register_llm_stream_execution_intercept( + qualified_name, + priority, + wrap_py_llm_stream_exec_intercept_fn(callback), + ) + }, + deregister_llm_stream_execution_intercept, + "llm stream execution intercept", ) - .map_err(to_py_err)?; - - let name_owned = qualified_name; - let mut guard = self.registrations.lock().map_err(|e| { - pyo3::exceptions::PyRuntimeError::new_err(format!("plugin context lock poisoned: {e}")) - })?; - guard.push(PluginRegistration::new( - "plugin", - name_owned.clone(), - Box::new(move || { - deregister_llm_stream_execution_intercept(&name_owned) - .map(|_| ()) - .map_err(|e| { - PluginError::RegistrationFailed(format!( - "llm stream execution intercept deregistration failed: {e}" - )) - }) - }), - )); - Ok(()) } #[pyo3(signature = ( @@ -661,33 +527,19 @@ impl PyPluginContext { break_chain: bool, callback: Py, ) -> PyResult<()> { - let qualified_name = self.qualify_name(name); - register_tool_request_intercept( - &qualified_name, - priority, - break_chain, - wrap_py_tool_request_intercept_fn(callback), + self.register_callback( + name, + |qualified_name| { + register_tool_request_intercept( + qualified_name, + priority, + break_chain, + wrap_py_tool_request_intercept_fn(callback), + ) + }, + deregister_tool_request_intercept, + "tool request intercept", ) - .map_err(to_py_err)?; - - let name_owned = qualified_name; - let mut guard = self.registrations.lock().map_err(|e| { - pyo3::exceptions::PyRuntimeError::new_err(format!("plugin context lock poisoned: {e}")) - })?; - guard.push(PluginRegistration::new( - "plugin", - name_owned.clone(), - Box::new(move || { - deregister_tool_request_intercept(&name_owned) - .map(|_| ()) - .map_err(|e| { - PluginError::RegistrationFailed(format!( - "tool request intercept deregistration failed: {e}" - )) - }) - }), - )); - Ok(()) } #[pyo3(signature = (name: "str", priority: "int", callback: "object") -> "None", text_signature = "(name: str, priority: int, callback: object) -> None")] @@ -697,32 +549,18 @@ impl PyPluginContext { priority: i32, callback: Py, ) -> PyResult<()> { - let qualified_name = self.qualify_name(name); - register_tool_execution_intercept( - &qualified_name, - priority, - wrap_py_tool_exec_intercept_fn(callback), + self.register_callback( + name, + |qualified_name| { + register_tool_execution_intercept( + qualified_name, + priority, + wrap_py_tool_exec_intercept_fn(callback), + ) + }, + deregister_tool_execution_intercept, + "tool execution intercept", ) - .map_err(to_py_err)?; - - let name_owned = qualified_name; - let mut guard = self.registrations.lock().map_err(|e| { - pyo3::exceptions::PyRuntimeError::new_err(format!("plugin context lock poisoned: {e}")) - })?; - guard.push(PluginRegistration::new( - "plugin", - name_owned.clone(), - Box::new(move || { - deregister_tool_execution_intercept(&name_owned) - .map(|_| ()) - .map_err(|e| { - PluginError::RegistrationFailed(format!( - "tool execution intercept deregistration failed: {e}" - )) - }) - }), - )); - Ok(()) } fn __repr__(&self) -> String { diff --git a/python/plugin/src/nemo_relay_plugin/_api.py b/python/plugin/src/nemo_relay_plugin/_api.py index a5471c37f..15c94b027 100644 --- a/python/plugin/src/nemo_relay_plugin/_api.py +++ b/python/plugin/src/nemo_relay_plugin/_api.py @@ -432,7 +432,9 @@ def register(self, ctx: PluginContext, config: Json) -> None | Awaitable[None]: class _Handlers: registrations: list[Any] subscribers: dict[str, SubscriberCallback] - event_sanitizers: dict[str, EventSanitizeCallback] + mark_sanitizers: dict[str, EventSanitizeCallback] + scope_start_sanitizers: dict[str, EventSanitizeCallback] + scope_end_sanitizers: dict[str, EventSanitizeCallback] tool_sanitize_requests: dict[str, ToolSanitizeCallback] tool_sanitize_responses: dict[str, ToolSanitizeCallback] tool_conditionals: dict[str, ToolConditionalCallback] @@ -450,7 +452,9 @@ def empty(cls) -> _Handlers: return cls( registrations=[], subscribers={}, - event_sanitizers={}, + mark_sanitizers={}, + scope_start_sanitizers={}, + scope_end_sanitizers={}, tool_sanitize_requests={}, tool_sanitize_responses={}, tool_conditionals={}, @@ -521,27 +525,76 @@ def _register_event_sanitizer( callback: EventSanitizeCallback, surface: int, priority: int, + handlers: dict[str, EventSanitizeCallback], ) -> None: self._push_registration(name, surface, priority, False) - self._handlers.event_sanitizers[name] = callback + handlers[name] = callback def register_mark_sanitize_guardrail( self, name: str, callback: EventSanitizeCallback, *, priority: int = 0 ) -> None: - """Register a sanitizer for mark event observability fields.""" - self._register_event_sanitizer(name, callback, pb.MARK_SANITIZE_GUARDRAIL, priority) + """Register a sanitizer for mark event observability fields. + + Args: + name: Component-local registration name. + callback: Function receiving the event and its observability fields. + It can return sanitized fields directly or through an awaitable. + priority: Execution order. Lower values run first. + + Callback errors: + An exception becomes a structured worker invocation error. + """ + self._register_event_sanitizer( + name, + callback, + pb.MARK_SANITIZE_GUARDRAIL, + priority, + self._handlers.mark_sanitizers, + ) def register_scope_sanitize_start_guardrail( self, name: str, callback: EventSanitizeCallback, *, priority: int = 0 ) -> None: - """Register a sanitizer for scope start event observability fields.""" - self._register_event_sanitizer(name, callback, pb.SCOPE_SANITIZE_START_GUARDRAIL, priority) + """Register a sanitizer for scope start event observability fields. + + Args: + name: Component-local registration name. + callback: Function receiving the event and its observability fields. + It can return sanitized fields directly or through an awaitable. + priority: Execution order. Lower values run first. + + Callback errors: + An exception becomes a structured worker invocation error. + """ + self._register_event_sanitizer( + name, + callback, + pb.SCOPE_SANITIZE_START_GUARDRAIL, + priority, + self._handlers.scope_start_sanitizers, + ) def register_scope_sanitize_end_guardrail( self, name: str, callback: EventSanitizeCallback, *, priority: int = 0 ) -> None: - """Register a sanitizer for scope end event observability fields.""" - self._register_event_sanitizer(name, callback, pb.SCOPE_SANITIZE_END_GUARDRAIL, priority) + """Register a sanitizer for scope end event observability fields. + + Args: + name: Component-local registration name. + callback: Function receiving the event and its observability fields. + It can return sanitized fields directly or through an awaitable. + priority: Execution order. Lower values run first. + + Callback errors: + An exception becomes a structured worker invocation error. + """ + self._register_event_sanitizer( + name, + callback, + pb.SCOPE_SANITIZE_END_GUARDRAIL, + priority, + self._handlers.scope_end_sanitizers, + ) def register_tool_sanitize_request_guardrail( self, @@ -1466,22 +1519,20 @@ async def _invoke_result(self, request: Any) -> Any: event = _decode_required_envelope(request.event, "event", EVENT_SCHEMA) await _maybe_await(self._handler(self._handlers.subscribers, request.registration_name)(event)) return pb.InvokeResponse(empty=pb.EmptyResult()) - if request.surface in { - pb.MARK_SANITIZE_GUARDRAIL, - pb.SCOPE_SANITIZE_START_GUARDRAIL, - pb.SCOPE_SANITIZE_END_GUARDRAIL, - }: + event_sanitizer_handlers = { + pb.MARK_SANITIZE_GUARDRAIL: self._handlers.mark_sanitizers, + pb.SCOPE_SANITIZE_START_GUARDRAIL: self._handlers.scope_start_sanitizers, + pb.SCOPE_SANITIZE_END_GUARDRAIL: self._handlers.scope_end_sanitizers, + } + if request.surface in event_sanitizer_handlers: event = _decode_required_envelope(request.event, "event", EVENT_SCHEMA) fields: EventSanitizeFields = { "data": event.get("data"), "category_profile": event.get("category_profile"), "metadata": event.get("metadata"), } - return _json_response( - await _maybe_await( - self._handler(self._handlers.event_sanitizers, request.registration_name)(event, fields) - ) - ) + handler = self._handler(event_sanitizer_handlers[request.surface], request.registration_name) + return _json_response(await _maybe_await(handler(event, fields))) if request.surface == pb.TOOL_SANITIZE_REQUEST_GUARDRAIL: return _json_response( await _maybe_await( diff --git a/python/tests/plugin/test_worker_sdk.py b/python/tests/plugin/test_worker_sdk.py index da37a5662..279874d06 100644 --- a/python/tests/plugin/test_worker_sdk.py +++ b/python/tests/plugin/test_worker_sdk.py @@ -184,8 +184,14 @@ def register(self, ctx: PluginContext, config: Json) -> None: async def subscriber(event: Json) -> None: await ctx.runtime.emit_mark("tests.subscriber", event) - async def event_sanitize(event: Json, fields: Json) -> Json: - return {**fields, "data": {"sanitized": event["name"]}, "metadata": None} + async def mark_sanitize(event: Json, fields: Json) -> Json: + return {**fields, "data": {"sanitized": f"mark:{event['name']}"}, "metadata": None} + + async def scope_start_sanitize(event: Json, fields: Json) -> Json: + return {**fields, "data": {"sanitized": f"scope-start:{event['name']}"}, "metadata": None} + + async def scope_end_sanitize(event: Json, fields: Json) -> Json: + return {**fields, "data": {"sanitized": f"scope-end:{event['name']}"}, "metadata": None} def tool_sanitize(name: str, value: Json) -> Json: return _tag(value, f"sanitize_{name}") @@ -232,9 +238,9 @@ async def llm_stream_execution(name: str, request: Json, next_call: Any) -> Asyn yield _tag(chunk, "llm_stream_execution") ctx.register_subscriber("subscriber", subscriber) - ctx.register_mark_sanitize_guardrail("mark_sanitize", event_sanitize, priority=1) - ctx.register_scope_sanitize_start_guardrail("scope_start_sanitize", event_sanitize, priority=2) - ctx.register_scope_sanitize_end_guardrail("scope_end_sanitize", event_sanitize, priority=3) + ctx.register_mark_sanitize_guardrail("event_sanitize", mark_sanitize, priority=1) + ctx.register_scope_sanitize_start_guardrail("event_sanitize", scope_start_sanitize, priority=2) + ctx.register_scope_sanitize_end_guardrail("scope_end_sanitize", scope_end_sanitize, priority=3) ctx.register_tool_sanitize_request_guardrail("tool_sanitize", tool_sanitize, priority=1) ctx.register_tool_sanitize_response_guardrail("tool_sanitize", tool_sanitize, priority=2) ctx.register_tool_conditional_execution_guardrail("tool_conditional", tool_block, priority=3) @@ -317,8 +323,8 @@ async def test_health_handshake_validate_register_and_all_surfaces(service: _Wor ] assert registrations == [ ("subscriber", pb.SUBSCRIBER, 0, False), - ("mark_sanitize", pb.MARK_SANITIZE_GUARDRAIL, 1, False), - ("scope_start_sanitize", pb.SCOPE_SANITIZE_START_GUARDRAIL, 2, False), + ("event_sanitize", pb.MARK_SANITIZE_GUARDRAIL, 1, False), + ("event_sanitize", pb.SCOPE_SANITIZE_START_GUARDRAIL, 2, False), ("scope_end_sanitize", pb.SCOPE_SANITIZE_END_GUARDRAIL, 3, False), ("tool_sanitize", pb.TOOL_SANITIZE_REQUEST_GUARDRAIL, 1, False), ("tool_sanitize", pb.TOOL_SANITIZE_RESPONSE_GUARDRAIL, 2, False), @@ -615,10 +621,15 @@ async def test_event_sanitizer_surfaces_receive_context_and_return_all_fields( ) -> None: await _register(service) registration = { - pb.MARK_SANITIZE_GUARDRAIL: "mark_sanitize", - pb.SCOPE_SANITIZE_START_GUARDRAIL: "scope_start_sanitize", + pb.MARK_SANITIZE_GUARDRAIL: "event_sanitize", + pb.SCOPE_SANITIZE_START_GUARDRAIL: "event_sanitize", pb.SCOPE_SANITIZE_END_GUARDRAIL: "scope_end_sanitize", }[surface] + sanitizer = { + pb.MARK_SANITIZE_GUARDRAIL: "mark", + pb.SCOPE_SANITIZE_START_GUARDRAIL: "scope-start", + pb.SCOPE_SANITIZE_END_GUARDRAIL: "scope-end", + }[surface] response = await service.Invoke( _invoke_request( registration, @@ -637,7 +648,7 @@ async def test_event_sanitizer_surfaces_receive_context_and_return_all_fields( AbortContext(), ) assert _envelope_value(response.json.value) == { - "data": {"sanitized": "worker-event"}, + "data": {"sanitized": f"{sanitizer}:worker-event"}, "category_profile": {"subtype": "test"}, "metadata": None, } diff --git a/python/tests/test_event_sanitizers.py b/python/tests/test_event_sanitizers.py index 2949a48b5..b104c911b 100644 --- a/python/tests/test_event_sanitizers.py +++ b/python/tests/test_event_sanitizers.py @@ -3,6 +3,7 @@ from __future__ import annotations +from collections.abc import Iterator from typing import cast import pytest @@ -11,15 +12,17 @@ from nemo_relay import EventSanitizeFields, guardrails, plugin, scope, scope_local, subscribers -def _capture_events(): - events = [] +@pytest.fixture(name="capture_events") +def capture_events_fixture() -> Iterator[tuple[str, list[nemo_relay.Event]]]: + events: list[nemo_relay.Event] = [] name = "test-event-sanitizer-capture" subscribers.register(name, events.append) - return name, events + yield name, events + subscribers.deregister(name) -def test_global_mark_sanitizers_order_convert_fields_and_remove_values() -> None: - capture_name, events = _capture_events() +def test_global_mark_sanitizers_order_convert_fields_and_remove_values(capture_events): + _capture_name, events = capture_events calls: list[tuple[str, object]] = [] def first(event: nemo_relay.Event, fields: EventSanitizeFields) -> EventSanitizeFields: @@ -46,7 +49,6 @@ def second(event: nemo_relay.Event, fields: EventSanitizeFields) -> EventSanitiz finally: guardrails.deregister_mark_sanitize("python-mark-first") guardrails.deregister_mark_sanitize("python-mark-second") - subscribers.deregister(capture_name) mark = events[-1] assert mark.data == {"stage": "second"} @@ -54,8 +56,8 @@ def second(event: nemo_relay.Event, fields: EventSanitizeFields) -> EventSanitiz assert calls == [("checkpoint", {"secret": "raw"}), ("mark", {"stage": "first"})] -def test_invalid_mark_sanitizer_result_fails_open() -> None: - capture_name, events = _capture_events() +def test_invalid_mark_sanitizer_result_fails_open(capture_events): + _capture_name, events = capture_events guardrails.register_mark_sanitize( "python-mark-invalid", 0, @@ -66,13 +68,12 @@ def test_invalid_mark_sanitizer_result_fails_open() -> None: subscribers.flush() finally: guardrails.deregister_mark_sanitize("python-mark-invalid") - subscribers.deregister(capture_name) assert events[-1].data == {"kept": True} -def test_scope_start_and_end_sanitizers_cover_category_profile() -> None: - capture_name, events = _capture_events() +def test_scope_start_and_end_sanitizers_cover_category_profile(capture_events): + _capture_name, events = capture_events def sanitize(_event: nemo_relay.Event, fields: EventSanitizeFields) -> EventSanitizeFields: profile = dict(fields["category_profile"] or {}) @@ -94,7 +95,6 @@ def sanitize(_event: nemo_relay.Event, fields: EventSanitizeFields) -> EventSani finally: guardrails.deregister_scope_sanitize_start("python-scope-start") guardrails.deregister_scope_sanitize_end("python-scope-end") - subscribers.deregister(capture_name) lifecycle = [event for event in events if event.name == "generic"] assert len(lifecycle) == 2 @@ -103,8 +103,8 @@ def sanitize(_event: nemo_relay.Event, fields: EventSanitizeFields) -> EventSani assert all(event.category_profile["subtype"] == "sanitized" for event in lifecycle) -def test_scope_local_event_sanitizers_are_inherited_and_cleaned_up() -> None: - capture_name, events = _capture_events() +def test_scope_local_event_sanitizers_are_inherited_and_cleaned_up(capture_events): + _capture_name, events = capture_events def sanitize(_event: nemo_relay.Event, fields: EventSanitizeFields) -> EventSanitizeFields: return { @@ -114,15 +114,18 @@ def sanitize(_event: nemo_relay.Event, fields: EventSanitizeFields) -> EventSani } owner = scope.push("owner", nemo_relay.ScopeType.Agent) - scope_local.register_mark_sanitize(owner, "python-local-mark", 0, sanitize) - scope.event("inside", data={"raw": True}) - child = scope.push("child", nemo_relay.ScopeType.Function) - scope.event("inherited", data={"raw": True}) - scope.pop(child) - scope.pop(owner) + try: + scope_local.register_mark_sanitize(owner, "python-local-mark", 0, sanitize) + scope.event("inside", data={"raw": True}) + child = scope.push("child", nemo_relay.ScopeType.Function) + try: + scope.event("inherited", data={"raw": True}) + finally: + scope.pop(child) + finally: + scope.pop(owner) scope.event("outside", data={"raw": True}) subscribers.flush() - subscribers.deregister(capture_name) marks = {event.name: event for event in events if event.kind == "mark"} assert marks["inside"].data == {"scope_local": True} @@ -130,7 +133,7 @@ def sanitize(_event: nemo_relay.Event, fields: EventSanitizeFields) -> EventSani assert marks["outside"].data == {"raw": True} -async def test_in_process_plugin_event_sanitizers_are_removed_on_clear() -> None: +async def test_in_process_plugin_event_sanitizers_are_removed_on_clear(capture_events): class EventPlugin: def validate(self, _config): return None @@ -146,7 +149,7 @@ def sanitize(_event: nemo_relay.Event, fields: EventSanitizeFields) -> EventSani context.register_mark_sanitize_guardrail("mark", 0, sanitize) kind = "python.test_event_sanitizer" - capture_name, events = _capture_events() + _capture_name, events = capture_events plugin.register(kind, cast(plugin.Plugin, EventPlugin())) try: await plugin.initialize(plugin.PluginConfig(components=[plugin.ComponentSpec(kind=kind)])) @@ -158,14 +161,13 @@ def sanitize(_event: nemo_relay.Event, fields: EventSanitizeFields) -> EventSani finally: plugin.clear() plugin.deregister(kind) - subscribers.deregister(capture_name) marks = {event.name: event for event in events if event.kind == "mark"} assert marks["configured"].data == {"plugin": True} assert marks["cleared"].data == {"raw": True} -async def test_in_process_plugin_rolls_back_event_sanitizer_when_registration_fails() -> None: +async def test_in_process_plugin_rolls_back_event_sanitizer_when_registration_fails(capture_events): class FailingPlugin: def validate(self, _config): return None @@ -183,15 +185,12 @@ def sanitize(_event: nemo_relay.Event, fields: EventSanitizeFields) -> EventSani kind = "python.test_event_sanitizer_rollback" plugin.register(kind, cast(plugin.Plugin, FailingPlugin())) + _capture_name, events = capture_events try: with pytest.raises(RuntimeError, match="registration failed"): await plugin.initialize(plugin.PluginConfig(components=[plugin.ComponentSpec(kind=kind)])) - capture_name, events = _capture_events() - try: - scope.event("after-failure", data={"raw": True}) - subscribers.flush() - finally: - subscribers.deregister(capture_name) + scope.event("after-failure", data={"raw": True}) + subscribers.flush() assert events[-1].data == {"raw": True} finally: plugin.clear() From c7c3e5e92f2e6b42ac5b65a84a7c5891a58e07e8 Mon Sep 17 00:00:00 2001 From: Will Killian Date: Thu, 9 Jul 2026 10:58:08 -0400 Subject: [PATCH 3/3] refactor: centralize event sanitizer handlers Signed-off-by: Will Killian --- python/plugin/src/nemo_relay_plugin/_api.py | 29 ++++++++++++--------- 1 file changed, 17 insertions(+), 12 deletions(-) diff --git a/python/plugin/src/nemo_relay_plugin/_api.py b/python/plugin/src/nemo_relay_plugin/_api.py index 15c94b027..de730eaf8 100644 --- a/python/plugin/src/nemo_relay_plugin/_api.py +++ b/python/plugin/src/nemo_relay_plugin/_api.py @@ -69,7 +69,7 @@ from enum import Enum from importlib import metadata from pathlib import Path -from typing import Any, Protocol, TypeAlias, TypedDict +from typing import Any, ClassVar, Protocol, TypeAlias, TypedDict from urllib.parse import urlsplit grpc: Any = importlib.import_module("grpc") @@ -487,6 +487,12 @@ class PluginContext: context without host access, primarily for tests. """ + _EVENT_SANITIZER_HANDLER_ATTRIBUTES: ClassVar[dict[int, str]] = { + pb.MARK_SANITIZE_GUARDRAIL: "mark_sanitizers", + pb.SCOPE_SANITIZE_START_GUARDRAIL: "scope_start_sanitizers", + pb.SCOPE_SANITIZE_END_GUARDRAIL: "scope_end_sanitizers", + } + def __init__(self, runtime: PluginRuntime | None = None) -> None: self._runtime = runtime self._handlers = _Handlers.empty() @@ -525,9 +531,12 @@ def _register_event_sanitizer( callback: EventSanitizeCallback, surface: int, priority: int, - handlers: dict[str, EventSanitizeCallback], ) -> None: self._push_registration(name, surface, priority, False) + handlers: dict[str, EventSanitizeCallback] = getattr( + self._handlers, + self._EVENT_SANITIZER_HANDLER_ATTRIBUTES[surface], + ) handlers[name] = callback def register_mark_sanitize_guardrail( @@ -549,7 +558,6 @@ def register_mark_sanitize_guardrail( callback, pb.MARK_SANITIZE_GUARDRAIL, priority, - self._handlers.mark_sanitizers, ) def register_scope_sanitize_start_guardrail( @@ -571,7 +579,6 @@ def register_scope_sanitize_start_guardrail( callback, pb.SCOPE_SANITIZE_START_GUARDRAIL, priority, - self._handlers.scope_start_sanitizers, ) def register_scope_sanitize_end_guardrail( @@ -593,7 +600,6 @@ def register_scope_sanitize_end_guardrail( callback, pb.SCOPE_SANITIZE_END_GUARDRAIL, priority, - self._handlers.scope_end_sanitizers, ) def register_tool_sanitize_request_guardrail( @@ -1519,19 +1525,18 @@ async def _invoke_result(self, request: Any) -> Any: event = _decode_required_envelope(request.event, "event", EVENT_SCHEMA) await _maybe_await(self._handler(self._handlers.subscribers, request.registration_name)(event)) return pb.InvokeResponse(empty=pb.EmptyResult()) - event_sanitizer_handlers = { - pb.MARK_SANITIZE_GUARDRAIL: self._handlers.mark_sanitizers, - pb.SCOPE_SANITIZE_START_GUARDRAIL: self._handlers.scope_start_sanitizers, - pb.SCOPE_SANITIZE_END_GUARDRAIL: self._handlers.scope_end_sanitizers, - } - if request.surface in event_sanitizer_handlers: + if request.surface in PluginContext._EVENT_SANITIZER_HANDLER_ATTRIBUTES: event = _decode_required_envelope(request.event, "event", EVENT_SCHEMA) fields: EventSanitizeFields = { "data": event.get("data"), "category_profile": event.get("category_profile"), "metadata": event.get("metadata"), } - handler = self._handler(event_sanitizer_handlers[request.surface], request.registration_name) + handlers = getattr( + self._handlers, + PluginContext._EVENT_SANITIZER_HANDLER_ATTRIBUTES[request.surface], + ) + handler = self._handler(handlers, request.registration_name) return _json_response(await _maybe_await(handler(event, fields))) if request.surface == pb.TOOL_SANITIZE_REQUEST_GUARDRAIL: return _json_response(