diff --git a/crates/core/src/api/runtime/state.rs b/crates/core/src/api/runtime/state.rs index e5878e35e..bb0aa6a16 100644 --- a/crates/core/src/api/runtime/state.rs +++ b/crates/core/src/api/runtime/state.rs @@ -1785,19 +1785,38 @@ fn validate_event_metadata_attributes(attributes: &BTreeMap) -> Re Ok(()) } -fn is_otel_compatible_attribute_number(value: &serde_json::Number) -> bool { +#[derive(Clone, Copy, PartialEq, Eq)] +enum OtelAttributePrimitiveKind { + Boolean, + Integer, + Double, + String, +} + +fn otel_compatible_attribute_number_kind( + value: &serde_json::Number, +) -> Option { if let Some(value) = value.as_u64() { - return i64::try_from(value).is_ok(); + return i64::try_from(value) + .is_ok() + .then_some(OtelAttributePrimitiveKind::Integer); + } + if value.as_i64().is_some() { + return Some(OtelAttributePrimitiveKind::Integer); } - value.as_i64().is_some() || value.as_f64().is_some() + value.as_f64().map(|_| OtelAttributePrimitiveKind::Double) +} + +fn is_otel_compatible_attribute_number(value: &serde_json::Number) -> bool { + otel_compatible_attribute_number_kind(value).is_some() } fn is_otel_compatible_attribute_value(value: &Json) -> bool { - fn primitive_kind(value: &Json) -> Option { + fn primitive_kind(value: &Json) -> Option { match value { - Json::Bool(_) => Some(0), - Json::Number(value) if is_otel_compatible_attribute_number(value) => Some(1), - Json::String(_) => Some(2), + Json::Bool(_) => Some(OtelAttributePrimitiveKind::Boolean), + Json::Number(value) => otel_compatible_attribute_number_kind(value), + Json::String(_) => Some(OtelAttributePrimitiveKind::String), _ => None, } } diff --git a/crates/core/tests/unit/runtime_state_tests.rs b/crates/core/tests/unit/runtime_state_tests.rs index a304d339e..08403222a 100644 --- a/crates/core/tests/unit/runtime_state_tests.rs +++ b/crates/core/tests/unit/runtime_state_tests.rs @@ -39,7 +39,8 @@ async fn event_metadata_injection_accepts_flat_otel_values_and_empty_output() { ), ("nv.test.strings".into(), json!(["a", "b"])), ("nv.test.booleans".into(), json!([true, false])), - ("nv.test.numbers".into(), json!([1, 2])), + ("nv.test.integers".into(), json!([1, 2])), + ("nv.test.doubles".into(), json!([1.0, 2.5])), ("nv.test.empty".into(), json!([])), ])) }) @@ -68,7 +69,8 @@ async fn event_metadata_injection_accepts_flat_otel_values_and_empty_output() { ); assert_eq!(metadata["nv.test.strings"], json!(["a", "b"])); assert_eq!(metadata["nv.test.booleans"], json!([true, false])); - assert_eq!(metadata["nv.test.numbers"], json!([1, 2])); + assert_eq!(metadata["nv.test.integers"], json!([1, 2])); + assert_eq!(metadata["nv.test.doubles"], json!([1.0, 2.5])); assert_eq!(metadata["nv.test.empty"], json!([])); } @@ -92,6 +94,7 @@ async fn event_metadata_injection_rejects_invalid_output_atomically() { BTreeMap::from([("nv.test.object".into(), json!({"nested": true}))]), BTreeMap::from([("nv.test.nested_list".into(), json!([[1]]))]), BTreeMap::from([("nv.test.mixed_list".into(), json!([1, "two"]))]), + BTreeMap::from([("nv.test.mixed_numbers".into(), json!([1, 2.5]))]), BTreeMap::from([("nv.test.oversized_number".into(), json!(u64::MAX))]), BTreeMap::from([("nv.test.oversized_list".into(), json!([u64::MAX]))]), ]; diff --git a/crates/node/plugin.d.ts b/crates/node/plugin.d.ts index 8f1d689e0..5226cc510 100644 --- a/crates/node/plugin.d.ts +++ b/crates/node/plugin.d.ts @@ -204,6 +204,18 @@ export interface ToolExecutionInterceptOutcome { pendingMarks?: PendingMarkSpec[]; } +/** Scalar value accepted in event metadata additions. */ +export type EventMetadataScalar = string | number | boolean; + +/** + * Flat value accepted in event metadata additions. After JSON conversion, + * numeric arrays must contain only integer values or only floating-point values. + */ +export type EventMetadataValue = EventMetadataScalar | string[] | number[] | boolean[]; + +/** Metadata additions returned by an event metadata injector. */ +export type EventMetadata = Record; + /** Component-scoped registration context passed to plugin handlers. */ export interface PluginContext { /** @@ -211,6 +223,12 @@ export interface PluginContext { * through the Node binding's callback-error channel; flushSubscribers waits for returned promises. */ registerSubscriber(name: string, callback: (event: Json) => void | Promise): void; + /** Register an event metadata injector for this component. */ + registerEventMetadataInjector( + name: string, + priority: number, + callback: (event: Json) => EventMetadata | Promise, + ): void; /** Register a mark event sanitizer for this component. */ registerMarkSanitizeGuardrail( name: string, diff --git a/crates/node/src/api/mod.rs b/crates/node/src/api/mod.rs index 0935a3f40..9d1b5d506 100644 --- a/crates/node/src/api/mod.rs +++ b/crates/node/src/api/mod.rs @@ -35,8 +35,9 @@ use nemo_relay::api::runtime::subscriber_dispatcher::{ PublicationBuffer, capture_nested_publication_buffer, with_nested_publication_buffer, }; use nemo_relay::api::runtime::{ - EventSanitizeFn, LlmExecutionNextFn, LlmJsonStream, LlmStreamExecutionNextFn, LlmStreamInner, - ScopeStackHandle as CoreScopeStackHandle, ToolExecutionNextFn, + EventMetadataInjectorFn, EventSanitizeFn, LlmExecutionNextFn, LlmJsonStream, + LlmStreamExecutionNextFn, LlmStreamInner, ScopeStackHandle as CoreScopeStackHandle, + ToolExecutionNextFn, }; use nemo_relay::api::runtime::{ TASK_SCOPE_STACK, capture_propagation_context as capture_propagation_context_handle, @@ -820,6 +821,45 @@ fn add_plugin_event_sanitizer( context.set_named_property(property, function) } +fn add_plugin_event_metadata_injector( + env: &Env, + context: &mut JsObject, + namespace_prefix: String, + registrations: Arc>>, +) -> napi::Result<()> { + let function = env.create_function_from_closure( + "__nemo_relay_plugin_register_event_metadata_injector", + move |ctx| { + let name = format!("{}{}", namespace_prefix, ctx.get::(0)?); + let priority = ctx.get::(1)?; + let callback = ctx.get::(2)?; + core_registry_api::register_event_metadata_injector( + &name, + priority, + node_event_metadata_injector_fn(ctx.env, &callback)?, + ) + .map_err(to_napi_err)?; + + let name_clone = name.clone(); + registrations.lock().unwrap().push(PluginRegistration::new( + "plugin", + name_clone.clone(), + Box::new(move || { + core_registry_api::deregister_event_metadata_injector(&name_clone) + .map(|_| ()) + .map_err(|error| { + PluginError::RegistrationFailed(format!( + "event metadata injector deregistration failed: {error}" + )) + }) + }), + )); + ctx.env.get_undefined() + }, + )?; + context.set_named_property("registerEventMetadataInjector", function) +} + fn build_plugin_context( env: &Env, namespace_prefix: String, @@ -862,6 +902,13 @@ fn build_plugin_context( )?; context.set_named_property("registerSubscriber", register_subscriber)?; + add_plugin_event_metadata_injector( + env, + &mut context, + namespace_prefix.clone(), + registrations.clone(), + )?; + add_plugin_event_sanitizer( env, &mut context, @@ -1563,6 +1610,16 @@ fn node_event_sanitize_fn(env: &Env, func: &JsFunction) -> napi::Result napi::Result { + let callback = Arc::new(crate::promise_call::PromiseAwareFn::new(env, func)?); + Ok(callable::wrap_js_event_metadata_injector_promise_fn( + callback, + )) +} + type NodeLlmCodec = ( Arc, Vec>, @@ -3083,9 +3140,38 @@ pub fn llm_stream_call_execute( } // --------------------------------------------------------------------------- -// Tool guardrail registrations +// Event metadata injector and event guardrail registrations // --------------------------------------------------------------------------- +/// Register a global event metadata injector. +/// +/// The callback receives an immutable event snapshot and may return additions +/// directly or in a Promise. Relay validates and merges accepted additions +/// before event sanitizers run. +#[napi] +pub fn register_event_metadata_injector( + env: Env, + name: String, + priority: i32, + #[napi( + ts_arg_type = "(event: Json) => import('./plugin').EventMetadata | Promise" + )] + injector: JsFunction, +) -> Result<()> { + core_registry_api::register_event_metadata_injector( + &name, + priority, + node_event_metadata_injector_fn(&env, &injector)?, + ) + .map_err(to_napi_err) +} + +/// Deregister a global event metadata injector by name. +#[napi] +pub fn deregister_event_metadata_injector(name: String) -> Result { + core_registry_api::deregister_event_metadata_injector(&name).map_err(to_napi_err) +} + macro_rules! napi_event_guardrail_api { ($register_name:ident, $deregister_name:ident, $core_register:path, $core_deregister:path) => { /// Register an event sanitize guardrail. @@ -3628,9 +3714,40 @@ pub fn flush_subscribers(env: Env) -> Result { } // --------------------------------------------------------------------------- -// Scope-local guardrail registrations — Tool +// Scope-local event metadata injector and event guardrail registrations // --------------------------------------------------------------------------- +/// Register an event metadata injector owned by an active scope. +#[napi] +pub fn scope_register_event_metadata_injector( + env: Env, + scope_uuid: String, + name: String, + priority: i32, + #[napi( + ts_arg_type = "(event: Json) => import('./plugin').EventMetadata | Promise" + )] + injector: JsFunction, +) -> Result<()> { + let uuid = uuid::Uuid::parse_str(&scope_uuid) + .map_err(|error| napi::Error::from_reason(format!("invalid UUID: {error}")))?; + core_registry_api::scope_register_event_metadata_injector( + &uuid, + &name, + priority, + node_event_metadata_injector_fn(&env, &injector)?, + ) + .map_err(to_napi_err) +} + +/// Deregister a scope-local event metadata injector by name. +#[napi] +pub fn scope_deregister_event_metadata_injector(scope_uuid: String, name: String) -> Result { + let uuid = uuid::Uuid::parse_str(&scope_uuid) + .map_err(|error| napi::Error::from_reason(format!("invalid UUID: {error}")))?; + core_registry_api::scope_deregister_event_metadata_injector(&uuid, &name).map_err(to_napi_err) +} + macro_rules! napi_scope_event_guardrail_api { ($register_name:ident, $deregister_name:ident, $core_register:path, $core_deregister:path) => { /// Register a scope-local event sanitize guardrail. diff --git a/crates/node/src/callable.rs b/crates/node/src/callable.rs index 55e4bbcef..987a13173 100644 --- a/crates/node/src/callable.rs +++ b/crates/node/src/callable.rs @@ -9,7 +9,7 @@ //! handles serialization of arguments to/from JSON and manages cross-thread communication //! between the Rust async runtime and the Node.js event loop. -use std::collections::BTreeSet; +use std::collections::{BTreeMap, BTreeSet}; use std::future::Future; use std::pin::Pin; use std::sync::atomic::{AtomicBool, Ordering}; @@ -22,10 +22,11 @@ use napi::threadsafe_function::{ use napi::{Env, JsFunction, JsObject, JsUnknown, NapiRaw, NapiValue}; use napi_derive::napi; use nemo_relay::api::runtime::{ - EventSanitizeFn, EventSubscriberFn, LlmCodecIdentity, LlmConditionalFn, LlmExecutionNextFn, - LlmJsonStream, LlmRequestInterceptFn, LlmSanitizeRequestContext, LlmSanitizeRequestFn, - LlmSanitizeResponseContext, LlmSanitizeResponseFn, LlmStreamExecutionNextFn, ToolConditionalFn, - ToolExecutionNextFn, ToolInterceptFn, ToolSanitizeFn, + EventMetadataInjectorFn, EventSanitizeFn, EventSubscriberFn, LlmCodecIdentity, + LlmConditionalFn, LlmExecutionNextFn, LlmJsonStream, LlmRequestInterceptFn, + LlmSanitizeRequestContext, LlmSanitizeRequestFn, LlmSanitizeResponseContext, + LlmSanitizeResponseFn, LlmStreamExecutionNextFn, ToolConditionalFn, ToolExecutionNextFn, + ToolInterceptFn, ToolSanitizeFn, }; use serde::{Deserialize, Serialize}; use serde_json::Value as Json; @@ -619,6 +620,41 @@ pub fn wrap_js_event_sanitize_promise_fn(func: Arc) -> EventSani }) } +/// Wrap a Promise-aware JavaScript event metadata injector. +pub fn wrap_js_event_metadata_injector_promise_fn( + func: Arc, +) -> EventMetadataInjectorFn { + Arc::new(move |event: Arc| { + let func = func.clone(); + Box::pin(async move { + let event_json = JsEvent::try_from_event(&event) + .map(JsEvent::into_json) + .map_err(|error| { + let error = FlowError::Internal(format!( + "failed to serialize JavaScript event metadata injector context: {error}" + )); + record_callback_error(error.to_string()); + error + })?; + let publication = + nemo_relay::api::runtime::subscriber_dispatcher::in_dispatcher_callback(); + let value = if publication { + func.call_spread_for_publication(vec![event_json]).await + } else { + func.call(event_json).await + } + .inspect_err(|error| record_callback_error(error.to_string()))?; + serde_json::from_value::>(value).map_err(|error| { + let error = FlowError::Internal(format!( + "invalid JavaScript event metadata injector result: {error}" + )); + record_callback_error(error.to_string()); + error + }) + }) + }) +} + fn recv_json_or_null(rx: std::sync::mpsc::Receiver, error_prefix: &str) -> Json { rx.recv().unwrap_or_else(|e| { record_callback_error(format!("{error_prefix}: {e}")); diff --git a/crates/node/tests/event_metadata_injection_tests.mjs b/crates/node/tests/event_metadata_injection_tests.mjs new file mode 100644 index 000000000..45132a877 --- /dev/null +++ b/crates/node/tests/event_metadata_injection_tests.mjs @@ -0,0 +1,185 @@ +// SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +import assert from 'node:assert/strict'; +import { createRequire } from 'node:module'; +import { mkdtempSync, rmSync } from 'node:fs'; +import { tmpdir } from 'node:os'; +import path from 'node:path'; +import { describe, it } from 'node:test'; + +const configHome = mkdtempSync(path.join(tmpdir(), 'nemo-relay-node-config-')); +process.env.XDG_CONFIG_HOME = configHome; +process.on('exit', () => rmSync(configHome, { recursive: true, force: true })); + +const require = createRequire(import.meta.url); +const lib = require('../index.js'); +const plugin = require('../plugin.js'); + +function capture(name) { + const events = []; + lib.registerSubscriber(name, (event) => events.push(event)); + return events; +} + +async function waitFor(events, count) { + for (let attempt = 0; attempt < 100 && events.length < count; attempt += 1) { + await new Promise((resolve) => setTimeout(resolve, 10)); + } + assert.ok(events.length >= count, `expected ${count} events, received ${events.length}`); +} + +async function initializeWithoutDiscoveredPluginConfig(config) { + const previousDirectory = process.cwd(); + const directory = mkdtempSync(path.join(tmpdir(), 'nemo-relay-node-')); + try { + process.chdir(directory); + return await plugin.initialize(config); + } finally { + process.chdir(previousDirectory); + rmSync(directory, { recursive: true, force: true }); + } +} + +describe('event metadata injector bindings', () => { + it('preserves insertion order while isolating callback failures and invalid output', async () => { + const events = capture('node-event-metadata-global-sub'); + lib.clearLastCallbackError(); + lib.registerEventMetadataInjector('node-event-metadata-sync', 10, (event) => ({ + 'node.existing': 'ignored', + 'node.injector.shared': 'sync-first', + 'node.injector.sync': event.name, + })); + lib.registerEventMetadataInjector('node-event-metadata-async', 20, async () => ({ + 'node.injector.async': true, + 'node.injector.shared': 'async-later', + })); + lib.registerEventMetadataInjector('node-event-metadata-failure', 30, () => { + throw new Error('node injector failure'); + }); + lib.registerEventMetadataInjector('node-event-metadata-after-failure', 40, () => ({ + 'node.injector.after_failure': 'added', + })); + lib.registerEventMetadataInjector('node-event-metadata-invalid', 50, () => ['not', 'a', 'mapping']); + try { + lib.event('node-event-metadata-global', null, null, { 'node.existing': 'preserved' }); + await lib.flushSubscribers(); + await waitFor(events, 1); + } finally { + lib.deregisterEventMetadataInjector('node-event-metadata-invalid'); + lib.deregisterEventMetadataInjector('node-event-metadata-after-failure'); + lib.deregisterEventMetadataInjector('node-event-metadata-failure'); + lib.deregisterEventMetadataInjector('node-event-metadata-async'); + lib.deregisterEventMetadataInjector('node-event-metadata-sync'); + lib.deregisterSubscriber('node-event-metadata-global-sub'); + } + + assert.deepEqual(events.at(-1).metadata, { + 'node.existing': 'preserved', + 'node.injector.after_failure': 'added', + 'node.injector.async': true, + 'node.injector.shared': 'sync-first', + 'node.injector.sync': 'node-event-metadata-global', + }); + assert.match(lib.getLastCallbackError() ?? '', /invalid JavaScript event metadata injector result/i); + lib.clearLastCallbackError(); + }); + + it('applies and deregisters scope-local callbacks', async () => { + const events = capture('node-event-metadata-scope-sub'); + const owner = lib.pushScope('node-event-metadata-owner', lib.ScopeType.Agent); + lib.scopeRegisterEventMetadataInjector(owner.uuid, 'node-event-metadata-local-first', 10, (event) => ({ + 'node.injector.scope_local': event.name, + 'node.injector.scope_order': 'first', + })); + lib.scopeRegisterEventMetadataInjector(owner.uuid, 'node-event-metadata-local-later', 20, () => ({ + 'node.injector.scope_order': 'later', + })); + lib.event('node-event-metadata-before-deregister', owner); + assert.equal(lib.scopeDeregisterEventMetadataInjector(owner.uuid, 'node-event-metadata-local-first'), true); + assert.equal(lib.scopeDeregisterEventMetadataInjector(owner.uuid, 'node-event-metadata-local-first'), false); + assert.equal(lib.scopeDeregisterEventMetadataInjector(owner.uuid, 'node-event-metadata-local-later'), true); + lib.event('node-event-metadata-after-deregister', owner); + lib.popScope(owner); + await lib.flushSubscribers(); + await waitFor(events, 4); + lib.deregisterSubscriber('node-event-metadata-scope-sub'); + + const marks = Object.fromEntries( + events.filter((event) => event.kind === 'mark').map((event) => [event.name, event]), + ); + assert.deepEqual(marks['node-event-metadata-before-deregister'].metadata, { + 'node.injector.scope_local': 'node-event-metadata-before-deregister', + 'node.injector.scope_order': 'first', + }); + assert.equal(marks['node-event-metadata-after-deregister'].metadata, null); + }); + + it('requires homogeneous numeric arrays at runtime', async () => { + const events = capture('node-event-metadata-numeric-sub'); + lib.registerEventMetadataInjector('node-event-metadata-integers', 10, () => ({ + 'node.injector.integers': [1, 2], + })); + lib.registerEventMetadataInjector('node-event-metadata-doubles', 20, () => ({ + 'node.injector.doubles': [1.25, 2.5], + })); + lib.registerEventMetadataInjector('node-event-metadata-mixed-numbers', 30, () => ({ + 'node.injector.mixed_numbers': [1, 2.5], + })); + try { + lib.event('node-event-metadata-homogeneous-numeric-arrays'); + await lib.flushSubscribers(); + await waitFor(events, 1); + } finally { + lib.deregisterEventMetadataInjector('node-event-metadata-mixed-numbers'); + lib.deregisterEventMetadataInjector('node-event-metadata-doubles'); + lib.deregisterEventMetadataInjector('node-event-metadata-integers'); + lib.deregisterSubscriber('node-event-metadata-numeric-sub'); + } + + assert.deepEqual(events.at(-1).metadata, { + 'node.injector.doubles': [1.25, 2.5], + 'node.injector.integers': [1, 2], + }); + }); + + it('cleans up plugin-owned callbacks', async () => { + const kind = `node.test.event-metadata.${Date.now()}`; + const events = capture('node-event-metadata-plugin-sub'); + plugin.register(kind, { + register(config, context) { + context.registerEventMetadataInjector('configured', 10, () => config.metadata); + }, + }); + try { + await initializeWithoutDiscoveredPluginConfig({ + version: 1, + components: [ + plugin.ComponentSpec(kind, { + metadata: { 'node.injector.plugin': 'configured' }, + }), + ], + }); + lib.event('node-event-metadata-plugin-configured'); + await lib.flushSubscribers(); + await waitFor(events, 1); + + plugin.clear(); + lib.event('node-event-metadata-plugin-cleared'); + await lib.flushSubscribers(); + await waitFor(events, 2); + } finally { + plugin.clear(); + plugin.deregister(kind); + lib.deregisterSubscriber('node-event-metadata-plugin-sub'); + } + + const marks = Object.fromEntries( + events.filter((event) => event.kind === 'mark').map((event) => [event.name, event]), + ); + assert.deepEqual(marks['node-event-metadata-plugin-configured'].metadata, { + 'node.injector.plugin': 'configured', + }); + assert.equal(marks['node-event-metadata-plugin-cleared'].metadata, null); + }); +}); diff --git a/crates/node/tests/public_event_metadata_api_fixture.ts b/crates/node/tests/public_event_metadata_api_fixture.ts new file mode 100644 index 000000000..56a17b816 --- /dev/null +++ b/crates/node/tests/public_event_metadata_api_fixture.ts @@ -0,0 +1,40 @@ +// SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +import { registerEventMetadataInjector, scopeRegisterEventMetadataInjector } from '../index.js'; +import type { EventMetadata, PluginContext } from '../plugin.js'; + +const acceptedMetadata: EventMetadata = { + 'fixture.string': 'value', + 'fixture.number': 42, + 'fixture.boolean': true, + 'fixture.strings': ['alpha', 'beta'], + 'fixture.numbers': [1, 2.5], + 'fixture.booleans': [true, false], + 'fixture.empty': [], +}; + +registerEventMetadataInjector('fixture-global', 0, () => acceptedMetadata); +scopeRegisterEventMetadataInjector( + '00000000-0000-0000-0000-000000000001', + 'fixture-scope', + 0, + async () => acceptedMetadata, +); + +declare const context: PluginContext; +context.registerEventMetadataInjector('fixture-plugin', 0, () => acceptedMetadata); + +// @ts-expect-error Global injector callbacks cannot return object values. +registerEventMetadataInjector('fixture-object', 0, () => ({ 'fixture.object': { nested: true } })); + +// @ts-expect-error Global injector callbacks cannot return nested arrays. +registerEventMetadataInjector('fixture-nested', 0, () => ({ 'fixture.nested': [[1]] })); + +scopeRegisterEventMetadataInjector('00000000-0000-0000-0000-000000000001', 'fixture-null', 0, () => ({ + // @ts-expect-error Scope-local injector callbacks cannot return null values. + 'fixture.null': null, +})); + +// @ts-expect-error Plugin injector arrays must contain one primitive type. +context.registerEventMetadataInjector('fixture-mixed', 0, () => ({ 'fixture.mixed': [1, 'two'] })); diff --git a/crates/node/tests/types_tests.mjs b/crates/node/tests/types_tests.mjs index 74083f2d5..706152b12 100644 --- a/crates/node/tests/types_tests.mjs +++ b/crates/node/tests/types_tests.mjs @@ -43,7 +43,7 @@ describe('Type constants', () => { assert.equal(ScopeType.Unknown, 10); }); - it('type-checks the public observability API fixture', () => { + it('type-checks the public API fixtures', () => { execFileSync( process.execPath, [ @@ -58,6 +58,7 @@ describe('Type constants', () => { '--moduleResolution', 'NodeNext', 'tests/public_observability_api_fixture.ts', + 'tests/public_event_metadata_api_fixture.ts', ], { cwd: fileURLToPath(new URL('..', import.meta.url)), diff --git a/crates/python/src/py_api/mod.rs b/crates/python/src/py_api/mod.rs index fcc0f22d6..63768c459 100644 --- a/crates/python/src/py_api/mod.rs +++ b/crates/python/src/py_api/mod.rs @@ -1315,9 +1315,27 @@ fn llm_stream_call_execute<'py>( } // --------------------------------------------------------------------------- -// Guardrail registrations (macro-generated) +// Event metadata injector and guardrail registrations // --------------------------------------------------------------------------- +#[pyfunction] +fn register_event_metadata_injector( + name: &str, + priority: i32, + injector: Py, +) -> PyResult<()> { + core_registry_api::register_event_metadata_injector( + name, + priority, + py_callable::wrap_py_event_metadata_injector_fn(injector), + ) + .map_err(to_py_err) +} + +#[pyfunction] +fn deregister_event_metadata_injector(name: &str) -> PyResult { + core_registry_api::deregister_event_metadata_injector(name).map_err(to_py_err) +} macro_rules! py_event_guardrail_api { ($register_name:ident, $deregister_name:ident, $core_register:path, $core_deregister:path) => { #[pyfunction] @@ -1806,6 +1824,28 @@ fn parse_uuid(scope_uuid: &str) -> PyResult { .map_err(|e| PyErr::new::(format!("invalid UUID: {e}"))) } +#[pyfunction] +fn scope_register_event_metadata_injector( + scope_uuid: &str, + name: &str, + priority: i32, + injector: Py, +) -> PyResult<()> { + let uuid = parse_uuid(scope_uuid)?; + core_registry_api::scope_register_event_metadata_injector( + &uuid, + name, + priority, + py_callable::wrap_py_event_metadata_injector_fn(injector), + ) + .map_err(to_py_err) +} + +#[pyfunction] +fn scope_deregister_event_metadata_injector(scope_uuid: &str, name: &str) -> PyResult { + let uuid = parse_uuid(scope_uuid)?; + core_registry_api::scope_deregister_event_metadata_injector(&uuid, name).map_err(to_py_err) +} macro_rules! py_scope_event_guardrail_api { ($register_name:ident, $deregister_name:ident, $core_register:path, $core_deregister:path) => { #[pyfunction] @@ -2241,6 +2281,8 @@ pub fn register(m: &Bound<'_, PyModule>) -> PyResult<()> { )?)?; // Mark and scope event guardrails + m.add_function(wrap_pyfunction!(register_event_metadata_injector, m)?)?; + m.add_function(wrap_pyfunction!(deregister_event_metadata_injector, m)?)?; 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!( @@ -2339,6 +2381,11 @@ pub fn register(m: &Bound<'_, PyModule>) -> PyResult<()> { scope_deregister_tool_conditional_execution_guardrail, m )?)?; + m.add_function(wrap_pyfunction!(scope_register_event_metadata_injector, m)?)?; + m.add_function(wrap_pyfunction!( + scope_deregister_event_metadata_injector, + m + )?)?; m.add_function(wrap_pyfunction!(scope_register_mark_sanitize_guardrail, m)?)?; m.add_function(wrap_pyfunction!( scope_deregister_mark_sanitize_guardrail, diff --git a/crates/python/src/py_callable.rs b/crates/python/src/py_callable.rs index 266238f47..38daa3082 100644 --- a/crates/python/src/py_callable.rs +++ b/crates/python/src/py_callable.rs @@ -20,6 +20,7 @@ #![allow(clippy::type_complexity)] +use std::collections::BTreeMap; use std::future::Future; use std::pin::Pin; use std::sync::{Arc, Mutex}; @@ -29,12 +30,12 @@ use nemo_relay::api::runtime::subscriber_dispatcher::{ PublicationBuffer, PublicationContext, capture_nested_publication_buffer, publication_context, }; use nemo_relay::api::runtime::{ - EventSanitizeFn, EventSubscriberFn, LlmConditionalFn, LlmExecutionNextFn, LlmJsonStream, - LlmRequestInterceptFn, LlmSanitizeRequestContext, LlmSanitizeRequestFn, - LlmSanitizeResponseContext, LlmSanitizeResponseFn, LlmStreamExecutionNextFn, LlmStreamInner, - MiddlewareContinuationContext, ScopeStackHandle, ToolConditionalFn, ToolExecutionNextFn, - ToolInterceptFn, ToolSanitizeFn, capture_propagation_context, capture_traceparent, - current_scope_stack, + EventMetadataInjectorFn, EventSanitizeFn, EventSubscriberFn, LlmConditionalFn, + LlmExecutionNextFn, LlmJsonStream, LlmRequestInterceptFn, LlmSanitizeRequestContext, + LlmSanitizeRequestFn, LlmSanitizeResponseContext, LlmSanitizeResponseFn, + LlmStreamExecutionNextFn, LlmStreamInner, MiddlewareContinuationContext, ScopeStackHandle, + ToolConditionalFn, ToolExecutionNextFn, ToolInterceptFn, ToolSanitizeFn, + capture_propagation_context, capture_traceparent, current_scope_stack, }; use nemo_relay::error::{FlowError, Result as FlowResult}; use pyo3::exceptions::PyRuntimeError; @@ -1927,6 +1928,109 @@ pub fn wrap_py_event_sanitize_fn(py_fn: Py) -> EventSanitizeFn { }) } +fn call_event_metadata_injector( + py: Python<'_>, + invoke: &Bound<'_, PyAny>, + callback: &Py, + invocation_context: Option<&Bound<'_, PyAny>>, + loop_affine: bool, + py_event: Py, +) -> PyResult> { + let result = match (invocation_context, loop_affine) { + (Some(context), false) => { + context.call_method1("run", (invoke, callback.bind(py), py_event)) + } + (None, false) => invoke.call1((callback.bind(py), py_event)), + (Some(context), true) => context.call_method1("run", (callback.bind(py), py_event)), + (None, true) => callback.bind(py).call1((py_event,)), + }?; + Ok(result.unbind()) +} + +fn start_py_event_metadata_injector( + py: Python<'_>, + py_fn: &Py, + event: &Event, + publication_context: Option<&PythonPublicationContext>, + task_locals: Option, + publication_buffer: Option, +) -> FlowResult, PyValueFuture>> { + let (invocation_context, task_locals) = prepare_event_sanitizer_invocation( + py, + publication_context, + task_locals, + publication_buffer, + )?; + let py_event = + py_event_object(py, event).map_err(|error| FlowError::Internal(error.to_string()))?; + let invoke = py + .import("nemo_relay._event_sanitizer_context") + .and_then(|module| module.getattr("invoke")) + .map_err(|error| FlowError::Internal(error.to_string()))?; + let loop_affine = task_locals.is_some(); + let callback = loop_affine_callback(py, py_fn.bind(py), task_locals.as_ref(), true) + .map_err(|error| FlowError::Internal(error.to_string()))?; + let result = call_event_metadata_injector( + py, + &invoke, + &callback, + invocation_context.as_ref(), + loop_affine, + py_event, + ) + .map_err(python_callback_error)?; + split_py_object_or_future_with_locals( + py, + result, + task_locals.as_ref(), + invocation_context.as_ref(), + ) +} + +/// Wrap a Python callable ``(Event) -> dict[str, Json]``. +pub fn wrap_py_event_metadata_injector_fn(py_fn: Py) -> EventMetadataInjectorFn { + let py_fn = Arc::new(py_fn); + let task_locals = capture_python_task_locals(); + Arc::new(move |event: Arc| { + let py_fn = py_fn.clone(); + let task_locals = task_locals_with_running_loop(task_locals.as_ref()); + let publication_context = publication_context::(); + let publication_buffer = capture_nested_publication_buffer(); + Box::pin(async move { + let result = Python::attach(|py| { + start_py_event_metadata_injector( + py, + py_fn.as_ref(), + event.as_ref(), + publication_context.as_deref(), + task_locals, + publication_buffer, + ) + }); + let result = resolve_py_object_or_future(result) + .await + .and_then(|result| { + Python::attach(|py| { + py_to_json(result.bind(py)) + .map_err(|error| FlowError::Internal(error.to_string())) + .and_then(|value| { + serde_json::from_value::>(value).map_err( + |error| { + FlowError::Internal(format!( + "invalid event metadata injector result: {error}" + )) + }, + ) + }) + }) + }); + if let Err(error) = &result { + eprintln!("nemo_relay: Python event metadata injector failed: {error}"); + } + result + }) + }) +} // --------------------------------------------------------------------------- // LLM Codec wrapper // --------------------------------------------------------------------------- diff --git a/crates/python/src/py_plugin.rs b/crates/python/src/py_plugin.rs index ec11ce32e..ad5c4e3c2 100644 --- a/crates/python/src/py_plugin.rs +++ b/crates/python/src/py_plugin.rs @@ -13,13 +13,14 @@ use pyo3::prelude::*; use serde_json::{Map, Value as Json}; 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_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, + deregister_event_metadata_injector, 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_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_event_metadata_injector, 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, @@ -39,8 +40,8 @@ use nemo_relay::plugin::{ use crate::convert::{json_to_py, py_to_json}; use crate::py_callable::{ - 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_event_metadata_injector_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, @@ -240,6 +241,26 @@ impl PyPluginContext { #[pymethods] impl PyPluginContext { + #[pyo3(signature = (name: "str", priority: "int", callback: "object") -> "None", text_signature = "(name: str, priority: int, callback: object) -> None")] + fn register_event_metadata_injector( + &self, + name: &str, + priority: i32, + callback: Py, + ) -> PyResult<()> { + self.register_callback( + name, + |qualified_name| { + register_event_metadata_injector( + qualified_name, + priority, + wrap_py_event_metadata_injector_fn(callback), + ) + }, + deregister_event_metadata_injector, + "event metadata injector", + ) + } #[pyo3(signature = (name: "str", priority: "int", callback: "object") -> "None", text_signature = "(name: str, priority: int, callback: object) -> None")] fn register_mark_sanitize_guardrail( &self, diff --git a/crates/python/tests/coverage/coverage_tests.rs b/crates/python/tests/coverage/coverage_tests.rs index 8c3db4757..d2d1dcdc2 100644 --- a/crates/python/tests/coverage/coverage_tests.rs +++ b/crates/python/tests/coverage/coverage_tests.rs @@ -278,6 +278,8 @@ fn test_register_exposes_all_native_api_functions() { "llm_call_end", "llm_call_execute", "llm_stream_call_execute", + "register_event_metadata_injector", + "deregister_event_metadata_injector", "register_mark_sanitize_guardrail", "deregister_mark_sanitize_guardrail", "register_scope_sanitize_start_guardrail", @@ -314,6 +316,8 @@ 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_event_metadata_injector", + "scope_deregister_event_metadata_injector", "scope_register_mark_sanitize_guardrail", "scope_deregister_mark_sanitize_guardrail", "scope_register_scope_sanitize_start_guardrail", diff --git a/crates/python/tests/coverage/py_callable_coverage_tests.rs b/crates/python/tests/coverage/py_callable_coverage_tests.rs index c94b28a4d..411ba034d 100644 --- a/crates/python/tests/coverage/py_callable_coverage_tests.rs +++ b/crates/python/tests/coverage/py_callable_coverage_tests.rs @@ -873,6 +873,57 @@ def invalid(event, fields): ); } +#[test] +fn event_metadata_injector_wrapper_covers_sync_async_and_invalid_results() { + use nemo_relay::api::event::{BaseEvent, MarkEvent}; + + let _python = crate::test_support::init_python_test(); + let (_context_module, sync_injector, async_injector, invalid_injector) = Python::attach(|py| { + let context_module = install_event_sanitizer_context_module(py); + let module = load_module( + py, + r#" +import asyncio + +def inject(event): + return {"python.sync": event.name} + +async def inject_async(event): + await asyncio.sleep(0) + return {"python.async": True} + +def invalid(event): + return [event.name] +"#, + ); + ( + context_module, + wrap_py_event_metadata_injector_fn(module.getattr("inject").unwrap().unbind()), + wrap_py_event_metadata_injector_fn(module.getattr("inject_async").unwrap().unbind()), + wrap_py_event_metadata_injector_fn(module.getattr("invalid").unwrap().unbind()), + ) + }); + let event = Arc::new(Event::Mark(MarkEvent::new( + BaseEvent::builder().name("checkpoint").build(), + None, + None, + ))); + let runtime = tokio::runtime::Runtime::new().unwrap(); + + let sync_values = runtime.block_on(sync_injector(event.clone())).unwrap(); + assert_eq!(sync_values.get("python.sync"), Some(&json!("checkpoint"))); + + let async_values = runtime.block_on(async_injector(event.clone())).unwrap(); + assert_eq!(async_values.get("python.async"), Some(&json!(true))); + + let invalid = runtime.block_on(invalid_injector(event)).unwrap_err(); + assert!( + invalid + .to_string() + .contains("invalid event metadata injector result") + ); +} + #[test] fn awaitable_middleware_wrappers_cover_success_and_failure() { let _python = crate::test_support::init_python_test(); diff --git a/crates/python/tests/coverage/py_plugin_coverage_tests.rs b/crates/python/tests/coverage/py_plugin_coverage_tests.rs index e3973e587..c50296cee 100644 --- a/crates/python/tests/coverage/py_plugin_coverage_tests.rs +++ b/crates/python/tests/coverage/py_plugin_coverage_tests.rs @@ -215,6 +215,7 @@ fn stale_async_clear_completion_keeps_the_newer_state() { } #[test] +#[allow(clippy::cognitive_complexity)] fn plugin_context_registers_all_runtime_hooks_and_drains_registrations() { let _python = crate::test_support::init_python_test(); Python::attach(|py| { @@ -227,6 +228,9 @@ def subscriber(event): def event_sanitize(event, fields): return fields +def event_metadata(event): + return {"python.plugin": event.name} + def tool_fn(name, value): return value @@ -271,6 +275,13 @@ async def tool_execution_intercept(name, value, next): helpers.getattr("subscriber").unwrap().unbind(), ) .unwrap(); + context + .register_event_metadata_injector( + "event_metadata", + 1, + helpers.getattr("event_metadata").unwrap().unbind(), + ) + .unwrap(); context .register_mark_sanitize_guardrail( "mark_sanitize", @@ -379,7 +390,7 @@ async def tool_execution_intercept(name, value, next): .unwrap(); let registrations = context.drain_registrations().unwrap(); - assert_eq!(registrations.len(), 15); + assert_eq!(registrations.len(), 16); assert!( registrations .iter() @@ -387,6 +398,7 @@ async def tool_execution_intercept(name, value, next): ); assert!(deregister_subscriber("demo.subscriber").unwrap()); + assert!(deregister_event_metadata_injector("demo.event_metadata").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()); diff --git a/python/nemo_relay/__init__.py b/python/nemo_relay/__init__.py index ffbad037a..e32715100 100644 --- a/python/nemo_relay/__init__.py +++ b/python/nemo_relay/__init__.py @@ -11,6 +11,7 @@ - ``nemo_relay.tools`` for tool lifecycle management - ``nemo_relay.llm`` for non-streaming and streaming LLM lifecycle management - ``nemo_relay.guardrails`` and ``nemo_relay.intercepts`` for global middleware +- ``nemo_relay.event_metadata`` for global event metadata injection - ``nemo_relay.scope_local`` for middleware scoped to a specific ``ScopeHandle`` - ``nemo_relay.typed`` for codec-based typed wrappers - ``nemo_relay.plugin`` for global plugin configuration and custom plugin registration @@ -199,6 +200,17 @@ class EventSanitizeFields(TypedDict): EventSanitizeGuardrail: TypeAlias = Callable[ ["Event", EventSanitizeFields], EventSanitizeFields | Awaitable[EventSanitizeFields] ] +#: Primitive values accepted from event metadata injectors. +EventMetadataScalar: TypeAlias = str | int | float | bool +#: Flat event metadata value accepted from an injector. Lists must be empty or +#: contain values of one primitive type; Relay validates this at runtime. +EventMetadataValue: TypeAlias = EventMetadataScalar | list[str] | list[int] | list[float] | list[bool] +#: Flat metadata additions returned by an event metadata injector. +EventMetadata: TypeAlias = dict[str, EventMetadataValue] +#: Additive middleware callback that inspects an immutable event and returns +#: proposed metadata additions. Both synchronous and asynchronous callbacks are +#: supported. +EventMetadataInjectorCallback: TypeAlias = Callable[["Event"], EventMetadata | Awaitable[EventMetadata]] #: 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] | Awaitable[Optional[str]]] @@ -256,6 +268,7 @@ class EventSanitizeFields(TypedDict): from nemo_relay import ( # noqa: E402 adaptive, codecs, + event_metadata, guardrails, intercepts, llm, @@ -601,6 +614,7 @@ def worker() -> None: "tools", "llm", "guardrails", + "event_metadata", "intercepts", "subscribers", "scope_local", @@ -669,6 +683,10 @@ def worker() -> None: "UnsupportedBehavior", "EventSanitizeFields", "EventSanitizeGuardrail", + "EventMetadataScalar", + "EventMetadataValue", + "EventMetadata", + "EventMetadataInjectorCallback", "ToolSanitizeGuardrail", "ToolConditionalExecutionGuardrail", "LlmSanitizeRequestGuardrail", diff --git a/python/nemo_relay/__init__.pyi b/python/nemo_relay/__init__.pyi index 3b466799a..de631a980 100644 --- a/python/nemo_relay/__init__.pyi +++ b/python/nemo_relay/__init__.pyi @@ -27,6 +27,7 @@ from typing import Literal, Optional, TypeAlias, TypedDict from nemo_relay import adaptive as adaptive from nemo_relay import codecs as codecs +from nemo_relay import event_metadata as event_metadata from nemo_relay import guardrails as guardrails from nemo_relay import intercepts as intercepts from nemo_relay import llm as llm @@ -232,6 +233,14 @@ Exceptional flow: Exceptions fail closed, clear the mutable observability fields, and stop the remaining sanitizer chain. """ +EventMetadataScalar: TypeAlias = str | int | float | bool +EventMetadataValue: TypeAlias = EventMetadataScalar | list[str] | list[int] | list[float] | list[bool] +EventMetadata: TypeAlias = dict[str, EventMetadataValue] +EventMetadataInjectorCallback: TypeAlias = Callable[ + [Event], + EventMetadata | Awaitable[EventMetadata], +] +"""Additive middleware callback that returns flat event metadata additions.""" ToolConditionalExecutionGuardrail: TypeAlias = Callable[[str, Json], Optional[str] | Awaitable[Optional[str]]] """Guardrail callback that can block tool execution. diff --git a/python/nemo_relay/_native.pyi b/python/nemo_relay/_native.pyi index ee4a46ae4..ce5acab33 100644 --- a/python/nemo_relay/_native.pyi +++ b/python/nemo_relay/_native.pyi @@ -58,6 +58,13 @@ _EventSanitizeGuardrail: TypeAlias = Callable[ [ScopeEvent | MarkEvent, _EventSanitizeFields], _EventSanitizeFields | Awaitable[_EventSanitizeFields], ] +_EventMetadataScalar: TypeAlias = str | int | float | bool +_EventMetadataValue: TypeAlias = _EventMetadataScalar | list[str] | list[int] | list[float] | list[bool] +_EventMetadata: TypeAlias = dict[str, _EventMetadataValue] +_EventMetadataInjector: TypeAlias = Callable[ + [ScopeEvent | MarkEvent], + _EventMetadata | Awaitable[_EventMetadata], +] _LlmConditionalExecutionGuardrail: TypeAlias = Callable[["LLMRequest"], Optional[str] | Awaitable[Optional[str]]] _ToolRequestIntercept: TypeAlias = Callable[[str, _Json], _Json | Awaitable[_Json]] _ToolExecutionIntercept: TypeAlias = Callable[ @@ -1451,6 +1458,7 @@ class PluginContext: Python plugin protocols expose the public shape. The native class exists for runtime registration callbacks. """ + def register_event_metadata_injector(self, name: str, priority: int, callback: _EventMetadataInjector) -> 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 @@ -1459,6 +1467,8 @@ class PluginContext: self, name: str, priority: int, callback: _EventSanitizeGuardrail ) -> None: ... +def register_event_metadata_injector(name: str, priority: int, injector: _EventMetadataInjector) -> None: ... +def deregister_event_metadata_injector(name: str) -> bool: ... 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: ... @@ -1477,6 +1487,10 @@ 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 scope_register_event_metadata_injector( + scope_uuid: str, name: str, priority: int, injector: _EventMetadataInjector +) -> None: ... +def scope_deregister_event_metadata_injector(scope_uuid: str, name: str) -> bool: ... def create_scope_stack() -> ScopeStack: """Create a fresh native scope stack. diff --git a/python/nemo_relay/event_metadata.py b/python/nemo_relay/event_metadata.py new file mode 100644 index 000000000..380164183 --- /dev/null +++ b/python/nemo_relay/event_metadata.py @@ -0,0 +1,36 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Global event metadata injector registration. + +Injectors receive an immutable ``ScopeEvent`` or ``MarkEvent`` and return +metadata additions. Relay validates and merges accepted additions before the +event sanitizer chain runs. +""" + +from __future__ import annotations + +from nemo_relay import EventMetadataInjectorCallback +from nemo_relay._native import ( + deregister_event_metadata_injector as _deregister_event_metadata_injector, +) +from nemo_relay._native import ( + register_event_metadata_injector as _register_event_metadata_injector, +) + + +def register_injector(name: str, priority: int, injector: EventMetadataInjectorCallback) -> None: + """Register a global event metadata injector. + + The registration applies to every subsequently published canonical Relay + event until it is deregistered. Lower numeric priorities run first. + """ + _register_event_metadata_injector(name, priority, injector) + + +def deregister_injector(name: str) -> bool: + """Remove a global event metadata injector by registration name.""" + return _deregister_event_metadata_injector(name) + + +__all__ = ["register_injector", "deregister_injector"] diff --git a/python/nemo_relay/plugin.py b/python/nemo_relay/plugin.py index 6afe38b64..17f68de5a 100644 --- a/python/nemo_relay/plugin.py +++ b/python/nemo_relay/plugin.py @@ -21,6 +21,7 @@ from typing import TYPE_CHECKING, AsyncIterator, Callable, Literal, Protocol, Self, TypedDict, cast from nemo_relay import ( + EventMetadataInjectorCallback, EventSanitizeGuardrail, Json, JsonObject, @@ -117,6 +118,12 @@ def register_subscriber(self, name: str, callback: Callable[[Event], None]) -> N """Register an infallible event subscriber for this component.""" ... + def register_event_metadata_injector( + self, name: str, priority: int, callback: EventMetadataInjectorCallback + ) -> None: + """Register an event metadata injector for this component.""" + ... + def register_mark_sanitize_guardrail(self, name: str, priority: int, callback: EventSanitizeGuardrail) -> None: """Register a mark event sanitizer for this component.""" ... diff --git a/python/nemo_relay/plugin.pyi b/python/nemo_relay/plugin.pyi index ee8f165dd..5a5f8a72b 100644 --- a/python/nemo_relay/plugin.pyi +++ b/python/nemo_relay/plugin.pyi @@ -8,6 +8,7 @@ from typing import AsyncContextManager, Literal, Protocol, Self, TypedDict from nemo_relay import ( Event, + EventMetadataInjectorCallback, EventSanitizeGuardrail, JsonObject, LlmConditionalExecutionGuardrail, @@ -50,6 +51,9 @@ class ConfigReport(TypedDict): class PluginContext(Protocol): def register_subscriber(self, name: str, callback: Callable[[Event], None]) -> None: ... + def register_event_metadata_injector( + self, name: str, priority: int, callback: EventMetadataInjectorCallback + ) -> 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 diff --git a/python/nemo_relay/scope_local.py b/python/nemo_relay/scope_local.py index 7b4da72a7..09f1ad763 100644 --- a/python/nemo_relay/scope_local.py +++ b/python/nemo_relay/scope_local.py @@ -19,6 +19,12 @@ def redact(tool_name, args): nemo_relay.scope_local.register_tool_sanitize_request(handle, "redact", 10, redact) """ +from __future__ import annotations + +from nemo_relay import EventMetadataInjectorCallback, ScopeHandle +from nemo_relay._native import ( + scope_deregister_event_metadata_injector as _deregister_event_metadata_injector, +) from nemo_relay._native import ( scope_deregister_llm_conditional_execution_guardrail as _deregister_llm_conditional_execution, ) @@ -64,6 +70,9 @@ def redact(tool_name, args): from nemo_relay._native import ( scope_deregister_tool_sanitize_response_guardrail as _deregister_tool_sanitize_response, ) +from nemo_relay._native import ( + scope_register_event_metadata_injector as _register_event_metadata_injector, +) from nemo_relay._native import ( scope_register_llm_conditional_execution_guardrail as _register_llm_conditional_execution, ) @@ -115,6 +124,21 @@ def redact(tool_name, args): # --------------------------------------------------------------------------- +def register_event_metadata_injector( + scope_handle: ScopeHandle, + name: str, + priority: int, + injector: EventMetadataInjectorCallback, +) -> None: + """Register an event metadata injector owned by an active scope.""" + _register_event_metadata_injector(scope_handle.uuid, name, priority, injector) + + +def deregister_event_metadata_injector(scope_handle: ScopeHandle, name: str) -> bool: + """Remove a scope-local event metadata injector by registration name.""" + return _deregister_event_metadata_injector(scope_handle.uuid, name) + + 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) @@ -671,6 +695,9 @@ def deregister_subscriber(scope_handle, name): __all__ = [ + # Injector registrations + "register_event_metadata_injector", + "deregister_event_metadata_injector", # Mark and scope event guardrails "register_mark_sanitize", "deregister_mark_sanitize", diff --git a/python/tests/test_event_metadata_injection.py b/python/tests/test_event_metadata_injection.py new file mode 100644 index 000000000..86927bbfb --- /dev/null +++ b/python/tests/test_event_metadata_injection.py @@ -0,0 +1,200 @@ +# 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 +from uuid import uuid4 + +import nemo_relay +from nemo_relay import event_metadata, plugin, scope, scope_local, subscribers + + +async def test_global_python_injectors_support_sync_async_and_failure_safe_output(subscribed_events): + sync_name = f"python-sync-{uuid4()}" + async_name = f"python-async-{uuid4()}" + failure_name = f"python-failure-{uuid4()}" + after_failure_name = f"python-after-failure-{uuid4()}" + invalid_name = f"python-invalid-{uuid4()}" + + async def inject_async(_event: nemo_relay.Event) -> nemo_relay.EventMetadata: + return { + "python.injector.async": True, + "python.injector.shared": "async-later", + } + + event_metadata.register_injector( + sync_name, + 10, + lambda _event: { + "python.existing": "ignored", + "python.injector.shared": "sync-first", + "python.injector.sync": "added", + }, + ) + event_metadata.register_injector(async_name, 20, inject_async) + invalid_injector = cast( + nemo_relay.EventMetadataInjectorCallback, + lambda _event: ["not", "a", "mapping"], + ) + + def fail(_event: nemo_relay.Event) -> nemo_relay.EventMetadata: + raise RuntimeError("python injector failure") + + event_metadata.register_injector(failure_name, 30, fail) + event_metadata.register_injector( + after_failure_name, + 40, + lambda _event: {"python.injector.after_failure": "added"}, + ) + event_metadata.register_injector(invalid_name, 50, invalid_injector) + try: + scope.event("python-global-injection", metadata={"python.existing": "preserved"}) + await subscribers.flush_async() + finally: + event_metadata.deregister_injector(invalid_name) + event_metadata.deregister_injector(after_failure_name) + event_metadata.deregister_injector(failure_name) + event_metadata.deregister_injector(async_name) + event_metadata.deregister_injector(sync_name) + + event = next(event for event in subscribed_events if event.name == "python-global-injection") + assert event.metadata == { + "python.existing": "preserved", + "python.injector.after_failure": "added", + "python.injector.async": True, + "python.injector.shared": "sync-first", + "python.injector.sync": "added", + } + + +async def test_python_injectors_require_homogeneous_numeric_lists(subscribed_events): + integers_name = f"python-integers-{uuid4()}" + doubles_name = f"python-doubles-{uuid4()}" + mixed_name = f"python-mixed-numbers-{uuid4()}" + mixed_injector = cast( + nemo_relay.EventMetadataInjectorCallback, + lambda _event: {"python.injector.mixed_numbers": [1, 2.5]}, + ) + + event_metadata.register_injector( + integers_name, + 10, + lambda _event: {"python.injector.integers": [1, 2]}, + ) + event_metadata.register_injector( + doubles_name, + 20, + lambda _event: {"python.injector.doubles": [1.0, 2.5]}, + ) + event_metadata.register_injector(mixed_name, 30, mixed_injector) + try: + scope.event("python-homogeneous-numeric-lists") + await subscribers.flush_async() + finally: + event_metadata.deregister_injector(mixed_name) + event_metadata.deregister_injector(doubles_name) + event_metadata.deregister_injector(integers_name) + + event = next(event for event in subscribed_events if event.name == "python-homogeneous-numeric-lists") + assert event.metadata == { + "python.injector.doubles": [1.0, 2.5], + "python.injector.integers": [1, 2], + } + + +async def test_scope_local_python_injector_applies_only_to_owned_events(subscribed_events): + with scope.scope("python-scope-owner", nemo_relay.ScopeType.Agent) as owner: + scope_local.register_event_metadata_injector( + owner, + "python-scope-local-first", + 10, + lambda _event: { + "python.injector.scope_local": "active", + "python.injector.scope_order": "first", + }, + ) + scope_local.register_event_metadata_injector( + owner, + "python-scope-local-later", + 20, + lambda _event: {"python.injector.scope_order": "later"}, + ) + scope.event("python-scope-inside") + with scope.scope("python-scope-child", nemo_relay.ScopeType.Function): + pass + + scope.event("python-scope-outside") + await subscribers.flush_async() + + events = {event.name: event for event in subscribed_events} + assert events["python-scope-inside"].metadata["python.injector.scope_local"] == "active" + assert events["python-scope-inside"].metadata["python.injector.scope_order"] == "first" + assert events["python-scope-child"].metadata["python.injector.scope_local"] == "active" + assert events["python-scope-owner"].metadata["python.injector.scope_local"] == "active" + assert events["python-scope-outside"].metadata is None + + +async def test_scope_local_python_injector_can_be_deregistered_while_owner_is_active(subscribed_events): + with scope.scope("python-scope-deregister-owner", nemo_relay.ScopeType.Agent) as owner: + scope_local.register_event_metadata_injector( + owner, + "python-scope-deregister", + 10, + lambda _event: {"python.injector.scope_local": "active"}, + ) + scope.event("python-scope-before-deregister") + assert scope_local.deregister_event_metadata_injector(owner, "python-scope-deregister") is True + assert scope_local.deregister_event_metadata_injector(owner, "python-scope-deregister") is False + scope.event("python-scope-after-deregister") + + await subscribers.flush_async() + + events = {event.name: event for event in subscribed_events} + assert events["python-scope-before-deregister"].metadata == {"python.injector.scope_local": "active"} + assert events["python-scope-after-deregister"].metadata is None + + +async def test_in_process_python_plugin_registers_configured_injector_and_cleans_up(subscribed_events): + class ConfiguredMetadataPlugin: + def validate(self, _config: nemo_relay.JsonObject): + return None + + def register( + self, + config: nemo_relay.JsonObject, + context: plugin.PluginContext, + ) -> None: + configured = cast(nemo_relay.EventMetadata, config["metadata"]) + context.register_event_metadata_injector( + "configured", + 10, + lambda _event: configured, + ) + + kind = f"python.test_event_metadata.{uuid4()}" + plugin.register(kind, cast(plugin.Plugin, ConfiguredMetadataPlugin())) + try: + await plugin.initialize( + plugin.PluginConfig( + components=[ + plugin.ComponentSpec( + kind=kind, + config={"metadata": {"python.injector.plugin": "configured"}}, + ) + ] + ) + ) + scope.event("python-plugin-configured") + await subscribers.flush_async() + + await plugin.clear_async() + scope.event("python-plugin-cleared") + await subscribers.flush_async() + finally: + await plugin.clear_async() + plugin.deregister(kind) + + events = {event.name: event for event in subscribed_events} + assert events["python-plugin-configured"].metadata == {"python.injector.plugin": "configured"} + assert events["python-plugin-cleared"].metadata is None