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..3178471dc 100644 --- a/crates/python/src/py_plugin.rs +++ b/crates/python/src/py_plugin.rs @@ -16,16 +16,20 @@ 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::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 +39,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,18 +210,16 @@ impl PyPluginContext { fn qualify_name(&self, name: &str) -> String { format!("{}{}", self.namespace_prefix, name) } -} -#[pymethods] -impl PyPluginContext { - #[pyo3( - signature = (name: "str", callback: "object") -> "None", - text_signature = "(name: str, callback: object) -> None" - )] - fn register_subscriber(&self, name: &str, callback: Py) -> PyResult<()> { + fn register_callback( + &self, + name: &str, + register: impl FnOnce(&str) -> FlowResult<()>, + deregister: fn(&str) -> FlowResult, + label: &'static str, + ) -> PyResult<()> { let qualified_name = self.qualify_name(name); - register_subscriber(&qualified_name, wrap_py_event_subscriber(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| { @@ -227,154 +229,177 @@ impl PyPluginContext { "plugin", name_owned.clone(), Box::new(move || { - deregister_subscriber(&name_owned).map(|_| ()).map_err(|e| { - PluginError::RegistrationFailed(format!( - "subscriber deregistration failed: {e}" - )) + 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_tool_sanitize_request_guardrail( + fn register_mark_sanitize_guardrail( &self, name: &str, 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_mark_sanitize_guardrail( + qualified_name, + priority, + wrap_py_event_sanitize_fn(callback), + ) + }, + deregister_mark_sanitize_guardrail, + "mark sanitize 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")] + fn register_scope_sanitize_start_guardrail( + &self, + name: &str, + priority: i32, + callback: Py, + ) -> PyResult<()> { + self.register_callback( + name, + |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", + ) } #[pyo3(signature = (name: "str", priority: "int", callback: "object") -> "None", text_signature = "(name: str, priority: int, callback: object) -> None")] - fn register_tool_sanitize_response_guardrail( + fn register_scope_sanitize_end_guardrail( &self, name: &str, 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_scope_sanitize_end_guardrail( + qualified_name, + priority, + wrap_py_event_sanitize_fn(callback), + ) + }, + deregister_scope_sanitize_end_guardrail, + "scope end sanitize 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", callback: "object") -> "None", + text_signature = "(name: str, callback: object) -> None" + )] + fn register_subscriber(&self, name: &str, callback: Py) -> PyResult<()> { + 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")] - fn register_tool_conditional_execution_guardrail( + fn register_tool_sanitize_request_guardrail( &self, name: &str, 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_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_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")] + fn register_tool_sanitize_response_guardrail( + &self, + name: &str, + priority: i32, + callback: Py, + ) -> PyResult<()> { + 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", + ) } #[pyo3(signature = (name: "str", priority: "int", callback: "object") -> "None", text_signature = "(name: str, priority: int, callback: object) -> None")] - fn register_llm_sanitize_request_guardrail( + fn register_tool_conditional_execution_guardrail( &self, name: &str, 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_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_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")] + fn register_llm_sanitize_request_guardrail( + &self, + name: &str, + priority: i32, + callback: Py, + ) -> PyResult<()> { + 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", + ) } #[pyo3(signature = (name: "str", priority: "int", callback: "object") -> "None", text_signature = "(name: str, priority: int, callback: object) -> None")] @@ -384,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")] @@ -419,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 = ( @@ -460,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")] @@ -496,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")] @@ -531,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 = ( @@ -572,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")] @@ -608,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/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 90a725b90..400307416 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[ @@ -1189,7 +1196,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 cf44e8e6e..f3052e42a 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. @@ -38,6 +39,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. @@ -66,6 +68,8 @@ ConfigDiagnostic, DiagnosticLevel, Event, + EventSanitizeCallback, + EventSanitizeFields, Json, LlmConditionalCallback, LlmExecutionCallback, @@ -105,6 +109,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 99d89b390..b688addc8 100644 --- a/python/plugin/src/nemo_relay_plugin/_api.py +++ b/python/plugin/src/nemo_relay_plugin/_api.py @@ -76,7 +76,7 @@ from enum import Enum from importlib import metadata from pathlib import Path -from typing import Any, Protocol, TypeAlias +from typing import Any, ClassVar, Protocol, TypeAlias, TypedDict from urllib.parse import urlsplit grpc: Any = importlib.import_module("grpc") @@ -87,6 +87,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. @@ -709,6 +719,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]] @@ -734,6 +748,9 @@ def register(self, ctx: PluginContext, config: Json) -> None | Awaitable[None]: class _Handlers: registrations: list[Any] subscribers: dict[str, SubscriberCallback] + 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] @@ -751,6 +768,9 @@ def empty(cls) -> _Handlers: return cls( registrations=[], subscribers={}, + mark_sanitizers={}, + scope_start_sanitizers={}, + scope_end_sanitizers={}, tool_sanitize_requests={}, tool_sanitize_responses={}, tool_conditionals={}, @@ -783,6 +803,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() @@ -815,6 +841,83 @@ 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) + handlers: dict[str, EventSanitizeCallback] = getattr( + self._handlers, + self._EVENT_SANITIZER_HANDLER_ATTRIBUTES[surface], + ) + 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. + + 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, + ) + + def register_scope_sanitize_start_guardrail( + self, name: str, callback: EventSanitizeCallback, *, priority: int = 0 + ) -> None: + """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, + ) + + def register_scope_sanitize_end_guardrail( + self, name: str, callback: EventSanitizeCallback, *, priority: int = 0 + ) -> None: + """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, + ) + def register_tool_sanitize_request_guardrail( self, name: str, @@ -1738,6 +1841,19 @@ 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 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"), + } + 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( await _maybe_await( @@ -1900,6 +2016,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 bdb70028a..7ecafbeea 100644 --- a/python/tests/plugin/test_worker_sdk.py +++ b/python/tests/plugin/test_worker_sdk.py @@ -275,6 +275,15 @@ def register(self, ctx: PluginContext, config: Json) -> None: async def subscriber(event: Json) -> None: await ctx.runtime.emit_mark("tests.subscriber", event) + 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}") @@ -320,6 +329,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("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) @@ -361,6 +373,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 @@ -399,6 +414,9 @@ async def test_health_handshake_validate_register_and_all_surfaces(service: _Wor ] assert registrations == [ ("subscriber", pb.SUBSCRIBER, 0, 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), ("tool_conditional", pb.TOOL_CONDITIONAL_EXECUTION_GUARDRAIL, 3, False), @@ -681,6 +699,52 @@ 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: "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, + 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": f"{sanitizer}: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" @@ -2453,6 +2517,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..b104c911b --- /dev/null +++ b/python/tests/test_event_sanitizers.py @@ -0,0 +1,197 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from collections.abc import Iterator +from typing import cast + +import pytest + +import nemo_relay +from nemo_relay import EventSanitizeFields, guardrails, plugin, scope, scope_local, subscribers + + +@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) + yield name, events + subscribers.deregister(name) + + +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: + 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") + + 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(capture_events): + _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") + + assert events[-1].data == {"kept": True} + + +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 {}) + 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") + + 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(capture_events): + _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) + 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() + + 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(capture_events): + 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) + + 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(capture_events): + 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())) + _capture_name, events = capture_events + try: + with pytest.raises(RuntimeError, match="registration failed"): + await plugin.initialize(plugin.PluginConfig(components=[plugin.ComponentSpec(kind=kind)])) + scope.event("after-failure", data={"raw": True}) + subscribers.flush() + 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(