diff --git a/Cargo.toml b/Cargo.toml index c62ce5393..4d711cf0a 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -45,3 +45,6 @@ uuid = "=1.18.1" [workspace.lints.rust] missing_docs = "deny" + +[workspace.lints.clippy] +cognitive_complexity = "deny" diff --git a/clippy.toml b/clippy.toml new file mode 100644 index 000000000..e7c65218e --- /dev/null +++ b/clippy.toml @@ -0,0 +1,4 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +cognitive-complexity-threshold = 18 diff --git a/codecov.yml b/codecov.yml index 55b332a21..c52334c14 100644 --- a/codecov.yml +++ b/codecov.yml @@ -135,11 +135,12 @@ ignore: - "**/examples/**" - "**/tests/**" - "crates/cli/tests/" - # CLI TTY shells are exercised by smoke tests, but their prompt loops are - # intentionally split away from testable model modules. - - "crates/cli/src/plugins/mod.rs" - - "crates/cli/src/plugins/dynamic_editor.rs" - - "crates/cli/src/commands/configure/wizard.rs" + # CLI TTY shells are exercised by smoke tests, while deterministic state, + # validation, and persistence behavior remains in the CLI component. + - "crates/cli/src/commands/configure/editor/prompt.rs" + - "crates/cli/src/commands/configure/wizard/prompt.rs" + - "crates/cli/src/plugins/prompt.rs" + - "crates/cli/src/plugins/dynamic_editor/prompt.rs" - "**/tests-js/**" # The Node binding currently reports JS package coverage separately; exclude the # native Rust bridge until we have direct Rust-side coverage for this crate. diff --git a/crates/adaptive/src/plugin_component.rs b/crates/adaptive/src/plugin_component.rs index eae7759d1..1d1b530a9 100644 --- a/crates/adaptive/src/plugin_component.rs +++ b/crates/adaptive/src/plugin_component.rs @@ -215,37 +215,7 @@ fn validate_adaptive_plugin_config_with_policy( ); } - if let Some(state_json) = plugin_config.get("state").and_then(Json::as_object) { - validate_unknown_fields( - &mut diagnostics, - &config.policy, - Some("state".to_string()), - state_json, - &["backend"], - ); - if let Some(backend_json) = state_json.get("backend").and_then(Json::as_object) { - validate_unknown_fields( - &mut diagnostics, - &config.policy, - Some("backend".to_string()), - backend_json, - &["kind", "config"], - ); - let backend_kind = backend_json - .get("kind") - .and_then(Json::as_str) - .unwrap_or_default(); - if let Some(backend_config_json) = backend_json.get("config").and_then(Json::as_object) - { - validate_backend_config_fields( - &mut diagnostics, - &config.policy, - backend_kind, - backend_config_json, - ); - } - } - } + validate_adaptive_state_section(&mut diagnostics, &config.policy, plugin_config); if let Some(telemetry_json) = plugin_config.get("telemetry").and_then(Json::as_object) { validate_unknown_fields( @@ -303,52 +273,94 @@ fn validate_adaptive_plugin_config_with_policy( ); } - if let Some(response_cache_json) = plugin_config + validate_response_cache_section(&mut diagnostics, &config.policy, plugin_config); + + diagnostics.extend(AdaptiveRuntime::validate_config(&config).diagnostics); + diagnostics +} + +fn validate_adaptive_state_section( + diagnostics: &mut Vec, + policy: &ConfigPolicy, + plugin_config: &Map, +) { + let Some(state_json) = plugin_config.get("state").and_then(Json::as_object) else { + return; + }; + validate_unknown_fields( + diagnostics, + policy, + Some("state".to_string()), + state_json, + &["backend"], + ); + let Some(backend_json) = state_json.get("backend").and_then(Json::as_object) else { + return; + }; + validate_unknown_fields( + diagnostics, + policy, + Some("backend".to_string()), + backend_json, + &["kind", "config"], + ); + let backend_kind = backend_json + .get("kind") + .and_then(Json::as_str) + .unwrap_or_default(); + if let Some(backend_config_json) = backend_json.get("config").and_then(Json::as_object) { + validate_backend_config_fields(diagnostics, policy, backend_kind, backend_config_json); + } +} + +fn validate_response_cache_section( + diagnostics: &mut Vec, + policy: &ConfigPolicy, + plugin_config: &Map, +) { + let Some(response_cache_json) = plugin_config .get("response_cache") .and_then(Json::as_object) - { + else { + return; + }; + validate_unknown_fields( + diagnostics, + policy, + Some("response_cache".to_string()), + response_cache_json, + &[ + "ttl_seconds", + "namespace", + "priority", + "bypass_rate", + "cache_nondeterministic", + "key_strategy", + "header_allowlist", + "backend", + ], + ); + if let Some(backend_json) = response_cache_json.get("backend").and_then(Json::as_object) { validate_unknown_fields( - &mut diagnostics, - &config.policy, - Some("response_cache".to_string()), - response_cache_json, - &[ - "ttl_seconds", - "namespace", - "priority", - "bypass_rate", - "cache_nondeterministic", - "key_strategy", - "header_allowlist", - "backend", - ], + diagnostics, + policy, + Some("response_cache.backend".to_string()), + backend_json, + &["kind", "config"], ); - if let Some(backend_json) = response_cache_json.get("backend").and_then(Json::as_object) { - validate_unknown_fields( - &mut diagnostics, - &config.policy, - Some("response_cache.backend".to_string()), - backend_json, - &["kind", "config"], + let backend_kind = backend_json + .get("kind") + .and_then(Json::as_str) + .unwrap_or("in_memory"); + if let Some(backend_config_json) = backend_json.get("config").and_then(Json::as_object) { + validate_response_cache_backend_config_fields( + diagnostics, + policy, + backend_kind, + backend_config_json, ); - let backend_kind = backend_json - .get("kind") - .and_then(Json::as_str) - .unwrap_or("in_memory"); - if let Some(backend_config_json) = backend_json.get("config").and_then(Json::as_object) - { - validate_response_cache_backend_config_fields( - &mut diagnostics, - &config.policy, - backend_kind, - backend_config_json, - ); - } } } - - diagnostics.extend(AdaptiveRuntime::validate_config(&config).diagnostics); - diagnostics } fn validate_response_cache_backend_config_fields( diff --git a/crates/adaptive/src/response_cache/key.rs b/crates/adaptive/src/response_cache/key.rs index 895d54102..36eeca3d1 100644 --- a/crates/adaptive/src/response_cache/key.rs +++ b/crates/adaptive/src/response_cache/key.rs @@ -56,48 +56,8 @@ pub fn build_cache_key( request: &LlmRequest, config: &ResponseCacheConfig, ) -> KeyOutcome { - // Unparseable bodies arrive as null; they would all share one key. - if request.content.is_null() { - return KeyOutcome::Bypass("unparseable_body"); - } - // Cacheability gates run on the RAW request, so they are correct regardless - // of which codec (if any) decodes the body — a chat codec may park `store` - // in `extra` rather than the typed field, so we must not rely on the decode. - if let Some(object) = request.content.as_object() { - // Any present, non-`false` `store` opts into server-side persistence — - // bypass even a malformed non-boolean rather than risk caching a stateful - // call (whose result is otherwise keyed with `store` stripped). - if object - .get("store") - .is_some_and(|value| !matches!(value, Json::Bool(false) | Json::Null)) - { - return KeyOutcome::Bypass("stateful_store"); - } - if object.contains_key("previous_response_id") { - return KeyOutcome::Bypass("stateful_previous_response_id"); - } - // Server-side conversation state the key cannot see. - if object.contains_key("conversation") || object.contains_key("container") { - return KeyOutcome::Bypass("stateful_conversation"); - } - // Responses persists by default; only an explicit opt-out is stateless. - // A `prompt` object is the Responses prompt-template reference; a bare - // string `prompt` is a completions body with no server-side state. - if (object.contains_key("input") - || object.contains_key("instructions") - || object.get("prompt").is_some_and(Json::is_object)) - && !object - .get("store") - .is_some_and(|store| store == &Json::Bool(false)) - { - return KeyOutcome::Bypass("stateful_store"); - } - } - // Toggle off = explicit temperature 0 only; absent defaults to sampling. - if !config.cache_nondeterministic - && request_temperature(&request.content).is_none_or(|temperature| temperature > 0.0) - { - return KeyOutcome::Bypass("nondeterministic_temperature"); + if let Some(reason) = cache_bypass_reason(request, config) { + return KeyOutcome::Bypass(reason); } // Body to fingerprint: the decoded/normalized form when a surface resolves @@ -136,6 +96,44 @@ pub fn build_cache_key( } } +fn cache_bypass_reason(request: &LlmRequest, config: &ResponseCacheConfig) -> Option<&'static str> { + if request.content.is_null() { + return Some("unparseable_body"); + } + if let Some(reason) = request + .content + .as_object() + .and_then(stateful_request_bypass_reason) + { + return Some(reason); + } + (!config.cache_nondeterministic + && request_temperature(&request.content).is_none_or(|temperature| temperature > 0.0)) + .then_some("nondeterministic_temperature") +} + +fn stateful_request_bypass_reason(object: &Map) -> Option<&'static str> { + if object + .get("store") + .is_some_and(|value| !matches!(value, Json::Bool(false) | Json::Null)) + { + return Some("stateful_store"); + } + if object.contains_key("previous_response_id") { + return Some("stateful_previous_response_id"); + } + if object.contains_key("conversation") || object.contains_key("container") { + return Some("stateful_conversation"); + } + let responses_surface = object.contains_key("input") + || object.contains_key("instructions") + || object.get("prompt").is_some_and(Json::is_object); + let explicitly_stateless = object + .get("store") + .is_some_and(|store| store == &Json::Bool(false)); + (responses_surface && !explicitly_stateless).then_some("stateful_store") +} + /// Preserves which OpenAI Chat token-cap field the caller sent. /// /// The Chat codec normalizes both spellings into `GenerationParams.max_tokens`, diff --git a/crates/adaptive/src/response_cache/replay.rs b/crates/adaptive/src/response_cache/replay.rs index ed3d1afe6..02a8bf0d4 100644 --- a/crates/adaptive/src/response_cache/replay.rs +++ b/crates/adaptive/src/response_cache/replay.rs @@ -192,44 +192,7 @@ fn synthesize_chat_chunks(aggregate: &Json) -> Vec { .cloned() .unwrap_or_default(); for (position, choice) in choices.iter().enumerate() { - let index = choice - .get("index") - .and_then(Json::as_u64) - .unwrap_or(position as u64); - let message = choice.get("message").cloned().unwrap_or(json!({})); - if let Some(role) = message.get("role") { - chunks.push(base( - json!([{"index": index, "delta": {"role": role}, "finish_reason": null}]), - )); - } - if let Some(content) = message.get("content").and_then(Json::as_str) - && !content.is_empty() - { - chunks.push(base( - json!([{"index": index, "delta": {"content": content}, "finish_reason": null}]), - )); - } - if let Some(tool_calls) = message.get("tool_calls").and_then(Json::as_array) { - let deltas: Vec = tool_calls - .iter() - .enumerate() - .map(|(call_index, call)| { - let mut delta = call.clone(); - if let Some(map) = delta.as_object_mut() { - map.entry("index".to_string()) - .or_insert(json!(call_index as u64)); - } - delta - }) - .collect(); - chunks.push(base( - json!([{"index": index, "delta": {"tool_calls": deltas}, "finish_reason": null}]), - )); - } - let finish = choice.get("finish_reason").cloned().unwrap_or(Json::Null); - chunks.push(base( - json!([{"index": index, "delta": {}, "finish_reason": finish}]), - )); + chunks.extend(synthesize_chat_choice_chunks(&base, position, choice)); } if let Some(usage) = aggregate.get("usage") { let mut usage_chunk = base(json!([])); @@ -241,6 +204,53 @@ fn synthesize_chat_chunks(aggregate: &Json) -> Vec { chunks } +fn synthesize_chat_choice_chunks( + base: &impl Fn(Json) -> Json, + position: usize, + choice: &Json, +) -> Vec { + let index = choice + .get("index") + .and_then(Json::as_u64) + .unwrap_or(position as u64); + let message = choice.get("message").cloned().unwrap_or(json!({})); + let mut chunks = Vec::new(); + if let Some(role) = message.get("role") { + chunks.push(base( + json!([{"index": index, "delta": {"role": role}, "finish_reason": null}]), + )); + } + if let Some(content) = message.get("content").and_then(Json::as_str) + && !content.is_empty() + { + chunks.push(base( + json!([{"index": index, "delta": {"content": content}, "finish_reason": null}]), + )); + } + if let Some(tool_calls) = message.get("tool_calls").and_then(Json::as_array) { + let deltas = tool_calls + .iter() + .enumerate() + .map(|(call_index, call)| { + let mut delta = call.clone(); + if let Some(map) = delta.as_object_mut() { + map.entry("index".to_string()) + .or_insert(json!(call_index as u64)); + } + delta + }) + .collect::>(); + chunks.push(base( + json!([{"index": index, "delta": {"tool_calls": deltas}, "finish_reason": null}]), + )); + } + let finish = choice.get("finish_reason").cloned().unwrap_or(Json::Null); + chunks.push(base( + json!([{"index": index, "delta": {}, "finish_reason": finish}]), + )); + chunks +} + /// OpenAI Responses: a `response.created` snapshot, one `response.output_item.done` /// per output item, and a `response.completed` carrying the full stored aggregate /// (the collector keeps the last snapshot wholesale, so reassembly is exact). diff --git a/crates/adaptive/tests/unit/plugin_component_tests.rs b/crates/adaptive/tests/unit/plugin_component_tests.rs index 993688271..e051bb301 100644 --- a/crates/adaptive/tests/unit/plugin_component_tests.rs +++ b/crates/adaptive/tests/unit/plugin_component_tests.rs @@ -271,6 +271,11 @@ fn response_cache_backend_validation_uses_the_default_backend_kind() { #[test] fn adaptive_to_plugin_error_maps_all_non_redis_variants() { + assert_adaptive_config_and_lookup_errors(); + assert_adaptive_internal_and_serialization_errors(); +} + +fn assert_adaptive_config_and_lookup_errors() { assert!(matches!( adaptive_to_plugin_error(AdaptiveError::InvalidConfig("bad".into())), nemo_relay::plugin::PluginError::InvalidConfig(message) if message == "bad" @@ -283,6 +288,9 @@ fn adaptive_to_plugin_error_maps_all_non_redis_variants() { adaptive_to_plugin_error(AdaptiveError::Storage("store".into())), nemo_relay::plugin::PluginError::Internal(message) if message == "store" )); +} + +fn assert_adaptive_internal_and_serialization_errors() { assert!(matches!( adaptive_to_plugin_error(AdaptiveError::Internal("internal".into())), nemo_relay::plugin::PluginError::Internal(message) if message == "internal" diff --git a/crates/adaptive/tests/unit/response_cache/intercept_tests.rs b/crates/adaptive/tests/unit/response_cache/intercept_tests.rs index 67e39a8a3..51eda4ccf 100644 --- a/crates/adaptive/tests/unit/response_cache/intercept_tests.rs +++ b/crates/adaptive/tests/unit/response_cache/intercept_tests.rs @@ -86,3 +86,89 @@ async fn write_behind_returns_eof_before_cache_commit_completes() { .expect("detached cache commit must resume after release") .expect("detached cache commit must run to completion"); } + +#[test] +fn response_cache_fidelity_helpers_cover_all_rejection_shapes() { + assert_uncollected_response_field_shapes(); + assert_aggregate_replay_and_stream_completion(); + assert_inband_error_and_content_detection(); + assert_error_response_and_bypass_detection(); +} + +fn assert_uncollected_response_field_shapes() { + for chunk in [ + json!(false), + json!({"choices": false}), + json!({"choices": [false]}), + json!({"choices": [{"index": "zero"}]}), + json!({"choices": [{"finish_reason": false}]}), + json!({"choices": [{"delta": false}]}), + json!({"choices": [{"logprobs": {}}]}), + json!({"choices": [{"extension": true}]}), + json!({"choices": [{"delta": {"tool_calls": [false]}}]}), + json!({"choices": [{"delta": {"tool_calls": [{"index": "zero"}]}}]}), + json!({"choices": [{"delta": {"tool_calls": [{"id": false}]}}]}), + json!({"choices": [{"delta": {"tool_calls": [{"function": false}]}}]}), + json!({"choices": [{"delta": {"tool_calls": [{"extension": true}]}}]}), + ] { + assert!(chunk_has_uncollected_response_fields(&chunk), "{chunk}"); + } + assert!(!chunk_has_uncollected_response_fields(&json!({ + "choices": null, + "metadata": null + }))); +} + +fn assert_aggregate_replay_and_stream_completion() { + assert!(aggregate_replay_lossy(&json!({ + "content": [{"type": "thinking"}] + }))); + assert!(aggregate_replay_lossy(&json!({ + "choices": [{"message": {"content": null, "tool_calls": []}}] + }))); + assert!(!aggregate_replay_lossy(&json!({ + "choices": [{"message": {"content": "answer"}}] + }))); + + let mut completion = StreamCompletion::default(); + completion.observe(&json!({ + "choices": [ + {"index": 0, "finish_reason": "stop"}, + {"index": 1, "finish_reason": null} + ] + })); + assert!(!completion.is_terminal()); + completion.observe(&json!({"choices": [{"index": 1, "finish_reason": "stop"}]})); + assert!(completion.is_terminal()); + + let mut stopped = StreamCompletion::default(); + stopped.observe(&json!({"type": "response.completed"})); + assert!(stopped.is_terminal()); +} + +fn assert_inband_error_and_content_detection() { + assert!(chunk_is_inband_error(&json!({"error": "bad"}))); + assert!(chunk_is_inband_error(&json!({"type": "response.failed"}))); + assert!(!chunk_is_inband_error(&json!({"error": null}))); + assert!(aggregate_has_no_content(&json!({}))); + assert!(!aggregate_has_no_content(&json!({"output": [1]}))); +} + +fn assert_error_response_and_bypass_detection() { + assert!(!is_error_response(&json!(false))); + assert!(is_error_response(&json!({"error": "bad"}))); + for status in [ + "failed", + "cancelled", + "canceled", + "incomplete", + "in_progress", + "queued", + ] { + assert!(is_error_response(&json!({"status": status}))); + } + assert!(!is_error_response(&json!({"status": "completed"}))); + assert!(!should_bypass(0.0)); + assert!(should_bypass(1.0)); + assert!(next_unit_f64().is_finite()); +} diff --git a/crates/adaptive/tests/unit/response_cache/store_tests.rs b/crates/adaptive/tests/unit/response_cache/store_tests.rs index 43a90188f..50446d0b9 100644 --- a/crates/adaptive/tests/unit/response_cache/store_tests.rs +++ b/crates/adaptive/tests/unit/response_cache/store_tests.rs @@ -214,3 +214,23 @@ async fn an_entry_larger_than_the_budget_is_not_cached_and_keeps_existing_entrie ); assert_eq!(store.total_bytes(), 0); } + +#[tokio::test] +async fn repeated_replacement_compacts_stale_insertion_order_nodes() { + let store = InMemoryCacheStore::new(BIG); + for generation in 0..70 { + store + .set( + "stable-key", + entry("stable-key", generation, u64::MAX), + Duration::MAX, + ) + .await + .unwrap(); + } + + let guard = store.inner.lock().unwrap(); + assert_eq!(guard.map.len(), 1); + assert_eq!(guard.order.len(), 4); + assert_eq!(guard.next_generation, 70); +} diff --git a/crates/adaptive/tests/unit/trie/builder_tests.rs b/crates/adaptive/tests/unit/trie/builder_tests.rs index d47abc613..9d5ac1e94 100644 --- a/crates/adaptive/tests/unit/trie/builder_tests.rs +++ b/crates/adaptive/tests/unit/trie/builder_tests.rs @@ -348,6 +348,36 @@ fn test_score_to_sensitivity_no_samples() { assert_eq!(score_to_sensitivity(&acc, 5), None); } +#[test] +fn sensitivity_helpers_cover_empty_zero_duration_and_parallel_edges() { + let mut contexts = Vec::new(); + compute_sensitivity_scores(&mut contexts, &SensitivityConfig::default()); + assert!(compute_logical_positions(&contexts).is_empty()); + + let context = LlmCallContext { + path: vec!["edge".into()], + call_index: 0, + remaining_calls: 0, + time_to_next_ms: None, + output_tokens: 0, + call_duration_s: 0.0, + workflow_duration_s: 0.0, + parallel_slack_ratio: 0.4, + sensitivity_score: 0.0, + span_start_time: 0.0, + span_end_time: 0.0, + }; + assert_eq!(critical_path_weight(&context), 1.0); + assert_eq!(fanout_score(0, 0), 0.0); + assert_eq!(position_score(0, 1), 1.0); + assert_eq!(parallel_penalty(0.4, &HashMap::from([(0, 2)]), 0), 0.45); + + let mut run = make_test_run(1, 0); + run.ended_at = None; + run.calls[0].ended_at = None; + assert!(extract_llm_contexts(&run).is_empty()); +} + // ----------------------------------------------------------------------- // PredictionTrieBuilder integration tests // ----------------------------------------------------------------------- diff --git a/crates/cli/src/commands/configure/editor.rs b/crates/cli/src/commands/configure/editor.rs index 5be154746..93504a700 100644 --- a/crates/cli/src/commands/configure/editor.rs +++ b/crates/cli/src/commands/configure/editor.rs @@ -3,21 +3,19 @@ //! Interactive editor for the non-agent sections of Relay's `config.toml`. -use std::io::IsTerminal; use std::path::{Path, PathBuf}; -use dialoguer::theme::ColorfulTheme; -use dialoguer::{Input, Password, Select}; use nemo_relay::logging::MAX_FILE_SINK_QUEUE_ENTRIES; use toml_edit::{ArrayOfTables, DocumentMut, Item, Table, Value, value}; use super::ConfigEditCommand; use crate::error::CliError; -const EDIT_CANCELLED_MESSAGE: &str = "configuration edit cancelled — no config saved"; const LOG_LEVELS: &[&str] = &["error", "warn", "info", "debug", "trace"]; const LOG_FORMATS: &[&str] = &["human", "jsonl"]; +mod prompt; + #[derive(Debug, Clone, Copy, PartialEq, Eq)] enum TargetScope { User, @@ -41,39 +39,7 @@ pub(super) fn edit( command: ConfigEditCommand, explicit_path: Option, ) -> Result<(), CliError> { - ensure_tty()?; - let (scope, path) = resolve_edit_target(&command, explicit_path)?; - let mut document = ConfigDocument::read(path)?; - let theme = ColorfulTheme::default(); - - crate::banner::print_intro(); - println!(" Editing config at {}", document.path().display()); - println!(" Secrets are never displayed. Choose Save to write changes."); - println!(); - - loop { - let choices = [ - format!("Gateway limits ({})", document.gateway_summary()), - format!("Provider upstreams ({})", document.upstream_summary()), - format!("Operational logging ({})", document.logging_summary()), - "Preview".into(), - "Save".into(), - "Cancel".into(), - ]; - match select(&theme, "config.toml", &choices)? { - 0 => edit_gateway(&theme, &mut document)?, - 1 => edit_upstream(&theme, &mut document)?, - 2 => edit_logging(&theme, &mut document)?, - 3 => print_preview(&document), - 4 => { - document.write(scope)?; - println!(" ✓ Saved {}", document.path().display()); - return Ok(()); - } - 5 => return Err(CliError::Config(EDIT_CANCELLED_MESSAGE.into())), - _ => unreachable!("select returns an in-range index"), - } - } + prompt::edit(command, explicit_path) } fn resolve_edit_target( @@ -92,10 +58,6 @@ fn resolve_edit_target( Ok((scope, path)) } -fn ensure_tty() -> Result<(), CliError> { - ensure_tty_with(std::io::stdin().is_terminal()) -} - fn ensure_tty_with(stdin_is_terminal: bool) -> Result<(), CliError> { if stdin_is_terminal { Ok(()) @@ -106,407 +68,6 @@ fn ensure_tty_with(stdin_is_terminal: bool) -> Result<(), CliError> { } } -fn select(theme: &ColorfulTheme, prompt: &str, choices: &[String]) -> Result { - Select::with_theme(theme) - .with_prompt(prompt) - .items(choices) - .default(0) - .interact() - .map_err(prompt_error) -} - -fn choose_action(theme: &ColorfulTheme, configured: bool) -> Result { - let choices = if configured { - vec!["Set or replace".into(), "Clear".into(), "Back".into()] - } else { - vec!["Set".into(), "Back".into()] - }; - select(theme, "Action", &choices) -} - -fn edit_gateway(theme: &ColorfulTheme, document: &mut ConfigDocument) -> Result<(), CliError> { - loop { - let choices = [ - format!( - "Maximum hook payload bytes: {}", - document.integer_summary("gateway", "max_hook_payload_bytes") - ), - format!( - "Maximum passthrough body bytes: {}", - document.integer_summary("gateway", "max_passthrough_body_bytes") - ), - "Back".into(), - ]; - match select(theme, "Gateway limits", &choices)? { - 0 => edit_positive_integer(theme, document, "gateway", "max_hook_payload_bytes")?, - 1 => edit_positive_integer(theme, document, "gateway", "max_passthrough_body_bytes")?, - 2 => return Ok(()), - _ => unreachable!(), - } - } -} - -fn edit_upstream(theme: &ColorfulTheme, document: &mut ConfigDocument) -> Result<(), CliError> { - loop { - let choices = [ - format!( - "OpenAI base URL: {}", - document.string_summary("upstream", "openai_base_url") - ), - format!( - "OpenAI authorization header: {}", - document.secret_summary("openai_auth_header") - ), - format!( - "Anthropic base URL: {}", - document.string_summary("upstream", "anthropic_base_url") - ), - format!( - "Anthropic authorization header: {}", - document.secret_summary("anthropic_auth_header") - ), - "Back".into(), - ]; - match select(theme, "Provider upstreams", &choices)? { - 0 => edit_string(theme, document, "upstream", "openai_base_url")?, - 1 => edit_secret(theme, document, "openai_auth_header")?, - 2 => edit_string(theme, document, "upstream", "anthropic_base_url")?, - 3 => edit_secret(theme, document, "anthropic_auth_header")?, - 4 => return Ok(()), - _ => unreachable!(), - } - } -} - -fn edit_logging(theme: &ColorfulTheme, document: &mut ConfigDocument) -> Result<(), CliError> { - loop { - let choices = [ - format!("Level: {}", document.string_summary("logging", "level")), - format!( - "Stderr format: {}", - document.string_summary("logging", "stderr_format") - ), - format!( - "Flush interval (ms): {}", - document.integer_summary("logging", "flush_interval_millis") - ), - format!("File sinks ({})", document.sink_count()), - "Back".into(), - ]; - match select(theme, "Operational logging", &choices)? { - 0 => edit_enum(theme, document, "logging", "level", LOG_LEVELS)?, - 1 => edit_enum(theme, document, "logging", "stderr_format", LOG_FORMATS)?, - 2 => edit_nonnegative_integer(theme, document, "logging", "flush_interval_millis")?, - 3 => edit_sinks(theme, document)?, - 4 => return Ok(()), - _ => unreachable!(), - } - } -} - -fn edit_positive_integer( - theme: &ColorfulTheme, - document: &mut ConfigDocument, - section: &str, - key: &str, -) -> Result<(), CliError> { - let configured = document.has_key(section, key); - match choose_action(theme, configured)? { - 0 => { - let value = prompt_u64(theme, "Value in bytes", document.integer(section, key))?; - document.set_positive_integer(section, key, value)?; - } - 1 if configured => document.clear_key(section, key)?, - _ => {} - } - Ok(()) -} - -fn edit_nonnegative_integer( - theme: &ColorfulTheme, - document: &mut ConfigDocument, - section: &str, - key: &str, -) -> Result<(), CliError> { - let configured = document.has_key(section, key); - match choose_action(theme, configured)? { - 0 => { - let value = prompt_u64( - theme, - "Milliseconds (0 flushes on shutdown)", - document.integer(section, key), - )?; - document.set_integer(section, key, value)?; - } - 1 if configured => document.clear_key(section, key)?, - _ => {} - } - Ok(()) -} - -fn edit_string( - theme: &ColorfulTheme, - document: &mut ConfigDocument, - section: &str, - key: &str, -) -> Result<(), CliError> { - let configured = document.has_key(section, key); - match choose_action(theme, configured)? { - 0 => { - let default = document.string(section, key).unwrap_or_default(); - let value = Input::::with_theme(theme) - .with_prompt("Value") - .with_initial_text(default) - .validate_with(|value: &String| { - if value.trim().is_empty() { - Err("value must not be empty; use Clear to remove it") - } else { - Ok(()) - } - }) - .interact_text() - .map_err(prompt_error)?; - document.set_string(section, key, value)?; - } - 1 if configured => document.clear_key(section, key)?, - _ => {} - } - Ok(()) -} - -fn edit_secret( - theme: &ColorfulTheme, - document: &mut ConfigDocument, - key: &str, -) -> Result<(), CliError> { - let configured = document.has_key("upstream", key); - match choose_action(theme, configured)? { - 0 => { - let value = Password::with_theme(theme) - .with_prompt("Authorization header value") - .allow_empty_password(false) - .interact() - .map_err(prompt_error)?; - document.set_auth_header(key, value)?; - } - 1 if configured => document.clear_key("upstream", key)?, - _ => {} - } - Ok(()) -} - -fn edit_enum( - theme: &ColorfulTheme, - document: &mut ConfigDocument, - section: &str, - key: &str, - values: &[&str], -) -> Result<(), CliError> { - let configured = document.has_key(section, key); - match choose_action(theme, configured)? { - 0 => { - let current = document.string(section, key); - let default = current - .as_deref() - .and_then(|current| values.iter().position(|value| *value == current)) - .unwrap_or(0); - let selected = Select::with_theme(theme) - .with_prompt("Value") - .items(values) - .default(default) - .interact() - .map_err(prompt_error)?; - document.set_enum(section, key, values[selected], values)?; - } - 1 if configured => document.clear_key(section, key)?, - _ => {} - } - Ok(()) -} - -fn edit_sinks(theme: &ColorfulTheme, document: &mut ConfigDocument) -> Result<(), CliError> { - loop { - let mut choices = document - .sink_labels() - .into_iter() - .map(|label| format!("Edit {label}")) - .collect::>(); - let sink_count = choices.len(); - choices.push("Add file sink".into()); - choices.push("Back".into()); - match select(theme, "File sinks", &choices)? { - index if index < sink_count => edit_sink(theme, document, index)?, - index if index == sink_count => { - let path = Input::::with_theme(theme) - .with_prompt("File path") - .validate_with(|value: &String| { - if value.trim().is_empty() { - Err("value must not be empty".to_owned()) - } else { - Ok(()) - } - }) - .interact_text() - .map_err(prompt_error)?; - document.add_sink(path)?; - } - _ => return Ok(()), - } - } -} - -fn edit_sink( - theme: &ColorfulTheme, - document: &mut ConfigDocument, - index: usize, -) -> Result<(), CliError> { - loop { - let choices = [ - format!("Path: {}", document.sink_string_summary(index, "path")), - format!("Level: {}", document.sink_string_summary(index, "level")), - format!("Format: {}", document.sink_string_summary(index, "format")), - format!( - "Queue capacity: {}", - document.sink_integer_summary(index, "queue_capacity") - ), - format!("Rotation: {}", document.sink_rotation_summary(index)), - "Remove sink".into(), - "Back".into(), - ]; - match select(theme, "File sink", &choices)? { - 0 => edit_sink_path(theme, document, index)?, - 1 => edit_sink_enum(theme, document, index, "level", LOG_LEVELS)?, - 2 => edit_sink_enum(theme, document, index, "format", LOG_FORMATS)?, - 3 => edit_sink_queue_capacity(theme, document, index)?, - 4 => edit_sink_rotation(theme, document, index)?, - 5 => { - document.remove_sink(index)?; - return Ok(()); - } - _ => return Ok(()), - } - } -} - -fn edit_sink_path( - theme: &ColorfulTheme, - document: &mut ConfigDocument, - index: usize, -) -> Result<(), CliError> { - let current = document.sink_string(index, "path").unwrap_or_default(); - let value = Input::::with_theme(theme) - .with_prompt("File path") - .with_initial_text(current) - .validate_with(|value: &String| { - if value.trim().is_empty() { - Err("value must not be empty".to_owned()) - } else { - Ok(()) - } - }) - .interact_text() - .map_err(prompt_error)?; - document.set_sink_string(index, "path", value) -} - -fn edit_sink_enum( - theme: &ColorfulTheme, - document: &mut ConfigDocument, - index: usize, - key: &str, - values: &[&str], -) -> Result<(), CliError> { - let configured = document.sink_has_key(index, key)?; - match choose_action(theme, configured)? { - 0 => { - let default = document - .sink_string(index, key) - .as_deref() - .and_then(|current| values.iter().position(|value| *value == current)) - .unwrap_or(0); - let selected = Select::with_theme(theme) - .with_prompt("Value") - .items(values) - .default(default) - .interact() - .map_err(prompt_error)?; - document.set_sink_enum(index, key, values[selected], values)?; - } - 1 if configured => document.clear_sink_key(index, key)?, - _ => {} - } - Ok(()) -} - -fn edit_sink_queue_capacity( - theme: &ColorfulTheme, - document: &mut ConfigDocument, - index: usize, -) -> Result<(), CliError> { - let configured = document.sink_has_key(index, "queue_capacity")?; - match choose_action(theme, configured)? { - 0 => { - let value = prompt_u64( - theme, - "Queue entries", - document.sink_integer(index, "queue_capacity"), - )?; - document.set_sink_queue_capacity(index, value)?; - } - 1 if configured => document.clear_sink_key(index, "queue_capacity")?, - _ => {} - } - Ok(()) -} - -fn edit_sink_rotation( - theme: &ColorfulTheme, - document: &mut ConfigDocument, - index: usize, -) -> Result<(), CliError> { - let configured = document.sink_has_key(index, "max_file_size_bytes")? - || document.sink_has_key(index, "retained_files")?; - match choose_action(theme, configured)? { - 0 => { - let size = prompt_u64( - theme, - "Maximum file size in bytes", - document.sink_integer(index, "max_file_size_bytes"), - )?; - let retained = prompt_u64( - theme, - "Retained backup files", - document.sink_integer(index, "retained_files"), - )?; - document.set_sink_rotation(index, size, retained)?; - } - 1 if configured => document.clear_sink_rotation(index)?, - _ => {} - } - Ok(()) -} - -fn prompt_u64(theme: &ColorfulTheme, prompt: &str, current: Option) -> Result { - let mut input = Input::::with_theme(theme).with_prompt(prompt); - if let Some(current) = current { - input = input.with_initial_text(current.to_string()); - } - input.interact_text().map_err(prompt_error) -} - -fn prompt_error(error: dialoguer::Error) -> CliError { - CliError::Config(format!("configuration edit error: {error}")) -} - -fn print_preview(document: &ConfigDocument) { - println!(); - println!(" ─── Preview ─────────────────────────────────────────────"); - for line in document.preview().lines() { - println!(" {line}"); - } - println!(); -} - struct ConfigDocument { path: PathBuf, document: DocumentMut, diff --git a/crates/cli/src/commands/configure/editor/prompt.rs b/crates/cli/src/commands/configure/editor/prompt.rs new file mode 100644 index 000000000..3f843c9d4 --- /dev/null +++ b/crates/cli/src/commands/configure/editor/prompt.rs @@ -0,0 +1,461 @@ +// SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +//! Terminal-only prompt adapter for the interactive config editor. + +use std::io::IsTerminal; +use std::path::PathBuf; + +use dialoguer::theme::ColorfulTheme; +use dialoguer::{Input, Password, Select}; + +use super::{ + ConfigDocument, ConfigEditCommand, LOG_FORMATS, LOG_LEVELS, ensure_tty_with, + resolve_edit_target, +}; +use crate::error::CliError; + +const EDIT_CANCELLED_MESSAGE: &str = "configuration edit cancelled — no config saved"; +pub(super) fn edit( + command: ConfigEditCommand, + explicit_path: Option, +) -> Result<(), CliError> { + ensure_tty()?; + let (scope, path) = resolve_edit_target(&command, explicit_path)?; + let mut document = ConfigDocument::read(path)?; + let theme = ColorfulTheme::default(); + + crate::banner::print_intro(); + println!(" Editing config at {}", document.path().display()); + println!(" Secrets are never displayed. Choose Save to write changes."); + println!(); + + loop { + let choices = [ + format!("Gateway limits ({})", document.gateway_summary()), + format!("Provider upstreams ({})", document.upstream_summary()), + format!("Operational logging ({})", document.logging_summary()), + "Preview".into(), + "Save".into(), + "Cancel".into(), + ]; + match select(&theme, "config.toml", &choices)? { + 0 => edit_gateway(&theme, &mut document)?, + 1 => edit_upstream(&theme, &mut document)?, + 2 => edit_logging(&theme, &mut document)?, + 3 => print_preview(&document), + 4 => { + document.write(scope)?; + println!(" ✓ Saved {}", document.path().display()); + return Ok(()); + } + 5 => return Err(CliError::Config(EDIT_CANCELLED_MESSAGE.into())), + _ => unreachable!("select returns an in-range index"), + } + } +} + +fn ensure_tty() -> Result<(), CliError> { + ensure_tty_with(std::io::stdin().is_terminal()) +} + +fn select(theme: &ColorfulTheme, prompt: &str, choices: &[String]) -> Result { + Select::with_theme(theme) + .with_prompt(prompt) + .items(choices) + .default(0) + .interact() + .map_err(prompt_error) +} + +fn choose_action(theme: &ColorfulTheme, configured: bool) -> Result { + let choices = if configured { + vec!["Set or replace".into(), "Clear".into(), "Back".into()] + } else { + vec!["Set".into(), "Back".into()] + }; + select(theme, "Action", &choices) +} + +fn edit_gateway(theme: &ColorfulTheme, document: &mut ConfigDocument) -> Result<(), CliError> { + loop { + let choices = [ + format!( + "Maximum hook payload bytes: {}", + document.integer_summary("gateway", "max_hook_payload_bytes") + ), + format!( + "Maximum passthrough body bytes: {}", + document.integer_summary("gateway", "max_passthrough_body_bytes") + ), + "Back".into(), + ]; + match select(theme, "Gateway limits", &choices)? { + 0 => edit_positive_integer(theme, document, "gateway", "max_hook_payload_bytes")?, + 1 => edit_positive_integer(theme, document, "gateway", "max_passthrough_body_bytes")?, + 2 => return Ok(()), + _ => unreachable!(), + } + } +} + +fn edit_upstream(theme: &ColorfulTheme, document: &mut ConfigDocument) -> Result<(), CliError> { + loop { + let choices = [ + format!( + "OpenAI base URL: {}", + document.string_summary("upstream", "openai_base_url") + ), + format!( + "OpenAI authorization header: {}", + document.secret_summary("openai_auth_header") + ), + format!( + "Anthropic base URL: {}", + document.string_summary("upstream", "anthropic_base_url") + ), + format!( + "Anthropic authorization header: {}", + document.secret_summary("anthropic_auth_header") + ), + "Back".into(), + ]; + match select(theme, "Provider upstreams", &choices)? { + 0 => edit_string(theme, document, "upstream", "openai_base_url")?, + 1 => edit_secret(theme, document, "openai_auth_header")?, + 2 => edit_string(theme, document, "upstream", "anthropic_base_url")?, + 3 => edit_secret(theme, document, "anthropic_auth_header")?, + 4 => return Ok(()), + _ => unreachable!(), + } + } +} + +fn edit_logging(theme: &ColorfulTheme, document: &mut ConfigDocument) -> Result<(), CliError> { + loop { + let choices = [ + format!("Level: {}", document.string_summary("logging", "level")), + format!( + "Stderr format: {}", + document.string_summary("logging", "stderr_format") + ), + format!( + "Flush interval (ms): {}", + document.integer_summary("logging", "flush_interval_millis") + ), + format!("File sinks ({})", document.sink_count()), + "Back".into(), + ]; + match select(theme, "Operational logging", &choices)? { + 0 => edit_enum(theme, document, "logging", "level", LOG_LEVELS)?, + 1 => edit_enum(theme, document, "logging", "stderr_format", LOG_FORMATS)?, + 2 => edit_nonnegative_integer(theme, document, "logging", "flush_interval_millis")?, + 3 => edit_sinks(theme, document)?, + 4 => return Ok(()), + _ => unreachable!(), + } + } +} + +fn edit_positive_integer( + theme: &ColorfulTheme, + document: &mut ConfigDocument, + section: &str, + key: &str, +) -> Result<(), CliError> { + let configured = document.has_key(section, key); + match choose_action(theme, configured)? { + 0 => { + let value = prompt_u64(theme, "Value in bytes", document.integer(section, key))?; + document.set_positive_integer(section, key, value)?; + } + 1 if configured => document.clear_key(section, key)?, + _ => {} + } + Ok(()) +} + +fn edit_nonnegative_integer( + theme: &ColorfulTheme, + document: &mut ConfigDocument, + section: &str, + key: &str, +) -> Result<(), CliError> { + let configured = document.has_key(section, key); + match choose_action(theme, configured)? { + 0 => { + let value = prompt_u64( + theme, + "Milliseconds (0 flushes on shutdown)", + document.integer(section, key), + )?; + document.set_integer(section, key, value)?; + } + 1 if configured => document.clear_key(section, key)?, + _ => {} + } + Ok(()) +} + +fn edit_string( + theme: &ColorfulTheme, + document: &mut ConfigDocument, + section: &str, + key: &str, +) -> Result<(), CliError> { + let configured = document.has_key(section, key); + match choose_action(theme, configured)? { + 0 => { + let default = document.string(section, key).unwrap_or_default(); + let value = Input::::with_theme(theme) + .with_prompt("Value") + .with_initial_text(default) + .validate_with(|value: &String| { + if value.trim().is_empty() { + Err("value must not be empty; use Clear to remove it") + } else { + Ok(()) + } + }) + .interact_text() + .map_err(prompt_error)?; + document.set_string(section, key, value)?; + } + 1 if configured => document.clear_key(section, key)?, + _ => {} + } + Ok(()) +} + +fn edit_secret( + theme: &ColorfulTheme, + document: &mut ConfigDocument, + key: &str, +) -> Result<(), CliError> { + let configured = document.has_key("upstream", key); + match choose_action(theme, configured)? { + 0 => { + let value = Password::with_theme(theme) + .with_prompt("Authorization header value") + .allow_empty_password(false) + .interact() + .map_err(prompt_error)?; + document.set_auth_header(key, value)?; + } + 1 if configured => document.clear_key("upstream", key)?, + _ => {} + } + Ok(()) +} + +fn edit_enum( + theme: &ColorfulTheme, + document: &mut ConfigDocument, + section: &str, + key: &str, + values: &[&str], +) -> Result<(), CliError> { + let configured = document.has_key(section, key); + match choose_action(theme, configured)? { + 0 => { + let current = document.string(section, key); + let default = current + .as_deref() + .and_then(|current| values.iter().position(|value| *value == current)) + .unwrap_or(0); + let selected = Select::with_theme(theme) + .with_prompt("Value") + .items(values) + .default(default) + .interact() + .map_err(prompt_error)?; + document.set_enum(section, key, values[selected], values)?; + } + 1 if configured => document.clear_key(section, key)?, + _ => {} + } + Ok(()) +} + +fn edit_sinks(theme: &ColorfulTheme, document: &mut ConfigDocument) -> Result<(), CliError> { + loop { + let mut choices = document + .sink_labels() + .into_iter() + .map(|label| format!("Edit {label}")) + .collect::>(); + let sink_count = choices.len(); + choices.push("Add file sink".into()); + choices.push("Back".into()); + match select(theme, "File sinks", &choices)? { + index if index < sink_count => edit_sink(theme, document, index)?, + index if index == sink_count => { + let path = Input::::with_theme(theme) + .with_prompt("File path") + .validate_with(|value: &String| { + if value.trim().is_empty() { + Err("value must not be empty".to_owned()) + } else { + Ok(()) + } + }) + .interact_text() + .map_err(prompt_error)?; + document.add_sink(path)?; + } + _ => return Ok(()), + } + } +} + +fn edit_sink( + theme: &ColorfulTheme, + document: &mut ConfigDocument, + index: usize, +) -> Result<(), CliError> { + loop { + let choices = [ + format!("Path: {}", document.sink_string_summary(index, "path")), + format!("Level: {}", document.sink_string_summary(index, "level")), + format!("Format: {}", document.sink_string_summary(index, "format")), + format!( + "Queue capacity: {}", + document.sink_integer_summary(index, "queue_capacity") + ), + format!("Rotation: {}", document.sink_rotation_summary(index)), + "Remove sink".into(), + "Back".into(), + ]; + match select(theme, "File sink", &choices)? { + 0 => edit_sink_path(theme, document, index)?, + 1 => edit_sink_enum(theme, document, index, "level", LOG_LEVELS)?, + 2 => edit_sink_enum(theme, document, index, "format", LOG_FORMATS)?, + 3 => edit_sink_queue_capacity(theme, document, index)?, + 4 => edit_sink_rotation(theme, document, index)?, + 5 => { + document.remove_sink(index)?; + return Ok(()); + } + _ => return Ok(()), + } + } +} + +fn edit_sink_path( + theme: &ColorfulTheme, + document: &mut ConfigDocument, + index: usize, +) -> Result<(), CliError> { + let current = document.sink_string(index, "path").unwrap_or_default(); + let value = Input::::with_theme(theme) + .with_prompt("File path") + .with_initial_text(current) + .validate_with(|value: &String| { + if value.trim().is_empty() { + Err("value must not be empty".to_owned()) + } else { + Ok(()) + } + }) + .interact_text() + .map_err(prompt_error)?; + document.set_sink_string(index, "path", value) +} + +fn edit_sink_enum( + theme: &ColorfulTheme, + document: &mut ConfigDocument, + index: usize, + key: &str, + values: &[&str], +) -> Result<(), CliError> { + let configured = document.sink_has_key(index, key)?; + match choose_action(theme, configured)? { + 0 => { + let default = document + .sink_string(index, key) + .as_deref() + .and_then(|current| values.iter().position(|value| *value == current)) + .unwrap_or(0); + let selected = Select::with_theme(theme) + .with_prompt("Value") + .items(values) + .default(default) + .interact() + .map_err(prompt_error)?; + document.set_sink_enum(index, key, values[selected], values)?; + } + 1 if configured => document.clear_sink_key(index, key)?, + _ => {} + } + Ok(()) +} + +fn edit_sink_queue_capacity( + theme: &ColorfulTheme, + document: &mut ConfigDocument, + index: usize, +) -> Result<(), CliError> { + let configured = document.sink_has_key(index, "queue_capacity")?; + match choose_action(theme, configured)? { + 0 => { + let value = prompt_u64( + theme, + "Queue entries", + document.sink_integer(index, "queue_capacity"), + )?; + document.set_sink_queue_capacity(index, value)?; + } + 1 if configured => document.clear_sink_key(index, "queue_capacity")?, + _ => {} + } + Ok(()) +} + +fn edit_sink_rotation( + theme: &ColorfulTheme, + document: &mut ConfigDocument, + index: usize, +) -> Result<(), CliError> { + let configured = document.sink_has_key(index, "max_file_size_bytes")? + || document.sink_has_key(index, "retained_files")?; + match choose_action(theme, configured)? { + 0 => { + let size = prompt_u64( + theme, + "Maximum file size in bytes", + document.sink_integer(index, "max_file_size_bytes"), + )?; + let retained = prompt_u64( + theme, + "Retained backup files", + document.sink_integer(index, "retained_files"), + )?; + document.set_sink_rotation(index, size, retained)?; + } + 1 if configured => document.clear_sink_rotation(index)?, + _ => {} + } + Ok(()) +} + +fn prompt_u64(theme: &ColorfulTheme, prompt: &str, current: Option) -> Result { + let mut input = Input::::with_theme(theme).with_prompt(prompt); + if let Some(current) = current { + input = input.with_initial_text(current.to_string()); + } + input.interact_text().map_err(prompt_error) +} + +fn prompt_error(error: dialoguer::Error) -> CliError { + CliError::Config(format!("configuration edit error: {error}")) +} + +fn print_preview(document: &ConfigDocument) { + println!(); + println!(" ─── Preview ─────────────────────────────────────────────"); + for line in document.preview().lines() { + println!(" {line}"); + } + println!(); +} diff --git a/crates/cli/src/commands/configure/wizard.rs b/crates/cli/src/commands/configure/wizard.rs index 7a26335d6..3b3a418d3 100644 --- a/crates/cli/src/commands/configure/wizard.rs +++ b/crates/cli/src/commands/configure/wizard.rs @@ -3,90 +3,32 @@ //! First-run setup for `nemo-relay` configuration. //! -//! Drives the required scope and agent prompts, then writes a `config.toml` to the chosen scope. Pure -//! helpers (`detect_installed_agents`, `build_config`, `save_config`) are split out from the -//! `dialoguer`-driven orchestrator so the data path can be unit-tested without a TTY. -//! -//! Keep this module focused on TTY and `dialoguer` orchestration. New testable setup behavior -//! should live in `setup/model.rs`, with focused unit tests, so Codecov does not depend on -//! exercising interactive prompt loops. +//! Coordinates first-run setup while terminal-only interaction lives in `wizard/prompt.rs`. -use std::io::IsTerminal; use std::path::PathBuf; -use dialoguer::theme::ColorfulTheme; -use dialoguer::{Confirm, MultiSelect, Select}; +#[cfg(test)] use toml_edit::DocumentMut; +#[cfg(test)] use self::model::{ - ConfigScope, SetupAnswers, agent_key_and_command, build_config, detect_installed_agents, - home_dir, plugins_edit_command_for_scope, plugins_resume_command, preview_paths, - read_existing_defaults, save_config, + ConfigScope, build_config, plugins_edit_command_for_scope, plugins_resume_command, + preview_paths, save_config, }; use super::model; use crate::agents::CodingAgent; use crate::error::CliError; #[cfg(test)] -use self::model::{Defaults, global_config_dir, read_agents_from_doc, reset, write_or_merge}; +use self::model::{ + Defaults, SetupAnswers, global_config_dir, read_agents_from_doc, read_existing_defaults, reset, + write_or_merge, +}; #[cfg(test)] use self::model::detect_installed_agents_in; -/// -/// When `agent_hint` is `Some`, the agent multi-select is skipped — the user already declared -/// intent by typing `nemo-relay claude` (or another agent name), so respect that and only ask -/// scope and agents. To set up multiple agents, the user re-runs `nemo-relay config` later. -pub(crate) fn prompt_user( - detected_agents: &[CodingAgent], - agent_hint: Option, -) -> Result { - ensure_tty()?; - let defaults = read_existing_defaults().unwrap_or_default(); - crate::banner::print_intro(); - match agent_hint { - Some(agent) => { - let (name, _) = agent_key_and_command(agent); - println!(" Setting up {name}."); - println!(" Re-run `nemo-relay config` later to configure additional agents."); - } - None => { - println!(" Let's set up your coding agent."); - println!(" This runs once. Re-run later with `nemo-relay config`."); - } - } - // Only print the detected-agents listing for the unscoped wizard (`nemo-relay config`), - // where the user is about to pick from the multi-select. When the agent was already chosen - // via the easy-path shortcut (`nemo-relay codex`), listing the other two agents is noise. - if agent_hint.is_none() { - println!(); - print_detected_agents(detected_agents); - } - if defaults.has_any() { - println!(); - println!(" Existing config detected — current values are pre-selected."); - } - println!(); - // Keybinding hint shown once: dialoguer's MultiSelect needs SPACE to toggle and ENTER to - // confirm, but doesn't surface that itself. Without this line, users hit Enter expecting - // to check a box and the prompt confirms with the wrong selection. - println!( - " Tip: ↑/↓ to move, SPACE to toggle a checkbox, ENTER to confirm. Defaults are pre-selected." - ); - println!(); - - let theme = ColorfulTheme::default(); - let scope = ask_scope(&theme, defaults.scope)?; - let agents = match agent_hint { - Some(agent) => vec![agent], - None => ask_agents(&theme, detected_agents, &defaults.agents)?, - }; - if agents.contains(&CodingAgent::Codex) { - print_codex_api_key_guide(); - } - - Ok(SetupAnswers { scope, agents }) -} +mod prompt; /// Top-level setup entry point used by `nemo-relay config` and the easy-path fallback. /// Detects agents, prompts the user, writes the config, prints a final summary. @@ -98,77 +40,7 @@ pub(crate) async fn run( agent_hint: Option, explicit_plugin_path: Option, ) -> Result<(), CliError> { - let detected = detect_installed_agents(); - let answers = prompt_user(&detected, agent_hint)?; - - let cwd = std::env::current_dir()?; - let home = home_dir().ok_or_else(|| { - CliError::Config("cannot determine home directory (set $HOME or $USERPROFILE)".into()) - })?; - let doc = build_config(&answers); - let preview_paths = preview_paths(answers.scope, &cwd, &home); - - if !confirm_summary(&preview_paths, &doc)? { - return Err(CliError::Config("setup cancelled — no config saved".into())); - } - - let written = save_config(&doc, answers.scope, &cwd, &home, agent_hint)?; - println!(); - println!(" ✓ Saved:"); - for path in &written { - println!(" {}", path.display()); - } - println!(); - continue_to_plugins(answers.scope, explicit_plugin_path) -} - -/// After the base config is saved, offers to continue into plugin configuration in-process. -/// -/// Prompts once. On acceptance it runs the existing plugin editor targeting an explicit runtime -/// plugin path when present, otherwise the scope derived from base setup (project for -/// `Project`/`Both`, user for `Global`). On decline it reports that the base config was saved, -/// that plugin setup was skipped, and prints the command to resume later. Prompt interruption is -/// treated as a skip; other prompt or editor failures surface an error that makes clear the base -/// config remains saved. The saved `config.toml` is never rolled back here. -fn continue_to_plugins( - scope: ConfigScope, - explicit_plugin_path: Option, -) -> Result<(), CliError> { - let resume_command = plugins_resume_command(scope, explicit_plugin_path.as_deref()); - let proceed = match Confirm::with_theme(&ColorfulTheme::default()) - .with_prompt("Configure Relay plugins now?") - .default(true) - .interact() - { - Ok(proceed) => proceed, - Err(error) if plugin_prompt_was_interrupted(&error) => { - print_plugins_skipped(&resume_command); - return Ok(()); - } - Err(error) => { - return Err(CliError::Config(format!( - "plugin setup did not complete; base configuration remains saved. \ - Resume with `{}`. Cause: {error}", - resume_command - ))); - } - }; - if !proceed { - print_plugins_skipped(&resume_command); - return Ok(()); - } - let result = crate::plugins::edit(plugins_edit_command_for_scope(scope, explicit_plugin_path)); - result.map_err(|error| { - let cause = match error { - CliError::Config(message) => message, - other => other.to_string(), - }; - CliError::Config(format!( - "plugin setup did not complete; base configuration remains saved. \ - Resume with `{}`. Cause: {cause}", - resume_command - )) - }) + prompt::run(agent_hint, explicit_plugin_path).await } fn plugin_prompt_was_interrupted(error: &dialoguer::Error) -> bool { @@ -182,151 +54,6 @@ fn plugin_prompt_was_interrupted(error: &dialoguer::Error) -> bool { ) } -fn print_plugins_skipped(resume_command: &str) { - println!(); - println!(" Base configuration saved. Plugin configuration skipped."); - println!(" Configure plugins later with `{resume_command}`."); - println!(); -} - -fn print_codex_api_key_guide() { - // Codex supports two auth flows (see `codex-rs/login/src/auth/manager.rs`): - // 1. ChatGPT-Plus PKCE OAuth via `codex --login` → tokens stored in `~/.codex/auth.json` - // 2. OpenAI API key via `OPENAI_API_KEY` env var - // The gateway routes to the correct upstream automatically: ChatGPT OAuth goes to - // `chatgpt.com/backend-api/codex`, API key goes to `api.openai.com`. - println!(); - println!(" ℹ Codex sends Responses-API requests through the gateway."); - println!(" Authentication (pick one):"); - println!(" • ChatGPT-Plus login: codex --login (uses ~/.codex/auth.json)"); - println!(" • OpenAI API key: export OPENAI_API_KEY=sk-..."); - println!(" When OPENAI_API_KEY is set the gateway uses it; otherwise the"); - println!(" ChatGPT-Plus OAuth token is forwarded to the ChatGPT backend."); - println!(); -} - -fn ensure_tty() -> Result<(), CliError> { - if !std::io::stdin().is_terminal() { - return Err(CliError::Config( - "interactive setup requires a TTY; pass `--config ` or set up \ - `.nemo-relay/config.toml` manually" - .into(), - )); - } - Ok(()) -} - -fn print_detected_agents(detected: &[CodingAgent]) { - println!(" Detected agents on $PATH:"); - for agent in detected { - let (name, _) = agent_key_and_command(*agent); - println!(" ✓ {name}"); - } - if detected.is_empty() { - println!(" (none — you can still add agents later)"); - } -} - -fn ask_scope( - theme: &ColorfulTheme, - existing: Option, -) -> Result { - let options = [ConfigScope::Project, ConfigScope::Global, ConfigScope::Both]; - let labels: Vec<&str> = options.iter().map(|s| s.label()).collect(); - // Start on the user's existing scope if there is one (so re-running the wizard doesn't - // accidentally relocate their config), else `Project` per the design default. - let default_idx = existing - .and_then(|s| options.iter().position(|opt| *opt == s)) - .unwrap_or(0); - let idx = Select::with_theme(theme) - .with_prompt("Save config where?") - .items(&labels) - .default(default_idx) - .interact() - .map_err(setup_error)?; - Ok(options[idx]) -} - -fn ask_agents( - theme: &ColorfulTheme, - detected: &[CodingAgent], - configured: &[CodingAgent], -) -> Result, CliError> { - let all_supported = [ - CodingAgent::ClaudeCode, - CodingAgent::Codex, - CodingAgent::Hermes, - ]; - let labels: Vec = all_supported - .iter() - .map(|a| { - let (name, _) = agent_key_and_command(*a); - name.to_string() - }) - .collect(); - // Pre-check: union of "already in the existing config" and "detected on $PATH". The existing - // entries take precedence — if the user previously deselected an agent that's on PATH, we - // shouldn't re-check it for them. On first run (no existing config), this falls back to - // pre-checking everything detected. - let defaults: Vec = if configured.is_empty() { - all_supported.iter().map(|a| detected.contains(a)).collect() - } else { - all_supported - .iter() - .map(|a| configured.contains(a)) - .collect() - }; - let selected_idx = MultiSelect::with_theme(theme) - .with_prompt("Which agents to observe?") - .items(&labels) - .defaults(&defaults) - .interact() - .map_err(setup_error)?; - Ok(selected_idx.into_iter().map(|i| all_supported[i]).collect()) -} - -/// Confirms the summary with the user before writing the file. Returns true if the user accepted. -/// Shows both the destination path(s) and the exact TOML body about to be written so the user -/// can verify what they're committing to instead of confirming a path blind. -pub(crate) fn confirm_summary( - written_paths: &[PathBuf], - doc: &DocumentMut, -) -> Result { - println!(); - println!(" ─── Summary ─────────────────────────────────────────────"); - println!(" Will write to:"); - for path in written_paths { - println!(" {}", path.display()); - } - println!(); - println!(" Contents:"); - for line in doc.to_string().lines() { - println!(" {line}"); - } - println!(); - Confirm::with_theme(&ColorfulTheme::default()) - .with_prompt("Looks good?") - .default(true) - .interact() - .map_err(setup_error) -} - -fn setup_error(err: dialoguer::Error) -> CliError { - // dialoguer errors are mostly IO. Translate cancellation (Ctrl-C, EOF on stdin) into a - // friendly "cancelled" message; surface anything else as the raw error. - match err { - dialoguer::Error::IO(io_err) - if matches!( - io_err.kind(), - std::io::ErrorKind::Interrupted | std::io::ErrorKind::UnexpectedEof - ) => - { - CliError::Config("setup cancelled — no config saved".into()) - } - other => CliError::Config(format!("setup error: {other}")), - } -} - #[cfg(test)] #[path = "../../../tests/coverage/shared/setup_tests.rs"] mod tests; diff --git a/crates/cli/src/commands/configure/wizard/prompt.rs b/crates/cli/src/commands/configure/wizard/prompt.rs new file mode 100644 index 000000000..4c9d01192 --- /dev/null +++ b/crates/cli/src/commands/configure/wizard/prompt.rs @@ -0,0 +1,299 @@ +// SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +//! Terminal-only prompt adapter for first-run configuration. + +use std::io::IsTerminal; +use std::path::PathBuf; + +use dialoguer::theme::ColorfulTheme; +use dialoguer::{Confirm, MultiSelect, Select}; +use toml_edit::DocumentMut; + +use super::model::{ + ConfigScope, SetupAnswers, agent_key_and_command, build_config, detect_installed_agents, + home_dir, plugins_edit_command_for_scope, plugins_resume_command, preview_paths, + read_existing_defaults, save_config, +}; +use crate::agents::CodingAgent; +use crate::error::CliError; + +/// Prompts for the configuration scope and agents selected by the user. +/// +/// When `agent_hint` is present, the agent picker is skipped because the command already +/// identified the requested agent. +pub(crate) fn prompt_user( + detected_agents: &[CodingAgent], + agent_hint: Option, +) -> Result { + ensure_tty()?; + let defaults = read_existing_defaults().unwrap_or_default(); + crate::banner::print_intro(); + match agent_hint { + Some(agent) => { + let (name, _) = agent_key_and_command(agent); + println!(" Setting up {name}."); + println!(" Re-run `nemo-relay config` later to configure additional agents."); + } + None => { + println!(" Let's set up your coding agent."); + println!(" This runs once. Re-run later with `nemo-relay config`."); + } + } + // Only print the detected-agents listing for the unscoped wizard (`nemo-relay config`), + // where the user is about to pick from the multi-select. When the agent was already chosen + // via the easy-path shortcut (`nemo-relay codex`), listing the other two agents is noise. + if agent_hint.is_none() { + println!(); + print_detected_agents(detected_agents); + } + if defaults.has_any() { + println!(); + println!(" Existing config detected — current values are pre-selected."); + } + println!(); + // Keybinding hint shown once: dialoguer's MultiSelect needs SPACE to toggle and ENTER to + // confirm, but doesn't surface that itself. Without this line, users hit Enter expecting + // to check a box and the prompt confirms with the wrong selection. + println!( + " Tip: ↑/↓ to move, SPACE to toggle a checkbox, ENTER to confirm. Defaults are pre-selected." + ); + println!(); + + let theme = ColorfulTheme::default(); + let scope = ask_scope(&theme, defaults.scope)?; + let agents = match agent_hint { + Some(agent) => vec![agent], + None => ask_agents(&theme, detected_agents, &defaults.agents)?, + }; + if agents.contains(&CodingAgent::Codex) { + print_codex_api_key_guide(); + } + + Ok(SetupAnswers { scope, agents }) +} + +pub(super) async fn run( + agent_hint: Option, + explicit_plugin_path: Option, +) -> Result<(), CliError> { + let detected = detect_installed_agents(); + let answers = prompt_user(&detected, agent_hint)?; + + let cwd = std::env::current_dir()?; + let home = home_dir().ok_or_else(|| { + CliError::Config("cannot determine home directory (set $HOME or $USERPROFILE)".into()) + })?; + let doc = build_config(&answers); + let preview_paths = preview_paths(answers.scope, &cwd, &home); + + if !confirm_summary(&preview_paths, &doc)? { + return Err(CliError::Config("setup cancelled — no config saved".into())); + } + + let written = save_config(&doc, answers.scope, &cwd, &home, agent_hint)?; + println!(); + println!(" ✓ Saved:"); + for path in &written { + println!(" {}", path.display()); + } + println!(); + continue_to_plugins(answers.scope, explicit_plugin_path) +} + +/// After the base config is saved, offers to continue into plugin configuration in-process. +/// +/// Prompts once. On acceptance it runs the existing plugin editor targeting an explicit runtime +/// plugin path when present, otherwise the scope derived from base setup (project for +/// `Project`/`Both`, user for `Global`). On decline it reports that the base config was saved, +/// that plugin setup was skipped, and prints the command to resume later. Prompt interruption is +/// treated as a skip; other prompt or editor failures surface an error that makes clear the base +/// config remains saved. The saved `config.toml` is never rolled back here. +fn continue_to_plugins( + scope: ConfigScope, + explicit_plugin_path: Option, +) -> Result<(), CliError> { + let resume_command = plugins_resume_command(scope, explicit_plugin_path.as_deref()); + let proceed = match confirm_plugin_setup() { + Ok(proceed) => proceed, + Err(error) if super::plugin_prompt_was_interrupted(&error) => { + print_plugins_skipped(&resume_command); + return Ok(()); + } + Err(error) => { + return Err(CliError::Config(format!( + "plugin setup did not complete; base configuration remains saved. \ + Resume with `{}`. Cause: {error}", + resume_command + ))); + } + }; + if !proceed { + print_plugins_skipped(&resume_command); + return Ok(()); + } + let result = crate::plugins::edit(plugins_edit_command_for_scope(scope, explicit_plugin_path)); + result.map_err(|error| { + let cause = match error { + CliError::Config(message) => message, + other => other.to_string(), + }; + CliError::Config(format!( + "plugin setup did not complete; base configuration remains saved. \ + Resume with `{}`. Cause: {cause}", + resume_command + )) + }) +} + +pub(super) fn confirm_plugin_setup() -> Result { + Confirm::with_theme(&ColorfulTheme::default()) + .with_prompt("Configure Relay plugins now?") + .default(true) + .interact() +} + +pub(super) fn print_plugins_skipped(resume_command: &str) { + println!(); + println!(" Base configuration saved. Plugin configuration skipped."); + println!(" Configure plugins later with `{resume_command}`."); + println!(); +} + +fn print_codex_api_key_guide() { + // Codex supports two auth flows (see `codex-rs/login/src/auth/manager.rs`): + // 1. ChatGPT-Plus PKCE OAuth via `codex --login` → tokens stored in `~/.codex/auth.json` + // 2. OpenAI API key via `OPENAI_API_KEY` env var + // The gateway routes to the correct upstream automatically: ChatGPT OAuth goes to + // `chatgpt.com/backend-api/codex`, API key goes to `api.openai.com`. + println!(); + println!(" ℹ Codex sends Responses-API requests through the gateway."); + println!(" Authentication (pick one):"); + println!(" • ChatGPT-Plus login: codex --login (uses ~/.codex/auth.json)"); + println!(" • OpenAI API key: export OPENAI_API_KEY=sk-..."); + println!(" When OPENAI_API_KEY is set the gateway uses it; otherwise the"); + println!(" ChatGPT-Plus OAuth token is forwarded to the ChatGPT backend."); + println!(); +} + +fn ensure_tty() -> Result<(), CliError> { + if !std::io::stdin().is_terminal() { + return Err(CliError::Config( + "interactive setup requires a TTY; pass `--config ` or set up \ + `.nemo-relay/config.toml` manually" + .into(), + )); + } + Ok(()) +} + +fn print_detected_agents(detected: &[CodingAgent]) { + println!(" Detected agents on $PATH:"); + for agent in detected { + let (name, _) = agent_key_and_command(*agent); + println!(" ✓ {name}"); + } + if detected.is_empty() { + println!(" (none — you can still add agents later)"); + } +} + +fn ask_scope( + theme: &ColorfulTheme, + existing: Option, +) -> Result { + let options = [ConfigScope::Project, ConfigScope::Global, ConfigScope::Both]; + let labels: Vec<&str> = options.iter().map(|s| s.label()).collect(); + // Start on the user's existing scope if there is one (so re-running the wizard doesn't + // accidentally relocate their config), else `Project` per the design default. + let default_idx = existing + .and_then(|s| options.iter().position(|opt| *opt == s)) + .unwrap_or(0); + let idx = Select::with_theme(theme) + .with_prompt("Save config where?") + .items(&labels) + .default(default_idx) + .interact() + .map_err(setup_error)?; + Ok(options[idx]) +} + +fn ask_agents( + theme: &ColorfulTheme, + detected: &[CodingAgent], + configured: &[CodingAgent], +) -> Result, CliError> { + let all_supported = [ + CodingAgent::ClaudeCode, + CodingAgent::Codex, + CodingAgent::Hermes, + ]; + let labels: Vec = all_supported + .iter() + .map(|a| { + let (name, _) = agent_key_and_command(*a); + name.to_string() + }) + .collect(); + // Pre-check: union of "already in the existing config" and "detected on $PATH". The existing + // entries take precedence — if the user previously deselected an agent that's on PATH, we + // shouldn't re-check it for them. On first run (no existing config), this falls back to + // pre-checking everything detected. + let defaults: Vec = if configured.is_empty() { + all_supported.iter().map(|a| detected.contains(a)).collect() + } else { + all_supported + .iter() + .map(|a| configured.contains(a)) + .collect() + }; + let selected_idx = MultiSelect::with_theme(theme) + .with_prompt("Which agents to observe?") + .items(&labels) + .defaults(&defaults) + .interact() + .map_err(setup_error)?; + Ok(selected_idx.into_iter().map(|i| all_supported[i]).collect()) +} + +/// Confirms the summary with the user before writing the file. Returns true if the user accepted. +/// Shows both the destination path(s) and the exact TOML body about to be written so the user +/// can verify what they're committing to instead of confirming a path blind. +pub(crate) fn confirm_summary( + written_paths: &[PathBuf], + doc: &DocumentMut, +) -> Result { + println!(); + println!(" ─── Summary ─────────────────────────────────────────────"); + println!(" Will write to:"); + for path in written_paths { + println!(" {}", path.display()); + } + println!(); + println!(" Contents:"); + for line in doc.to_string().lines() { + println!(" {line}"); + } + println!(); + Confirm::with_theme(&ColorfulTheme::default()) + .with_prompt("Looks good?") + .default(true) + .interact() + .map_err(setup_error) +} + +fn setup_error(err: dialoguer::Error) -> CliError { + // dialoguer errors are mostly IO. Translate cancellation (Ctrl-C, EOF on stdin) into a + // friendly "cancelled" message; surface anything else as the raw error. + match err { + dialoguer::Error::IO(io_err) + if matches!( + io_err.kind(), + std::io::ErrorKind::Interrupted | std::io::ErrorKind::UnexpectedEof + ) => + { + CliError::Config("setup cancelled — no config saved".into()) + } + other => CliError::Config(format!("setup error: {other}")), + } +} diff --git a/crates/cli/src/commands/mod.rs b/crates/cli/src/commands/mod.rs index 1a5e158f8..8ce414493 100644 --- a/crates/cli/src/commands/mod.rs +++ b/crates/cli/src/commands/mod.rs @@ -53,55 +53,58 @@ pub(crate) async fn run(bootstrap_shutdown_token: Option) -> ExitCode { // Dispatches CLI subcommands while keeping the no-subcommand path as server mode. `run` inherits // top-level server flags so transparent launch can share config parsing with daemon startup. -async fn dispatch(bootstrap_shutdown_token: Option) -> Result { - let cli = Cli::parse(); - let command_name = cli - .command - .as_ref() - .map(Command::log_name) - .unwrap_or("default"); - - let initialize_logging = match cli.command.as_ref() { +fn configure_logging( + cli: &Cli, +) -> Result, error::CliError> { + let initialize = match cli.command.as_ref() { Some(command) => !command.skips_logging(), None => { cli.server.to_runtime().requested_daemon_mode() || runtime_configuration::any_config_file_exists() } }; - let _logging = if initialize_logging { - let user_only = matches!(cli.command.as_ref(), Some(Command::Mcp)); - let explicit_config = if user_only { - None - } else { - match cli.command.as_ref() { - Some(Command::Run(command)) => { - command.config.as_deref().or(cli.server.config.as_deref()) - } - _ => cli.server.config.as_deref(), - } - }; - let mut logging_fallback_error = None; - let config = match cli.logging.resolve(explicit_config, user_only) { - Ok(config) => config, - Err(error) if matches!(cli.command.as_ref(), Some(Command::Doctor(_))) => { - logging_fallback_error = Some(error); - nemo_relay::logging::LoggingConfig::default() - } - Err(error) => return Err(error), - }; - let runtime = nemo_relay::logging::LoggingRuntime::configure(config)?; - if let Some(error) = logging_fallback_error { - log::warn!( - target: "nemo_relay.cli", - event = "doctor_logging_fallback", - error_kind = error.log_kind(); - "Doctor fell back to default logging after resolution failure" - ); + if !initialize { + return Ok(None); + } + + let user_only = matches!(cli.command.as_ref(), Some(Command::Mcp)); + let explicit_config = match (user_only, cli.command.as_ref()) { + (true, _) => None, + (false, Some(Command::Run(command))) => { + command.config.as_deref().or(cli.server.config.as_deref()) } - Some(runtime) - } else { - None + (false, _) => cli.server.config.as_deref(), }; + let mut fallback_error = None; + let config = match cli.logging.resolve(explicit_config, user_only) { + Ok(config) => config, + Err(error) if matches!(cli.command.as_ref(), Some(Command::Doctor(_))) => { + fallback_error = Some(error); + nemo_relay::logging::LoggingConfig::default() + } + Err(error) => return Err(error), + }; + let runtime = nemo_relay::logging::LoggingRuntime::configure(config)?; + if let Some(error) = fallback_error { + log::warn!( + target: "nemo_relay.cli", + event = "doctor_logging_fallback", + error_kind = error.log_kind(); + "Doctor fell back to default logging after resolution failure" + ); + } + Ok(Some(runtime)) +} + +async fn dispatch(bootstrap_shutdown_token: Option) -> Result { + let cli = Cli::parse(); + let command_name = cli + .command + .as_ref() + .map(Command::log_name) + .unwrap_or("default"); + + let _logging = configure_logging(&cli)?; log::info!( target: "nemo_relay.cli", diff --git a/crates/cli/src/plugins/dynamic_editor.rs b/crates/cli/src/plugins/dynamic_editor.rs index e31df9673..1d04a50e6 100644 --- a/crates/cli/src/plugins/dynamic_editor.rs +++ b/crates/cli/src/plugins/dynamic_editor.rs @@ -6,9 +6,8 @@ use std::collections::HashSet; use dialoguer::theme::ColorfulTheme; -use dialoguer::{Input, Password, Select}; use nemo_relay::plugin::dynamic::DynamicPluginManifest; -use serde_json::{Map, Number, Value}; +use serde_json::{Map, Value}; use crate::error::CliError; @@ -23,6 +22,8 @@ use super::{ const REDACTED: &str = ""; +mod prompt; + #[derive(Debug)] pub(super) struct DynamicPluginEditorState { document_index: usize, @@ -149,6 +150,11 @@ impl DynamicPluginEditorState { .unwrap_or_default() } + #[cfg(test)] + pub(super) fn editor_fields(&self) -> &[DynamicConfigField] { + self.schema.as_ref().map_or(&[], |schema| schema.fields()) + } + #[cfg(test)] pub(super) fn reset_top_level_field(&mut self, key: &str) -> Result<(), CliError> { let field = self @@ -354,7 +360,7 @@ fn load_config_schema( } #[derive(Debug, Clone, Copy)] -enum DynamicMenuAction { +pub(super) enum DynamicMenuAction { EditField(usize), EditRawConfig, ResetPlugin, @@ -365,45 +371,10 @@ pub(super) fn edit_dynamic_plugin( theme: &ColorfulTheme, state: &mut DynamicPluginEditorState, ) -> Result<(), CliError> { - if let Some(description) = &state.description { - println!(" {}", super::single_line_text(description)); - } - let fields = state - .schema - .as_ref() - .map(|schema| schema.fields().to_vec()) - .unwrap_or_default(); - if state.schema.is_none() || fields.is_empty() { - edit_dynamic_root_menu(theme, state, &fields) - } else { - let prompt = state - .editor_title - .clone() - .unwrap_or_else(|| state.label.clone()); - edit_dynamic_fields_menu(theme, state, &fields, &[], prompt) - } + prompt::edit_dynamic_plugin(theme, state) } -fn edit_dynamic_root_menu( - theme: &ColorfulTheme, - state: &mut DynamicPluginEditorState, - fields: &[DynamicConfigField], -) -> Result<(), CliError> { - let mut selected_index = 0; - loop { - let (items, actions) = dynamic_root_menu_items(state, fields); - - let selection = prompt_menu(theme, state.label(), &items, selected_index)?; - if let Some(selected) = menu_response_index(&selection) { - selected_index = selected; - } - if handle_dynamic_root_menu_response(theme, state, &actions, selection)? { - return Ok(()); - } - } -} - -fn dynamic_root_menu_items( +pub(super) fn dynamic_root_menu_items( state: &DynamicPluginEditorState, fields: &[DynamicConfigField], ) -> (Vec, Vec) { @@ -426,96 +397,7 @@ fn dynamic_root_menu_items( (items, actions) } -fn handle_dynamic_root_menu_response( - theme: &ColorfulTheme, - state: &mut DynamicPluginEditorState, - actions: &[DynamicMenuAction], - selection: MenuResponse, -) -> Result { - match selection { - MenuResponse::Selected(selected) => match actions.get(selected).copied() { - Some(DynamicMenuAction::EditRawConfig) => { - prompt_raw_config(theme, state)?; - Ok(false) - } - Some(DynamicMenuAction::ResetPlugin) => { - state.reset(); - Ok(false) - } - Some(DynamicMenuAction::Back) | None => Ok(true), - Some(DynamicMenuAction::EditField(_)) => { - println!(" Select Edit raw configuration to modify settings."); - Ok(false) - } - }, - MenuResponse::Shortcut(MenuShortcut::Reset, selected) => { - if matches!(actions.get(selected), Some(DynamicMenuAction::ResetPlugin)) { - state.reset(); - } else { - println!(" Select Reset plugin configuration to remove config."); - } - Ok(false) - } - MenuResponse::Shortcut(MenuShortcut::Clear, selected) => { - if matches!( - actions.get(selected), - Some(DynamicMenuAction::EditRawConfig) - ) { - state.set_raw_config(Map::new()); - } - Ok(false) - } - MenuResponse::Shortcut(MenuShortcut::Help, _) => { - super::print_editor_help(); - Ok(false) - } - MenuResponse::Shortcut(MenuShortcut::Preview | MenuShortcut::Save, _) => { - println!(" Preview and save are available from the main plugins.toml menu."); - Ok(false) - } - MenuResponse::Cancel => Ok(true), - } -} - -fn edit_dynamic_fields_menu( - theme: &ColorfulTheme, - state: &mut DynamicPluginEditorState, - fields: &[DynamicConfigField], - parent_path: &[String], - prompt: String, -) -> Result<(), CliError> { - let mut selected_index = 0; - loop { - let (items, actions) = dynamic_field_menu_items(state, fields, parent_path); - let selection = prompt_menu(theme, &prompt, &items, selected_index)?; - if let Some(selected) = menu_response_index(&selection) { - selected_index = selected; - } - match selection { - MenuResponse::Selected(selected) => match actions.get(selected).copied() { - Some(DynamicMenuAction::EditField(index)) => { - edit_dynamic_field(theme, state, &fields[index], parent_path)?; - } - Some(DynamicMenuAction::ResetPlugin) => state.reset(), - Some(DynamicMenuAction::Back) | None => return Ok(()), - Some(DynamicMenuAction::EditRawConfig) => unreachable!(), - }, - MenuResponse::Shortcut(MenuShortcut::Reset, selected) => { - reset_dynamic_selection(state, fields, parent_path, &actions, selected); - } - MenuResponse::Shortcut(MenuShortcut::Clear, selected) => { - clear_dynamic_selection(state, fields, parent_path, &actions, selected); - } - MenuResponse::Shortcut(MenuShortcut::Help, _) => super::print_editor_help(), - MenuResponse::Shortcut(MenuShortcut::Preview | MenuShortcut::Save, _) => { - println!(" Preview and save are available from the main plugins.toml menu."); - } - MenuResponse::Cancel => return Ok(()), - } - } -} - -fn dynamic_field_menu_items( +pub(super) fn dynamic_field_menu_items( state: &DynamicPluginEditorState, fields: &[DynamicConfigField], parent_path: &[String], @@ -558,265 +440,7 @@ fn dynamic_field_menu_items( (items, actions) } -fn edit_dynamic_field( - theme: &ColorfulTheme, - state: &mut DynamicPluginEditorState, - field: &DynamicConfigField, - parent_path: &[String], -) -> Result<(), CliError> { - if let Some(description) = &field.description { - println!(" {}", super::single_line_text(description)); - } - let path = field_path(parent_path, field); - if let DynamicConfigFieldKind::Object { fields } = &field.kind { - return edit_dynamic_fields_menu(theme, state, fields, &path, field.title.clone()); - } - if let Some(value) = prompt_dynamic_value(theme, state, field, &path)? { - state.set_field(&path, value); - } - Ok(()) -} - -fn prompt_dynamic_value( - theme: &ColorfulTheme, - state: &DynamicPluginEditorState, - field: &DynamicConfigField, - path: &[String], -) -> Result, CliError> { - let current = state.field_value(path); - match &field.kind { - DynamicConfigFieldKind::Boolean => { - let values = ["false", "true"]; - let default = current - .and_then(Value::as_bool) - .or_else(|| field.default.as_ref().and_then(Value::as_bool)) - .map(usize::from) - .unwrap_or(0); - let selected = Select::with_theme(theme) - .with_prompt(super::single_line_text(&field.title)) - .items(&values) - .default(default) - .interact() - .map_err(editor_error)?; - Ok(Some(Value::Bool(selected == 1))) - } - DynamicConfigFieldKind::String { secret } => { - prompt_dynamic_string(theme, field, current, *secret, None) - } - DynamicConfigFieldKind::StringEnum { options, secret } => { - if *secret { - prompt_dynamic_string(theme, field, current, true, Some(options)) - } else { - let default = current - .and_then(Value::as_str) - .or_else(|| field.default.as_ref().and_then(Value::as_str)) - .and_then(|value| options.iter().position(|option| option == value)) - .unwrap_or(0); - let selected = Select::with_theme(theme) - .with_prompt(super::single_line_text(&field.title)) - .items(options) - .default(default) - .interact() - .map_err(editor_error)?; - Ok(Some(Value::String(options[selected].clone()))) - } - } - DynamicConfigFieldKind::Integer => { - let initial = current - .or(field.default.as_ref()) - .map(json_text) - .unwrap_or_default(); - let value: String = Input::with_theme(theme) - .with_prompt(super::single_line_text(&field.title)) - .with_initial_text(initial) - .interact_text() - .map_err(editor_error)?; - let value = value.trim().parse::().map_err(|error| { - CliError::Config(format!("{} must be an integer: {error}", field.key)) - })?; - Ok(Some(Value::Number(value.into()))) - } - DynamicConfigFieldKind::Number => { - let initial = current - .or(field.default.as_ref()) - .map(json_text) - .unwrap_or_default(); - let value: String = Input::with_theme(theme) - .with_prompt(super::single_line_text(&field.title)) - .with_initial_text(initial) - .interact_text() - .map_err(editor_error)?; - let parsed = value.trim().parse::().map_err(|error| { - CliError::Config(format!("{} must be a number: {error}", field.key)) - })?; - let number = Number::from_f64(parsed).ok_or_else(|| { - CliError::Config(format!("{} must be a finite number", field.key)) - })?; - Ok(Some(Value::Number(number))) - } - DynamicConfigFieldKind::StringMap => { - let (current, redacted_config, secrets, hidden) = state.field_value_for_raw_edit(path); - let Some(value) = prompt_json_value( - theme, - field, - current.as_ref(), - Value::Object(Map::new()), - hidden, - )? - else { - return Ok(None); - }; - let value = state.restore_raw_field_edit(path, value, redacted_config, &secrets)?; - let object = value - .as_object() - .ok_or_else(|| CliError::Config(format!("{} must be a JSON object", field.key)))?; - if object.values().any(|value| !value.is_string()) { - return Err(CliError::Config(format!( - "{} must contain only string values", - field.key - ))); - } - Ok(Some(value)) - } - DynamicConfigFieldKind::RawJson => { - let fallback = field.default.clone().unwrap_or(Value::Null); - let (current, redacted_config, secrets, hidden) = state.field_value_for_raw_edit(path); - let Some(value) = prompt_json_value(theme, field, current.as_ref(), fallback, hidden)? - else { - return Ok(None); - }; - let value = state.restore_raw_field_edit(path, value, redacted_config, &secrets)?; - Ok(Some(value)) - } - DynamicConfigFieldKind::Object { .. } => unreachable!(), - } -} - -fn prompt_dynamic_string( - theme: &ColorfulTheme, - field: &DynamicConfigField, - current: Option<&Value>, - secret: bool, - options: Option<&[String]>, -) -> Result, CliError> { - if secret { - let title = super::single_line_text(&field.title); - let value = Password::with_theme(theme) - .with_prompt(format!("New {} (blank preserves the current value)", title)) - .allow_empty_password(true) - .report(false) - .interact() - .map_err(editor_error)?; - if value.is_empty() { - return Ok(None); - } - if options.is_some_and(|options| !options.iter().any(|option| option == &value)) { - return Err(CliError::Config(format!( - "{} must be one of the schema enum values", - field.key - ))); - } - return Ok(Some(Value::String(value))); - } - let initial = current - .and_then(Value::as_str) - .or_else(|| field.default.as_ref().and_then(Value::as_str)) - .unwrap_or_default(); - let value: String = Input::with_theme(theme) - .with_prompt(super::single_line_text(&field.title)) - .with_initial_text(initial) - .interact_text() - .map_err(editor_error)?; - Ok(Some(Value::String(value))) -} - -fn prompt_json_value( - theme: &ColorfulTheme, - field: &DynamicConfigField, - current: Option<&Value>, - fallback: Value, - hidden: bool, -) -> Result, CliError> { - let initial = current.or(field.default.as_ref()).unwrap_or(&fallback); - let prompt = format!("{} as JSON", super::single_line_text(&field.title)); - let value = if hidden { - if current.is_some() { - println!(" Current redacted JSON: {}", json_text(initial)); - } - let value = Password::with_theme(theme) - .with_prompt(format!("New {prompt} (blank preserves the current value)")) - .allow_empty_password(true) - .report(false) - .interact() - .map_err(editor_error)?; - if value.is_empty() { - return Ok(None); - } - value - } else { - Input::with_theme(theme) - .with_prompt(prompt) - .with_initial_text(json_text(initial)) - .interact_text() - .map_err(editor_error)? - }; - serde_json::from_str(value.trim()) - .map_err(|error| CliError::Config(format!("invalid JSON for {}: {error}", field.key))) - .map(Some) -} - -fn prompt_raw_config( - theme: &ColorfulTheme, - state: &mut DynamicPluginEditorState, -) -> Result<(), CliError> { - let original = Value::Object(state.config.clone().unwrap_or_default()); - let (initial, secrets, hidden) = state - .schema - .as_ref() - .map(|schema| { - let (redacted, secrets) = schema.redact_for_edit(&original); - (redacted, secrets, schema.has_secrets()) - }) - .unwrap_or_else(|| (original, SecretEditValues::new(), false)); - let value = if hidden { - println!(" Current redacted JSON: {}", json_text(&initial)); - let value = Password::with_theme(theme) - .with_prompt("New configuration as JSON object (blank preserves the current value)") - .allow_empty_password(true) - .report(false) - .interact() - .map_err(editor_error)?; - if value.is_empty() { - return Ok(()); - } - value - } else { - Input::with_theme(theme) - .with_prompt("Configuration as JSON object") - .with_initial_text(json_text(&initial)) - .interact_text() - .map_err(editor_error)? - }; - let value: Value = serde_json::from_str(value.trim()) - .map_err(|error| CliError::Config(format!("invalid JSON configuration: {error}")))?; - let value = match &state.schema { - Some(schema) => schema.restore_edit_secrets(&value, &secrets)?, - None => value, - }; - let object = value.as_object().cloned().ok_or_else(|| { - CliError::Config(format!( - "dynamic plugin '{}' configuration must be a JSON object", - state.plugin_id - )) - })?; - if let Some(schema) = &state.schema { - schema.validate(&value)?; - } - state.set_raw_config(object); - Ok(()) -} - -fn reset_dynamic_selection( +pub(super) fn reset_dynamic_selection( state: &mut DynamicPluginEditorState, fields: &[DynamicConfigField], parent_path: &[String], @@ -833,7 +457,7 @@ fn reset_dynamic_selection( } } -fn clear_dynamic_selection( +pub(super) fn clear_dynamic_selection( state: &mut DynamicPluginEditorState, fields: &[DynamicConfigField], parent_path: &[String], @@ -848,7 +472,7 @@ fn clear_dynamic_selection( } } -fn field_path(parent_path: &[String], field: &DynamicConfigField) -> Vec { +pub(super) fn field_path(parent_path: &[String], field: &DynamicConfigField) -> Vec { let mut path = parent_path.to_vec(); path.push(field.key.clone()); path @@ -862,7 +486,10 @@ fn field_is_secret(field: &DynamicConfigField) -> bool { ) } -fn value_at_path<'a>(config: Option<&'a Map>, path: &[String]) -> Option<&'a Value> { +pub(super) fn value_at_path<'a>( + config: Option<&'a Map>, + path: &[String], +) -> Option<&'a Value> { let (first, rest) = path.split_first()?; let mut value = config?.get(first)?; for segment in rest { @@ -871,7 +498,11 @@ fn value_at_path<'a>(config: Option<&'a Map>, path: &[String]) -> Some(value) } -fn set_value_at_path(config: &mut Option>, path: &[String], value: Value) { +pub(super) fn set_value_at_path( + config: &mut Option>, + path: &[String], + value: Value, +) { let Some((last, parents)) = path.split_last() else { return; }; @@ -890,7 +521,7 @@ fn set_value_at_path(config: &mut Option>, path: &[String], v object.insert(last.clone(), value); } -fn remove_value_at_path(config: &mut Map, path: &[String]) -> bool { +pub(super) fn remove_value_at_path(config: &mut Map, path: &[String]) -> bool { let Some((first, rest)) = path.split_first() else { return config.is_empty(); }; diff --git a/crates/cli/src/plugins/dynamic_editor/prompt.rs b/crates/cli/src/plugins/dynamic_editor/prompt.rs new file mode 100644 index 000000000..ca43e99c6 --- /dev/null +++ b/crates/cli/src/plugins/dynamic_editor/prompt.rs @@ -0,0 +1,401 @@ +// SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +//! Terminal-only prompt adapter for dynamic plugin configuration. + +use dialoguer::theme::ColorfulTheme; +use dialoguer::{Input, Password, Select}; +use serde_json::{Map, Number, Value}; + +use super::*; +use crate::error::CliError; +use crate::plugins::{print_editor_help, single_line_text}; + +pub(super) fn edit_dynamic_plugin( + theme: &ColorfulTheme, + state: &mut DynamicPluginEditorState, +) -> Result<(), CliError> { + if let Some(description) = &state.description { + println!(" {}", single_line_text(description)); + } + let fields = state + .schema + .as_ref() + .map(|schema| schema.fields().to_vec()) + .unwrap_or_default(); + if state.schema.is_none() || fields.is_empty() { + edit_dynamic_root_menu(theme, state, &fields) + } else { + let prompt = state + .editor_title + .clone() + .unwrap_or_else(|| state.label.clone()); + edit_dynamic_fields_menu(theme, state, &fields, &[], prompt) + } +} + +fn edit_dynamic_root_menu( + theme: &ColorfulTheme, + state: &mut DynamicPluginEditorState, + fields: &[DynamicConfigField], +) -> Result<(), CliError> { + let mut selected_index = 0; + loop { + let (items, actions) = dynamic_root_menu_items(state, fields); + + let selection = prompt_menu(theme, state.label(), &items, selected_index)?; + if let Some(selected) = menu_response_index(&selection) { + selected_index = selected; + } + if handle_dynamic_root_menu_response(theme, state, &actions, selection)? { + return Ok(()); + } + } +} + +fn handle_dynamic_root_menu_response( + theme: &ColorfulTheme, + state: &mut DynamicPluginEditorState, + actions: &[DynamicMenuAction], + selection: MenuResponse, +) -> Result { + match selection { + MenuResponse::Selected(selected) => match actions.get(selected).copied() { + Some(DynamicMenuAction::EditRawConfig) => { + prompt_raw_config(theme, state)?; + Ok(false) + } + Some(DynamicMenuAction::ResetPlugin) => { + state.reset(); + Ok(false) + } + Some(DynamicMenuAction::Back) | None => Ok(true), + Some(DynamicMenuAction::EditField(_)) => { + println!(" Select Edit raw configuration to modify settings."); + Ok(false) + } + }, + MenuResponse::Shortcut(MenuShortcut::Reset, selected) => { + if matches!(actions.get(selected), Some(DynamicMenuAction::ResetPlugin)) { + state.reset(); + } else { + println!(" Select Reset plugin configuration to remove config."); + } + Ok(false) + } + MenuResponse::Shortcut(MenuShortcut::Clear, selected) => { + if matches!( + actions.get(selected), + Some(DynamicMenuAction::EditRawConfig) + ) { + state.set_raw_config(Map::new()); + } + Ok(false) + } + MenuResponse::Shortcut(MenuShortcut::Help, _) => { + print_editor_help(); + Ok(false) + } + MenuResponse::Shortcut(MenuShortcut::Preview | MenuShortcut::Save, _) => { + println!(" Preview and save are available from the main plugins.toml menu."); + Ok(false) + } + MenuResponse::Cancel => Ok(true), + } +} + +fn edit_dynamic_fields_menu( + theme: &ColorfulTheme, + state: &mut DynamicPluginEditorState, + fields: &[DynamicConfigField], + parent_path: &[String], + prompt: String, +) -> Result<(), CliError> { + let mut selected_index = 0; + loop { + let (items, actions) = dynamic_field_menu_items(state, fields, parent_path); + let selection = prompt_menu(theme, &prompt, &items, selected_index)?; + if let Some(selected) = menu_response_index(&selection) { + selected_index = selected; + } + match selection { + MenuResponse::Selected(selected) => match actions.get(selected).copied() { + Some(DynamicMenuAction::EditField(index)) => { + edit_dynamic_field(theme, state, &fields[index], parent_path)?; + } + Some(DynamicMenuAction::ResetPlugin) => state.reset(), + Some(DynamicMenuAction::Back) | None => return Ok(()), + Some(DynamicMenuAction::EditRawConfig) => unreachable!(), + }, + MenuResponse::Shortcut(MenuShortcut::Reset, selected) => { + reset_dynamic_selection(state, fields, parent_path, &actions, selected); + } + MenuResponse::Shortcut(MenuShortcut::Clear, selected) => { + clear_dynamic_selection(state, fields, parent_path, &actions, selected); + } + MenuResponse::Shortcut(MenuShortcut::Help, _) => print_editor_help(), + MenuResponse::Shortcut(MenuShortcut::Preview | MenuShortcut::Save, _) => { + println!(" Preview and save are available from the main plugins.toml menu."); + } + MenuResponse::Cancel => return Ok(()), + } + } +} + +fn edit_dynamic_field( + theme: &ColorfulTheme, + state: &mut DynamicPluginEditorState, + field: &DynamicConfigField, + parent_path: &[String], +) -> Result<(), CliError> { + if let Some(description) = &field.description { + println!(" {}", single_line_text(description)); + } + let path = field_path(parent_path, field); + if let DynamicConfigFieldKind::Object { fields } = &field.kind { + return edit_dynamic_fields_menu(theme, state, fields, &path, field.title.clone()); + } + if let Some(value) = prompt_dynamic_value(theme, state, field, &path)? { + state.set_field(&path, value); + } + Ok(()) +} + +fn prompt_dynamic_value( + theme: &ColorfulTheme, + state: &DynamicPluginEditorState, + field: &DynamicConfigField, + path: &[String], +) -> Result, CliError> { + let current = state.field_value(path); + match &field.kind { + DynamicConfigFieldKind::Boolean => { + let values = ["false", "true"]; + let default = current + .and_then(Value::as_bool) + .or_else(|| field.default.as_ref().and_then(Value::as_bool)) + .map(usize::from) + .unwrap_or(0); + let selected = Select::with_theme(theme) + .with_prompt(single_line_text(&field.title)) + .items(&values) + .default(default) + .interact() + .map_err(editor_error)?; + Ok(Some(Value::Bool(selected == 1))) + } + DynamicConfigFieldKind::String { secret } => { + prompt_dynamic_string(theme, field, current, *secret, None) + } + DynamicConfigFieldKind::StringEnum { options, secret } => { + if *secret { + prompt_dynamic_string(theme, field, current, true, Some(options)) + } else { + let default = current + .and_then(Value::as_str) + .or_else(|| field.default.as_ref().and_then(Value::as_str)) + .and_then(|value| options.iter().position(|option| option == value)) + .unwrap_or(0); + let selected = Select::with_theme(theme) + .with_prompt(single_line_text(&field.title)) + .items(options) + .default(default) + .interact() + .map_err(editor_error)?; + Ok(Some(Value::String(options[selected].clone()))) + } + } + DynamicConfigFieldKind::Integer => { + let initial = current + .or(field.default.as_ref()) + .map(json_text) + .unwrap_or_default(); + let value: String = Input::with_theme(theme) + .with_prompt(single_line_text(&field.title)) + .with_initial_text(initial) + .interact_text() + .map_err(editor_error)?; + let value = value.trim().parse::().map_err(|error| { + CliError::Config(format!("{} must be an integer: {error}", field.key)) + })?; + Ok(Some(Value::Number(value.into()))) + } + DynamicConfigFieldKind::Number => { + let initial = current + .or(field.default.as_ref()) + .map(json_text) + .unwrap_or_default(); + let value: String = Input::with_theme(theme) + .with_prompt(single_line_text(&field.title)) + .with_initial_text(initial) + .interact_text() + .map_err(editor_error)?; + let parsed = value.trim().parse::().map_err(|error| { + CliError::Config(format!("{} must be a number: {error}", field.key)) + })?; + let number = Number::from_f64(parsed).ok_or_else(|| { + CliError::Config(format!("{} must be a finite number", field.key)) + })?; + Ok(Some(Value::Number(number))) + } + DynamicConfigFieldKind::StringMap => { + let (current, redacted_config, secrets, hidden) = state.field_value_for_raw_edit(path); + let Some(value) = prompt_json_value( + theme, + field, + current.as_ref(), + Value::Object(Map::new()), + hidden, + )? + else { + return Ok(None); + }; + let value = state.restore_raw_field_edit(path, value, redacted_config, &secrets)?; + let object = value + .as_object() + .ok_or_else(|| CliError::Config(format!("{} must be a JSON object", field.key)))?; + if object.values().any(|value| !value.is_string()) { + return Err(CliError::Config(format!( + "{} must contain only string values", + field.key + ))); + } + Ok(Some(value)) + } + DynamicConfigFieldKind::RawJson => { + let fallback = field.default.clone().unwrap_or(Value::Null); + let (current, redacted_config, secrets, hidden) = state.field_value_for_raw_edit(path); + let Some(value) = prompt_json_value(theme, field, current.as_ref(), fallback, hidden)? + else { + return Ok(None); + }; + let value = state.restore_raw_field_edit(path, value, redacted_config, &secrets)?; + Ok(Some(value)) + } + DynamicConfigFieldKind::Object { .. } => unreachable!(), + } +} + +fn prompt_dynamic_string( + theme: &ColorfulTheme, + field: &DynamicConfigField, + current: Option<&Value>, + secret: bool, + options: Option<&[String]>, +) -> Result, CliError> { + if secret { + let title = single_line_text(&field.title); + let value = Password::with_theme(theme) + .with_prompt(format!("New {} (blank preserves the current value)", title)) + .allow_empty_password(true) + .report(false) + .interact() + .map_err(editor_error)?; + if value.is_empty() { + return Ok(None); + } + if options.is_some_and(|options| !options.iter().any(|option| option == &value)) { + return Err(CliError::Config(format!( + "{} must be one of the schema enum values", + field.key + ))); + } + return Ok(Some(Value::String(value))); + } + let initial = current + .and_then(Value::as_str) + .or_else(|| field.default.as_ref().and_then(Value::as_str)) + .unwrap_or_default(); + let value: String = Input::with_theme(theme) + .with_prompt(single_line_text(&field.title)) + .with_initial_text(initial) + .interact_text() + .map_err(editor_error)?; + Ok(Some(Value::String(value))) +} + +fn prompt_json_value( + theme: &ColorfulTheme, + field: &DynamicConfigField, + current: Option<&Value>, + fallback: Value, + hidden: bool, +) -> Result, CliError> { + let initial = current.or(field.default.as_ref()).unwrap_or(&fallback); + let prompt = format!("{} as JSON", single_line_text(&field.title)); + let value = if hidden { + if current.is_some() { + println!(" Current redacted JSON: {}", json_text(initial)); + } + let value = Password::with_theme(theme) + .with_prompt(format!("New {prompt} (blank preserves the current value)")) + .allow_empty_password(true) + .report(false) + .interact() + .map_err(editor_error)?; + if value.is_empty() { + return Ok(None); + } + value + } else { + Input::with_theme(theme) + .with_prompt(prompt) + .with_initial_text(json_text(initial)) + .interact_text() + .map_err(editor_error)? + }; + serde_json::from_str(value.trim()) + .map_err(|error| CliError::Config(format!("invalid JSON for {}: {error}", field.key))) + .map(Some) +} + +fn prompt_raw_config( + theme: &ColorfulTheme, + state: &mut DynamicPluginEditorState, +) -> Result<(), CliError> { + let original = Value::Object(state.config.clone().unwrap_or_default()); + let (initial, secrets, hidden) = state + .schema + .as_ref() + .map(|schema| { + let (redacted, secrets) = schema.redact_for_edit(&original); + (redacted, secrets, schema.has_secrets()) + }) + .unwrap_or_else(|| (original, SecretEditValues::new(), false)); + let value = if hidden { + println!(" Current redacted JSON: {}", json_text(&initial)); + let value = Password::with_theme(theme) + .with_prompt("New configuration as JSON object (blank preserves the current value)") + .allow_empty_password(true) + .report(false) + .interact() + .map_err(editor_error)?; + if value.is_empty() { + return Ok(()); + } + value + } else { + Input::with_theme(theme) + .with_prompt("Configuration as JSON object") + .with_initial_text(json_text(&initial)) + .interact_text() + .map_err(editor_error)? + }; + let value: Value = serde_json::from_str(value.trim()) + .map_err(|error| CliError::Config(format!("invalid JSON configuration: {error}")))?; + let value = match &state.schema { + Some(schema) => schema.restore_edit_secrets(&value, &secrets)?, + None => value, + }; + let object = value.as_object().cloned().ok_or_else(|| { + CliError::Config(format!( + "dynamic plugin '{}' configuration must be a JSON object", + state.plugin_id + )) + })?; + if let Some(schema) = &state.schema { + schema.validate(&value)?; + } + state.set_raw_config(object); + Ok(()) +} diff --git a/crates/cli/src/plugins/mod.rs b/crates/cli/src/plugins/mod.rs index 61b542e2d..c76012c26 100644 --- a/crates/cli/src/plugins/mod.rs +++ b/crates/cli/src/plugins/mod.rs @@ -1,18 +1,14 @@ // SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. // SPDX-License-Identifier: Apache-2.0 -//! Interactive plugin configuration editing. +//! Plugin configuration state and deterministic editor behavior. //! -//! Keep this module focused on TTY and `dialoguer` orchestration. New testable plugin config -//! behavior should live in `plugins/config_io.rs` or `plugins/editor_model.rs`, with focused unit -//! tests, so Codecov does not depend on exercising interactive prompt loops. +//! Terminal-only interaction lives in `plugins/prompt.rs`. -use std::io::IsTerminal; use std::path::{Path, PathBuf}; -use console::{Key, Term, style, truncate_str}; +use console::{Key, style, truncate_str}; use dialoguer::theme::ColorfulTheme; -use dialoguer::{Input, Select}; use nemo_relay::config_editor::{EditorFieldKind, EditorFieldSpec}; use serde_json::{Value, json}; @@ -24,6 +20,7 @@ mod editor_model; pub(crate) mod lifecycle; pub(crate) mod policy; pub(crate) mod pricing; +mod prompt; pub(crate) mod schema; mod types; @@ -32,6 +29,10 @@ pub(crate) use types::*; use self::config_io::*; use self::dynamic_editor::*; use self::editor_model::*; +use self::prompt::{editor_error, print_editor_help, prompt_menu}; + +#[cfg(test)] +use self::prompt::menu_error; const PLUGIN_EDIT_CANCELLED_MESSAGE: &str = "plugin edit cancelled; no plugin changes saved"; @@ -103,46 +104,7 @@ fn print_save_success(path: &Path) { } pub(crate) fn edit(command: PluginsEditRequest) -> Result<(), CliError> { - ensure_tty()?; - let (scope, path) = resolve_edit_target(command)?; - let mut document = PluginConfigDocument::read(&path)?; - ensure_observability_component(document.config_mut())?; - ensure_adaptive_component(document.config_mut())?; - let mut components = editable_components(document.config())?; - let mut dynamic_plugins = load_dynamic_plugin_states(&document)?; - - let theme = ColorfulTheme::default(); - crate::banner::print_intro(); - println!( - " Editing plugin config at {}", - single_line_text(&path.display().to_string()) - ); - println!(" Tip: ↑/↓ or j/k to move, PageUp/PageDown to scroll, SPACE/ENTER to select."); - println!(); - let mut selected_index = 0; - loop { - let dynamic_rows = dynamic_plugins - .iter() - .map(|plugin| (plugin.label().to_owned(), plugin.menu_summary())) - .collect::>(); - let (items, actions) = plugin_menu_items(&components, &dynamic_rows, &path); - let selection = prompt_menu(&theme, "plugins.toml", &items, selected_index)?; - if let Some(selected) = menu_response_index(&selection) { - selected_index = selected; - } - if handle_menu_response( - &theme, - &mut document, - &mut components, - &mut dynamic_plugins, - &actions, - selection, - scope, - )? == EditLoopControl::Finish - { - return Ok(()); - } - } + prompt::edit(command) } pub(crate) fn resolve_edit_target( @@ -156,109 +118,6 @@ pub(crate) fn resolve_edit_target( Ok((scope, path)) } -fn handle_menu_response( - theme: &ColorfulTheme, - document: &mut PluginConfigDocument, - components: &mut [EditableComponent], - dynamic_plugins: &mut [DynamicPluginEditorState], - actions: &[MenuAction], - selection: MenuResponse, - scope: TargetScope, -) -> Result { - match selection { - MenuResponse::Selected(selection) => handle_menu_action( - theme, - document, - components, - dynamic_plugins, - actions.get(selection).copied(), - scope, - ), - MenuResponse::Shortcut(MenuShortcut::Preview, _) => { - preview_document(document, components, dynamic_plugins)?; - Ok(EditLoopControl::Continue) - } - MenuResponse::Shortcut(MenuShortcut::Save, _) => { - save_document(document, components, dynamic_plugins, scope) - } - MenuResponse::Shortcut(MenuShortcut::Help, _) => { - print_editor_help(); - Ok(EditLoopControl::Continue) - } - MenuResponse::Shortcut( - shortcut @ (MenuShortcut::Reset | MenuShortcut::Clear), - selected, - ) => handle_reset_or_clear_shortcut(components, actions.get(selected).copied(), shortcut), - MenuResponse::Cancel => Err(cancelled_error()), - } -} - -fn handle_menu_action( - theme: &ColorfulTheme, - document: &mut PluginConfigDocument, - components: &mut [EditableComponent], - dynamic_plugins: &mut [DynamicPluginEditorState], - action: Option, - scope: TargetScope, -) -> Result { - match action { - Some(MenuAction::EditComponent(component_index)) => { - if let Some(component) = components.get_mut(component_index) { - edit_component(theme, component)?; - } - Ok(EditLoopControl::Continue) - } - Some(MenuAction::EditDynamic(dynamic_index)) => { - if let Some(plugin) = dynamic_plugins.get_mut(dynamic_index) { - edit_dynamic_plugin(theme, plugin)?; - } - Ok(EditLoopControl::Continue) - } - Some(MenuAction::Preview) => { - preview_document(document, components, dynamic_plugins)?; - Ok(EditLoopControl::Continue) - } - Some(MenuAction::Save) => save_document(document, components, dynamic_plugins, scope), - Some(MenuAction::Cancel) | None => Err(cancelled_error()), - } -} - -fn edit_component( - theme: &ColorfulTheme, - component: &mut EditableComponent, -) -> Result<(), CliError> { - let mut selected_index = 0; - loop { - let (items, actions) = component_menu_items(component); - let selection = prompt_menu(theme, component.label(), &items, selected_index)?; - if let Some(selected) = menu_response_index(&selection) { - selected_index = selected; - } - match selection { - MenuResponse::Selected(selected) => match actions.get(selected).copied() { - Some(ComponentMenuAction::Toggle) => component.toggle_enabled(), - Some(ComponentMenuAction::EditField(field_index)) => { - if let Some(field) = component.fields().get(field_index) { - edit_component_field(theme, component, *field)?; - } - } - Some(ComponentMenuAction::Back) | None => return Ok(()), - }, - MenuResponse::Shortcut(MenuShortcut::Reset, selected) => { - reset_component_menu_item(component, actions.get(selected).copied())?; - } - MenuResponse::Shortcut(MenuShortcut::Clear, selected) => { - clear_component_menu_item(component, actions.get(selected).copied())?; - } - MenuResponse::Shortcut(MenuShortcut::Help, _) => print_editor_help(), - MenuResponse::Shortcut(MenuShortcut::Preview | MenuShortcut::Save, _) => { - println!(" Preview and save are available from the main plugins.toml menu."); - } - MenuResponse::Cancel => return Ok(()), - } - } -} - fn preview_document( document: &PluginConfigDocument, components: &[EditableComponent], @@ -358,37 +217,6 @@ fn cancelled_error() -> CliError { CliError::Config(PLUGIN_EDIT_CANCELLED_MESSAGE.into()) } -fn edit_component_field( - theme: &ColorfulTheme, - component: &mut EditableComponent, - field: EditorFieldSpec, -) -> Result<(), CliError> { - match component { - EditableComponent::Observability(state) => { - edit_section(theme, &mut state.config, field)?; - state.mark_config_touched(); - } - EditableComponent::Adaptive(state) => { - edit_config_field(theme, &mut state.config, field)?; - state.mark_config_touched(); - } - EditableComponent::NemoGuardrails(state) => { - edit_config_field(theme, &mut state.config, field)?; - state.mark_config_touched(); - } - EditableComponent::PiiRedaction(state) => { - edit_config_field(theme, &mut state.config, field)?; - state.mark_config_touched(); - } - #[cfg(feature = "switchyard")] - EditableComponent::Switchyard(state) => { - edit_config_field(theme, &mut state.config, field)?; - state.mark_config_touched(); - } - } - Ok(()) -} - fn menu_response_index(response: &MenuResponse) -> Option { match response { MenuResponse::Selected(index) @@ -404,75 +232,16 @@ fn menu_response_index(response: &MenuResponse) -> Option { } } -fn prompt_menu( - theme: &ColorfulTheme, - prompt: &str, - items: &[MenuItem], - default: usize, -) -> Result { - if items.is_empty() { - return Err(CliError::Config(format!("{prompt} menu has no items"))); - } - let term = Term::stderr(); - let mut selected = default.min(items.len() - 1); - let mut rendered_lines = 0; - loop { - if rendered_lines > 0 { - term.clear_last_lines(rendered_lines).map_err(menu_error)?; - } - let (rows, columns) = term.size(); - let viewport = menu_viewport(items.len(), selected, usize::from(rows)); - let lines = render_menu_for_size( - theme, - prompt, - items, - selected, - usize::from(rows), - usize::from(columns), - ); - rendered_lines = lines.len(); - for line in &lines { - term.write_line(line).map_err(menu_error)?; - } - term.flush().map_err(menu_error)?; - let key = term.read_key().map_err(menu_error)?; - if let Some(next) = - menu_selection_after_key(&key, selected, items.len(), viewport.page_size) - { - selected = next; - continue; - } - match key { - Key::Enter | Key::Char(' ') => { - clear_menu(&term, rendered_lines)?; - return Ok(MenuResponse::Selected(selected)); - } - Key::Char('p') => { - clear_menu(&term, rendered_lines)?; - return Ok(MenuResponse::Shortcut(MenuShortcut::Preview, selected)); - } - Key::Char('s') => { - clear_menu(&term, rendered_lines)?; - return Ok(MenuResponse::Shortcut(MenuShortcut::Save, selected)); - } - Key::Char('r') => { - clear_menu(&term, rendered_lines)?; - return Ok(MenuResponse::Shortcut(MenuShortcut::Reset, selected)); - } - Key::Backspace | Key::Del => { - clear_menu(&term, rendered_lines)?; - return Ok(MenuResponse::Shortcut(MenuShortcut::Clear, selected)); - } - Key::Char('?') => { - clear_menu(&term, rendered_lines)?; - return Ok(MenuResponse::Shortcut(MenuShortcut::Help, selected)); - } - Key::Escape | Key::CtrlC | Key::Char('q') => { - clear_menu(&term, rendered_lines)?; - return Ok(MenuResponse::Cancel); - } - _ => {} - } +fn menu_response_for_key(key: &Key, selected: usize) -> Option { + match key { + Key::Enter | Key::Char(' ') => Some(MenuResponse::Selected(selected)), + Key::Char('p') => Some(MenuResponse::Shortcut(MenuShortcut::Preview, selected)), + Key::Char('s') => Some(MenuResponse::Shortcut(MenuShortcut::Save, selected)), + Key::Char('r') => Some(MenuResponse::Shortcut(MenuShortcut::Reset, selected)), + Key::Backspace | Key::Del => Some(MenuResponse::Shortcut(MenuShortcut::Clear, selected)), + Key::Char('?') => Some(MenuResponse::Shortcut(MenuShortcut::Help, selected)), + Key::Escape | Key::CtrlC | Key::Char('q') => Some(MenuResponse::Cancel), + _ => None, } } @@ -634,119 +403,6 @@ fn single_line_text(value: &str) -> String { .collect() } -fn clear_menu(term: &Term, rendered_lines: usize) -> Result<(), CliError> { - if rendered_lines > 0 { - term.clear_last_lines(rendered_lines).map_err(menu_error)?; - } - Ok(()) -} - -fn menu_error(error: std::io::Error) -> CliError { - if matches!( - error.kind(), - std::io::ErrorKind::Interrupted | std::io::ErrorKind::UnexpectedEof - ) { - CliError::Config(PLUGIN_EDIT_CANCELLED_MESSAGE.into()) - } else { - CliError::Config(format!("plugin editor terminal error: {error}")) - } -} - -fn print_editor_help() { - println!(); - println!( - "{} {}", - style("?").yellow(), - style("Plugin editor keys").bold() - ); - println!(" {} move", style("↑/↓ or j/k").cyan()); - println!( - " {} move by page or jump to an end", - style("PageUp/PageDown, Home/End").cyan() - ); - println!( - " {} select/toggle the highlighted item", - style("Enter/Space").cyan() - ); - println!( - " {} reset the highlighted field or section", - style("r").cyan() - ); - println!( - " {} clear the highlighted optional field", - style("Backspace/Del").cyan() - ); - println!( - " {} preview TOML from the main menu", - style("p").cyan() - ); - println!( - " {} save from the main menu", - style("s").cyan() - ); - println!(" {} go back/cancel", style("q or Esc").cyan()); -} - -fn ensure_tty() -> Result<(), CliError> { - if !std::io::stdin().is_terminal() - || !std::io::stdout().is_terminal() - || !std::io::stderr().is_terminal() - { - return Err(CliError::Config( - "interactive plugin editing requires a TTY".into(), - )); - } - Ok(()) -} - -fn edit_section( - theme: &ColorfulTheme, - config: &mut T, - section: EditorFieldSpec, -) -> Result<(), CliError> -where - T: SerializeConfig, -{ - let fields = section - .schema() - .ok_or_else(|| CliError::Config(format!("{} is not an editable section", section.name)))? - .fields; - let mut selected_index = 0; - loop { - let items = section_menu_items(config, section, fields)?; - let selection = prompt_menu(theme, section.name, &items, selected_index)?; - if let Some(selected) = menu_response_index(&selection) { - selected_index = selected; - } - let selection = match selection { - MenuResponse::Selected(selection) => selection, - MenuResponse::Shortcut(MenuShortcut::Help, _) => { - print_editor_help(); - continue; - } - MenuResponse::Shortcut(MenuShortcut::Reset, selected) => { - reset_selected_item(config, section, fields, selected)?; - continue; - } - MenuResponse::Shortcut(MenuShortcut::Clear, selected) => { - if reset_selected_field(config, section, fields, selected)? { - continue; - } - println!(" Select a field to clear."); - continue; - } - MenuResponse::Shortcut(MenuShortcut::Preview | MenuShortcut::Save, _) => { - println!(" Preview and save are available from the main plugins.toml menu."); - continue; - } - MenuResponse::Cancel => return Ok(()), - }; - if !edit_selected_section_item(theme, config, section, fields, selection)? { - return Ok(()); - } - } -} - fn section_menu_items( config: &T, section: EditorFieldSpec, @@ -820,393 +476,12 @@ where Ok(()) } -fn edit_selected_section_item( - theme: &ColorfulTheme, - config: &mut T, - section: EditorFieldSpec, - fields: &[EditorFieldSpec], - selection: usize, -) -> Result -where - T: SerializeConfig, -{ - if section_has_enabled_toggle(section) && selection == 0 { - toggle_section(config, section); - return Ok(true); - } - let index = selected_field_index(section, selection); - if let Some(field) = fields.get(index) { - edit_field(theme, config, section, field)?; - return Ok(true); - } - if index == fields.len() { - reset_section(config, section); - return Ok(true); - } - Ok(false) -} - -fn edit_field( - theme: &ColorfulTheme, - config: &mut T, - section: EditorFieldSpec, - field: &EditorFieldSpec, -) -> Result<(), CliError> -where - T: SerializeConfig, -{ - if field.kind == EditorFieldKind::Section { - edit_nested_section(theme, config, section, *field)?; - return Ok(()); - } - let current = section_field_value(config, section, field.name)?; - if field.kind == EditorFieldKind::List { - let item = field.list_item.ok_or_else(|| { - CliError::Config(format!("{} does not describe its list entries", field.name)) - })?; - let default = section_field_default(section, *field); - let mut items = current - .or_else(|| default.clone()) - .unwrap_or_else(|| json!([])); - if edit_list_value( - theme, - &format!("{}.{}", section.name, field.name), - &mut items, - default, - item, - )? { - set_section_field(config, section, field.name, items)?; - } - return Ok(()); - } - if field.kind == EditorFieldKind::StringMap { - let default = section_field_default(section, *field); - let mut entries = current - .or_else(|| default.clone()) - .unwrap_or_else(|| json!({})); - if edit_string_map_value( - theme, - &format!("{}.{}", section.name, field.name), - &mut entries, - default, - )? { - set_section_field(config, section, field.name, entries)?; - } - return Ok(()); - } - if field.kind == EditorFieldKind::TaggedUnion { - let tagged_union = field.tagged_union.ok_or_else(|| { - CliError::Config(format!("{} does not describe its variants", field.name)) - })?; - let default = section_field_default(section, *field); - match edit_tagged_union_field( - theme, - &format!("{}.{}", section.name, field.name), - current, - default, - tagged_union, - )? { - TaggedUnionFieldEdit::Set(value) => { - set_section_field(config, section, field.name, value)?; - } - TaggedUnionFieldEdit::Reset => remove_section_field(config, section, field.name)?, - TaggedUnionFieldEdit::Unchanged => {} - } - return Ok(()); - } - let actions = [ - MenuItem::new("Set value"), - MenuItem::new(shortcut_label( - "Reset to default/none", - "r, Backspace, Delete", - )), - MenuItem::new(shortcut_label("Back", "q")), - ]; - let action = prompt_menu( - theme, - &format!( - "{}.{}, current {}", - section.name, - field.name, - current - .as_ref() - .map(|value| display_field_value(section, *field, value)) - .unwrap_or_else(|| "(default)".to_string()) - ), - &actions, - 0, - )?; - match action { - MenuResponse::Selected(0) => { - let value = prompt_value(theme, field, current.as_ref())?; - set_section_field(config, section, field.name, value)?; - } - MenuResponse::Selected(1) - | MenuResponse::Shortcut(MenuShortcut::Reset | MenuShortcut::Clear, _) => { - remove_section_field(config, section, field.name)? - } - MenuResponse::Shortcut(MenuShortcut::Help, _) => print_editor_help(), - MenuResponse::Shortcut(MenuShortcut::Preview | MenuShortcut::Save, _) => { - println!(" Preview and save are available from the main plugins.toml menu."); - } - _ => {} - } - Ok(()) -} - -fn edit_list_value( - theme: &ColorfulTheme, - prompt: &str, - value: &mut Value, - default: Option, - item: &nemo_relay::config_editor::EditorListItemSpec, -) -> Result { - if !value.is_array() { - *value = default.clone().unwrap_or_else(|| json!([])); - } - let original = value.clone(); - let mut selected_index = 0; - loop { - let entries = value.as_array().expect("list value is an array"); - let mut menu = vec![MenuItem::new("Add item")]; - menu.extend(entries.iter().enumerate().map(|(index, entry)| { - MenuItem::new(format!( - "Edit item {}: {}", - index + 1, - editor_item_label(entry, item) - )) - })); - menu.push(MenuItem::new(shortcut_label("Back", "q"))); - let selection = prompt_menu(theme, prompt, &menu, selected_index)?; - if let Some(selected) = menu_response_index(&selection) { - selected_index = selected; - } - match selection { - MenuResponse::Selected(0) => { - let mut entry = new_editor_item(theme, item)?; - edit_editor_item( - theme, - &format!("{prompt}[{}]", entries.len()), - &mut entry, - item, - )?; - value - .as_array_mut() - .expect("list value is an array") - .push(entry); - } - MenuResponse::Selected(index) if index <= entries.len() => { - edit_existing_list_item(theme, prompt, value, index - 1, item)?; - } - MenuResponse::Cancel | MenuResponse::Selected(_) => return Ok(*value != original), - MenuResponse::Shortcut(MenuShortcut::Help, _) => print_editor_help(), - MenuResponse::Shortcut(shortcut @ (MenuShortcut::Reset | MenuShortcut::Clear), _) => { - *value = collection_shortcut_value(default.as_ref(), json!([]), shortcut) - } - MenuResponse::Shortcut(MenuShortcut::Preview | MenuShortcut::Save, _) => { - println!(" Preview and save are available from the main plugins.toml menu."); - } - } - } -} - -fn edit_string_map_value( - theme: &ColorfulTheme, - prompt: &str, - value: &mut Value, - default: Option, -) -> Result { - if !value.is_object() { - *value = default.clone().unwrap_or_else(|| json!({})); - } - let original = value.clone(); - let mut selected_index = 0; - loop { - let entries = value.as_object().expect("string map value is an object"); - let keys = entries.keys().cloned().collect::>(); - let mut menu = vec![MenuItem::new("Add entry")]; - menu.extend(keys.iter().map(|key| { - MenuItem::new(format!( - "Edit {key}: {}", - entries.get(key).map(display_value).unwrap_or_default() - )) - })); - menu.push(MenuItem::new(shortcut_label("Back", "q"))); - let selection = prompt_menu(theme, prompt, &menu, selected_index)?; - if let Some(selected) = menu_response_index(&selection) { - selected_index = selected; - } - match selection { - MenuResponse::Selected(0) => { - let key: String = Input::with_theme(theme) - .with_prompt("Entry key") - .interact_text() - .map_err(editor_error)?; - if key.trim().is_empty() { - println!(" Entry key must not be empty."); - continue; - } - let key = key.trim().to_owned(); - if string_map_entry_exists(value, &key) { - println!(" Entry already exists; select it to edit."); - continue; - } - let entry: String = Input::with_theme(theme) - .with_prompt("Entry value") - .interact_text() - .map_err(editor_error)?; - value - .as_object_mut() - .expect("string map value is an object") - .insert(key, Value::String(entry)); - } - MenuResponse::Selected(index) if index <= keys.len() => { - edit_existing_string_map_entry(theme, prompt, value, &keys[index - 1])?; - } - MenuResponse::Cancel | MenuResponse::Selected(_) => return Ok(*value != original), - MenuResponse::Shortcut(MenuShortcut::Help, _) => print_editor_help(), - MenuResponse::Shortcut(shortcut @ (MenuShortcut::Reset | MenuShortcut::Clear), _) => { - *value = collection_shortcut_value(default.as_ref(), json!({}), shortcut) - } - MenuResponse::Shortcut(MenuShortcut::Preview | MenuShortcut::Save, _) => { - println!(" Preview and save are available from the main plugins.toml menu."); - } - } - } -} - fn string_map_entry_exists(value: &Value, key: &str) -> bool { value .as_object() .is_some_and(|entries| entries.contains_key(key.trim())) } -fn edit_existing_string_map_entry( - theme: &ColorfulTheme, - prompt: &str, - value: &mut Value, - key: &str, -) -> Result<(), CliError> { - let actions = [ - MenuItem::new("Edit value"), - MenuItem::new("Remove entry"), - MenuItem::new(shortcut_label("Back", "q")), - ]; - match prompt_menu(theme, &format!("{prompt}.{key}"), &actions, 0)? { - MenuResponse::Selected(0) => { - let current = value - .as_object() - .and_then(|entries| entries.get(key)) - .and_then(Value::as_str) - .unwrap_or_default(); - let entry: String = Input::with_theme(theme) - .with_prompt("Entry value") - .with_initial_text(current) - .interact_text() - .map_err(editor_error)?; - value - .as_object_mut() - .expect("string map value is an object") - .insert(key.to_owned(), Value::String(entry)); - } - MenuResponse::Selected(1) | MenuResponse::Shortcut(MenuShortcut::Clear, _) => { - value - .as_object_mut() - .expect("string map value is an object") - .remove(key); - } - _ => {} - } - Ok(()) -} - -fn edit_existing_list_item( - theme: &ColorfulTheme, - prompt: &str, - value: &mut Value, - index: usize, - item: &nemo_relay::config_editor::EditorListItemSpec, -) -> Result<(), CliError> { - let actions = [ - MenuItem::new("Edit item"), - MenuItem::new("Remove item"), - MenuItem::new(shortcut_label("Back", "q")), - ]; - match prompt_menu(theme, &format!("{prompt}[{}]", index + 1), &actions, 0)? { - MenuResponse::Selected(0) => { - if let Some(entry) = value - .as_array_mut() - .and_then(|entries| entries.get_mut(index)) - { - edit_editor_item(theme, &format!("{prompt}[{}]", index + 1), entry, item)?; - } - } - MenuResponse::Selected(1) | MenuResponse::Shortcut(MenuShortcut::Clear, _) => { - value - .as_array_mut() - .expect("list value is an array") - .remove(index); - } - _ => {} - } - Ok(()) -} - -fn new_editor_item( - theme: &ColorfulTheme, - item: &nemo_relay::config_editor::EditorListItemSpec, -) -> Result { - if let Some(tagged_union) = item.tagged_union { - return new_tagged_union_value(theme, tagged_union); - } - Ok(item.default.map(|default| default()).unwrap_or(Value::Null)) -} - -fn edit_editor_item( - theme: &ColorfulTheme, - prompt: &str, - value: &mut Value, - item: &nemo_relay::config_editor::EditorListItemSpec, -) -> Result<(), CliError> { - if let Some(tagged_union) = item.tagged_union { - return edit_tagged_union_payload(theme, prompt, value, tagged_union); - } - - match item.kind { - EditorFieldKind::Section => { - let schema = item - .schema - .ok_or_else(|| CliError::Config("list item has no schema".into()))?( - ); - edit_value_section(theme, prompt, value, schema, None)?; - } - EditorFieldKind::List => { - let nested = item.list_item.ok_or_else(|| { - CliError::Config("nested list item has no entry description".into()) - })?; - let _ = edit_list_value(theme, prompt, value, None, nested)?; - } - EditorFieldKind::StringMap => { - let _ = edit_string_map_value(theme, prompt, value, None)?; - } - kind => { - let field = EditorFieldSpec { - name: "item", - label: "item", - kind, - enum_values: &[], - optional: false, - nested_schema: None, - nested_default: None, - list_item: None, - tagged_union: None, - }; - *value = prompt_value(theme, &field, Some(value))?; - } - } - Ok(()) -} - fn editor_item_label( value: &Value, item: &nemo_relay::config_editor::EditorListItemSpec, @@ -1232,59 +507,6 @@ fn tagged_union_variant_value( .ok_or_else(|| CliError::Config("tagged union variant does not exist".into())) } -fn select_tagged_union_variant( - theme: &ColorfulTheme, - tagged_union: &nemo_relay::config_editor::EditorTaggedUnionSpec, -) -> Result { - if tagged_union.variants.is_empty() { - return Err(CliError::Config("tagged union has no variants".into())); - } - Select::with_theme(theme) - .with_prompt("Variant type") - .items( - &tagged_union - .variants - .iter() - .map(|variant| variant.label) - .collect::>(), - ) - .default(0) - .interact() - .map_err(editor_error) -} - -fn new_tagged_union_value( - theme: &ColorfulTheme, - tagged_union: &nemo_relay::config_editor::EditorTaggedUnionSpec, -) -> Result { - tagged_union_variant_value( - tagged_union, - select_tagged_union_variant(theme, tagged_union)?, - ) -} - -fn edit_tagged_union_payload( - theme: &ColorfulTheme, - prompt: &str, - value: &mut Value, - tagged_union: &nemo_relay::config_editor::EditorTaggedUnionSpec, -) -> Result<(), CliError> { - if !value.is_object() { - *value = new_tagged_union_value(theme, tagged_union)?; - } - let tag = value - .get(tagged_union.discriminator) - .and_then(Value::as_str) - .ok_or_else(|| CliError::Config("tagged union has no discriminator value".into()))?; - let variant = tagged_union - .variants - .iter() - .find(|variant| variant.tag == tag) - .ok_or_else(|| CliError::Config(format!("unknown tagged union type {tag:?}")))?; - edit_value_section(theme, prompt, value, (variant.schema)(), None)?; - Ok(()) -} - #[derive(Debug, PartialEq)] enum TaggedUnionFieldEdit { Set(Value), @@ -1336,55 +558,6 @@ impl TaggedUnionFieldState { } } -fn edit_tagged_union_field( - theme: &ColorfulTheme, - prompt: &str, - current: Option, - default: Option, - tagged_union: &nemo_relay::config_editor::EditorTaggedUnionSpec, -) -> Result { - let mut state = TaggedUnionFieldState::new(current, default); - loop { - let actions = [ - MenuItem::new("Edit fields"), - MenuItem::new("Change variant"), - MenuItem::new(shortcut_label( - "Reset to default/none", - "r, Backspace, Delete", - )), - MenuItem::new(shortcut_label("Back", "q")), - ]; - match prompt_menu( - theme, - &format!("{prompt}, current {}", display_value(state.value())), - &actions, - 0, - )? { - MenuResponse::Selected(0) => { - edit_tagged_union_payload(theme, prompt, state.value_mut(), tagged_union)?; - } - MenuResponse::Selected(1) => { - state.change_variant( - tagged_union, - select_tagged_union_variant(theme, tagged_union)?, - )?; - edit_tagged_union_payload(theme, prompt, state.value_mut(), tagged_union)?; - } - MenuResponse::Selected(2) - | MenuResponse::Shortcut(MenuShortcut::Reset | MenuShortcut::Clear, _) => { - return Ok(state.reset()); - } - MenuResponse::Shortcut(MenuShortcut::Help, _) => print_editor_help(), - MenuResponse::Shortcut(MenuShortcut::Preview | MenuShortcut::Save, _) => { - println!(" Preview and save are available from the main plugins.toml menu."); - } - MenuResponse::Cancel | MenuResponse::Selected(_) => { - return Ok(state.finish()); - } - } - } -} - fn collection_shortcut_value( default: Option<&Value>, empty: Value, @@ -1397,141 +570,6 @@ fn collection_shortcut_value( } } -fn edit_config_field( - theme: &ColorfulTheme, - config: &mut T, - field: EditorFieldSpec, -) -> Result<(), CliError> -where - T: Default + SerializeConfig, -{ - if field.kind == EditorFieldKind::Section { - let mut value = config_field_value(config, field.name)? - .or_else(|| field.default_value()) - .unwrap_or_else(|| json!({})); - let schema = field.schema().ok_or_else(|| { - CliError::Config(format!("{} is not an editable section", field.name)) - })?; - if edit_value_section(theme, field.name, &mut value, schema, field.default_value())? { - store_edited_config_section(config, field, value)?; - } - return Ok(()); - } - - if field.kind == EditorFieldKind::List { - let item = field.list_item.ok_or_else(|| { - CliError::Config(format!("{} does not describe its list entries", field.name)) - })?; - let default = default_config_field_value::(field).or_else(|| field.default_value()); - let mut items = config_field_value(config, field.name)? - .or_else(|| default.clone()) - .unwrap_or_else(|| json!([])); - if edit_list_value(theme, field.name, &mut items, default, item)? { - set_struct_field(config, field.name, items)?; - } - return Ok(()); - } - - if field.kind == EditorFieldKind::StringMap { - let default = default_config_field_value::(field).or_else(|| field.default_value()); - let mut entries = config_field_value(config, field.name)? - .or_else(|| default.clone()) - .unwrap_or_else(|| json!({})); - if edit_string_map_value(theme, field.name, &mut entries, default)? { - set_struct_field(config, field.name, entries)?; - } - return Ok(()); - } - - if field.kind == EditorFieldKind::TaggedUnion { - let tagged_union = field.tagged_union.ok_or_else(|| { - CliError::Config(format!("{} does not describe its variants", field.name)) - })?; - let default = default_config_field_value::(field).or_else(|| field.default_value()); - match edit_tagged_union_field( - theme, - field.name, - config_field_value(config, field.name)?, - default, - tagged_union, - )? { - TaggedUnionFieldEdit::Set(value) => set_struct_field(config, field.name, value)?, - TaggedUnionFieldEdit::Reset => reset_config_field(config, field)?, - TaggedUnionFieldEdit::Unchanged => {} - } - return Ok(()); - } - - let current = config_field_value(config, field.name)?; - let actions = [ - MenuItem::new("Set value"), - MenuItem::new(shortcut_label( - "Reset to default/none", - "r, Backspace, Delete", - )), - MenuItem::new(shortcut_label("Back", "q")), - ]; - let action = prompt_menu( - theme, - &format!( - "{}, current {}", - field.label, - current - .as_ref() - .map(display_value) - .or_else(|| default_config_field_value::(field) - .map(|value| { format!("{} (default)", display_value(&value)) })) - .unwrap_or_else(|| "(default)".to_string()) - ), - &actions, - 0, - )?; - match action { - MenuResponse::Selected(0) => { - let value = prompt_value(theme, &field, current.as_ref())?; - set_struct_field(config, field.name, value)?; - } - MenuResponse::Selected(1) - | MenuResponse::Shortcut(MenuShortcut::Reset | MenuShortcut::Clear, _) => { - reset_config_field(config, field)? - } - MenuResponse::Shortcut(MenuShortcut::Help, _) => print_editor_help(), - MenuResponse::Shortcut(MenuShortcut::Preview | MenuShortcut::Save, _) => { - println!(" Preview and save are available from the main plugins.toml menu."); - } - _ => {} - } - Ok(()) -} - -fn edit_nested_section( - theme: &ColorfulTheme, - config: &mut T, - section: EditorFieldSpec, - field: EditorFieldSpec, -) -> Result<(), CliError> -where - T: SerializeConfig, -{ - let mut value = section_field_value(config, section, field.name)? - .or_else(|| section_field_default(section, field)) - .unwrap_or_else(|| json!({})); - let schema = field - .schema() - .ok_or_else(|| CliError::Config(format!("{} is not an editable section", field.name)))?; - let default = section_field_default(section, field); - if edit_value_section( - theme, - &format!("{}.{}", section.name, field.name), - &mut value, - schema, - default, - )? { - store_edited_section_field(config, section, field, value)?; - } - Ok(()) -} - fn section_field_default(section: EditorFieldSpec, field: EditorFieldSpec) -> Option { default_field_value(section, field).or_else(|| field.default_value()) } @@ -1567,51 +605,6 @@ where } } -fn edit_value_section( - theme: &ColorfulTheme, - prompt: &str, - value: &mut Value, - schema: &nemo_relay::config_editor::EditorSchema, - default: Option, -) -> Result { - ensure_object(value); - let original = value.clone(); - let mut selected_index = 0; - loop { - let items = value_section_menu_items(value, schema, default.as_ref())?; - let selection = prompt_menu(theme, prompt, &items, selected_index)?; - if let Some(selected) = menu_response_index(&selection) { - selected_index = selected; - } - let selection = match selection { - MenuResponse::Selected(selection) => selection, - MenuResponse::Shortcut(MenuShortcut::Help, _) => { - print_editor_help(); - continue; - } - MenuResponse::Shortcut(MenuShortcut::Reset, selected) => { - reset_value_section_item(value, schema, default.as_ref(), selected); - continue; - } - MenuResponse::Shortcut(MenuShortcut::Clear, selected) => { - if clear_value_field(value, schema, selected) { - continue; - } - println!(" Select a field to clear."); - continue; - } - MenuResponse::Shortcut(MenuShortcut::Preview | MenuShortcut::Save, _) => { - println!(" Preview and save are available from the main plugins.toml menu."); - continue; - } - MenuResponse::Cancel => return Ok(*value != original), - }; - if !edit_selected_value_item(theme, prompt, value, schema, default.as_ref(), selection)? { - return Ok(*value != original); - } - } -} - fn value_section_menu_items( value: &Value, schema: &nemo_relay::config_editor::EditorSchema, @@ -1647,156 +640,6 @@ fn value_field_menu_item( ))) } -fn edit_selected_value_item( - theme: &ColorfulTheme, - prompt: &str, - value: &mut Value, - schema: &nemo_relay::config_editor::EditorSchema, - default: Option<&Value>, - selection: usize, -) -> Result { - if let Some(field) = schema.fields.get(selection) { - edit_value_field(theme, prompt, value, *field, default)?; - return Ok(true); - } - if selection == schema.fields.len() { - *value = default.cloned().unwrap_or_else(|| json!({})); - ensure_object(value); - return Ok(true); - } - Ok(false) -} - -fn edit_value_field( - theme: &ColorfulTheme, - prompt: &str, - value: &mut Value, - field: EditorFieldSpec, - default: Option<&Value>, -) -> Result<(), CliError> { - if field.kind == EditorFieldKind::Section { - let nested_default = value_field_default(default, field); - let mut nested_value = value_field_value(value, field.name) - .or_else(|| nested_default.clone()) - .unwrap_or_else(|| json!({})); - let nested_schema = field.schema().ok_or_else(|| { - CliError::Config(format!("{} is not an editable section", field.name)) - })?; - if edit_value_section( - theme, - &format!("{prompt}.{}", field.name), - &mut nested_value, - nested_schema, - nested_default, - )? { - store_edited_value_section(value, field, nested_value); - } - return Ok(()); - } - - if field.kind == EditorFieldKind::List { - let item = field.list_item.ok_or_else(|| { - CliError::Config(format!("{} does not describe its list entries", field.name)) - })?; - let field_default = value_field_default(default, field); - let mut items = value_field_value(value, field.name) - .or_else(|| field_default.clone()) - .unwrap_or_else(|| json!([])); - if edit_list_value( - theme, - &format!("{prompt}.{}", field.name), - &mut items, - field_default, - item, - )? { - set_value_field(value, field.name, items); - } - return Ok(()); - } - - if field.kind == EditorFieldKind::StringMap { - let field_default = value_field_default(default, field); - let mut entries = value_field_value(value, field.name) - .or_else(|| field_default.clone()) - .unwrap_or_else(|| json!({})); - if edit_string_map_value( - theme, - &format!("{prompt}.{}", field.name), - &mut entries, - field_default, - )? { - set_value_field(value, field.name, entries); - } - return Ok(()); - } - - if field.kind == EditorFieldKind::TaggedUnion { - let tagged_union = field.tagged_union.ok_or_else(|| { - CliError::Config(format!("{} does not describe its variants", field.name)) - })?; - let field_default = value_field_default(default, field); - match edit_tagged_union_field( - theme, - &format!("{prompt}.{}", field.name), - value_field_value(value, field.name), - field_default.clone(), - tagged_union, - )? { - TaggedUnionFieldEdit::Set(tagged_value) => { - set_value_field(value, field.name, tagged_value); - } - TaggedUnionFieldEdit::Reset => reset_value_field(value, field, default), - TaggedUnionFieldEdit::Unchanged => {} - } - return Ok(()); - } - - let current = value_field_value(value, field.name); - let actions = [ - MenuItem::new("Set value"), - MenuItem::new(shortcut_label( - "Reset to default/none", - "r, Backspace, Delete", - )), - MenuItem::new(shortcut_label("Back", "q")), - ]; - let action = prompt_menu( - theme, - &format!( - "{prompt}.{}, current {}", - field.name, - current - .as_ref() - .map(|value| { - display_value_with_default(value, value_field_default(default, field)) - }) - .or_else(|| { - value_field_default(default, field) - .map(|value| format!("{} (default)", display_value(&value))) - }) - .unwrap_or_else(|| "(default)".to_string()) - ), - &actions, - 0, - )?; - match action { - MenuResponse::Selected(0) => { - let field_value = prompt_value(theme, &field, current.as_ref())?; - set_value_field(value, field.name, field_value); - } - MenuResponse::Selected(1) - | MenuResponse::Shortcut(MenuShortcut::Reset | MenuShortcut::Clear, _) => { - reset_value_field(value, field, default) - } - MenuResponse::Shortcut(MenuShortcut::Help, _) => print_editor_help(), - MenuResponse::Shortcut(MenuShortcut::Preview | MenuShortcut::Save, _) => { - println!(" Preview and save are available from the main plugins.toml menu."); - } - _ => {} - } - Ok(()) -} - fn reset_value_section_item( value: &mut Value, schema: &nemo_relay::config_editor::EditorSchema, @@ -1896,98 +739,6 @@ trait SerializeConfig: serde::Serialize + serde::de::DeserializeOwned {} impl SerializeConfig for T where T: serde::Serialize + serde::de::DeserializeOwned {} -fn prompt_value( - theme: &ColorfulTheme, - field: &EditorFieldSpec, - current: Option<&Value>, -) -> Result { - match field.kind { - EditorFieldKind::Boolean => { - let values = ["false", "true"]; - let default_idx = current - .and_then(Value::as_bool) - .map(usize::from) - .unwrap_or(0); - let idx = Select::with_theme(theme) - .with_prompt(field.label) - .items(&values) - .default(default_idx) - .interact() - .map_err(editor_error)?; - Ok(json!(idx == 1)) - } - EditorFieldKind::Integer => { - let initial = current.map(display_value).unwrap_or_default(); - let value: String = Input::with_theme(theme) - .with_prompt(field.label) - .with_initial_text(initial) - .interact_text() - .map_err(editor_error)?; - let parsed = value.trim().parse::().map_err(|error| { - CliError::Config(format!("{} must be an integer: {error}", field.name)) - })?; - Ok(json!(parsed)) - } - EditorFieldKind::Float => { - let initial = current.map(display_value).unwrap_or_default(); - let value: String = Input::with_theme(theme) - .with_prompt(field.label) - .with_initial_text(initial) - .interact_text() - .map_err(editor_error)?; - parse_float_value(field, &value) - } - EditorFieldKind::StringMap | EditorFieldKind::Json => { - let initial = current.map(display_value).unwrap_or_else(|| { - if matches!(field.name, "tool_definitions" | "learners") { - "[]".to_string() - } else { - "{}".to_string() - } - }); - let value: String = Input::with_theme(theme) - .with_prompt(format!("{} as JSON", field.label)) - .with_initial_text(initial) - .interact_text() - .map_err(editor_error)?; - serde_json::from_str(value.trim()).map_err(|error| { - CliError::Config(format!("invalid JSON for {}: {error}", field.name)) - }) - } - EditorFieldKind::Enum => { - let values = field.enum_values; - let default_idx = current - .and_then(Value::as_str) - .and_then(|value| values.iter().position(|candidate| *candidate == value)) - .unwrap_or(0); - let idx = Select::with_theme(theme) - .with_prompt(field.label) - .items(values) - .default(default_idx) - .interact() - .map_err(editor_error)?; - Ok(json!(values[idx])) - } - EditorFieldKind::String => { - let initial = current.and_then(Value::as_str).unwrap_or_default(); - let value: String = Input::with_theme(theme) - .with_prompt(field.label) - .with_initial_text(initial) - .interact_text() - .map_err(editor_error)?; - Ok(json!(value)) - } - EditorFieldKind::Section => Err(CliError::Config(format!( - "{} is a nested section and cannot be edited as a scalar", - field.name - ))), - EditorFieldKind::List | EditorFieldKind::TaggedUnion => Err(CliError::Config(format!( - "{} is a structured value and cannot be edited as a scalar", - field.name - ))), - } -} - fn parse_float_value(field: &EditorFieldSpec, value: &str) -> Result { let value = value.trim(); let parsed = value @@ -2002,20 +753,6 @@ fn parse_float_value(field: &EditorFieldSpec, value: &str) -> Result CliError { - match err { - dialoguer::Error::IO(io_err) - if matches!( - io_err.kind(), - std::io::ErrorKind::Interrupted | std::io::ErrorKind::UnexpectedEof - ) => - { - CliError::Config(PLUGIN_EDIT_CANCELLED_MESSAGE.into()) - } - other => CliError::Config(format!("plugin edit error: {other}")), - } -} - #[cfg(test)] #[path = "../../tests/coverage/shared/plugins_tests.rs"] mod tests; diff --git a/crates/cli/src/plugins/prompt.rs b/crates/cli/src/plugins/prompt.rs new file mode 100644 index 000000000..be1c4e64c --- /dev/null +++ b/crates/cli/src/plugins/prompt.rs @@ -0,0 +1,1266 @@ +// SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +//! Terminal-only prompt adapter for plugin configuration. + +use std::io::IsTerminal; + +use console::Term; +use dialoguer::theme::ColorfulTheme; +use dialoguer::{Input, Select}; + +use super::*; + +pub(crate) fn edit(command: PluginsEditRequest) -> Result<(), CliError> { + ensure_tty()?; + let (scope, path) = resolve_edit_target(command)?; + let mut document = PluginConfigDocument::read(&path)?; + ensure_observability_component(document.config_mut())?; + ensure_adaptive_component(document.config_mut())?; + let mut components = editable_components(document.config())?; + let mut dynamic_plugins = load_dynamic_plugin_states(&document)?; + + let theme = ColorfulTheme::default(); + crate::banner::print_intro(); + println!( + " Editing plugin config at {}", + single_line_text(&path.display().to_string()) + ); + println!(" Tip: ↑/↓ or j/k to move, PageUp/PageDown to scroll, SPACE/ENTER to select."); + println!(); + let mut selected_index = 0; + loop { + let dynamic_rows = dynamic_plugins + .iter() + .map(|plugin| (plugin.label().to_owned(), plugin.menu_summary())) + .collect::>(); + let (items, actions) = plugin_menu_items(&components, &dynamic_rows, &path); + let selection = prompt_menu(&theme, "plugins.toml", &items, selected_index)?; + if let Some(selected) = menu_response_index(&selection) { + selected_index = selected; + } + if handle_menu_response( + &theme, + &mut document, + &mut components, + &mut dynamic_plugins, + &actions, + selection, + scope, + )? == EditLoopControl::Finish + { + return Ok(()); + } + } +} + +fn handle_menu_response( + theme: &ColorfulTheme, + document: &mut PluginConfigDocument, + components: &mut [EditableComponent], + dynamic_plugins: &mut [DynamicPluginEditorState], + actions: &[MenuAction], + selection: MenuResponse, + scope: TargetScope, +) -> Result { + match selection { + MenuResponse::Selected(selection) => handle_menu_action( + theme, + document, + components, + dynamic_plugins, + actions.get(selection).copied(), + scope, + ), + MenuResponse::Shortcut(MenuShortcut::Preview, _) => { + preview_document(document, components, dynamic_plugins)?; + Ok(EditLoopControl::Continue) + } + MenuResponse::Shortcut(MenuShortcut::Save, _) => { + save_document(document, components, dynamic_plugins, scope) + } + MenuResponse::Shortcut(MenuShortcut::Help, _) => { + print_editor_help(); + Ok(EditLoopControl::Continue) + } + MenuResponse::Shortcut( + shortcut @ (MenuShortcut::Reset | MenuShortcut::Clear), + selected, + ) => handle_reset_or_clear_shortcut(components, actions.get(selected).copied(), shortcut), + MenuResponse::Cancel => Err(cancelled_error()), + } +} + +fn handle_menu_action( + theme: &ColorfulTheme, + document: &mut PluginConfigDocument, + components: &mut [EditableComponent], + dynamic_plugins: &mut [DynamicPluginEditorState], + action: Option, + scope: TargetScope, +) -> Result { + match action { + Some(MenuAction::EditComponent(component_index)) => { + if let Some(component) = components.get_mut(component_index) { + edit_component(theme, component)?; + } + Ok(EditLoopControl::Continue) + } + Some(MenuAction::EditDynamic(dynamic_index)) => { + if let Some(plugin) = dynamic_plugins.get_mut(dynamic_index) { + edit_dynamic_plugin(theme, plugin)?; + } + Ok(EditLoopControl::Continue) + } + Some(MenuAction::Preview) => { + preview_document(document, components, dynamic_plugins)?; + Ok(EditLoopControl::Continue) + } + Some(MenuAction::Save) => save_document(document, components, dynamic_plugins, scope), + Some(MenuAction::Cancel) | None => Err(cancelled_error()), + } +} + +fn edit_component( + theme: &ColorfulTheme, + component: &mut EditableComponent, +) -> Result<(), CliError> { + let mut selected_index = 0; + loop { + let (items, actions) = component_menu_items(component); + let selection = prompt_menu(theme, component.label(), &items, selected_index)?; + if let Some(selected) = menu_response_index(&selection) { + selected_index = selected; + } + match selection { + MenuResponse::Selected(selected) => match actions.get(selected).copied() { + Some(ComponentMenuAction::Toggle) => component.toggle_enabled(), + Some(ComponentMenuAction::EditField(field_index)) => { + if let Some(field) = component.fields().get(field_index) { + edit_component_field(theme, component, *field)?; + } + } + Some(ComponentMenuAction::Back) | None => return Ok(()), + }, + MenuResponse::Shortcut(MenuShortcut::Reset, selected) => { + reset_component_menu_item(component, actions.get(selected).copied())?; + } + MenuResponse::Shortcut(MenuShortcut::Clear, selected) => { + clear_component_menu_item(component, actions.get(selected).copied())?; + } + MenuResponse::Shortcut(MenuShortcut::Help, _) => print_editor_help(), + MenuResponse::Shortcut(MenuShortcut::Preview | MenuShortcut::Save, _) => { + println!(" Preview and save are available from the main plugins.toml menu."); + } + MenuResponse::Cancel => return Ok(()), + } + } +} + +fn edit_component_field( + theme: &ColorfulTheme, + component: &mut EditableComponent, + field: EditorFieldSpec, +) -> Result<(), CliError> { + match component { + EditableComponent::Observability(state) => { + edit_section(theme, &mut state.config, field)?; + state.mark_config_touched(); + } + EditableComponent::Adaptive(state) => { + edit_config_field(theme, &mut state.config, field)?; + state.mark_config_touched(); + } + EditableComponent::NemoGuardrails(state) => { + edit_config_field(theme, &mut state.config, field)?; + state.mark_config_touched(); + } + EditableComponent::PiiRedaction(state) => { + edit_config_field(theme, &mut state.config, field)?; + state.mark_config_touched(); + } + #[cfg(feature = "switchyard")] + EditableComponent::Switchyard(state) => { + edit_config_field(theme, &mut state.config, field)?; + state.mark_config_touched(); + } + } + Ok(()) +} + +pub(super) fn prompt_menu( + theme: &ColorfulTheme, + prompt: &str, + items: &[MenuItem], + default: usize, +) -> Result { + if items.is_empty() { + return Err(CliError::Config(format!("{prompt} menu has no items"))); + } + let term = Term::stderr(); + let mut selected = default.min(items.len() - 1); + let mut rendered_lines = 0; + loop { + if rendered_lines > 0 { + term.clear_last_lines(rendered_lines).map_err(menu_error)?; + } + let (rows, columns) = term.size(); + let viewport = menu_viewport(items.len(), selected, usize::from(rows)); + let lines = render_menu_for_size( + theme, + prompt, + items, + selected, + usize::from(rows), + usize::from(columns), + ); + rendered_lines = lines.len(); + for line in &lines { + term.write_line(line).map_err(menu_error)?; + } + term.flush().map_err(menu_error)?; + let key = term.read_key().map_err(menu_error)?; + if let Some(next) = + menu_selection_after_key(&key, selected, items.len(), viewport.page_size) + { + selected = next; + continue; + } + if let Some(response) = menu_response_for_key(&key, selected) { + clear_menu(&term, rendered_lines)?; + return Ok(response); + } + } +} + +fn clear_menu(term: &Term, rendered_lines: usize) -> Result<(), CliError> { + if rendered_lines > 0 { + term.clear_last_lines(rendered_lines).map_err(menu_error)?; + } + Ok(()) +} + +pub(super) fn menu_error(error: std::io::Error) -> CliError { + if matches!( + error.kind(), + std::io::ErrorKind::Interrupted | std::io::ErrorKind::UnexpectedEof + ) { + CliError::Config(PLUGIN_EDIT_CANCELLED_MESSAGE.into()) + } else { + CliError::Config(format!("plugin editor terminal error: {error}")) + } +} + +pub(super) fn print_editor_help() { + println!(); + println!( + "{} {}", + style("?").yellow(), + style("Plugin editor keys").bold() + ); + println!(" {} move", style("↑/↓ or j/k").cyan()); + println!( + " {} move by page or jump to an end", + style("PageUp/PageDown, Home/End").cyan() + ); + println!( + " {} select/toggle the highlighted item", + style("Enter/Space").cyan() + ); + println!( + " {} reset the highlighted field or section", + style("r").cyan() + ); + println!( + " {} clear the highlighted optional field", + style("Backspace/Del").cyan() + ); + println!( + " {} preview TOML from the main menu", + style("p").cyan() + ); + println!( + " {} save from the main menu", + style("s").cyan() + ); + println!(" {} go back/cancel", style("q or Esc").cyan()); +} + +fn ensure_tty() -> Result<(), CliError> { + if !std::io::stdin().is_terminal() + || !std::io::stdout().is_terminal() + || !std::io::stderr().is_terminal() + { + return Err(CliError::Config( + "interactive plugin editing requires a TTY".into(), + )); + } + Ok(()) +} + +fn edit_section( + theme: &ColorfulTheme, + config: &mut T, + section: EditorFieldSpec, +) -> Result<(), CliError> +where + T: SerializeConfig, +{ + let fields = section + .schema() + .ok_or_else(|| CliError::Config(format!("{} is not an editable section", section.name)))? + .fields; + let mut selected_index = 0; + loop { + let items = section_menu_items(config, section, fields)?; + let selection = prompt_menu(theme, section.name, &items, selected_index)?; + if let Some(selected) = menu_response_index(&selection) { + selected_index = selected; + } + let selection = match selection { + MenuResponse::Selected(selection) => selection, + MenuResponse::Shortcut(MenuShortcut::Help, _) => { + print_editor_help(); + continue; + } + MenuResponse::Shortcut(MenuShortcut::Reset, selected) => { + reset_selected_item(config, section, fields, selected)?; + continue; + } + MenuResponse::Shortcut(MenuShortcut::Clear, selected) => { + if reset_selected_field(config, section, fields, selected)? { + continue; + } + println!(" Select a field to clear."); + continue; + } + MenuResponse::Shortcut(MenuShortcut::Preview | MenuShortcut::Save, _) => { + println!(" Preview and save are available from the main plugins.toml menu."); + continue; + } + MenuResponse::Cancel => return Ok(()), + }; + if !edit_selected_section_item(theme, config, section, fields, selection)? { + return Ok(()); + } + } +} + +fn edit_selected_section_item( + theme: &ColorfulTheme, + config: &mut T, + section: EditorFieldSpec, + fields: &[EditorFieldSpec], + selection: usize, +) -> Result +where + T: SerializeConfig, +{ + if section_has_enabled_toggle(section) && selection == 0 { + toggle_section(config, section); + return Ok(true); + } + let index = selected_field_index(section, selection); + if let Some(field) = fields.get(index) { + edit_field(theme, config, section, field)?; + return Ok(true); + } + if index == fields.len() { + reset_section(config, section); + return Ok(true); + } + Ok(false) +} + +fn edit_field( + theme: &ColorfulTheme, + config: &mut T, + section: EditorFieldSpec, + field: &EditorFieldSpec, +) -> Result<(), CliError> +where + T: SerializeConfig, +{ + if field.kind == EditorFieldKind::Section { + edit_nested_section(theme, config, section, *field)?; + return Ok(()); + } + let current = section_field_value(config, section, field.name)?; + if field.kind == EditorFieldKind::List { + let item = field.list_item.ok_or_else(|| { + CliError::Config(format!("{} does not describe its list entries", field.name)) + })?; + let default = section_field_default(section, *field); + let mut items = current + .or_else(|| default.clone()) + .unwrap_or_else(|| json!([])); + if edit_list_value( + theme, + &format!("{}.{}", section.name, field.name), + &mut items, + default, + item, + )? { + set_section_field(config, section, field.name, items)?; + } + return Ok(()); + } + if field.kind == EditorFieldKind::StringMap { + let default = section_field_default(section, *field); + let mut entries = current + .or_else(|| default.clone()) + .unwrap_or_else(|| json!({})); + if edit_string_map_value( + theme, + &format!("{}.{}", section.name, field.name), + &mut entries, + default, + )? { + set_section_field(config, section, field.name, entries)?; + } + return Ok(()); + } + if field.kind == EditorFieldKind::TaggedUnion { + let tagged_union = field.tagged_union.ok_or_else(|| { + CliError::Config(format!("{} does not describe its variants", field.name)) + })?; + let default = section_field_default(section, *field); + match edit_tagged_union_field( + theme, + &format!("{}.{}", section.name, field.name), + current, + default, + tagged_union, + )? { + TaggedUnionFieldEdit::Set(value) => { + set_section_field(config, section, field.name, value)?; + } + TaggedUnionFieldEdit::Reset => remove_section_field(config, section, field.name)?, + TaggedUnionFieldEdit::Unchanged => {} + } + return Ok(()); + } + let actions = [ + MenuItem::new("Set value"), + MenuItem::new(shortcut_label( + "Reset to default/none", + "r, Backspace, Delete", + )), + MenuItem::new(shortcut_label("Back", "q")), + ]; + let action = prompt_menu( + theme, + &format!( + "{}.{}, current {}", + section.name, + field.name, + current + .as_ref() + .map(|value| display_field_value(section, *field, value)) + .unwrap_or_else(|| "(default)".to_string()) + ), + &actions, + 0, + )?; + match action { + MenuResponse::Selected(0) => { + let value = prompt_value(theme, field, current.as_ref())?; + set_section_field(config, section, field.name, value)?; + } + MenuResponse::Selected(1) + | MenuResponse::Shortcut(MenuShortcut::Reset | MenuShortcut::Clear, _) => { + remove_section_field(config, section, field.name)? + } + MenuResponse::Shortcut(MenuShortcut::Help, _) => print_editor_help(), + MenuResponse::Shortcut(MenuShortcut::Preview | MenuShortcut::Save, _) => { + println!(" Preview and save are available from the main plugins.toml menu."); + } + _ => {} + } + Ok(()) +} + +fn edit_list_value( + theme: &ColorfulTheme, + prompt: &str, + value: &mut Value, + default: Option, + item: &nemo_relay::config_editor::EditorListItemSpec, +) -> Result { + if !value.is_array() { + *value = default.clone().unwrap_or_else(|| json!([])); + } + let original = value.clone(); + let mut selected_index = 0; + loop { + let entries = value.as_array().expect("list value is an array"); + let mut menu = vec![MenuItem::new("Add item")]; + menu.extend(entries.iter().enumerate().map(|(index, entry)| { + MenuItem::new(format!( + "Edit item {}: {}", + index + 1, + editor_item_label(entry, item) + )) + })); + menu.push(MenuItem::new(shortcut_label("Back", "q"))); + let selection = prompt_menu(theme, prompt, &menu, selected_index)?; + if let Some(selected) = menu_response_index(&selection) { + selected_index = selected; + } + match selection { + MenuResponse::Selected(0) => { + let mut entry = new_editor_item(theme, item)?; + edit_editor_item( + theme, + &format!("{prompt}[{}]", entries.len()), + &mut entry, + item, + )?; + value + .as_array_mut() + .expect("list value is an array") + .push(entry); + } + MenuResponse::Selected(index) if index <= entries.len() => { + edit_existing_list_item(theme, prompt, value, index - 1, item)?; + } + MenuResponse::Cancel | MenuResponse::Selected(_) => return Ok(*value != original), + MenuResponse::Shortcut(MenuShortcut::Help, _) => print_editor_help(), + MenuResponse::Shortcut(shortcut @ (MenuShortcut::Reset | MenuShortcut::Clear), _) => { + *value = collection_shortcut_value(default.as_ref(), json!([]), shortcut) + } + MenuResponse::Shortcut(MenuShortcut::Preview | MenuShortcut::Save, _) => { + println!(" Preview and save are available from the main plugins.toml menu."); + } + } + } +} + +fn edit_string_map_value( + theme: &ColorfulTheme, + prompt: &str, + value: &mut Value, + default: Option, +) -> Result { + if !value.is_object() { + *value = default.clone().unwrap_or_else(|| json!({})); + } + let original = value.clone(); + let mut selected_index = 0; + loop { + let entries = value.as_object().expect("string map value is an object"); + let keys = entries.keys().cloned().collect::>(); + let mut menu = vec![MenuItem::new("Add entry")]; + menu.extend(keys.iter().map(|key| { + MenuItem::new(format!( + "Edit {key}: {}", + entries.get(key).map(display_value).unwrap_or_default() + )) + })); + menu.push(MenuItem::new(shortcut_label("Back", "q"))); + let selection = prompt_menu(theme, prompt, &menu, selected_index)?; + if let Some(selected) = menu_response_index(&selection) { + selected_index = selected; + } + match selection { + MenuResponse::Selected(0) => { + let key: String = Input::with_theme(theme) + .with_prompt("Entry key") + .interact_text() + .map_err(editor_error)?; + if key.trim().is_empty() { + println!(" Entry key must not be empty."); + continue; + } + let key = key.trim().to_owned(); + if string_map_entry_exists(value, &key) { + println!(" Entry already exists; select it to edit."); + continue; + } + let entry: String = Input::with_theme(theme) + .with_prompt("Entry value") + .interact_text() + .map_err(editor_error)?; + value + .as_object_mut() + .expect("string map value is an object") + .insert(key, Value::String(entry)); + } + MenuResponse::Selected(index) if index <= keys.len() => { + edit_existing_string_map_entry(theme, prompt, value, &keys[index - 1])?; + } + MenuResponse::Cancel | MenuResponse::Selected(_) => return Ok(*value != original), + MenuResponse::Shortcut(MenuShortcut::Help, _) => print_editor_help(), + MenuResponse::Shortcut(shortcut @ (MenuShortcut::Reset | MenuShortcut::Clear), _) => { + *value = collection_shortcut_value(default.as_ref(), json!({}), shortcut) + } + MenuResponse::Shortcut(MenuShortcut::Preview | MenuShortcut::Save, _) => { + println!(" Preview and save are available from the main plugins.toml menu."); + } + } + } +} + +fn edit_existing_string_map_entry( + theme: &ColorfulTheme, + prompt: &str, + value: &mut Value, + key: &str, +) -> Result<(), CliError> { + let actions = [ + MenuItem::new("Edit value"), + MenuItem::new("Remove entry"), + MenuItem::new(shortcut_label("Back", "q")), + ]; + match prompt_menu(theme, &format!("{prompt}.{key}"), &actions, 0)? { + MenuResponse::Selected(0) => { + let current = value + .as_object() + .and_then(|entries| entries.get(key)) + .and_then(Value::as_str) + .unwrap_or_default(); + let entry: String = Input::with_theme(theme) + .with_prompt("Entry value") + .with_initial_text(current) + .interact_text() + .map_err(editor_error)?; + value + .as_object_mut() + .expect("string map value is an object") + .insert(key.to_owned(), Value::String(entry)); + } + MenuResponse::Selected(1) | MenuResponse::Shortcut(MenuShortcut::Clear, _) => { + value + .as_object_mut() + .expect("string map value is an object") + .remove(key); + } + _ => {} + } + Ok(()) +} + +fn edit_existing_list_item( + theme: &ColorfulTheme, + prompt: &str, + value: &mut Value, + index: usize, + item: &nemo_relay::config_editor::EditorListItemSpec, +) -> Result<(), CliError> { + let actions = [ + MenuItem::new("Edit item"), + MenuItem::new("Remove item"), + MenuItem::new(shortcut_label("Back", "q")), + ]; + match prompt_menu(theme, &format!("{prompt}[{}]", index + 1), &actions, 0)? { + MenuResponse::Selected(0) => { + if let Some(entry) = value + .as_array_mut() + .and_then(|entries| entries.get_mut(index)) + { + edit_editor_item(theme, &format!("{prompt}[{}]", index + 1), entry, item)?; + } + } + MenuResponse::Selected(1) | MenuResponse::Shortcut(MenuShortcut::Clear, _) => { + value + .as_array_mut() + .expect("list value is an array") + .remove(index); + } + _ => {} + } + Ok(()) +} + +fn new_editor_item( + theme: &ColorfulTheme, + item: &nemo_relay::config_editor::EditorListItemSpec, +) -> Result { + if let Some(tagged_union) = item.tagged_union { + return new_tagged_union_value(theme, tagged_union); + } + Ok(item.default.map(|default| default()).unwrap_or(Value::Null)) +} + +fn edit_editor_item( + theme: &ColorfulTheme, + prompt: &str, + value: &mut Value, + item: &nemo_relay::config_editor::EditorListItemSpec, +) -> Result<(), CliError> { + if let Some(tagged_union) = item.tagged_union { + return edit_tagged_union_payload(theme, prompt, value, tagged_union); + } + + match item.kind { + EditorFieldKind::Section => { + let schema = item + .schema + .ok_or_else(|| CliError::Config("list item has no schema".into()))?( + ); + edit_value_section(theme, prompt, value, schema, None)?; + } + EditorFieldKind::List => { + let nested = item.list_item.ok_or_else(|| { + CliError::Config("nested list item has no entry description".into()) + })?; + let _ = edit_list_value(theme, prompt, value, None, nested)?; + } + EditorFieldKind::StringMap => { + let _ = edit_string_map_value(theme, prompt, value, None)?; + } + kind => { + let field = EditorFieldSpec { + name: "item", + label: "item", + kind, + enum_values: &[], + optional: false, + nested_schema: None, + nested_default: None, + list_item: None, + tagged_union: None, + }; + *value = prompt_value(theme, &field, Some(value))?; + } + } + Ok(()) +} + +fn select_tagged_union_variant( + theme: &ColorfulTheme, + tagged_union: &nemo_relay::config_editor::EditorTaggedUnionSpec, +) -> Result { + if tagged_union.variants.is_empty() { + return Err(CliError::Config("tagged union has no variants".into())); + } + Select::with_theme(theme) + .with_prompt("Variant type") + .items( + &tagged_union + .variants + .iter() + .map(|variant| variant.label) + .collect::>(), + ) + .default(0) + .interact() + .map_err(editor_error) +} + +fn new_tagged_union_value( + theme: &ColorfulTheme, + tagged_union: &nemo_relay::config_editor::EditorTaggedUnionSpec, +) -> Result { + tagged_union_variant_value( + tagged_union, + select_tagged_union_variant(theme, tagged_union)?, + ) +} + +fn edit_tagged_union_payload( + theme: &ColorfulTheme, + prompt: &str, + value: &mut Value, + tagged_union: &nemo_relay::config_editor::EditorTaggedUnionSpec, +) -> Result<(), CliError> { + if !value.is_object() { + *value = new_tagged_union_value(theme, tagged_union)?; + } + let tag = value + .get(tagged_union.discriminator) + .and_then(Value::as_str) + .ok_or_else(|| CliError::Config("tagged union has no discriminator value".into()))?; + let variant = tagged_union + .variants + .iter() + .find(|variant| variant.tag == tag) + .ok_or_else(|| CliError::Config(format!("unknown tagged union type {tag:?}")))?; + edit_value_section(theme, prompt, value, (variant.schema)(), None)?; + Ok(()) +} + +fn edit_tagged_union_field( + theme: &ColorfulTheme, + prompt: &str, + current: Option, + default: Option, + tagged_union: &nemo_relay::config_editor::EditorTaggedUnionSpec, +) -> Result { + let mut state = TaggedUnionFieldState::new(current, default); + loop { + let actions = [ + MenuItem::new("Edit fields"), + MenuItem::new("Change variant"), + MenuItem::new(shortcut_label( + "Reset to default/none", + "r, Backspace, Delete", + )), + MenuItem::new(shortcut_label("Back", "q")), + ]; + match prompt_menu( + theme, + &format!("{prompt}, current {}", display_value(state.value())), + &actions, + 0, + )? { + MenuResponse::Selected(0) => { + edit_tagged_union_payload(theme, prompt, state.value_mut(), tagged_union)?; + } + MenuResponse::Selected(1) => { + state.change_variant( + tagged_union, + select_tagged_union_variant(theme, tagged_union)?, + )?; + edit_tagged_union_payload(theme, prompt, state.value_mut(), tagged_union)?; + } + MenuResponse::Selected(2) + | MenuResponse::Shortcut(MenuShortcut::Reset | MenuShortcut::Clear, _) => { + return Ok(state.reset()); + } + MenuResponse::Shortcut(MenuShortcut::Help, _) => print_editor_help(), + MenuResponse::Shortcut(MenuShortcut::Preview | MenuShortcut::Save, _) => { + println!(" Preview and save are available from the main plugins.toml menu."); + } + MenuResponse::Cancel | MenuResponse::Selected(_) => { + return Ok(state.finish()); + } + } + } +} + +fn edit_config_field( + theme: &ColorfulTheme, + config: &mut T, + field: EditorFieldSpec, +) -> Result<(), CliError> +where + T: Default + SerializeConfig, +{ + if field.kind == EditorFieldKind::Section { + let mut value = config_field_value(config, field.name)? + .or_else(|| field.default_value()) + .unwrap_or_else(|| json!({})); + let schema = field.schema().ok_or_else(|| { + CliError::Config(format!("{} is not an editable section", field.name)) + })?; + if edit_value_section(theme, field.name, &mut value, schema, field.default_value())? { + store_edited_config_section(config, field, value)?; + } + return Ok(()); + } + + if field.kind == EditorFieldKind::List { + let item = field.list_item.ok_or_else(|| { + CliError::Config(format!("{} does not describe its list entries", field.name)) + })?; + let default = default_config_field_value::(field).or_else(|| field.default_value()); + let mut items = config_field_value(config, field.name)? + .or_else(|| default.clone()) + .unwrap_or_else(|| json!([])); + if edit_list_value(theme, field.name, &mut items, default, item)? { + set_struct_field(config, field.name, items)?; + } + return Ok(()); + } + + if field.kind == EditorFieldKind::StringMap { + let default = default_config_field_value::(field).or_else(|| field.default_value()); + let mut entries = config_field_value(config, field.name)? + .or_else(|| default.clone()) + .unwrap_or_else(|| json!({})); + if edit_string_map_value(theme, field.name, &mut entries, default)? { + set_struct_field(config, field.name, entries)?; + } + return Ok(()); + } + + if field.kind == EditorFieldKind::TaggedUnion { + let tagged_union = field.tagged_union.ok_or_else(|| { + CliError::Config(format!("{} does not describe its variants", field.name)) + })?; + let default = default_config_field_value::(field).or_else(|| field.default_value()); + match edit_tagged_union_field( + theme, + field.name, + config_field_value(config, field.name)?, + default, + tagged_union, + )? { + TaggedUnionFieldEdit::Set(value) => set_struct_field(config, field.name, value)?, + TaggedUnionFieldEdit::Reset => reset_config_field(config, field)?, + TaggedUnionFieldEdit::Unchanged => {} + } + return Ok(()); + } + + let current = config_field_value(config, field.name)?; + let actions = [ + MenuItem::new("Set value"), + MenuItem::new(shortcut_label( + "Reset to default/none", + "r, Backspace, Delete", + )), + MenuItem::new(shortcut_label("Back", "q")), + ]; + let action = prompt_menu( + theme, + &format!( + "{}, current {}", + field.label, + current + .as_ref() + .map(display_value) + .or_else(|| default_config_field_value::(field) + .map(|value| { format!("{} (default)", display_value(&value)) })) + .unwrap_or_else(|| "(default)".to_string()) + ), + &actions, + 0, + )?; + match action { + MenuResponse::Selected(0) => { + let value = prompt_value(theme, &field, current.as_ref())?; + set_struct_field(config, field.name, value)?; + } + MenuResponse::Selected(1) + | MenuResponse::Shortcut(MenuShortcut::Reset | MenuShortcut::Clear, _) => { + reset_config_field(config, field)? + } + MenuResponse::Shortcut(MenuShortcut::Help, _) => print_editor_help(), + MenuResponse::Shortcut(MenuShortcut::Preview | MenuShortcut::Save, _) => { + println!(" Preview and save are available from the main plugins.toml menu."); + } + _ => {} + } + Ok(()) +} + +fn edit_nested_section( + theme: &ColorfulTheme, + config: &mut T, + section: EditorFieldSpec, + field: EditorFieldSpec, +) -> Result<(), CliError> +where + T: SerializeConfig, +{ + let mut value = section_field_value(config, section, field.name)? + .or_else(|| section_field_default(section, field)) + .unwrap_or_else(|| json!({})); + let schema = field + .schema() + .ok_or_else(|| CliError::Config(format!("{} is not an editable section", field.name)))?; + let default = section_field_default(section, field); + if edit_value_section( + theme, + &format!("{}.{}", section.name, field.name), + &mut value, + schema, + default, + )? { + store_edited_section_field(config, section, field, value)?; + } + Ok(()) +} + +fn edit_value_section( + theme: &ColorfulTheme, + prompt: &str, + value: &mut Value, + schema: &nemo_relay::config_editor::EditorSchema, + default: Option, +) -> Result { + ensure_object(value); + let original = value.clone(); + let mut selected_index = 0; + loop { + let items = value_section_menu_items(value, schema, default.as_ref())?; + let selection = prompt_menu(theme, prompt, &items, selected_index)?; + if let Some(selected) = menu_response_index(&selection) { + selected_index = selected; + } + let selection = match selection { + MenuResponse::Selected(selection) => selection, + MenuResponse::Shortcut(MenuShortcut::Help, _) => { + print_editor_help(); + continue; + } + MenuResponse::Shortcut(MenuShortcut::Reset, selected) => { + reset_value_section_item(value, schema, default.as_ref(), selected); + continue; + } + MenuResponse::Shortcut(MenuShortcut::Clear, selected) => { + if clear_value_field(value, schema, selected) { + continue; + } + println!(" Select a field to clear."); + continue; + } + MenuResponse::Shortcut(MenuShortcut::Preview | MenuShortcut::Save, _) => { + println!(" Preview and save are available from the main plugins.toml menu."); + continue; + } + MenuResponse::Cancel => return Ok(*value != original), + }; + if !edit_selected_value_item(theme, prompt, value, schema, default.as_ref(), selection)? { + return Ok(*value != original); + } + } +} + +fn edit_selected_value_item( + theme: &ColorfulTheme, + prompt: &str, + value: &mut Value, + schema: &nemo_relay::config_editor::EditorSchema, + default: Option<&Value>, + selection: usize, +) -> Result { + if let Some(field) = schema.fields.get(selection) { + edit_value_field(theme, prompt, value, *field, default)?; + return Ok(true); + } + if selection == schema.fields.len() { + *value = default.cloned().unwrap_or_else(|| json!({})); + ensure_object(value); + return Ok(true); + } + Ok(false) +} + +fn edit_value_field( + theme: &ColorfulTheme, + prompt: &str, + value: &mut Value, + field: EditorFieldSpec, + default: Option<&Value>, +) -> Result<(), CliError> { + if field.kind == EditorFieldKind::Section { + let nested_default = value_field_default(default, field); + let mut nested_value = value_field_value(value, field.name) + .or_else(|| nested_default.clone()) + .unwrap_or_else(|| json!({})); + let nested_schema = field.schema().ok_or_else(|| { + CliError::Config(format!("{} is not an editable section", field.name)) + })?; + if edit_value_section( + theme, + &format!("{prompt}.{}", field.name), + &mut nested_value, + nested_schema, + nested_default, + )? { + store_edited_value_section(value, field, nested_value); + } + return Ok(()); + } + + if field.kind == EditorFieldKind::List { + let item = field.list_item.ok_or_else(|| { + CliError::Config(format!("{} does not describe its list entries", field.name)) + })?; + let field_default = value_field_default(default, field); + let mut items = value_field_value(value, field.name) + .or_else(|| field_default.clone()) + .unwrap_or_else(|| json!([])); + if edit_list_value( + theme, + &format!("{prompt}.{}", field.name), + &mut items, + field_default, + item, + )? { + set_value_field(value, field.name, items); + } + return Ok(()); + } + + if field.kind == EditorFieldKind::StringMap { + let field_default = value_field_default(default, field); + let mut entries = value_field_value(value, field.name) + .or_else(|| field_default.clone()) + .unwrap_or_else(|| json!({})); + if edit_string_map_value( + theme, + &format!("{prompt}.{}", field.name), + &mut entries, + field_default, + )? { + set_value_field(value, field.name, entries); + } + return Ok(()); + } + + if field.kind == EditorFieldKind::TaggedUnion { + let tagged_union = field.tagged_union.ok_or_else(|| { + CliError::Config(format!("{} does not describe its variants", field.name)) + })?; + let field_default = value_field_default(default, field); + match edit_tagged_union_field( + theme, + &format!("{prompt}.{}", field.name), + value_field_value(value, field.name), + field_default.clone(), + tagged_union, + )? { + TaggedUnionFieldEdit::Set(tagged_value) => { + set_value_field(value, field.name, tagged_value); + } + TaggedUnionFieldEdit::Reset => reset_value_field(value, field, default), + TaggedUnionFieldEdit::Unchanged => {} + } + return Ok(()); + } + + let current = value_field_value(value, field.name); + let actions = [ + MenuItem::new("Set value"), + MenuItem::new(shortcut_label( + "Reset to default/none", + "r, Backspace, Delete", + )), + MenuItem::new(shortcut_label("Back", "q")), + ]; + let action = prompt_menu( + theme, + &format!( + "{prompt}.{}, current {}", + field.name, + current + .as_ref() + .map(|value| { + display_value_with_default(value, value_field_default(default, field)) + }) + .or_else(|| { + value_field_default(default, field) + .map(|value| format!("{} (default)", display_value(&value))) + }) + .unwrap_or_else(|| "(default)".to_string()) + ), + &actions, + 0, + )?; + match action { + MenuResponse::Selected(0) => { + let field_value = prompt_value(theme, &field, current.as_ref())?; + set_value_field(value, field.name, field_value); + } + MenuResponse::Selected(1) + | MenuResponse::Shortcut(MenuShortcut::Reset | MenuShortcut::Clear, _) => { + reset_value_field(value, field, default) + } + MenuResponse::Shortcut(MenuShortcut::Help, _) => print_editor_help(), + MenuResponse::Shortcut(MenuShortcut::Preview | MenuShortcut::Save, _) => { + println!(" Preview and save are available from the main plugins.toml menu."); + } + _ => {} + } + Ok(()) +} + +fn prompt_value( + theme: &ColorfulTheme, + field: &EditorFieldSpec, + current: Option<&Value>, +) -> Result { + match field.kind { + EditorFieldKind::Boolean => { + let values = ["false", "true"]; + let default_idx = current + .and_then(Value::as_bool) + .map(usize::from) + .unwrap_or(0); + let idx = Select::with_theme(theme) + .with_prompt(field.label) + .items(&values) + .default(default_idx) + .interact() + .map_err(editor_error)?; + Ok(json!(idx == 1)) + } + EditorFieldKind::Integer => { + let initial = current.map(display_value).unwrap_or_default(); + let value: String = Input::with_theme(theme) + .with_prompt(field.label) + .with_initial_text(initial) + .interact_text() + .map_err(editor_error)?; + let parsed = value.trim().parse::().map_err(|error| { + CliError::Config(format!("{} must be an integer: {error}", field.name)) + })?; + Ok(json!(parsed)) + } + EditorFieldKind::Float => { + let initial = current.map(display_value).unwrap_or_default(); + let value: String = Input::with_theme(theme) + .with_prompt(field.label) + .with_initial_text(initial) + .interact_text() + .map_err(editor_error)?; + parse_float_value(field, &value) + } + EditorFieldKind::StringMap | EditorFieldKind::Json => { + let initial = current.map(display_value).unwrap_or_else(|| { + if matches!(field.name, "tool_definitions" | "learners") { + "[]".to_string() + } else { + "{}".to_string() + } + }); + let value: String = Input::with_theme(theme) + .with_prompt(format!("{} as JSON", field.label)) + .with_initial_text(initial) + .interact_text() + .map_err(editor_error)?; + serde_json::from_str(value.trim()).map_err(|error| { + CliError::Config(format!("invalid JSON for {}: {error}", field.name)) + }) + } + EditorFieldKind::Enum => { + let values = field.enum_values; + let default_idx = current + .and_then(Value::as_str) + .and_then(|value| values.iter().position(|candidate| *candidate == value)) + .unwrap_or(0); + let idx = Select::with_theme(theme) + .with_prompt(field.label) + .items(values) + .default(default_idx) + .interact() + .map_err(editor_error)?; + Ok(json!(values[idx])) + } + EditorFieldKind::String => { + let initial = current.and_then(Value::as_str).unwrap_or_default(); + let value: String = Input::with_theme(theme) + .with_prompt(field.label) + .with_initial_text(initial) + .interact_text() + .map_err(editor_error)?; + Ok(json!(value)) + } + EditorFieldKind::Section => Err(CliError::Config(format!( + "{} is a nested section and cannot be edited as a scalar", + field.name + ))), + EditorFieldKind::List | EditorFieldKind::TaggedUnion => Err(CliError::Config(format!( + "{} is a structured value and cannot be edited as a scalar", + field.name + ))), + } +} + +pub(super) fn editor_error(err: dialoguer::Error) -> CliError { + match err { + dialoguer::Error::IO(io_err) + if matches!( + io_err.kind(), + std::io::ErrorKind::Interrupted | std::io::ErrorKind::UnexpectedEof + ) => + { + CliError::Config(PLUGIN_EDIT_CANCELLED_MESSAGE.into()) + } + other => CliError::Config(format!("plugin edit error: {other}")), + } +} diff --git a/crates/cli/src/process/launcher.rs b/crates/cli/src/process/launcher.rs index 81ed72569..e716d9da3 100644 --- a/crates/cli/src/process/launcher.rs +++ b/crates/cli/src/process/launcher.rs @@ -706,6 +706,13 @@ pub(crate) fn exporter_destinations(config: &GatewayConfig) -> Vec { fn observability_exporter_destinations(config: &ObservabilityConfig) -> Vec { let mut destinations = Vec::new(); + append_atof_destinations(&mut destinations, config); + append_atif_destinations(&mut destinations, config); + append_opentelemetry_destinations(&mut destinations, config); + destinations +} + +fn append_atof_destinations(destinations: &mut Vec, config: &ObservabilityConfig) { if let Some(section) = config.atof.as_ref().filter(|section| section.enabled) { for sink in §ion.sinks { match sink { @@ -727,6 +734,9 @@ fn observability_exporter_destinations(config: &ObservabilityConfig) -> Vec, config: &ObservabilityConfig) { if let Some(section) = config.atif.as_ref().filter(|section| section.enabled) { if section.storage.is_empty() { let directory = section @@ -746,6 +756,9 @@ fn observability_exporter_destinations(config: &ObservabilityConfig) -> Vec, config: &ObservabilityConfig) { if let Some(section) = config .opentelemetry .as_ref() @@ -763,7 +776,6 @@ fn observability_exporter_destinations(config: &ObservabilityConfig) -> Vec, + flush_result: &Result<(), CliError>, + clear_result: &Result<(), CliError>, + instance_id: &str, +) { + for (component, result) in [ + ("sessions", close_result), + ("subscribers", flush_result), + ("plugins", clear_result), + ] { + let Err(error) = result else { + continue; + }; + log::error!( + target: "nemo_relay.server", + event = "server_teardown_failed", + instance_id, + component, + error_kind = error.log_kind(); + "Gateway server teardown failed" + ); + } +} + async fn shutdown_signal() { #[cfg(unix)] { diff --git a/crates/cli/tests/cli_tests.rs b/crates/cli/tests/cli_tests.rs index b74e2e8dc..6ae8c4421 100644 --- a/crates/cli/tests/cli_tests.rs +++ b/crates/cli/tests/cli_tests.rs @@ -723,29 +723,7 @@ fn cli_internal_hermes_install_writes_mcp_hooks_trust_and_doctor_ready_state() { ); let config_path = hermes_home.join("config.yaml"); - let config: serde_json::Value = - serde_yaml::from_str(&std::fs::read_to_string(&config_path).unwrap()).unwrap(); - let server = &config["mcp_servers"]["nemo-relay"]; - assert_eq!(server["command"], gateway_bin()); - assert_eq!(server["args"], serde_json::json!(["mcp"])); - assert_eq!(server["env"]["NEMO_RELAY_GATEWAY_BIND"], "127.0.0.1:47632"); - assert_eq!(server["env"]["OPENAI_API_KEY"], "${OPENAI_API_KEY}"); - assert!( - !std::fs::read_to_string(&config_path) - .unwrap() - .contains("not-written-to-config") - ); - let command = config["hooks"]["on_session_start"][0]["command"] - .as_str() - .unwrap(); - assert!(command.contains("hook-forward hermes")); - let approvals: serde_json::Value = serde_json::from_str( - &std::fs::read_to_string(hermes_home.join("shell-hooks-allowlist.json")).unwrap(), - ) - .unwrap(); - let approvals = approvals["approvals"].as_array().unwrap(); - assert_eq!(approvals.len(), 13); - assert!(approvals.iter().all(|entry| entry["command"] == command)); + assert_hermes_install_config(&config_path, &hermes_home); let relay_config_dir = xdg.join("nemo-relay"); std::fs::create_dir_all(&relay_config_dir).unwrap(); @@ -774,14 +752,7 @@ fn cli_internal_hermes_install_writes_mcp_hooks_trust_and_doctor_ready_state() { String::from_utf8_lossy(&doctor.stderr) ); let report: serde_json::Value = serde_json::from_slice(&doctor.stdout).unwrap(); - assert_eq!(report["agents"][0]["name"], "hermes"); - assert_eq!(report["agents"][0]["status"], "pass"); - assert!( - report["agents"][0]["annotation"] - .as_str() - .unwrap() - .contains("MCP lifecycle") - ); + assert_hermes_doctor_report(&report); let uninstall = Command::new(gateway_bin()) .args(["uninstall", "hermes"]) @@ -801,6 +772,45 @@ fn cli_internal_hermes_install_writes_mcp_hooks_trust_and_doctor_ready_state() { assert!(!hermes_home.join(".nemo-relay-generation").exists()); } +#[cfg(unix)] +fn assert_hermes_install_config(config_path: &std::path::Path, hermes_home: &std::path::Path) { + let config: serde_json::Value = + serde_yaml::from_str(&std::fs::read_to_string(config_path).unwrap()).unwrap(); + let server = &config["mcp_servers"]["nemo-relay"]; + assert_eq!(server["command"], gateway_bin()); + assert_eq!(server["args"], serde_json::json!(["mcp"])); + assert_eq!(server["env"]["NEMO_RELAY_GATEWAY_BIND"], "127.0.0.1:47632"); + assert_eq!(server["env"]["OPENAI_API_KEY"], "${OPENAI_API_KEY}"); + assert!( + !std::fs::read_to_string(config_path) + .unwrap() + .contains("not-written-to-config") + ); + let command = config["hooks"]["on_session_start"][0]["command"] + .as_str() + .unwrap(); + assert!(command.contains("hook-forward hermes")); + let approvals: serde_json::Value = serde_json::from_str( + &std::fs::read_to_string(hermes_home.join("shell-hooks-allowlist.json")).unwrap(), + ) + .unwrap(); + let approvals = approvals["approvals"].as_array().unwrap(); + assert_eq!(approvals.len(), 13); + assert!(approvals.iter().all(|entry| entry["command"] == command)); +} + +#[cfg(unix)] +fn assert_hermes_doctor_report(report: &serde_json::Value) { + assert_eq!(report["agents"][0]["name"], "hermes"); + assert_eq!(report["agents"][0]["status"], "pass"); + assert!( + report["agents"][0]["annotation"] + .as_str() + .unwrap() + .contains("MCP lifecycle") + ); +} + fn start_mcp_client(temp: &std::path::Path, bind: SocketAddr) -> (Child, ChildStdin) { start_mcp_client_with_idle_timeout(temp, bind, "1") } @@ -1971,6 +1981,7 @@ fn cli_plugins_list_json_emits_empty_versioned_success_output() { let config_path = temp.path().join("config.toml"); std::fs::write(&config_path, "").unwrap(); let output = Command::new(gateway_bin()) + .current_dir(temp.path()) .env("XDG_CONFIG_HOME", temp.path().join("xdg")) .env("HOME", temp.path()) .args([ @@ -2088,6 +2099,7 @@ allowed = false .unwrap(); let output = Command::new(gateway_bin()) + .current_dir(temp.path()) .env("XDG_CONFIG_HOME", &xdg) .env("HOME", temp.path()) .args(["plugins", "validate"]) @@ -3513,6 +3525,7 @@ command = "hermes --yolo chat" .unwrap(); let output = Command::new(gateway_bin()) + .current_dir(temp.path()) .env("XDG_CONFIG_HOME", &xdg) .env("HOME", temp.path()) .args([ @@ -3891,6 +3904,7 @@ finally: .unwrap(); let output = Command::new("python3") + .current_dir(temp.path()) .arg(&driver) .arg(gateway_bin()) .arg(&config) @@ -3986,6 +4000,7 @@ fn assert_non_tty_signal_forwarding( ) { let pids = root.join(format!("agent-pids-{signal_name}")); let mut relay = Command::new(gateway_bin()) + .current_dir(root) .args([ "--config", config.to_str().unwrap(), diff --git a/crates/cli/tests/coverage/commands/configure_editor_tests.rs b/crates/cli/tests/coverage/commands/configure_editor_tests.rs index fa0648ce6..ce23eac42 100644 --- a/crates/cli/tests/coverage/commands/configure_editor_tests.rs +++ b/crates/cli/tests/coverage/commands/configure_editor_tests.rs @@ -300,4 +300,80 @@ fn noninteractive_editor_guard_is_deterministic() { error, "configuration error: interactive configuration editing requires a TTY" ); + assert!(ensure_tty_with(true).is_ok()); +} + +#[test] +fn document_accessors_cover_inline_values_defaults_and_invalid_shapes() { + let mut document = document( + "gateway = { max_hook_payload_bytes = 64, max_passthrough_body_bytes = -1 }\nupstream = { openai_base_url = \"https://example.test\", openai_auth_header = 7 }\nlogging = { level = 9 }\n", + ); + + assert_eq!(document.path(), Path::new("config.toml")); + assert_eq!(document.gateway_summary(), "configured"); + assert_eq!(document.upstream_summary(), "configured"); + assert_eq!(document.logging_summary(), "configured"); + assert_eq!( + document.integer_summary("gateway", "max_hook_payload_bytes"), + "64" + ); + assert_eq!( + document.integer_summary("gateway", "max_passthrough_body_bytes"), + "invalid" + ); + assert_eq!( + document.string_summary("upstream", "openai_base_url"), + "https://example.test" + ); + assert_eq!(document.string_summary("logging", "level"), "invalid"); + assert_eq!(document.string_summary("logging", "missing"), "unset"); + assert_eq!(document.secret_summary("anthropic_auth_header"), "unset"); + + document + .set_string("upstream", "openai_base_url", "https://changed.test".into()) + .unwrap(); + assert_eq!( + document.string("upstream", "openai_base_url").as_deref(), + Some("https://changed.test") + ); + document.clear_key("missing", "value").unwrap(); +} + +#[test] +fn sink_accessors_report_invalid_and_incomplete_entries() { + let mut document = document( + "[[logging.sinks]]\npath = 7\nlevel = \"debug\"\nqueue_capacity = -1\nmax_file_size_bytes = 1024\n", + ); + + assert_eq!(document.sink_labels(), ["sink 1 (invalid path)"]); + assert_eq!(document.sink_string_summary(0, "path"), "invalid"); + assert_eq!(document.sink_string_summary(0, "level"), "debug"); + assert_eq!(document.sink_string_summary(0, "format"), "unset"); + assert_eq!( + document.sink_integer_summary(0, "queue_capacity"), + "invalid" + ); + assert_eq!(document.sink_integer_summary(0, "retained_files"), "unset"); + assert_eq!(document.sink_rotation_summary(0), "incomplete"); + + document + .set_sink_string(0, "path", "relay.log".into()) + .unwrap(); + document.clear_sink_key(0, "level").unwrap(); + assert_eq!( + document.sink_string(0, "path").as_deref(), + Some("relay.log") + ); +} + +#[test] +fn project_path_defaults_to_start_when_no_ancestor_config_exists() { + let root = tempfile::tempdir().unwrap(); + let nested = root.path().join("a/b"); + std::fs::create_dir_all(&nested).unwrap(); + + assert_eq!( + project_config_path(&nested), + nested.join(".nemo-relay/config.toml") + ); } diff --git a/crates/cli/tests/coverage/shared/config_tests.rs b/crates/cli/tests/coverage/shared/config_tests.rs index cf74749b9..b72d5a956 100644 --- a/crates/cli/tests/coverage/shared/config_tests.rs +++ b/crates/cli/tests/coverage/shared/config_tests.rs @@ -4148,3 +4148,108 @@ level = "error" } } } + +#[test] +fn configuration_value_helpers_cover_empty_and_invalid_shapes() { + let source = Path::new("plugins.toml"); + let mut scalar = toml::Value::String("not a table".into()); + let resolved = + resolve_dynamic_plugin_refs(source, &mut scalar, &mut std::collections::HashSet::new()) + .unwrap(); + assert!(resolved.dynamic_plugins.is_empty()); + assert_eq!( + resolved.dynamic_plugin_policy, + DynamicPluginHostPolicy::default() + ); + + let value: toml::Value = toml::from_str( + r#" +[plugins] +dynamic = [] +[plugins.policy] +rules = [] +[other] +enabled = true +"#, + ) + .unwrap(); + let cleaned = remove_dynamic_plugin_sections(value); + assert!(cleaned.get("plugins").is_none()); + assert_eq!(cleaned["other"]["enabled"].as_bool(), Some(true)); + assert_eq!(plugin_toml_runtime_value(json!({})), None); + assert_eq!( + plugin_toml_runtime_value(json!({"enabled": true})), + Some(json!({"enabled": true})) + ); + + assert!(validate_auth_header("AUTH", " ".into()).is_err()); + assert!(parse_env_body_limit("LIMIT", "not-a-number").is_err()); + assert_eq!(parse_env_body_limit("LIMIT", "17").unwrap(), 17); +} + +#[test] +fn logging_path_and_sink_helpers_cover_lexical_fallbacks() { + let temp = tempfile::tempdir().unwrap(); + let existing_parent = temp.path().join("logs"); + std::fs::create_dir_all(&existing_parent).unwrap(); + let canonical = std::fs::canonicalize(&existing_parent) + .unwrap() + .join("relay.log"); + assert_eq!( + logging_path_identity(&existing_parent.join("relay.log")), + canonical + ); + + let lexical = logging_path_identity(&temp.path().join("missing/../nested/relay.log")); + assert!(lexical.ends_with("nested/relay.log")); + assert_eq!( + normalize_path_components(Path::new("a/./b/../c")), + PathBuf::from("a/c") + ); + assert_eq!( + normalize_path_components(Path::new("single")), + PathBuf::from("single") + ); + + let path = temp.path().join("coalesced.log"); + let lower = toml::Value::Table(toml::Table::from_iter([ + ( + "path".into(), + toml::Value::String(path.display().to_string()), + ), + ("level".into(), toml::Value::String("info".into())), + ])); + let higher = toml::Value::Table(toml::Table::from_iter([ + ( + "path".into(), + toml::Value::String(path.display().to_string()), + ), + ("format".into(), toml::Value::String("json".into())), + ])); + let coalesced = coalesce_logging_sinks(vec![lower, higher]); + assert_eq!(coalesced.len(), 1); + assert_eq!(coalesced[0]["level"].as_str(), Some("info")); + assert_eq!(coalesced[0]["format"].as_str(), Some("json")); + + let mut invalid_higher: toml::Value = toml::from_str("[logging]\nsinks = 'invalid'").unwrap(); + let lower = toml::Value::Table(toml::Table::new()); + merge_logging_sinks_by_path(&lower, &mut invalid_higher); + assert_eq!(invalid_higher["logging"]["sinks"].as_str(), Some("invalid")); +} + +#[test] +fn dynamic_plugin_identity_allows_worker_without_manifest() { + let component = ActiveDynamicPluginComponent { + plugin_id: "acme.manual-worker".into(), + kind: DynamicPluginKind::Worker, + lifecycle_generation: 7, + manifest_ref: None, + environment_ref: None, + config: serde_json::Map::new(), + activation_snapshot: None, + }; + let identity = dynamic_plugin_bootstrap_identity(&component).unwrap(); + assert_eq!(identity["plugin_id"], "acme.manual-worker"); + assert_eq!(identity["manifest"], Value::Null); + assert_eq!(identity["lifecycle_generation"], 7); +} diff --git a/crates/cli/tests/coverage/shared/installer_tests.rs b/crates/cli/tests/coverage/shared/installer_tests.rs index 96ff9475c..f77f7d74d 100644 --- a/crates/cli/tests/coverage/shared/installer_tests.rs +++ b/crates/cli/tests/coverage/shared/installer_tests.rs @@ -580,6 +580,13 @@ fn packaged_plugin_hooks_use_expected_forwarding_commands() { fn packaged_plugin_manifests_use_stable_plugin_name_and_version() { let root = std::path::PathBuf::from(env!("CARGO_MANIFEST_DIR")) .join("../../integrations/coding-agents"); + + assert_agent_plugin_manifests(&root); + assert_agent_mcp_manifests(&root); + assert_agent_marketplace_manifests(&root); +} + +fn assert_agent_plugin_manifests(root: &std::path::Path) { let claude_path = root.join("claude-code/.claude-plugin/plugin.json"); let claude = serde_json::from_str::(&std::fs::read_to_string(&claude_path).unwrap()).unwrap(); @@ -595,7 +602,9 @@ fn packaged_plugin_manifests_use_stable_plugin_name_and_version() { assert_eq!(codex["version"], json!(env!("CARGO_PKG_VERSION"))); assert!(codex.get("hooks").is_none()); assert_eq!(codex["mcpServers"], json!("./.mcp.json")); +} +fn assert_agent_mcp_manifests(root: &std::path::Path) { let codex_mcp_path = root.join("codex/.mcp.json"); let codex_mcp = serde_json::from_str::(&std::fs::read_to_string(&codex_mcp_path).unwrap()).unwrap(); @@ -628,7 +637,9 @@ fn packaged_plugin_manifests_use_stable_plugin_name_and_version() { json!({"NEMO_RELAY_GATEWAY_BIND": "127.0.0.1:47632"}) ); assert_eq!(claude_server["alwaysLoad"], json!(true)); +} +fn assert_agent_marketplace_manifests(root: &std::path::Path) { let codex_marketplace_path = root.join("../../.agents/plugins/marketplace.json"); let codex_marketplace = serde_json::from_str::(&std::fs::read_to_string(&codex_marketplace_path).unwrap()) diff --git a/crates/cli/tests/coverage/shared/plugins_lifecycle_tests.rs b/crates/cli/tests/coverage/shared/plugins_lifecycle_tests.rs index 4770cedc4..efcef6488 100644 --- a/crates/cli/tests/coverage/shared/plugins_lifecycle_tests.rs +++ b/crates/cli/tests/coverage/shared/plugins_lifecycle_tests.rs @@ -1918,21 +1918,7 @@ fn add_provisions_persists_and_removes_managed_python_environment() { .as_deref() .expect("managed environment should be persisted"); let environment_path = PathBuf::from(environment_ref); - assert!(environment_path.is_absolute()); - let expected_environment_name = Sha256::digest(b"acme.python") - .iter() - .map(|byte| format!("{byte:02x}")) - .collect::(); - assert_eq!( - environment_path.file_name(), - Some(OsStr::new(&expected_environment_name)) - ); - assert!( - environment_path - .parent() - .is_some_and(|parent| parent.ends_with(".dynamic-plugin-environments")) - ); - assert!(environment::environment_python_path(&environment_path).is_file()); + assert_managed_environment_path(&environment_path); assert_eq!( added.record.status.validation.environment, DynamicPluginCheckState::Valid @@ -1956,39 +1942,8 @@ fn add_provisions_persists_and_removes_managed_python_environment() { inspect["data"]["source"]["environment_ref"], serde_json::json!(environment_ref) ); - let calls = runner.calls(); - assert_eq!(calls.len(), 2); - assert_eq!( - calls[0].0, - OsString::from(if cfg!(windows) { "python" } else { "python3" }) - ); - assert_eq!( - calls[0].1, - vec![ - OsString::from("-m"), - OsString::from("venv"), - environment_path.as_os_str().to_owned(), - ] - ); - assert_eq!( - PathBuf::from(&calls[1].0), - environment::environment_python_path(&environment_path) - ); - assert_eq!( - calls[1].1, - vec![ - OsString::from("-m"), - OsString::from("pip"), - OsString::from("install"), - plugin_dir.canonicalize().unwrap().into_os_string(), - ] - ); - assert!( - !calls[1] - .1 - .iter() - .any(|arg| arg == "-e" || arg == "--editable") - ); + assert_python_environment_runner_calls(&runner.calls(), &environment_path, &plugin_dir); + enable( PluginsEnableRequest { id: "acme.python".into(), @@ -2031,6 +1986,63 @@ fn add_provisions_persists_and_removes_managed_python_environment() { assert!(!stale_marker.exists()); } +fn assert_managed_environment_path(environment_path: &Path) { + assert!(environment_path.is_absolute()); + let expected_environment_name = Sha256::digest(b"acme.python") + .iter() + .map(|byte| format!("{byte:02x}")) + .collect::(); + assert_eq!( + environment_path.file_name(), + Some(OsStr::new(&expected_environment_name)) + ); + assert!( + environment_path + .parent() + .is_some_and(|parent| parent.ends_with(".dynamic-plugin-environments")) + ); + assert!(environment::environment_python_path(environment_path).is_file()); +} + +fn assert_python_environment_runner_calls( + calls: &[(OsString, Vec)], + environment_path: &Path, + plugin_dir: &Path, +) { + assert_eq!(calls.len(), 2); + assert_eq!( + calls[0].0, + OsString::from(if cfg!(windows) { "python" } else { "python3" }) + ); + assert_eq!( + calls[0].1, + vec![ + OsString::from("-m"), + OsString::from("venv"), + environment_path.as_os_str().to_owned(), + ] + ); + assert_eq!( + PathBuf::from(&calls[1].0), + environment::environment_python_path(environment_path) + ); + assert_eq!( + calls[1].1, + vec![ + OsString::from("-m"), + OsString::from("pip"), + OsString::from("install"), + plugin_dir.canonicalize().unwrap().into_os_string(), + ] + ); + assert!( + !calls[1] + .1 + .iter() + .any(|arg| arg == "-e" || arg == "--editable") + ); +} + #[test] fn add_rolls_back_python_environment_when_installation_fails() { let temp = tempfile::tempdir().unwrap(); @@ -4244,3 +4256,279 @@ fn inspect_distinguishes_empty_host_config_from_missing_host_config() { 0 ); } + +fn required_lifecycle_record( + temp: &tempfile::TempDir, + plugin_id: &str, +) -> ScopedDynamicPluginRecord { + let plugin_dir = temp.path().join(plugin_id); + std::fs::create_dir_all(&plugin_dir).unwrap(); + write_dynamic_manifest(&plugin_dir, plugin_id); + let server = GatewayOverrides::default(); + add( + PluginsAddRequest { + scope: ConfigurationScope::Project, + path: plugin_dir, + }, + &server, + ) + .unwrap(); + let scopes = load_scoped_registries(None).unwrap(); + let mut entry = find_record_by_id(&scopes, plugin_id).unwrap().unwrap(); + entry.record.status.startup_class = + Some(nemo_relay::plugin::dynamic::DynamicPluginStartupClass::Required); + entry +} + +#[test] +fn lifecycle_helpers_cover_environment_manifest_scope_and_restore_paths() { + let temp = tempfile::tempdir().unwrap(); + let plugin_id = "acme.lifecycle-helpers"; + assert_eq!( + environment_last_error(plugin_id, DynamicPluginCheckState::Valid, None), + None + ); + let missing = environment_last_error(plugin_id, DynamicPluginCheckState::Invalid, None) + .expect("invalid environment should produce a diagnostic"); + assert_eq!(missing.code, "environment_failed"); + assert!(missing.message.contains("has no lifecycle-managed")); + let unavailable = environment_last_error( + plugin_id, + DynamicPluginCheckState::Invalid, + Some("managed/python"), + ) + .expect("invalid referenced environment should produce a diagnostic"); + assert!( + unavailable + .message + .contains("managed/python is unavailable") + ); + + let mut scopes = Vec::new(); + let plugins_path = temp.path().join("plugins.toml"); + let state_path = temp.path().join("state.json"); + let first = ensure_scope( + &mut scopes, + RegistryScope::Project, + plugins_path.clone(), + state_path.clone(), + ); + let existing = ensure_scope( + &mut scopes, + RegistryScope::Project, + plugins_path.clone(), + state_path, + ); + assert_eq!((first, existing, scopes.len()), (0, 0, 1)); + assert!(!scope_flags_selected(&ConfigurationScope::Default)); + assert!(scope_flags_selected(&ConfigurationScope::Project)); + + restore_plugins_toml(&plugins_path, Some(b"[plugins]\n")).unwrap(); + assert_eq!(std::fs::read(&plugins_path).unwrap(), b"[plugins]\n"); + restore_plugins_toml(&plugins_path, None).unwrap(); + assert!(!plugins_path.exists()); + restore_plugins_toml(&plugins_path, None).unwrap(); +} + +#[test] +fn manifest_helpers_report_missing_and_invalid_sources() { + let temp = tempfile::tempdir().unwrap(); + let _env = EnvScope::hermetic(&temp); + let _cwd = CurrentDirGuard::enter(temp.path()); + let mut entry = required_lifecycle_record(&temp, "acme.manifest-helpers"); + entry.record.source.manifest_ref = None; + let error = manifest_ref_from_record(&entry.record).unwrap_err(); + assert!(error.to_string().contains("has no manifest_ref")); + + let missing = temp.path().join("missing-manifest.toml"); + let error = load_manifest_for_action("inspect", &missing).unwrap_err(); + assert!(error.to_string().contains("dynamic plugin inspect failed")); +} + +#[test] +fn required_startup_failure_reports_policy_trust_and_environment_failures() { + let temp = tempfile::tempdir().unwrap(); + let _env = EnvScope::hermetic(&temp); + let _cwd = CurrentDirGuard::enter(temp.path()); + let mut entry = required_lifecycle_record(&temp, "acme.required-checks"); + + entry.record.status.validation.policy_satisfied = DynamicPluginCheckState::Invalid; + let failure = required_startup_failure(&entry, &[]).unwrap(); + assert!(failure.contains("blocked by host policy")); + + entry.record.status.validation.policy_satisfied = DynamicPluginCheckState::Valid; + entry.record.status.validation.integrity = DynamicPluginCheckState::Invalid; + let failure = required_startup_failure(&entry, &[]).unwrap(); + assert!(failure.contains("trust verification failed")); + + entry.record.status.validation.integrity = DynamicPluginCheckState::Valid; + entry.record.status.validation.environment = DynamicPluginCheckState::Invalid; + let failure = required_startup_failure(&entry, &[]).unwrap(); + assert!(failure.contains("environment is unavailable")); +} + +#[test] +fn required_startup_failure_reports_custom_and_manifest_diagnostics() { + let temp = tempfile::tempdir().unwrap(); + let _env = EnvScope::hermetic(&temp); + let _cwd = CurrentDirGuard::enter(temp.path()); + let mut entry = required_lifecycle_record(&temp, "acme.required-manifest"); + entry.record.status.validation.policy_satisfied = DynamicPluginCheckState::Invalid; + entry.record.status.last_error = Some(DynamicPluginFailure { + phase: DynamicPluginFailurePhase::Validation, + code: "custom_failure".into(), + message: "custom lifecycle diagnostic".into(), + }); + let failure = required_startup_failure(&entry, &[]).unwrap(); + assert!(failure.contains("custom lifecycle diagnostic")); + + entry.record.status.last_error = None; + entry.record.status.validation.policy_satisfied = DynamicPluginCheckState::Valid; + entry.record.source.manifest_ref = None; + let failure = required_startup_failure(&entry, &[]).unwrap(); + assert!(failure.contains("has no manifest_ref")); + + let missing = temp.path().join("removed-manifest.toml"); + entry.record.source.manifest_ref = Some(missing.display().to_string()); + let failure = required_startup_failure(&entry, &[]).unwrap(); + assert!(failure.contains("is no longer available")); + + let invalid = temp.path().join("invalid-manifest.toml"); + std::fs::write(&invalid, b"not valid TOML = [").unwrap(); + entry.record.source.manifest_ref = Some(invalid.display().to_string()); + let failure = required_startup_failure(&entry, &[]).unwrap(); + assert!(failure.contains("is unreadable")); +} + +#[test] +fn required_startup_failure_accepts_optional_and_readable_plugins() { + let temp = tempfile::tempdir().unwrap(); + let _env = EnvScope::hermetic(&temp); + let _cwd = CurrentDirGuard::enter(temp.path()); + let mut entry = required_lifecycle_record(&temp, "acme.required-ready"); + assert_eq!(required_startup_failure(&entry, &[]), None); + + entry.record.status.startup_class = + Some(nemo_relay::plugin::dynamic::DynamicPluginStartupClass::Optional); + entry.record.status.validation.policy_satisfied = DynamicPluginCheckState::Invalid; + assert_eq!(required_startup_failure(&entry, &[]), None); +} + +#[test] +fn lifecycle_commands_cover_json_and_human_output_paths() { + let temp = tempfile::tempdir().unwrap(); + let _env = EnvScope::hermetic(&temp); + let _cwd = CurrentDirGuard::enter(temp.path()); + let server = GatewayOverrides::default(); + + list( + PluginsListRequest { + all: false, + json: false, + }, + &server, + ) + .unwrap(); + list( + PluginsListRequest { + all: false, + json: true, + }, + &server, + ) + .unwrap(); + + let plugin_dir = temp.path().join("output-plugin"); + std::fs::create_dir_all(&plugin_dir).unwrap(); + let manifest_path = write_dynamic_manifest(&plugin_dir, "acme.output"); + add( + PluginsAddRequest { + scope: ConfigurationScope::Project, + path: plugin_dir, + }, + &server, + ) + .unwrap(); + + validate( + PluginsValidateRequest { + target: manifest_path.display().to_string(), + json: false, + }, + &server, + ) + .unwrap(); + validate( + PluginsValidateRequest { + target: manifest_path.display().to_string(), + json: true, + }, + &server, + ) + .unwrap(); + validate( + PluginsValidateRequest { + target: "acme.output".into(), + json: false, + }, + &server, + ) + .unwrap(); + validate( + PluginsValidateRequest { + target: "acme.output".into(), + json: true, + }, + &server, + ) + .unwrap(); + list( + PluginsListRequest { + all: false, + json: false, + }, + &server, + ) + .unwrap(); + list( + PluginsListRequest { + all: false, + json: true, + }, + &server, + ) + .unwrap(); + inspect( + PluginsInspectRequest { + id: "acme.output".into(), + json: false, + }, + &server, + ) + .unwrap(); + inspect( + PluginsInspectRequest { + id: "acme.output".into(), + json: true, + }, + &server, + ) + .unwrap(); +} + +#[test] +fn validate_rejects_a_missing_path_target() { + let temp = tempfile::tempdir().unwrap(); + let _env = EnvScope::hermetic(&temp); + let _cwd = CurrentDirGuard::enter(temp.path()); + let missing = temp.path().join("missing").join("relay-plugin.toml"); + let error = validate( + PluginsValidateRequest { + target: missing.display().to_string(), + json: false, + }, + &GatewayOverrides::default(), + ) + .unwrap_err(); + assert!(error.to_string().contains("does not exist")); +} diff --git a/crates/cli/tests/coverage/shared/plugins_schema_tests.rs b/crates/cli/tests/coverage/shared/plugins_schema_tests.rs index e126f62e7..9b1f9fbec 100644 --- a/crates/cli/tests/coverage/shared/plugins_schema_tests.rs +++ b/crates/cli/tests/coverage/shared/plugins_schema_tests.rs @@ -272,14 +272,21 @@ fn maps_native_nested_map_and_raw_controls() { "union": {"oneOf": [{"type": "string"}, {"type": "number"}]} } })); - let field = |key: &str| { - loaded - .fields() - .iter() - .find(|field| field.key == key) - .unwrap() - }; + assert_native_raw_and_scalar_fields(&loaded); + assert_native_choice_and_nested_fields(&loaded); + assert!(loaded.editor().title.is_none()); +} +fn native_config_field<'a>(schema: &'a PluginConfigSchema, key: &str) -> &'a DynamicConfigField { + schema + .fields() + .iter() + .find(|field| field.key == key) + .unwrap() +} + +fn assert_native_raw_and_scalar_fields(schema: &PluginConfigSchema) { + let field = |key| native_config_field(schema, key); assert!(matches!( field("array").kind, DynamicConfigFieldKind::RawJson @@ -307,6 +314,10 @@ fn maps_native_nested_map_and_raw_controls() { assert_eq!(field("enabled").title, "Enabled"); assert_eq!(field("enabled").default, Some(json!(true))); assert!(field("enabled").required); +} + +fn assert_native_choice_and_nested_fields(schema: &PluginConfigSchema) { + let field = |key| native_config_field(schema, key); assert!(matches!( field("choice").kind, DynamicConfigFieldKind::StringEnum { ref options, secret: false } @@ -320,7 +331,6 @@ fn maps_native_nested_map_and_raw_controls() { && fields[0].description.as_deref() == Some("Weight") && matches!(fields[0].kind, DynamicConfigFieldKind::Number) )); - assert!(loaded.editor().title.is_none()); } #[test] diff --git a/crates/cli/tests/coverage/shared/plugins_tests.rs b/crates/cli/tests/coverage/shared/plugins_tests.rs index cd6cc9abf..3deed9031 100644 --- a/crates/cli/tests/coverage/shared/plugins_tests.rs +++ b/crates/cli/tests/coverage/shared/plugins_tests.rs @@ -320,6 +320,12 @@ fn typed_editor_model_contains_nemo_guardrails_options() { #[test] fn typed_editor_model_contains_pii_redaction_options() { let schema = PiiRedactionConfig::editor_schema(); + assert_pii_root_editor_fields(schema); + assert_pii_builtin_editor_fields(schema.field("builtin").unwrap().schema().unwrap()); + assert_pii_local_editor_fields(schema.field("local").unwrap().schema().unwrap()); +} + +fn assert_pii_root_editor_fields(schema: &EditorSchema) { assert!(!schema.fields.iter().any(|field| field.name == "version")); assert_eq!( schema.field("mode").unwrap().enum_values, @@ -331,8 +337,9 @@ fn typed_editor_model_contains_pii_redaction_options() { schema.field("tool_output").unwrap().kind, EditorFieldKind::Boolean ); +} - let builtin = schema.field("builtin").unwrap().schema().unwrap(); +fn assert_pii_builtin_editor_fields(builtin: &EditorSchema) { assert_eq!(builtin.field("preset").unwrap().kind, EditorFieldKind::Enum); assert!( builtin @@ -398,8 +405,9 @@ fn typed_editor_model_contains_pii_redaction_options() { builtin.field("unmasked_suffix").unwrap().kind, EditorFieldKind::Integer ); +} - let local = schema.field("local").unwrap().schema().unwrap(); +fn assert_pii_local_editor_fields(local: &EditorSchema) { assert_eq!( local.field("backend").unwrap().kind, EditorFieldKind::String @@ -588,6 +596,216 @@ fn component_field_clear_only_removes_optional_fields() { assert!(!pii_redaction.field_configured(input)); } +#[test] +fn menu_keys_cover_selection_shortcuts_and_cancellation() { + assert_eq!( + menu_response_for_key(&Key::Enter, 2), + Some(MenuResponse::Selected(2)) + ); + assert_eq!( + menu_response_for_key(&Key::Char(' '), 3), + Some(MenuResponse::Selected(3)) + ); + assert_eq!( + menu_response_for_key(&Key::Char('p'), 1), + Some(MenuResponse::Shortcut(MenuShortcut::Preview, 1)) + ); + assert_eq!( + menu_response_for_key(&Key::Char('s'), 1), + Some(MenuResponse::Shortcut(MenuShortcut::Save, 1)) + ); + assert_eq!( + menu_response_for_key(&Key::Char('r'), 1), + Some(MenuResponse::Shortcut(MenuShortcut::Reset, 1)) + ); + assert_eq!( + menu_response_for_key(&Key::Backspace, 1), + Some(MenuResponse::Shortcut(MenuShortcut::Clear, 1)) + ); + assert_eq!( + menu_response_for_key(&Key::Char('?'), 1), + Some(MenuResponse::Shortcut(MenuShortcut::Help, 1)) + ); + assert_eq!( + menu_response_for_key(&Key::Escape, 1), + Some(MenuResponse::Cancel) + ); + assert_eq!(menu_response_for_key(&Key::Char('x'), 1), None); +} + +#[test] +fn section_menu_helpers_render_defaults_and_reset_sections() { + let mut config = ObservabilityConfig::default(); + let section = ObservabilityConfig::editor_schema().field("atof").unwrap(); + let fields = section.schema().unwrap().fields; + + let items = section_menu_items(&config, section, fields).unwrap(); + assert_eq!(items.len(), fields.len() + 3); + assert!(items.last().unwrap().label.contains("Back")); + assert_eq!(selected_field_index(section, 1), 0); + assert_eq!(reset_section_index(section, fields), fields.len() + 1); + + ensure_section(&mut config, section); + reset_selected_item( + &mut config, + section, + fields, + reset_section_index(section, fields), + ) + .unwrap(); + assert!(section_configured(&config, section)); +} + +#[test] +fn value_menu_helpers_store_reset_and_clear_nested_values() { + static FIELDS: [EditorFieldSpec; 2] = [ + EditorFieldSpec { + name: "optional", + label: "Optional", + kind: EditorFieldKind::String, + enum_values: &[], + optional: true, + nested_schema: None, + nested_default: None, + list_item: None, + tagged_union: None, + }, + EditorFieldSpec { + name: "required", + label: "Required", + kind: EditorFieldKind::String, + enum_values: &[], + optional: false, + nested_schema: None, + nested_default: None, + list_item: None, + tagged_union: None, + }, + ]; + static SCHEMA: EditorSchema = EditorSchema { fields: &FIELDS }; + let optional = FIELDS[0]; + let required = FIELDS[1]; + let mut value = json!({}); + set_value_field(&mut value, optional.name, json!("configured")); + set_value_field(&mut value, required.name, Value::Null); + + let items = value_section_menu_items(&value, &SCHEMA, None).unwrap(); + assert_eq!(items.len(), SCHEMA.fields.len() + 2); + assert!(value_field_configured(&value, optional, None)); + assert!(clear_value_field( + &mut value, + &SCHEMA, + SCHEMA + .fields + .iter() + .position(|field| field.name == optional.name) + .unwrap() + )); + assert!(!value_field_configured(&value, optional, None)); + assert!(!clear_value_field(&mut value, &SCHEMA, SCHEMA.fields.len())); + + reset_value_section_item( + &mut value, + &SCHEMA, + Some(&json!({"restored": true})), + SCHEMA.fields.len(), + ); + assert_eq!(value, json!({"restored": true})); +} + +#[test] +fn editor_storage_helpers_preserve_nonempty_values_and_prune_empty_sections() { + let field = EditorFieldSpec { + name: "section", + label: "Section", + kind: EditorFieldKind::Section, + enum_values: &[], + optional: true, + nested_schema: None, + nested_default: None, + list_item: None, + tagged_union: None, + }; + let mut config = json!({}); + + store_edited_config_section(&mut config, field, json!({"enabled": true})).unwrap(); + assert!(config_field_value(&config, field.name).unwrap().is_some()); + store_edited_config_section(&mut config, field, json!({})).unwrap(); + assert!(config_field_value(&config, field.name).unwrap().is_none()); + + let mut target = json!({"atof": {"enabled": true}}); + store_edited_value_section(&mut target, field, json!({})); + assert!(value_field_value(&target, field.name).is_none()); + store_edited_value_section(&mut target, field, json!({"enabled": true})); + assert_eq!( + value_field_value(&target, field.name), + Some(json!({"enabled": true})) + ); +} + +#[test] +fn component_shortcut_fallbacks_are_safe_noops() { + let mut config = PluginConfig::default(); + ensure_observability_component(&mut config).unwrap(); + let mut components = editable_components(&config).unwrap(); + + assert_eq!( + handle_reset_or_clear_shortcut(&mut components, None, MenuShortcut::Reset).unwrap(), + EditLoopControl::Continue + ); + reset_component_menu_item(&mut components[0], None).unwrap(); + clear_component_menu_item(&mut components[0], Some(ComponentMenuAction::Back)).unwrap(); +} + +#[test] +fn editable_component_dispatch_covers_every_component_variant() { + let config = PluginConfig::default(); + let mut components = editable_components(&config).unwrap(); + + for component in &mut components { + assert!(!component.label().is_empty()); + assert!(!component.fields().is_empty()); + assert!(!component.summary().is_empty()); + component.toggle_enabled(); + component.set_enabled(true); + component.reset_enabled(); + let optional = *component + .fields() + .iter() + .find(|field| field.optional) + .expect("every editable component exposes an optional field"); + component.reset_field(optional).unwrap(); + assert!(component.clear_field(optional).unwrap()); + } + + let rendered = config_with_editable_components(&config, &components).unwrap(); + let mut stored = config.clone(); + store_editable_components(&mut stored, &components).unwrap(); + assert_eq!( + serde_json::to_value(rendered).unwrap(), + serde_json::to_value(stored).unwrap() + ); +} + +#[test] +fn editor_model_object_and_schema_helpers_cover_fallbacks() { + let mut scalar = json!("replace me"); + ensure_object(&mut scalar).insert("enabled".into(), json!(true)); + assert_eq!(scalar, json!({"enabled": true})); + + let fields = observability_editor_fields_with_version(); + assert_eq!(fields.first(), Some(&"version")); + let nested = nested_editor_keys(ObservabilityConfig::editor_schema()); + assert!(nested.contains(&"atof")); + + let error = serde_json::from_str::("{").unwrap_err(); + assert!( + serde_error(error) + .to_string() + .contains("invalid plugin editor value") + ); +} + #[test] fn menu_viewport_keeps_selection_visible_and_pages() { let first = menu_viewport(20, 0, 8); @@ -1660,6 +1878,10 @@ value = "preserve-host-section" document.write().unwrap(); let rendered = std::fs::read_to_string(&path).unwrap(); let root = rendered.parse::().unwrap(); + assert_preserved_plugin_document(&root); +} + +fn assert_preserved_plugin_document(root: &toml::Table) { assert_eq!(root["host_setting"].as_str(), Some("preserve-me")); assert_eq!( root["host"]["extra"]["value"].as_str(), @@ -1846,6 +2068,131 @@ config = {} let document = PluginConfigDocument::read(&path).unwrap(); let mut states = load_dynamic_plugin_states(&document).unwrap(); + assert_dynamic_editor_initial_states(&states); + + states[2].clear_top_level_field("optional"); + states[2].reset_top_level_field("optional").unwrap(); + assert_eq!(states[2].config(), None); + + let mut preview = document.clone(); + for state in &states { + state.apply_to_document(&mut preview, true).unwrap(); + } + assert_dynamic_editor_redacted_preview(&preview.render().unwrap()); + + states[0].reset_top_level_field("retries").unwrap(); + assert_eq!(states[0].config().unwrap().get("retries"), Some(&json!(3))); + assert_eq!( + states[0].config().unwrap().get("unknown"), + Some(&json!({"nested": "keep"})) + ); + let mut touched = document.clone(); + states[0].apply_to_document(&mut touched, false).unwrap(); + assert_dynamic_editor_touched_document(&touched); + for field in ["token", "retries", "unknown", "observed_at", "records"] { + states[0].clear_top_level_field(field); + } + assert_eq!(states[0].config(), Some(&Map::new())); + + states[0].reset(); + let mut persisted = document.clone(); + for state in &states { + state.apply_to_document(&mut persisted, false).unwrap(); + } + assert_dynamic_editor_persisted_document(&persisted); +} + +#[test] +fn dynamic_editor_menu_actions_reset_clear_and_render_fields() { + let temp = tempfile::tempdir().unwrap(); + write_editor_dynamic_manifest( + &temp.path().join("plugin"), + "acme.menu", + Some("Menu Plugin"), + Some(&json!({ + "$schema": "https://json-schema.org/draft/2020-12/schema", + "type": "object", + "required": ["mode"], + "properties": { + "mode": {"type": "string", "default": "safe"}, + "token": {"type": "string", "writeOnly": true, "default": "secret"} + } + })), + ); + let path = temp.path().join("plugins.toml"); + std::fs::write( + &path, + "[[plugins.dynamic]]\nmanifest = \"./plugin/relay-plugin.toml\"\nconfig = { mode = \"fast\" }\n", + ) + .unwrap(); + let document = PluginConfigDocument::read(&path).unwrap(); + let mut states = load_dynamic_plugin_states(&document).unwrap(); + let state = &mut states[0]; + let fields = state.editor_fields().to_vec(); + + let (items, actions) = dynamic_field_menu_items(state, &fields, &[]); + assert_eq!(items.len(), fields.len() + 2); + assert!(items.iter().any(|item| item.label.contains("[required]"))); + assert!(items.iter().any(|item| item.label.contains(""))); + + let mode = fields.iter().position(|field| field.key == "mode").unwrap(); + reset_dynamic_selection(state, &fields, &[], &actions, mode); + assert_eq!(state.config().unwrap().get("mode"), Some(&json!("safe"))); + clear_dynamic_selection(state, &fields, &[], &actions, mode); + assert!(!state.config().unwrap().contains_key("mode")); + + let reset = actions + .iter() + .position(|action| matches!(action, DynamicMenuAction::ResetPlugin)) + .unwrap(); + reset_dynamic_selection(state, &fields, &[], &actions, reset); + assert_eq!(state.config(), None); +} + +#[test] +fn dynamic_editor_raw_menu_and_nested_value_paths_are_deterministic() { + let temp = tempfile::tempdir().unwrap(); + write_editor_dynamic_manifest(&temp.path().join("plugin"), "acme.raw-menu", None, None); + let path = temp.path().join("plugins.toml"); + std::fs::write( + &path, + "[[plugins.dynamic]]\nmanifest = \"./plugin/relay-plugin.toml\"\n", + ) + .unwrap(); + let document = PluginConfigDocument::read(&path).unwrap(); + let states = load_dynamic_plugin_states(&document).unwrap(); + let (items, actions) = dynamic_root_menu_items(&states[0], &[]); + assert_eq!(items.len(), 3); + assert!(matches!(actions[0], DynamicMenuAction::EditRawConfig)); + + let mut config = None; + let path = vec!["outer".to_owned(), "inner".to_owned()]; + set_value_at_path(&mut config, &path, json!(7)); + assert_eq!(value_at_path(config.as_ref(), &path), Some(&json!(7))); + assert_eq!(value_at_path(config.as_ref(), &[]), None); + assert!(remove_value_at_path(config.as_mut().unwrap(), &path)); + set_value_at_path(&mut config, &[], json!(9)); + assert!(config.as_ref().unwrap().is_empty()); +} + +#[test] +fn dynamic_editor_rejects_duplicate_plugin_ids() { + let temp = tempfile::tempdir().unwrap(); + write_editor_dynamic_manifest(&temp.path().join("plugin"), "acme.duplicate", None, None); + let path = temp.path().join("plugins.toml"); + std::fs::write( + &path, + "[[plugins.dynamic]]\nmanifest = \"./plugin/relay-plugin.toml\"\n\n[[plugins.dynamic]]\nmanifest = \"./plugin/relay-plugin.toml\"\n", + ) + .unwrap(); + let document = PluginConfigDocument::read(&path).unwrap(); + let error = load_dynamic_plugin_states(&document) + .unwrap_err() + .to_string(); + assert!(error.contains("declared more than once"), "{error}"); +} + +fn assert_dynamic_editor_initial_states(states: &[DynamicPluginEditorState]) { assert_eq!(states.len(), 4); assert_eq!(states[0].label(), "Structured Plugin (acme.structured)"); assert_eq!(states[1].label(), "acme.raw"); @@ -1863,16 +2210,9 @@ config = {} let labels = states[0].top_level_field_labels(); assert!(labels.iter().any(|label| label.contains(""))); assert!(labels.iter().all(|label| !label.contains("super-secret"))); +} - states[2].clear_top_level_field("optional"); - states[2].reset_top_level_field("optional").unwrap(); - assert_eq!(states[2].config(), None); - - let mut preview = document.clone(); - for state in &states { - state.apply_to_document(&mut preview, true).unwrap(); - } - let rendered = preview.render().unwrap(); +fn assert_dynamic_editor_redacted_preview(rendered: &str) { assert!(rendered.contains("")); assert!(!rendered.contains("super-secret")); assert!(!rendered.contains("nested-secret")); @@ -1890,15 +2230,9 @@ config = {} .get("config") .is_none() ); +} - states[0].reset_top_level_field("retries").unwrap(); - assert_eq!(states[0].config().unwrap().get("retries"), Some(&json!(3))); - assert_eq!( - states[0].config().unwrap().get("unknown"), - Some(&json!({"nested": "keep"})) - ); - let mut touched = document.clone(); - states[0].apply_to_document(&mut touched, false).unwrap(); +fn assert_dynamic_editor_touched_document(touched: &PluginConfigDocument) { assert_eq!( touched.dynamic_entries().unwrap()[0] .config @@ -1912,18 +2246,9 @@ config = {} touched_root["plugins"]["dynamic"].as_array().unwrap()[0]["config"]["observed_at"] .is_datetime() ); - states[0].clear_top_level_field("token"); - states[0].clear_top_level_field("retries"); - states[0].clear_top_level_field("unknown"); - states[0].clear_top_level_field("observed_at"); - states[0].clear_top_level_field("records"); - assert_eq!(states[0].config(), Some(&Map::new())); +} - states[0].reset(); - let mut persisted = document.clone(); - for state in &states { - state.apply_to_document(&mut persisted, false).unwrap(); - } +fn assert_dynamic_editor_persisted_document(persisted: &PluginConfigDocument) { let entries = persisted.dynamic_entries().unwrap(); assert_eq!(entries[0].config, None); assert_eq!(entries[2].config, None); @@ -1996,6 +2321,96 @@ fn plugin_config_document_reports_invalid_dynamic_entries_and_indexes() { ); } +#[test] +fn plugin_config_document_reports_each_dynamic_entry_shape_error() { + let temp = tempfile::tempdir().unwrap(); + let path = temp.path().join("plugins.toml"); + + std::fs::write(&path, "[plugins]\n").unwrap(); + let mut document = PluginConfigDocument::read(&path).unwrap(); + assert!(document.dynamic_entries().unwrap().is_empty()); + assert!(document.remove_dynamic_config(0).is_err()); + + std::fs::write(&path, "[plugins]\ndynamic = [1]\n").unwrap(); + let mut document = PluginConfigDocument::read(&path).unwrap(); + assert!( + document + .dynamic_entries() + .unwrap_err() + .to_string() + .contains("must be a table") + ); + assert!( + document + .remove_dynamic_config(0) + .unwrap_err() + .to_string() + .contains("must be a table") + ); + + std::fs::write(&path, "[[plugins.dynamic]]\nconfig = {}\n").unwrap(); + let document = PluginConfigDocument::read(&path).unwrap(); + assert!( + document + .dynamic_entries() + .unwrap_err() + .to_string() + .contains("manifest must be a string") + ); + + std::fs::write( + &path, + "[[plugins.dynamic]]\nmanifest = 'plugin.toml'\nconfig = 'invalid'\n", + ) + .unwrap(); + let document = PluginConfigDocument::read(&path).unwrap(); + assert!( + document + .dynamic_entries() + .unwrap_err() + .to_string() + .contains("config must be a table") + ); +} + +#[test] +fn dynamic_config_patching_covers_add_replace_remove_and_object_diff() { + let temp = tempfile::tempdir().unwrap(); + let path = temp.path().join("plugins.toml"); + std::fs::write(&path, "[[plugins.dynamic]]\nmanifest = 'plugin.toml'\n").unwrap(); + let mut document = PluginConfigDocument::read(&path).unwrap(); + + let first = Map::from_iter([("first".into(), json!(1))]); + document + .patch_dynamic_config(0, None, Some(first.clone())) + .unwrap(); + document.remove_dynamic_config(0).unwrap(); + document + .patch_dynamic_config(0, Some(&Map::new()), Some(first.clone())) + .unwrap(); + + let updated = Map::from_iter([("first".into(), json!(2)), ("second".into(), json!(true))]); + document + .patch_dynamic_config(0, Some(&first), Some(updated.clone())) + .unwrap(); + assert_eq!(document.dynamic_entries().unwrap()[0].config, Some(updated)); + document.patch_dynamic_config(0, None, None).unwrap(); + assert_eq!(document.dynamic_entries().unwrap()[0].config, None); +} + +#[test] +fn remove_dynamic_plugin_reference_covers_absent_container_paths() { + let temp = tempfile::tempdir().unwrap(); + let path = temp.path().join("plugins.toml"); + assert!(!remove_dynamic_plugin_reference(&path, "acme.missing", None).unwrap()); + + std::fs::write(&path, "version = 1\n").unwrap(); + assert!(!remove_dynamic_plugin_reference(&path, "acme.missing", None).unwrap()); + + std::fs::write(&path, "[plugins]\npolicy = {}\n").unwrap(); + assert!(!remove_dynamic_plugin_reference(&path, "acme.missing", None).unwrap()); +} + #[test] fn write_plugin_config_prunes_defaults_and_round_trips() { let temp = tempfile::tempdir().unwrap(); diff --git a/crates/cli/tests/coverage/shared/server_tests.rs b/crates/cli/tests/coverage/shared/server_tests.rs index b51cc44bf..e25c048c7 100644 --- a/crates/cli/tests/coverage/shared/server_tests.rs +++ b/crates/cli/tests/coverage/shared/server_tests.rs @@ -673,6 +673,36 @@ fn readiness_file_is_published_atomically_with_gateway_identity() { assert!(!path.with_extension("json.tmp").exists()); } +#[tokio::test] +async fn bind_listener_reports_an_actionable_address_conflict() { + let occupied = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let address = occupied.local_addr().unwrap(); + let error = bind_listener(address).await.unwrap_err(); + let message = error.to_string(); + assert!(message.contains("port is already in use")); + assert!(message.contains("ephemeral port")); +} + +#[test] +fn readiness_file_reports_write_and_publish_failures() { + let temp = tempfile::tempdir().unwrap(); + let address = "127.0.0.1:4040".parse().unwrap(); + + let missing_parent = temp.path().join("missing").join("ready.json"); + let error = write_ready_file(&missing_parent, address, "write-failure").unwrap_err(); + assert!(error.to_string().contains("failed to write readiness file")); + + let directory_target = temp.path().join("ready.json"); + std::fs::create_dir(&directory_target).unwrap(); + let error = write_ready_file(&directory_target, address, "publish-failure").unwrap_err(); + assert!( + error + .to_string() + .contains("failed to publish readiness file") + ); + assert!(!temp.path().join("ready.json.tmp").exists()); +} + #[tokio::test] async fn serve_listener_honors_plugin_idle_timeout_env() { let _guard = PLUGIN_CONFIG_TEST_LOCK.lock().await; @@ -2116,6 +2146,126 @@ async fn static_only_cli_configuration_keeps_the_legacy_lifecycle() { let _ = deregister_plugin(GENERIC_TEST_PLUGIN_KIND); } +#[test] +fn plugin_component_setup_errors_render_every_diagnostic_variant() { + let adaptive = PluginComponentSetupError::Adaptive("adaptive failure".into()); + assert_eq!(adaptive.check_name(), "Adaptive plugin"); + assert_eq!( + adaptive.diagnostic_details(), + "registration failed: adaptive failure" + ); + assert_eq!( + adaptive.to_string(), + "adaptive plugin registration failed: adaptive failure" + ); + + let pii = PluginComponentSetupError::PiiRedaction("pii failure".into()); + assert_eq!(pii.check_name(), "PII redaction plugin"); + assert_eq!(pii.diagnostic_details(), "registration failed: pii failure"); + assert_eq!( + pii.to_string(), + "PII redaction plugin registration failed: pii failure" + ); + + #[cfg(feature = "switchyard")] + { + let switchyard = PluginComponentSetupError::Switchyard("registration".into()); + assert_eq!(switchyard.check_name(), "Switchyard plugin"); + assert!(switchyard.to_string().contains("registration failed")); + + let atof = PluginComponentSetupError::SwitchyardAtof("atof ordering".into()); + assert_eq!(atof.check_name(), "Switchyard ATOF"); + assert_eq!(atof.diagnostic_details(), "atof ordering"); + assert!(atof.to_string().contains("ATOF validation failed")); + + let cache = PluginComponentSetupError::SwitchyardResponseCache("cache ordering".into()); + assert_eq!(cache.check_name(), "Switchyard response cache"); + assert_eq!(cache.diagnostic_details(), "cache ordering"); + assert!( + cache + .to_string() + .contains("response-cache validation failed") + ); + } +} + +fn dynamic_component_without_manifest( + plugin_id: &str, + kind: DynamicPluginKind, +) -> ActiveDynamicPluginComponent { + ActiveDynamicPluginComponent { + plugin_id: plugin_id.into(), + kind, + lifecycle_generation: 1, + manifest_ref: None, + environment_ref: None, + config: Map::new(), + activation_snapshot: None, + } +} + +#[tokio::test] +async fn plugin_activation_covers_empty_invalid_and_missing_manifest_paths() { + let _guard = PLUGIN_CONFIG_TEST_LOCK.lock().await; + let inactive = PluginActivation::initialize(None, Vec::new()) + .await + .unwrap(); + assert!(!inactive.active); + inactive.clear().unwrap(); + + let invalid = PluginActivation::initialize( + Some(json!("not a plugin config")), + vec![dynamic_component_without_manifest( + "acme.invalid-config", + DynamicPluginKind::Worker, + )], + ) + .await + .err() + .expect("invalid config should fail activation"); + assert!(invalid.to_string().contains("invalid plugin config")); + + let native = PluginActivation::initialize( + None, + vec![dynamic_component_without_manifest( + "acme.native-missing", + DynamicPluginKind::RustDynamic, + )], + ) + .await + .err() + .expect("native plugin without a manifest should fail activation"); + assert!(native.to_string().contains("native dynamic plugin")); + + let worker = PluginActivation::initialize( + None, + vec![dynamic_component_without_manifest( + "acme.worker-missing", + DynamicPluginKind::Worker, + )], + ) + .await + .err() + .expect("worker plugin without a manifest should fail activation"); + assert!(worker.to_string().contains("worker dynamic plugin")); +} + +#[tokio::test] +async fn shutdown_future_helpers_cover_receiver_combinations() { + let (shutdown_tx, shutdown_rx) = oneshot::channel(); + let shutdown = server_shutdown_future(Some(ShutdownMode::Receiver(shutdown_rx)), None).unwrap(); + shutdown_tx.send(()).unwrap(); + shutdown.await; + + let (bootstrap_tx, bootstrap_rx) = oneshot::channel(); + let shutdown = combine_shutdown_futures(None, Some(bootstrap_rx)).unwrap(); + bootstrap_tx.send(()).unwrap(); + shutdown.await; + + let ready: ShutdownFuture = Box::pin(async {}); + combine_shutdown_futures(Some(ready), None).unwrap().await; +} + #[cfg(feature = "switchyard")] #[test] fn switchyard_must_run_before_response_cache() { diff --git a/crates/core/src/api/runtime/scope_stack.rs b/crates/core/src/api/runtime/scope_stack.rs index 3712b4e6a..7ae5c12ea 100644 --- a/crates/core/src/api/runtime/scope_stack.rs +++ b/crates/core/src/api/runtime/scope_stack.rs @@ -452,7 +452,7 @@ pub fn create_scope_stack_from_propagation( /// # Examples /// /// ```no_run -/// # async fn example() -> nemo_relay::Result<()> { +/// # async fn example() -> nemo_relay::error::Result<()> { /// use nemo_relay::api::runtime::{TASK_SCOPE_STACK, fork_scope_stack}; /// /// let stack = fork_scope_stack()?; diff --git a/crates/core/src/logging/rotation.rs b/crates/core/src/logging/rotation.rs index 3314377c1..fa55780bb 100644 --- a/crates/core/src/logging/rotation.rs +++ b/crates/core/src/logging/rotation.rs @@ -146,3 +146,7 @@ pub(crate) fn rotated_log_path(base_path: &Path, index: usize) -> PathBuf { } base_path.with_file_name(file_name) } + +#[cfg(test)] +#[path = "../../tests/coverage/logging_rotation_tests.rs"] +mod tests; diff --git a/crates/core/src/observability/otel_genai.rs b/crates/core/src/observability/otel_genai.rs index aab9698e2..481dbf4e0 100644 --- a/crates/core/src/observability/otel_genai.rs +++ b/crates/core/src/observability/otel_genai.rs @@ -434,59 +434,31 @@ fn push_error_attributes(attributes: &mut Vec, event: &Event) { } fn scalar_string(event: &Event, keys: &[&str]) -> Option { - if let Some(profile) = event.category_profile() { - for key in keys { - if let Some(value) = profile.extra.get(*key) { - if let Some(value) = value.as_str() { - return Some(value.to_string()); - } - if value.is_number() || value.is_boolean() { - return Some(value.to_string()); - } - } - } - } - for object in event_objects(event) { - for key in keys { - if let Some(value) = object_value(object, key) { - if let Some(value) = value.as_str() { - return Some(value.to_string()); - } - if value.is_number() || value.is_boolean() { - return Some(value.to_string()); - } - } - } - } - None + find_scalar(event, keys, |value| { + value + .as_str() + .map(str::to_string) + .or_else(|| (value.is_number() || value.is_boolean()).then(|| value.to_string())) + }) } fn scalar_i64(event: &Event, keys: &[&str]) -> Option { - if let Some(profile) = event.category_profile() { - for key in keys { - if let Some(value) = profile.extra.get(*key) { - if let Some(value) = value.as_i64() { - return Some(value); - } - if let Some(value) = value.as_u64().and_then(to_i64) { - return Some(value); - } - } - } - } - for object in event_objects(event) { - for key in keys { - if let Some(value) = object_value(object, key) { - if let Some(value) = value.as_i64() { - return Some(value); - } - if let Some(value) = value.as_u64().and_then(to_i64) { - return Some(value); - } - } - } - } - None + find_scalar(event, keys, |value| { + value.as_i64().or_else(|| value.as_u64().and_then(to_i64)) + }) +} + +fn find_scalar(event: &Event, keys: &[&str], convert: impl Fn(&Json) -> Option) -> Option { + let profile_value = event.category_profile().and_then(|profile| { + keys.iter() + .find_map(|key| profile.extra.get(*key).and_then(&convert)) + }); + profile_value.or_else(|| { + event_objects(event).into_iter().find_map(|object| { + keys.iter() + .find_map(|key| object_value(object, key).and_then(&convert)) + }) + }) } fn object_value<'a>(object: &'a Map, key: &str) -> Option<&'a Json> { diff --git a/crates/core/src/observability/plugin_component.rs b/crates/core/src/observability/plugin_component.rs index 0dc46285e..cdf00e726 100644 --- a/crates/core/src/observability/plugin_component.rs +++ b/crates/core/src/observability/plugin_component.rs @@ -1807,22 +1807,7 @@ fn build_otel_config( ))); } }; - for (header, variable) in §ion.header_env { - if variable.trim().is_empty() || variable.trim() != variable { - return Err(PluginError::InvalidConfig(format!( - "OpenTelemetry endpoints[{index}].header_env.{header} must name a nonblank environment variable without surrounding whitespace" - ))); - } - if section - .headers - .keys() - .any(|configured| configured.eq_ignore_ascii_case(header)) - { - return Err(PluginError::InvalidConfig(format!( - "OpenTelemetry endpoints[{index}] header {header:?} cannot appear in both headers and header_env" - ))); - } - } + validate_otel_header_env(index, §ion)?; let mut config = CoreOpenTelemetryConfig::new(section.otel_type, section.endpoint) .with_transport(transport) .with_service_name(section.service_name) @@ -1840,7 +1825,42 @@ fn build_otel_config( for (key, value) in section.headers { config = config.with_header(key, value); } - for (key, variable) in section.header_env { + config = apply_otel_environment_headers(config, index, section.header_env)?; + for (key, value) in section.resource_attributes { + config = config.with_resource_attribute(key, value); + } + Ok(config) +} + +fn validate_otel_header_env( + index: usize, + section: &OpenTelemetryEndpointConfig, +) -> PluginResult<()> { + for (header, variable) in §ion.header_env { + if variable.trim().is_empty() || variable.trim() != variable { + return Err(PluginError::InvalidConfig(format!( + "OpenTelemetry endpoints[{index}].header_env.{header} must name a nonblank environment variable without surrounding whitespace" + ))); + } + if section + .headers + .keys() + .any(|configured| configured.eq_ignore_ascii_case(header)) + { + return Err(PluginError::InvalidConfig(format!( + "OpenTelemetry endpoints[{index}] header {header:?} cannot appear in both headers and header_env" + ))); + } + } + Ok(()) +} + +fn apply_otel_environment_headers( + mut config: CoreOpenTelemetryConfig, + index: usize, + header_env: HashMap, +) -> PluginResult { + for (key, variable) in header_env { let value = std::env::var(&variable).map_err(|error| { PluginError::InvalidConfig(format!( "OpenTelemetry endpoints[{index}].header_env.{key} could not read environment variable {variable:?}: {error}" @@ -1853,9 +1873,6 @@ fn build_otel_config( } config = config.with_header(key, value); } - for (key, value) in section.resource_attributes { - config = config.with_resource_attribute(key, value); - } Ok(config) } diff --git a/crates/core/src/plugin.rs b/crates/core/src/plugin.rs index 97bab10bb..b46d4c8d0 100644 --- a/crates/core/src/plugin.rs +++ b/crates/core/src/plugin.rs @@ -1634,33 +1634,54 @@ fn plugin_config_overlay_value(config: &PluginConfig) -> Result { root.remove("version"); } - if let Some(Json::Object(policy)) = root.get_mut("policy") { - let defaults = ConfigPolicy::default(); - if config.policy.unknown_component == defaults.unknown_component { - policy.remove("unknown_component"); - } - if config.policy.unknown_field == defaults.unknown_field { - policy.remove("unknown_field"); - } - if config.policy.unsupported_value == defaults.unsupported_value { - policy.remove("unsupported_value"); - } - if policy.is_empty() { - root.remove("policy"); + remove_default_policy_overlay(root, &config.policy); + remove_default_component_enabled_overlays(root, &config.components); + + Ok(overlay) +} + +fn remove_default_policy_overlay(root: &mut Map, config: &ConfigPolicy) { + let Some(Json::Object(policy)) = root.get_mut("policy") else { + return; + }; + let defaults = ConfigPolicy::default(); + for (field, is_default) in [ + ( + "unknown_component", + config.unknown_component == defaults.unknown_component, + ), + ( + "unknown_field", + config.unknown_field == defaults.unknown_field, + ), + ( + "unsupported_value", + config.unsupported_value == defaults.unsupported_value, + ), + ] { + if is_default { + policy.remove(field); } } + if policy.is_empty() { + root.remove("policy"); + } +} - if let Some(Json::Array(components)) = root.get_mut("components") { - for (component, typed) in components.iter_mut().zip(&config.components) { - if typed.enabled == default_enabled() - && let Json::Object(component) = component - { - component.remove("enabled"); - } +fn remove_default_component_enabled_overlays( + root: &mut Map, + configured: &[PluginComponentSpec], +) { + let Some(Json::Array(components)) = root.get_mut("components") else { + return; + }; + for (component, typed) in components.iter_mut().zip(configured) { + if typed.enabled == default_enabled() + && let Json::Object(component) = component + { + component.remove("enabled"); } } - - Ok(overlay) } /// Resolves the default `plugins.toml` layering into one JSON document, or an diff --git a/crates/core/src/plugin/dynamic/native.rs b/crates/core/src/plugin/dynamic/native.rs index d6fde6414..3b0124b5b 100644 --- a/crates/core/src/plugin/dynamic/native.rs +++ b/crates/core/src/plugin/dynamic/native.rs @@ -2218,80 +2218,11 @@ unsafe extern "C" fn native_async_next_invoke_stream( return; } continuation_context - .run(async move { - let mut callback_guard = callback_guard; - let result = AssertUnwindSafe(async { - match next_fn(request).await { - Ok(mut stream) => { - while let Some(item) = stream.next().await { - match item { - Ok(chunk) => { - if let Some(chunk) = native_string_from_json(&chunk) { - let keep_going = unsafe { - cb( - user_data as *mut c_void, - chunk, - ptr::null(), - false, - ) - }; - unsafe { - native_string_free(chunk); - } - if !keep_going { - callback_guard.finish(); - return; - } - } else { - break; - } - } - Err(error) => { - if let Some(message) = - native_string_from_str(&error.to_string()) - { - unsafe { - let _ = cb( - user_data as *mut c_void, - ptr::null(), - message, - false, - ); - native_string_free(message); - } - callback_guard.finish(); - } - return; - } - } - } - unsafe { - let _ = - cb(user_data as *mut c_void, ptr::null(), ptr::null(), true); - } - callback_guard.finish(); - } - Err(error) => { - if let Some(message) = native_string_from_str(&error.to_string()) { - unsafe { - let _ = - cb(user_data as *mut c_void, ptr::null(), message, false); - native_string_free(message); - } - callback_guard.finish(); - } - } - } - }) - .catch_unwind() - .await; - if let Err(payload) = result { - callback_guard.fail(&format!( - "native async stream continuation panicked: {}", - panic_payload_message(payload.as_ref()) - )); - } - }) + .run(deliver_native_async_next_stream( + next_fn, + request, + callback_guard, + )) .await; output_stream_for_cleanup .downstream_aborts @@ -2305,6 +2236,79 @@ unsafe extern "C" fn native_async_next_invoke_stream( NemoRelayStatus::Ok } +async fn deliver_native_async_next_stream( + next_fn: LlmStreamExecutionNextFn, + request: LlmRequest, + mut callback_guard: NativeAsyncStreamCallbackGuard, +) { + let result = AssertUnwindSafe(async { + match next_fn(request).await { + Ok(stream) => forward_native_async_next_stream(stream, &mut callback_guard).await, + Err(error) => callback_guard.fail(&error.to_string()), + } + }) + .catch_unwind() + .await; + if let Err(payload) = result { + callback_guard.fail(&format!( + "native async stream continuation panicked: {}", + panic_payload_message(payload.as_ref()) + )); + } +} + +async fn forward_native_async_next_stream( + stream: LlmJsonStream, + callback_guard: &mut NativeAsyncStreamCallbackGuard, +) { + forward_native_async_next_stream_with(stream, callback_guard, native_string_from_json).await; +} + +async fn forward_native_async_next_stream_with( + mut stream: LlmJsonStream, + callback_guard: &mut NativeAsyncStreamCallbackGuard, + to_native_string: impl Fn(&Json) -> Option<*mut NemoRelayNativeString>, +) { + while let Some(item) = stream.next().await { + match item { + Ok(chunk) => { + let Some(chunk) = to_native_string(&chunk) else { + callback_guard.fail( + "failed to serialize or allocate native async stream continuation chunk", + ); + return; + }; + let keep_going = unsafe { + (callback_guard.cb)( + callback_guard.user_data as *mut c_void, + chunk, + ptr::null(), + false, + ) + }; + unsafe { native_string_free(chunk) }; + if !keep_going { + callback_guard.finish(); + return; + } + } + Err(error) => { + callback_guard.fail(&error.to_string()); + return; + } + } + } + unsafe { + let _ = (callback_guard.cb)( + callback_guard.user_data as *mut c_void, + ptr::null(), + ptr::null(), + true, + ); + } + callback_guard.finish(); +} + fn wrap_native_async_tool_json( instance: Arc, cb: NemoRelayNativeAsyncMiddlewareCb, diff --git a/crates/core/src/plugins/nemo_guardrails/python.rs b/crates/core/src/plugins/nemo_guardrails/python.rs index 04b99c01c..504a798cb 100644 --- a/crates/core/src/plugins/nemo_guardrails/python.rs +++ b/crates/core/src/plugins/nemo_guardrails/python.rs @@ -1107,7 +1107,10 @@ impl LlmStreamInner for GuardedProviderStream { } } -#[allow(clippy::too_many_arguments)] +#[allow( + clippy::too_many_arguments, + reason = "stream cancellation, monitoring, delivery, and cleanup must remain ordered in one coordinator" +)] async fn forward_guarded_provider_stream( mut provider_stream: LlmJsonStream, codec: LocalGuardrailsCodec, @@ -1130,51 +1133,132 @@ async fn forward_guarded_provider_stream( let Some(item) = item else { break; }; - let chunk = match item { - Ok(chunk) => chunk, - Err(err) => { - let _ = chunk_tx.send(Err(err)).await; - let _ = text_tx.send(None).await; - let _ = monitor.take().expect("monitor available").await; - break; - } + let Some(chunk) = + receive_guarded_provider_chunk(item, &text_tx, &chunk_tx, &mut monitor).await + else { + break; }; - if let Some(message) = blocked_message(&blocked) { - let _ = chunk_tx.send(Err(streaming_output_blocked(message))).await; - let _ = text_tx.send(None).await; - let _ = monitor.take().expect("monitor available").await; + if stop_blocked_provider_stream(&text_tx, &chunk_tx, &blocked, &mut monitor).await { break; } - if let Some(text) = extract_stream_text(codec, &chunk) - && text_tx.send(Some(text)).await.is_err() + if !forward_guarded_stream_text(codec, &chunk, &text_tx, &chunk_tx, &blocked, &mut monitor) + .await { - send_stream_monitor_error( - monitor.take().expect("monitor available"), - &chunk_tx, - &blocked, - ) - .await; break; } - let sent = tokio::select! { - _ = cancel.changed() => break, - sent = chunk_tx.send(Ok(chunk)) => sent, - }; - if sent.is_err() { + if !send_guarded_provider_chunk(chunk, &text_tx, &chunk_tx, &mut monitor, &mut cancel).await + { + break; + } + } + finish_guarded_provider_stream( + &mut provider_stream, + &text_tx, + &chunk_tx, + &blocked, + &mut monitor, + &cancel, + &closed, + ) + .await; +} + +async fn receive_guarded_provider_chunk( + item: FlowResult, + text_tx: &mpsc::Sender>, + chunk_tx: &mpsc::Sender>, + monitor: &mut Option>>, +) -> Option { + match item { + Ok(chunk) => Some(chunk), + Err(err) => { + let _ = chunk_tx.send(Err(err)).await; let _ = text_tx.send(None).await; let _ = monitor.take().expect("monitor available").await; - break; + None } } +} + +async fn stop_blocked_provider_stream( + text_tx: &mpsc::Sender>, + chunk_tx: &mpsc::Sender>, + blocked: &Arc>>, + monitor: &mut Option>>, +) -> bool { + let Some(message) = blocked_message(blocked) else { + return false; + }; + let _ = chunk_tx.send(Err(streaming_output_blocked(message))).await; + let _ = text_tx.send(None).await; + let _ = monitor.take().expect("monitor available").await; + true +} + +async fn forward_guarded_stream_text( + codec: LocalGuardrailsCodec, + chunk: &Json, + text_tx: &mpsc::Sender>, + chunk_tx: &mpsc::Sender>, + blocked: &Arc>>, + monitor: &mut Option>>, +) -> bool { + let Some(text) = extract_stream_text(codec, chunk) else { + return true; + }; + if text_tx.send(Some(text)).await.is_ok() { + return true; + } + send_stream_monitor_error( + monitor.take().expect("monitor available"), + chunk_tx, + blocked, + ) + .await; + false +} + +async fn send_guarded_provider_chunk( + chunk: Json, + text_tx: &mpsc::Sender>, + chunk_tx: &mpsc::Sender>, + monitor: &mut Option>>, + cancel: &mut watch::Receiver, +) -> bool { + let sent = tokio::select! { + _ = cancel.changed() => return false, + sent = chunk_tx.send(Ok(chunk)) => sent, + }; + if sent.is_ok() { + return true; + } + let _ = text_tx.send(None).await; + let _ = monitor.take().expect("monitor available").await; + false +} + +#[allow( + clippy::too_many_arguments, + reason = "stream cleanup needs all channels and lifecycle handles" +)] +async fn finish_guarded_provider_stream( + provider_stream: &mut LlmJsonStream, + text_tx: &mpsc::Sender>, + chunk_tx: &mpsc::Sender>, + blocked: &Arc>>, + monitor: &mut Option>>, + cancel: &watch::Receiver, + closed: &watch::Sender>>, +) { let _ = text_tx.send(None).await; if *cancel.borrow() { if let Some(monitor) = monitor.take() { monitor.abort(); } } else if let Some(monitor) = monitor.take() { - let _ = send_stream_monitor_error(monitor, &chunk_tx, &blocked).await; + let _ = send_stream_monitor_error(monitor, chunk_tx, blocked).await; } closed.send_replace(Some(provider_stream.close().await)); } diff --git a/crates/core/src/stream.rs b/crates/core/src/stream.rs index dd99fde19..4679b2751 100644 --- a/crates/core/src/stream.rs +++ b/crates/core/src/stream.rs @@ -36,7 +36,7 @@ use crate::api::event::{BaseEvent, MarkEvent}; use crate::api::llm::LlmHandle; use crate::api::llm::emit_reserved_optimization_marks; use crate::api::optimization::finalize_optimization_summary; -use crate::api::runtime::LlmSanitizeResponseContext; +use crate::api::registry::Guardrail; use crate::api::runtime::NemoRelayContextState; use crate::api::runtime::global_context; use crate::api::runtime::subscriber_dispatcher; @@ -44,6 +44,7 @@ use crate::api::runtime::{ EventSubscriberFn, LlmJsonStream, LlmStreamInner, ScopeStackHandle, TASK_SCOPE_STACK, current_scope_stack, }; +use crate::api::runtime::{LlmSanitizeResponseContext, LlmSanitizeResponseFn}; use crate::api::shared::{ metadata_with_otel_error, metadata_with_otel_status, snapshot_event_sanitizers, }; @@ -249,31 +250,8 @@ impl LlmStreamWrapper { aggregated }; - let (entries, sanitizer_snapshot_failed) = match self.scope_stack.read() { - Ok(scope_guard) => { - let scope_locals = scope_guard - .collect_scope_local_registries(|r| &r.llm_sanitize_response_guardrails); - match global_context().read() { - Ok(state) => (state.llm_sanitize_response_entries(&scope_locals), false), - Err(error) => { - log::error!( - target: "nemo_relay.runtime", - event = "stream_end_sanitizer_snapshot_failed"; - "LLM stream END sanitizer snapshot failed; omitting the observability payload: {error}" - ); - (Vec::new(), true) - } - } - } - Err(error) => { - log::error!( - target: "nemo_relay.runtime", - event = "stream_end_sanitizer_snapshot_failed"; - "LLM stream END sanitizer snapshot failed; omitting the observability payload: {error}" - ); - (Vec::new(), true) - } - }; + let (entries, sanitizer_snapshot_failed) = + snapshot_stream_end_sanitizers(&self.scope_stack); let handle = self.handle.clone(); let scope_stack = self.scope_stack.clone(); let finalization_scope_stack = scope_stack.clone(); @@ -412,6 +390,30 @@ impl LlmStreamWrapper { } } +fn snapshot_stream_end_sanitizers( + scope_stack: &ScopeStackHandle, +) -> (Vec>, bool) { + let entries = scope_stack.read().ok().and_then(|scope_guard| { + let scope_locals = scope_guard + .collect_scope_local_registries(|registry| ®istry.llm_sanitize_response_guardrails); + global_context() + .read() + .ok() + .map(|state| state.llm_sanitize_response_entries(&scope_locals)) + }); + match entries { + Some(entries) => (entries, false), + None => { + log::error!( + target: "nemo_relay.runtime", + event = "stream_end_sanitizer_snapshot_failed"; + "LLM stream END sanitizer snapshot failed; omitting the observability payload" + ); + (Vec::new(), true) + } + } +} + impl Stream for LlmStreamWrapper { type Item = Result; diff --git a/crates/core/tests/coverage/logging_rotation_tests.rs b/crates/core/tests/coverage/logging_rotation_tests.rs new file mode 100644 index 000000000..1b9de3067 --- /dev/null +++ b/crates/core/tests/coverage/logging_rotation_tests.rs @@ -0,0 +1,57 @@ +// SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +use super::*; + +#[test] +fn rotating_writer_reports_missing_file_state() { + let temp = tempfile::tempdir().unwrap(); + let path = temp.path().join("events.log"); + let mut writer = SizeRotatingFileWriter::new(path, 1, 1).unwrap(); + writer.file = None; + writer.current_size = 1; + + let rotate_error = writer + .write(b"next") + .expect_err("rotation requires an open active file"); + assert_eq!(rotate_error.kind(), io::ErrorKind::Other); + + writer.current_size = 0; + let write_error = writer + .write(b"next") + .expect_err("writing requires an open active file"); + assert_eq!(write_error.kind(), io::ErrorKind::Other); + + let flush_error = writer + .flush() + .expect_err("flushing requires an open active file"); + assert_eq!(flush_error.kind(), io::ErrorKind::Other); +} + +#[test] +fn rotation_helpers_handle_relative_paths_and_missing_generations() { + let temp = tempfile::tempdir().unwrap(); + let base = temp.path().join("relay"); + assert_eq!(rotated_log_path(&base, 3), temp.path().join("relay.3")); + + rotate_files(&base, 3).unwrap(); + assert!(!rotated_log_path(&base, 1).exists()); + + create_parent_directory(Path::new("relay.log")).unwrap(); +} + +#[test] +fn failed_rotation_reopens_the_active_file() { + let temp = tempfile::tempdir().unwrap(); + let base = temp.path().join("relay.log"); + let backup = rotated_log_path(&base, 1); + fs::create_dir(&backup).unwrap(); + + let mut writer = SizeRotatingFileWriter::new(base.clone(), 1, 1).unwrap(); + writer.write_all(b"first").unwrap(); + let _error = writer + .write(b"second") + .expect_err("an existing backup directory prevents rotation"); + assert!(writer.file.is_some()); + assert_eq!(writer.current_size, fs::metadata(base).unwrap().len()); +} diff --git a/crates/core/tests/coverage/logging_sink_tests.rs b/crates/core/tests/coverage/logging_sink_tests.rs index 044ce2b47..e0c4e8f25 100644 --- a/crates/core/tests/coverage/logging_sink_tests.rs +++ b/crates/core/tests/coverage/logging_sink_tests.rs @@ -2,10 +2,15 @@ // SPDX-License-Identifier: Apache-2.0 use super::{ - DROP_REPORT_INTERVAL_MILLIS, DropNoticeRateLimiter, dropped_record_error_handler, - log_level_filter, now_millis, spdlog_level, stderr_error_handler, + DROP_REPORT_INTERVAL_MILLIS, DropNoticeRateLimiter, build_logger, dropped_record_error_handler, + log_level_filter, logging_path_identity, normalize_path_components, now_millis, + reserved_sink_paths, resolve_log_path, spdlog_level, stderr_error_handler, }; -use crate::logging::LogLevel; +use crate::logging::{ + FileLogRotationConfig, FileLogSinkConfig, LogLevel, LogSinkConfig, LoggingConfig, + MAX_FILE_SINK_QUEUE_ENTRIES, +}; +use std::path::{Path, PathBuf}; #[test] fn drop_notice_rate_limiter_reports_immediately_then_once_per_interval() { @@ -33,3 +38,114 @@ fn sink_helpers_cover_boundary_levels_time_and_emergency_handlers() { "expected test error", ))); } + +#[test] +fn sink_path_helpers_cover_rotation_and_normalization_edges() { + assert!(resolve_log_path(Path::new("")).is_err()); + assert_eq!( + normalize_path_components(Path::new("alpha/./beta/../gamma")), + PathBuf::from("alpha/gamma") + ); + assert_eq!( + logging_path_identity(Path::new("relay.log")), + PathBuf::from("relay.log") + ); + + let temp = tempfile::tempdir().unwrap(); + let base = temp.path().join("relay.log"); + std::fs::write(&base, "existing").unwrap(); + assert_eq!( + logging_path_identity(&base), + std::fs::canonicalize(&base).unwrap() + ); + + let rotation = FileLogRotationConfig::new(1_024, 2).unwrap(); + let paths = reserved_sink_paths(&base, Some(rotation)); + assert_eq!(paths.len(), 3); + assert_eq!(paths[0], base); + assert!(paths[1].ends_with("relay.1.log")); + assert!(paths[2].ends_with("relay.2.log")); + assert_eq!(reserved_sink_paths(&paths[0], None), vec![paths[0].clone()]); +} + +#[test] +fn sink_level_helpers_cover_all_intermediate_levels() { + for (level, spdlog_level_expected, log_level_expected) in [ + (LogLevel::Warn, spdlog::Level::Warn, log::LevelFilter::Warn), + (LogLevel::Info, spdlog::Level::Info, log::LevelFilter::Info), + ( + LogLevel::Debug, + spdlog::Level::Debug, + log::LevelFilter::Debug, + ), + ] { + assert_eq!(spdlog_level(level), spdlog_level_expected); + assert_eq!(log_level_filter(level), log_level_expected); + } +} + +fn file_sink(path: PathBuf) -> FileLogSinkConfig { + FileLogSinkConfig { + path, + ..FileLogSinkConfig::default() + } +} + +fn build_logger_error(config: &LoggingConfig) -> String { + match build_logger(config, "root".into()) { + Ok(_) => panic!("expected logger construction to fail"), + Err(error) => error.to_string(), + } +} + +#[test] +fn logger_builder_rejects_duplicate_reserved_and_invalid_queue_sinks() { + let temp = tempfile::tempdir().unwrap(); + let path = temp.path().join("relay.log"); + + let mut config = LoggingConfig { + sinks: vec![ + LogSinkConfig::File(file_sink(path.clone())), + LogSinkConfig::File(file_sink(path.clone())), + ], + ..LoggingConfig::default() + }; + assert!(build_logger_error(&config).contains("duplicate")); + + let mut rotating = file_sink(path.clone()); + rotating.rotation = Some(FileLogRotationConfig::new(1_024, 1).unwrap()); + config.sinks = vec![ + LogSinkConfig::File(rotating), + LogSinkConfig::File(file_sink(temp.path().join("relay.1.log"))), + ]; + assert!(build_logger_error(&config).contains("conflicts")); + + for (capacity, expected) in [ + (0, "must be greater than 0"), + (MAX_FILE_SINK_QUEUE_ENTRIES + 1, "exceeds maximum"), + ] { + let mut sink = file_sink(temp.path().join(format!("queue-{capacity}.log"))); + sink.queue_capacity = capacity; + config.sinks = vec![LogSinkConfig::File(sink)]; + assert!(build_logger_error(&config).contains(expected)); + } +} + +#[test] +fn logger_builder_reports_file_and_rotating_file_open_errors() { + let temp = tempfile::tempdir().unwrap(); + let blocked_parent = temp.path().join("not-a-directory"); + std::fs::write(&blocked_parent, "file").unwrap(); + let mut config = LoggingConfig { + sinks: vec![LogSinkConfig::File(file_sink( + blocked_parent.join("relay.log"), + ))], + ..LoggingConfig::default() + }; + assert!(build_logger_error(&config).contains("failed to open logging sink")); + + let mut rotating = file_sink(blocked_parent.join("rotating.log")); + rotating.rotation = Some(FileLogRotationConfig::new(1_024, 1).unwrap()); + config.sinks = vec![LogSinkConfig::File(rotating)]; + assert!(build_logger_error(&config).contains("failed to open rotating logging sink")); +} diff --git a/crates/core/tests/integration/api_surface_tests.rs b/crates/core/tests/integration/api_surface_tests.rs index 7b0f72ffc..e421ef436 100644 --- a/crates/core/tests/integration/api_surface_tests.rs +++ b/crates/core/tests/integration/api_surface_tests.rs @@ -894,6 +894,13 @@ fn test_global_registry_and_subscriber_wrappers_cover_success_and_duplicates() { reset_global(); setup_isolated_thread(); + assert_global_event_guardrail_registry(); + assert_global_tool_registry(); + assert_global_llm_registry(); + assert_global_subscriber_registry(); +} + +fn assert_global_event_guardrail_registry() { register_mark_sanitize_guardrail( "mark-sanitize", 1, @@ -929,7 +936,9 @@ fn test_global_registry_and_subscriber_wrappers_cover_success_and_duplicates() { .unwrap(); assert!(deregister_scope_sanitize_end_guardrail("scope-end-sanitize").unwrap()); assert!(!deregister_scope_sanitize_end_guardrail("scope-end-sanitize").unwrap()); +} +fn assert_global_tool_registry() { register_tool_sanitize_request_guardrail( "tool-sanitize-request", 1, @@ -980,7 +989,9 @@ fn test_global_registry_and_subscriber_wrappers_cover_success_and_duplicates() { ) .unwrap(); assert!(deregister_tool_execution_intercept("tool-execution").unwrap()); +} +fn assert_global_llm_registry() { register_llm_sanitize_request_guardrail( "llm-sanitize-request", 1, @@ -1039,7 +1050,9 @@ fn test_global_registry_and_subscriber_wrappers_cover_success_and_duplicates() { ) .unwrap(); assert!(deregister_llm_stream_execution_intercept("llm-stream").unwrap()); +} +fn assert_global_subscriber_registry() { register_subscriber("global-subscriber", Arc::new(|_event| {})).unwrap(); expect_already_exists( register_subscriber("global-subscriber", Arc::new(|_event| {})).unwrap_err(), @@ -1097,8 +1110,24 @@ fn test_scope_registry_and_subscriber_wrappers_cover_success_duplicates_and_miss ) .unwrap(); + assert_scope_event_guardrail_registry(&scope.uuid); + assert_scope_tool_registry(&scope.uuid); + assert_scope_llm_registry(&scope.uuid); + assert_scope_subscriber_registry(&scope.uuid); + + pop_scope( + nemo_relay::api::scope::PopScopeParams::builder() + .handle_uuid(&scope.uuid) + .build(), + ) + .unwrap(); + + assert_missing_scope_registry_errors(&scope.uuid); +} + +fn assert_scope_event_guardrail_registry(scope_uuid: &uuid::Uuid) { scope_register_mark_sanitize_guardrail( - &scope.uuid, + scope_uuid, "mark-sanitize", 1, Arc::new(|_, fields| Box::pin(async move { Ok(fields) })), @@ -1106,7 +1135,7 @@ fn test_scope_registry_and_subscriber_wrappers_cover_success_duplicates_and_miss .unwrap(); expect_already_exists( scope_register_mark_sanitize_guardrail( - &scope.uuid, + scope_uuid, "mark-sanitize", 1, Arc::new(|_, fields| Box::pin(async move { Ok(fields) })), @@ -1114,41 +1143,43 @@ fn test_scope_registry_and_subscriber_wrappers_cover_success_duplicates_and_miss .unwrap_err(), "mark-sanitize", ); - assert!(scope_deregister_mark_sanitize_guardrail(&scope.uuid, "mark-sanitize").unwrap()); - assert!(!scope_deregister_mark_sanitize_guardrail(&scope.uuid, "mark-sanitize").unwrap()); + assert!(scope_deregister_mark_sanitize_guardrail(scope_uuid, "mark-sanitize").unwrap()); + assert!(!scope_deregister_mark_sanitize_guardrail(scope_uuid, "mark-sanitize").unwrap()); scope_register_scope_sanitize_start_guardrail( - &scope.uuid, + scope_uuid, "scope-start-sanitize", 1, Arc::new(|_, fields| Box::pin(async move { Ok(fields) })), ) .unwrap(); assert!( - scope_deregister_scope_sanitize_start_guardrail(&scope.uuid, "scope-start-sanitize") + scope_deregister_scope_sanitize_start_guardrail(scope_uuid, "scope-start-sanitize") .unwrap() ); assert!( - !scope_deregister_scope_sanitize_start_guardrail(&scope.uuid, "scope-start-sanitize") + !scope_deregister_scope_sanitize_start_guardrail(scope_uuid, "scope-start-sanitize") .unwrap() ); scope_register_scope_sanitize_end_guardrail( - &scope.uuid, + scope_uuid, "scope-end-sanitize", 1, Arc::new(|_, fields| Box::pin(async move { Ok(fields) })), ) .unwrap(); assert!( - scope_deregister_scope_sanitize_end_guardrail(&scope.uuid, "scope-end-sanitize").unwrap() + scope_deregister_scope_sanitize_end_guardrail(scope_uuid, "scope-end-sanitize").unwrap() ); assert!( - !scope_deregister_scope_sanitize_end_guardrail(&scope.uuid, "scope-end-sanitize").unwrap() + !scope_deregister_scope_sanitize_end_guardrail(scope_uuid, "scope-end-sanitize").unwrap() ); +} +fn assert_scope_tool_registry(scope_uuid: &uuid::Uuid) { scope_register_tool_sanitize_request_guardrail( - &scope.uuid, + scope_uuid, "tool-sanitize-request", 1, Arc::new(|_name, args| Box::pin(async move { Ok(args) })), @@ -1156,7 +1187,7 @@ fn test_scope_registry_and_subscriber_wrappers_cover_success_duplicates_and_miss .unwrap(); expect_already_exists( scope_register_tool_sanitize_request_guardrail( - &scope.uuid, + scope_uuid, "tool-sanitize-request", 1, Arc::new(|_name, args| Box::pin(async move { Ok(args) })), @@ -1165,91 +1196,93 @@ fn test_scope_registry_and_subscriber_wrappers_cover_success_duplicates_and_miss "tool-sanitize-request", ); assert!( - scope_deregister_tool_sanitize_request_guardrail(&scope.uuid, "tool-sanitize-request") + scope_deregister_tool_sanitize_request_guardrail(scope_uuid, "tool-sanitize-request") .unwrap() ); scope_register_tool_sanitize_response_guardrail( - &scope.uuid, + scope_uuid, "tool-sanitize-response", 1, Arc::new(|_name, args| Box::pin(async move { Ok(args) })), ) .unwrap(); assert!( - scope_deregister_tool_sanitize_response_guardrail(&scope.uuid, "tool-sanitize-response") + scope_deregister_tool_sanitize_response_guardrail(scope_uuid, "tool-sanitize-response") .unwrap() ); scope_register_tool_conditional_execution_guardrail( - &scope.uuid, + scope_uuid, "tool-conditional", 1, Arc::new(|_name, _args| Box::pin(async { Ok(None) })), ) .unwrap(); assert!( - scope_deregister_tool_conditional_execution_guardrail(&scope.uuid, "tool-conditional") + scope_deregister_tool_conditional_execution_guardrail(scope_uuid, "tool-conditional") .unwrap() ); scope_register_tool_request_intercept( - &scope.uuid, + scope_uuid, "tool-request", 1, false, Arc::new(|_name, args| Box::pin(async move { Ok(args) })), ) .unwrap(); - assert!(scope_deregister_tool_request_intercept(&scope.uuid, "tool-request").unwrap()); + assert!(scope_deregister_tool_request_intercept(scope_uuid, "tool-request").unwrap()); scope_register_tool_execution_intercept( - &scope.uuid, + scope_uuid, "tool-execution", 1, Arc::new(|_name, args, _next| Box::pin(async move { Ok(args.into()) })), ) .unwrap(); - assert!(scope_deregister_tool_execution_intercept(&scope.uuid, "tool-execution").unwrap()); + assert!(scope_deregister_tool_execution_intercept(scope_uuid, "tool-execution").unwrap()); +} +fn assert_scope_llm_registry(scope_uuid: &uuid::Uuid) { scope_register_llm_sanitize_request_guardrail( - &scope.uuid, + scope_uuid, "llm-sanitize-request", 1, Arc::new(|request, _context| Box::pin(async move { Ok(Some(request)) })), ) .unwrap(); assert!( - scope_deregister_llm_sanitize_request_guardrail(&scope.uuid, "llm-sanitize-request") + scope_deregister_llm_sanitize_request_guardrail(scope_uuid, "llm-sanitize-request") .unwrap() ); scope_register_llm_sanitize_response_guardrail( - &scope.uuid, + scope_uuid, "llm-sanitize-response", 1, Arc::new(|response, _context| Box::pin(async move { Ok(Some(response)) })), ) .unwrap(); assert!( - scope_deregister_llm_sanitize_response_guardrail(&scope.uuid, "llm-sanitize-response") + scope_deregister_llm_sanitize_response_guardrail(scope_uuid, "llm-sanitize-response") .unwrap() ); scope_register_llm_conditional_execution_guardrail( - &scope.uuid, + scope_uuid, "llm-conditional", 1, Arc::new(|_request| Box::pin(async { Ok(None) })), ) .unwrap(); assert!( - scope_deregister_llm_conditional_execution_guardrail(&scope.uuid, "llm-conditional") + scope_deregister_llm_conditional_execution_guardrail(scope_uuid, "llm-conditional") .unwrap() ); scope_register_llm_request_intercept( - &scope.uuid, + scope_uuid, "llm-request", 1, false, @@ -1260,19 +1293,19 @@ fn test_scope_registry_and_subscriber_wrappers_cover_success_duplicates_and_miss }), ) .unwrap(); - assert!(scope_deregister_llm_request_intercept(&scope.uuid, "llm-request").unwrap()); + assert!(scope_deregister_llm_request_intercept(scope_uuid, "llm-request").unwrap()); scope_register_llm_execution_intercept( - &scope.uuid, + scope_uuid, "llm-execution", 1, Arc::new(|_name, request, _next| Box::pin(async move { Ok(request.content) })), ) .unwrap(); - assert!(scope_deregister_llm_execution_intercept(&scope.uuid, "llm-execution").unwrap()); + assert!(scope_deregister_llm_execution_intercept(scope_uuid, "llm-execution").unwrap()); scope_register_llm_stream_execution_intercept( - &scope.uuid, + scope_uuid, "llm-stream", 1, Arc::new(|_name, request, _next| { @@ -1284,27 +1317,24 @@ fn test_scope_registry_and_subscriber_wrappers_cover_success_duplicates_and_miss }), ) .unwrap(); - assert!(scope_deregister_llm_stream_execution_intercept(&scope.uuid, "llm-stream").unwrap()); + assert!(scope_deregister_llm_stream_execution_intercept(scope_uuid, "llm-stream").unwrap()); +} - scope_register_subscriber(&scope.uuid, "scope-subscriber", Arc::new(|_event| {})).unwrap(); +fn assert_scope_subscriber_registry(scope_uuid: &uuid::Uuid) { + scope_register_subscriber(scope_uuid, "scope-subscriber", Arc::new(|_event| {})).unwrap(); expect_already_exists( - scope_register_subscriber(&scope.uuid, "scope-subscriber", Arc::new(|_event| {})) + scope_register_subscriber(scope_uuid, "scope-subscriber", Arc::new(|_event| {})) .unwrap_err(), "scope-subscriber", ); - assert!(scope_deregister_subscriber(&scope.uuid, "scope-subscriber").unwrap()); - assert!(!scope_deregister_subscriber(&scope.uuid, "scope-subscriber").unwrap()); - - pop_scope( - nemo_relay::api::scope::PopScopeParams::builder() - .handle_uuid(&scope.uuid) - .build(), - ) - .unwrap(); + assert!(scope_deregister_subscriber(scope_uuid, "scope-subscriber").unwrap()); + assert!(!scope_deregister_subscriber(scope_uuid, "scope-subscriber").unwrap()); +} +fn assert_missing_scope_registry_errors(scope_uuid: &uuid::Uuid) { expect_not_found( scope_register_mark_sanitize_guardrail( - &scope.uuid, + scope_uuid, "missing-mark-sanitize", 1, Arc::new(|_, fields| Box::pin(async move { Ok(fields) })), @@ -1314,7 +1344,7 @@ fn test_scope_registry_and_subscriber_wrappers_cover_success_duplicates_and_miss ); expect_not_found( scope_register_scope_sanitize_start_guardrail( - &scope.uuid, + scope_uuid, "missing-scope-start-sanitize", 1, Arc::new(|_, fields| Box::pin(async move { Ok(fields) })), @@ -1324,7 +1354,7 @@ fn test_scope_registry_and_subscriber_wrappers_cover_success_duplicates_and_miss ); expect_not_found( scope_register_scope_sanitize_end_guardrail( - &scope.uuid, + scope_uuid, "missing-scope-end-sanitize", 1, Arc::new(|_, fields| Box::pin(async move { Ok(fields) })), @@ -1334,7 +1364,7 @@ fn test_scope_registry_and_subscriber_wrappers_cover_success_duplicates_and_miss ); expect_not_found( scope_register_tool_sanitize_request_guardrail( - &scope.uuid, + scope_uuid, "missing-tool-sanitize", 1, Arc::new(|_name, args| Box::pin(async move { Ok(args) })), @@ -1344,7 +1374,7 @@ fn test_scope_registry_and_subscriber_wrappers_cover_success_duplicates_and_miss ); expect_not_found( scope_register_tool_request_intercept( - &scope.uuid, + scope_uuid, "missing-tool-request", 1, false, @@ -1355,7 +1385,7 @@ fn test_scope_registry_and_subscriber_wrappers_cover_success_duplicates_and_miss ); expect_not_found( scope_register_tool_execution_intercept( - &scope.uuid, + scope_uuid, "missing-tool-exec", 1, Arc::new(|_name, args, _next| Box::pin(async move { Ok(args.into()) })), @@ -1364,7 +1394,7 @@ fn test_scope_registry_and_subscriber_wrappers_cover_success_duplicates_and_miss "scope", ); expect_not_found( - scope_register_subscriber(&scope.uuid, "missing-subscriber", Arc::new(|_event| {})) + scope_register_subscriber(scope_uuid, "missing-subscriber", Arc::new(|_event| {})) .unwrap_err(), "scope", ); diff --git a/crates/core/tests/unit/atif_tests.rs b/crates/core/tests/unit/atif_tests.rs index 050134f3a..ac71cd393 100644 --- a/crates/core/tests/unit/atif_tests.rs +++ b/crates/core/tests/unit/atif_tests.rs @@ -15,7 +15,7 @@ use crate::codec::anthropic::AnthropicMessagesCodec; use crate::codec::model_pricing::pricing_test_mutex; use crate::codec::openai_chat::OpenAIChatCodec; use crate::codec::openai_responses::OpenAIResponsesCodec; -use crate::codec::request::AnnotatedLlmRequest; +use crate::codec::request::{AnnotatedLlmRequest, ContentPart, Message, MessageContent}; use crate::codec::response::{ AnnotatedLlmResponse, CostEstimate, CostSource, PricingCatalog, PricingResolver, Usage, reset_active_pricing_resolver, set_active_pricing_resolver, @@ -368,6 +368,16 @@ fn set_event_timestamp(event: &mut Event, timestamp: chrono::DateTime, + step: chrono::Duration, +) { + for (offset, event) in events.iter_mut().enumerate() { + set_event_timestamp(event, base + step * offset as i32); + } +} + fn base_timestamp() -> chrono::DateTime { chrono::DateTime::parse_from_rfc3339("2026-01-01T00:00:00Z") .unwrap() @@ -721,8 +731,22 @@ fn test_exporter_llm_lifecycle() { let exporter = AtifExporter::new("session-1".to_string(), make_agent_info()); let llm_uuid = Uuid::now_v7(); - // Input wrapped in LlmRequest envelope — should be unwrapped. - let start = event_builder(llm_uuid, EventType::Start) + let start = llm_lifecycle_start_event(llm_uuid); + let end = llm_lifecycle_end_event(llm_uuid); + + { + let mut state = exporter.state.lock().unwrap(); + state.events.push(start); + state.events.push(end); + } + + let trajectory = exporter.export().unwrap(); + assert_eq!(trajectory.steps.len(), 2); + assert_llm_lifecycle_steps(&trajectory); +} + +fn llm_lifecycle_start_event(llm_uuid: Uuid) -> Event { + event_builder(llm_uuid, EventType::Start) .name("gpt-4") .scope_type(ScopeType::Llm) .input(json!({ @@ -745,10 +769,11 @@ fn test_exporter_llm_lifecycle() { "headers": {} })) .model_name("gpt-4") - .build(); + .build() +} - // Output with content, token_usage, and tool_calls. - let end = event_builder(llm_uuid, EventType::End) +fn llm_lifecycle_end_event(llm_uuid: Uuid) -> Event { + event_builder(llm_uuid, EventType::End) .name("gpt-4") .scope_type(ScopeType::Llm) .output(json!({ @@ -762,17 +787,10 @@ fn test_exporter_llm_lifecycle() { "tool_calls": [] })) .model_name("gpt-4") - .build(); - - { - let mut state = exporter.state.lock().unwrap(); - state.events.push(start); - state.events.push(end); - } - - let trajectory = exporter.export().unwrap(); - assert_eq!(trajectory.steps.len(), 2); + .build() +} +fn assert_llm_lifecycle_steps(trajectory: &AtifTrajectory) { // First step: user (LLM start — unwrapped LlmRequest, then messages extracted) let step1 = &trajectory.steps[0]; assert_eq!(step1.step_id, 1); @@ -2131,17 +2149,16 @@ fn test_exporter_openclaw_hook_only_fallbacks_preserve_stripped_content_and_expl .model_name("gpt-4") .build(); - for (offset, event) in [ - &mut stripped_start, - &mut stripped_end, - &mut partial_start, - &mut partial_end, - ] - .into_iter() - .enumerate() - { - set_event_timestamp(event, base + chrono::Duration::milliseconds(offset as i64)); - } + set_sequential_event_timestamps( + &mut [ + &mut stripped_start, + &mut stripped_end, + &mut partial_start, + &mut partial_end, + ], + base, + chrono::Duration::milliseconds(1), + ); { let mut state = exporter.state.lock().unwrap(); @@ -2153,7 +2170,17 @@ fn test_exporter_openclaw_hook_only_fallbacks_preserve_stripped_content_and_expl let trajectory = exporter.export().unwrap(); assert_atif_v17_shape(&trajectory); assert_eq!(trajectory.steps.len(), 4); + assert_openclaw_stripped_steps(&trajectory); + assert_openclaw_partial_steps(&trajectory); + assert_openclaw_combined_metrics(&trajectory); +} +fn assert_openclaw_stripped_steps(trajectory: &AtifTrajectory) { + assert_openclaw_stripped_user_step(trajectory); + assert_openclaw_stripped_agent_step(trajectory); +} + +fn assert_openclaw_stripped_user_step(trajectory: &AtifTrajectory) { let stripped_user = &trajectory.steps[0]; assert_eq!(stripped_user.source, "user"); let stripped_user_message: serde_json::Value = @@ -2175,7 +2202,9 @@ fn test_exporter_openclaw_hook_only_fallbacks_preserve_stripped_content_and_expl assert!(stripped_request.get("systemPrompt").is_none()); assert_eq!(stripped_request["messages"], json!([])); assert_eq!(stripped_request["imagesCount"], json!(1)); +} +fn assert_openclaw_stripped_agent_step(trajectory: &AtifTrajectory) { let stripped_agent = &trajectory.steps[1]; assert_eq!(stripped_agent.source, "agent"); let stripped_message: serde_json::Value = @@ -2192,7 +2221,9 @@ fn test_exporter_openclaw_hook_only_fallbacks_preserve_stripped_content_and_expl let stripped_response = stripped_agent_extra.llm_response.unwrap(); assert!(stripped_response.get("content").is_none()); assert_eq!(stripped_response["assistant_texts_count"], json!(1)); +} +fn assert_openclaw_partial_steps(trajectory: &AtifTrajectory) { let partial_user = &trajectory.steps[2]; assert_eq!(partial_user.source, "user"); assert_eq!(partial_user.message, json!("visible prompt")); @@ -2213,7 +2244,9 @@ fn test_exporter_openclaw_hook_only_fallbacks_preserve_stripped_content_and_expl assert_eq!(partial_metrics.completion_tokens, None); assert_eq!(partial_metrics.cached_tokens, None); assert_eq!(partial_metrics.cost_usd, None); +} +fn assert_openclaw_combined_metrics(trajectory: &AtifTrajectory) { let final_metrics = trajectory.final_metrics.as_ref().unwrap(); assert_eq!(final_metrics.total_prompt_tokens, Some(42)); assert_eq!(final_metrics.total_completion_tokens, None); @@ -2929,19 +2962,18 @@ fn test_exporter_embeds_nested_subagent_trajectory() { .scope_type(ScopeType::Agent) .build(); - for (offset, event) in [ - &mut root_start, - &mut child_start, - &mut llm_start, - &mut llm_end, - &mut child_end, - &mut root_end, - ] - .into_iter() - .enumerate() - { - set_event_timestamp(event, base + chrono::Duration::seconds(offset as i64)); - } + set_sequential_event_timestamps( + &mut [ + &mut root_start, + &mut child_start, + &mut llm_start, + &mut llm_end, + &mut child_end, + &mut root_end, + ], + base, + chrono::Duration::seconds(1), + ); { let mut state = exporter.state.lock().unwrap(); @@ -2961,7 +2993,10 @@ fn test_exporter_embeds_nested_subagent_trajectory() { assert_eq!(trajectory.session_id, root_uuid.to_string()); assert_eq!(trajectory.trajectory_id, Some(root_uuid.to_string())); assert_eq!(trajectory.steps.len(), 1); + assert_nested_subagent_projection(&trajectory, child_uuid); +} +fn assert_nested_subagent_projection(trajectory: &AtifTrajectory, child_uuid: Uuid) { let step = &trajectory.steps[0]; assert_eq!(step.source, "agent"); assert_eq!(step.llm_call_count, Some(0)); @@ -2993,7 +3028,7 @@ fn test_exporter_embeds_nested_subagent_trajectory() { assert_eq!(child.steps[0].source, "user"); assert_eq!(child.steps[1].source, "agent"); - let serialized = serde_json::to_value(&trajectory).unwrap(); + let serialized = serde_json::to_value(trajectory).unwrap(); assert!(serialized["steps"][0]["observation"]["results"][0]["content"].is_null()); } @@ -6036,3 +6071,423 @@ fn test_projected_duplicate_cleanup_preserves_partial_multi_call_step() { assert_eq!(steps.len(), 2); assert_eq!(step_tool_call_ids(&steps[0]), vec!["call_a", "call_b"]); } + +#[test] +fn atif_private_projection_helpers_cover_provider_edge_shapes() { + assert_provider_response_edge_shapes(); + assert_request_turn_state_edge_shapes(); + assert_annotated_turn_state_edge_shapes(); + assert_tool_argument_and_metric_edge_shapes(); +} + +fn assert_provider_response_edge_shapes() { + assert_eq!(atif_content_value(&Json::Null), empty_message()); + assert_eq!( + extract_llm_response_message(&json!({ + "assistant_message": {"content": "assistant"} + })), + json!("assistant") + ); + assert_eq!( + anthropic_messages_content_message( + &json!({"type": "message"}), + &json!([ + false, + {"type": "unknown"}, + {"type": "text", "text": "first"}, + {"type": "text", "text": "second"} + ]), + ), + Some(json!("first\nsecond")) + ); + assert_eq!( + anthropic_messages_content_message( + &json!({"type": "message"}), + &json!([{"type": "tool_use"}]), + ), + Some(empty_message()) + ); + assert_eq!( + anthropic_messages_content_message(&json!({"type": "message"}), &json!([])), + None + ); + + for invalid_parts in [ + json!(false), + json!([false]), + json!([{"type": "unknown"}]), + json!([{"type": "text"}]), + json!([{"type": "image", "source": {"media_type": "text/plain"}}]), + ] { + assert!(!is_atif_content_parts(&invalid_parts)); + } + assert_eq!( + openai_responses_output_message(&json!({"output_text": "direct"})), + Some(json!("direct")) + ); + assert_eq!( + openai_responses_output_message(&json!({ + "output": [ + false, + {"type": "unknown"}, + {"type": "output_text", "text": "first"}, + {"type": "message", "content": [ + false, + {"type": "output_text", "text": "second"} + ]} + ] + })), + Some(json!("first\nsecond")) + ); +} + +fn assert_request_turn_state_edge_shapes() { + assert_eq!( + request_turn_state_from_raw(&json!({"prompt": "hello"})), + Some(RequestTurnState::FreshUser) + ); + assert_eq!( + request_turn_state_from_raw(&json!({ + "messages": [{"role": "assistant", "content": "answer"}] + })), + Some(RequestTurnState::Continuation) + ); + assert_eq!( + request_turn_state_from_raw(&json!({ + "messages": [{"role": "user", "content": [{"type": "tool_result"}]}] + })), + Some(RequestTurnState::Continuation) + ); + assert_eq!( + request_turn_state_from_raw(&json!({ + "input": [{"type": "message", "role": "user", "content": "hello"}] + })), + Some(RequestTurnState::FreshUser) + ); + assert_eq!(provider_native_turn_state(&json!({"role": "system"})), None); + assert_eq!( + provider_native_turn_state(&json!({"role": "assistant"})), + Some(RequestTurnState::Continuation) + ); + assert_eq!(provider_native_turn_state(&json!({})), None); + assert!(raw_content_starts_new_turn(&json!({"type": "text"}))); + assert!(!raw_content_starts_new_turn( + &json!({"type": "tool_result"}) + )); + assert!(raw_content_starts_new_turn(&json!(false))); + assert_eq!( + openai_responses_turn_state(&json!([{"type": "function_call"}])), + Some(RequestTurnState::Continuation) + ); + assert_eq!( + openai_responses_input_content_message(&json!("direct input")), + Some(json!("direct input")) + ); + assert_eq!( + openai_responses_input_content_message(&json!([ + {"type": "text", "text": "first"}, + {"type": "text", "text": "second"} + ])), + Some(json!("first\nsecond")) + ); + assert_eq!(openai_responses_input_content_message(&json!(false)), None); +} + +fn assert_annotated_turn_state_edge_shapes() { + let text = MessageContent::Text("hello".into()); + assert_eq!( + annotated_message_turn_state(&Message::System { + content: text.clone(), + name: None, + }), + None + ); + assert_eq!( + annotated_message_turn_state(&Message::Assistant { + content: None, + tool_calls: None, + name: None, + }), + Some(RequestTurnState::Continuation) + ); + assert_eq!( + annotated_message_turn_state(&Message::Tool { + content: text.clone(), + tool_call_id: "call-1".into(), + }), + Some(RequestTurnState::Continuation) + ); + assert_eq!( + annotated_message_turn_state(&Message::ToolResultItem { + id: None, + call_id: "call-1".into(), + output: json!("done"), + extra: Default::default(), + }), + Some(RequestTurnState::Continuation) + ); + assert!(annotated_content_starts_new_turn(&MessageContent::Parts( + vec![ContentPart::Text { + text: "hello".into(), + extra: Default::default(), + },] + ))); + assert!(!annotated_content_starts_new_turn(&MessageContent::Parts( + vec![ContentPart::ToolResult { + tool_use_id: "call-1".into(), + content: json!("done"), + is_error: None, + extra: Default::default(), + },] + ))); + assert_eq!( + atif_message_from_annotated_response(&AnnotatedLlmResponse { + message: Some(MessageContent::Parts(vec![])), + ..Default::default() + }), + None + ); +} + +fn assert_tool_argument_and_metric_edge_shapes() { + assert_eq!(normalize_tool_arguments(None), json!({})); + assert_eq!(normalize_tool_arguments(Some(&json!(null))), json!({})); + assert_eq!( + normalize_tool_arguments(Some(&json!(42))), + json!({"value": 42}) + ); + assert_eq!( + normalize_tool_arguments(Some(&json!("not-json"))), + json!({"raw": "not-json"}) + ); + assert_eq!( + normalize_tool_arguments(Some(&json!("{\"key\":true}"))), + json!({"key": true}) + ); + + let supplemental = AtifMetrics { + prompt_tokens: Some(7), + extra: Some(json!({"supplemental": true})), + ..Default::default() + }; + let merged = merge_metrics(Some(AtifMetrics::default()), Some(&supplemental)).unwrap(); + assert_eq!(merged.prompt_tokens, Some(7)); + assert_eq!(merged.extra, supplemental.extra); + + let mut target = Some(json!({"preserved": true})); + merge_metrics_extra(&mut target, &Some(json!({"added": true}))); + assert_eq!(target, Some(json!({"preserved": true, "added": true}))); + let mut scalar = Some(json!(1)); + merge_metrics_extra(&mut scalar, &Some(json!(2))); + assert_eq!(scalar, Some(json!(1))); +} + +#[test] +fn atif_projection_helpers_cover_absent_and_fallback_values() { + assert!(!is_atif_image_source(None)); + assert_eq!( + extract_user_messages(&json!({"messages": [{"content": "implicit user"}]})), + json!("implicit user") + ); + assert_eq!( + raw_chat_message_turn_state(json!({"role": "user"}).as_object().expect("object fixture")), + Some(RequestTurnState::FreshUser) + ); + assert_eq!( + openai_responses_item_turn_state( + json!({"type": "message", "role": "assistant"}) + .as_object() + .expect("object fixture") + ), + Some(RequestTurnState::Continuation) + ); + assert_eq!( + atif_message_from_annotated_request(&AnnotatedLlmRequest::default()), + None + ); + assert_eq!( + annotated_message_turn_state(&Message::Developer { + content: MessageContent::Text("instruction".into()), + name: None, + }), + None + ); + assert_eq!( + normalize_tool_arguments(Some(&json!("[1,2]"))), + json!({"value": [1, 2]}) + ); + + let calls = extract_tool_calls(&json!({ + "tool_calls": [ + {}, + {"name": "lookup", "arguments": null} + ] + })) + .expect("one meaningful tool call"); + assert_eq!(calls.len(), 1); + assert_eq!(calls[0].tool_call_id, "lookup:2"); + + let mut metrics = Some(json!({"kept": true})); + merge_metrics_extra(&mut metrics, &None); + assert_eq!(metrics, Some(json!({"kept": true}))); +} + +#[test] +fn atif_observation_helpers_merge_and_prune_incomplete_references() { + let mut observation = AtifObservation { + results: vec![AtifObservationResult { + source_call_id: Some("call-1".into()), + content: None, + subagent_trajectory_ref: None, + extra: None, + }], + }; + merge_observation_result( + &mut observation, + AtifObservationResult { + source_call_id: Some("call-1".into()), + content: Some(json!("done")), + subagent_trajectory_ref: Some(vec![AtifSubagentTrajectoryRef { + trajectory_id: Some("child-1".into()), + session_id: Some("session-1".into()), + extra: None, + }]), + extra: Some(json!({"provider": "test"})), + }, + ); + assert_eq!(observation.results[0].content, Some(json!("done"))); + assert_eq!( + observation.results[0] + .subagent_trajectory_ref + .as_ref() + .map(Vec::len), + Some(1) + ); + assert_eq!( + observation.results[0].extra, + Some(json!({"provider": "test"})) + ); + + merge_observation_extra(&mut observation.results[0].extra, json!(false)); + assert_eq!( + observation.results[0].extra, + Some(json!({"provider": "test"})) + ); + + let mut step = cleanup_test_agent_step(empty_message(), &[], &[]); + step.source = "system".into(); + step.llm_call_count = None; + step.observation = Some(AtifObservation { + results: vec![AtifObservationResult { + source_call_id: None, + content: None, + subagent_trajectory_ref: Some(vec![AtifSubagentTrajectoryRef { + trajectory_id: Some("missing-child".into()), + session_id: None, + extra: None, + }]), + extra: None, + }], + }); + let mut steps = vec![step, cleanup_test_user_step()]; + prune_subagent_refs(&mut steps, &HashSet::new()); + assert_eq!(steps.len(), 1); + assert_eq!(steps[0].step_id, 1); + assert_eq!(steps[0].source, "user"); +} + +#[test] +fn atif_step_conversion_state_covers_tool_correlation_guard_paths() { + let tool_uuid = Uuid::now_v7(); + let tool_start = event_builder(tool_uuid, EventType::Start) + .name("lookup") + .scope_type(ScopeType::Tool) + .input(json!({"query": "relay"})) + .build(); + let mut lookups = EventLookupMaps::from_events(&[]); + + let mut state = StepConversionState::default(); + state.handle_tool_start(&tool_start, &lookups); + assert!(state.active_tool_call_id.is_none()); + assert!(!state.ensure_tool_call_on_current_agent(&tool_start, "call-1")); + + state.current_agent.step_idx = Some(99); + assert!(!state.ensure_tool_call_on_current_agent(&tool_start, "call-1")); + + state.steps = vec![cleanup_test_agent_step(empty_message(), &["call-1"], &[])]; + state.steps[0] + .tool_calls + .as_mut() + .expect("tool call fixture")[0] + .function_name = "lookup".into(); + state.current_agent.step_idx = Some(0); + assert!(state.ensure_tool_call_on_current_agent(&tool_start, "call-1")); + assert!(state.current_agent.has_tool_call_id("call-1")); + + lookups + .tool_call_ids + .insert(tool_uuid, "lookup-call".into()); + assert_eq!( + state.resolve_source_call_id(&tool_start, &lookups), + Some("lookup-call".into()) + ); + lookups.tool_call_ids.clear(); + assert_eq!( + state.resolve_source_call_id(&tool_start, &lookups), + Some("call-1".into()) + ); + + state.current_agent.tool_call_order.clear(); + state + .last_tool_call_map + .insert("lookup".into(), "call-1".into()); + assert_eq!( + state.resolve_source_call_id(&tool_start, &lookups), + Some("call-1".into()) + ); + state.last_tool_call_map.clear(); + state + .deferred_observations + .insert("call-1".into(), Vec::new()); + assert_eq!( + state.resolve_source_call_id(&tool_start, &lookups), + Some("call-1".into()) + ); + state.deferred_observations.clear(); + state + .deferred_tool_metadata + .insert("call-1".into(), Vec::new()); + assert_eq!( + state.resolve_source_call_id(&tool_start, &lookups), + Some("call-1".into()) + ); + state.deferred_tool_metadata.clear(); + assert_eq!(state.resolve_source_call_id(&tool_start, &lookups), None); + + let tool_end = event_builder(tool_uuid, EventType::End) + .name("lookup") + .scope_type(ScopeType::Tool) + .tool_call_id("call-1") + .build(); + state.handle_tool_end(&tool_end, &lookups); + assert!(state.pending_observations.is_empty()); + + let child = AgentScopeNode { + uuid: Uuid::now_v7(), + name: "child".into(), + session_id: None, + referenced_by_parent: false, + parent_agent: None, + children: Vec::new(), + start_timestamp: base_timestamp(), + }; + let mut empty_state = StepConversionState::default(); + assert!(!empty_state.attach_subagent_ref_to_agent_step(&child, &tool_start, "call-1")); + empty_state.current_agent.step_idx = Some(1); + assert!(!empty_state.attach_subagent_ref_to_agent_step(&child, &tool_start, "call-1")); + empty_state.steps.push(cleanup_test_user_step()); + empty_state.current_agent.step_idx = Some(0); + assert!(!empty_state.attach_subagent_ref_to_agent_step(&child, &tool_start, "call-1")); + empty_state.steps[0] = cleanup_test_agent_step(empty_message(), &["other"], &[]); + assert!(!empty_state.attach_subagent_ref_to_agent_step(&child, &tool_start, "call-1")); +} diff --git a/crates/core/tests/unit/codec/anthropic_tests.rs b/crates/core/tests/unit/codec/anthropic_tests.rs index 5a47b234c..6eb85d273 100644 --- a/crates/core/tests/unit/codec/anthropic_tests.rs +++ b/crates/core/tests/unit/codec/anthropic_tests.rs @@ -1008,6 +1008,11 @@ fn test_encode_max_tokens() { #[test] fn test_helper_and_error_paths_cover_remaining_anthropic_branches() { + assert_anthropic_helper_branches(); + assert_anthropic_codec_error_and_encode_paths(); +} + +fn assert_anthropic_helper_branches() { assert_eq!(json_f64(f64::NAN), Json::Null); assert_eq!( decode_anthropic_tool_choice(&json!({"type": "mystery"})), @@ -1019,7 +1024,9 @@ fn test_helper_and_error_paths_cover_remaining_anthropic_branches() { MessageContent::Parts(parts) if matches!(parts.as_slice(), [super::super::request::ContentPart::Image { .. }]) )); +} +fn assert_anthropic_codec_error_and_encode_paths() { let system_parts = MessageContent::Parts(vec![ super::super::request::ContentPart::Text { text: "First".into(), @@ -1126,6 +1133,11 @@ fn test_helper_and_error_paths_cover_remaining_anthropic_branches() { #[test] fn anthropic_request_component_branch_matrix() { + assert_anthropic_tool_choice_and_content_branches(); + assert_anthropic_message_and_tool_branches(); +} + +fn assert_anthropic_tool_choice_and_content_branches() { let specific = ToolChoice::Specific(ToolChoiceFunction { choice_type: "function".into(), function: ToolChoiceFunctionName { @@ -1233,7 +1245,9 @@ fn anthropic_request_component_branch_matrix() { }) .is_err() ); +} +fn assert_anthropic_message_and_tool_branches() { for invalid in [json!(42), json!({"content": "x"}), json!({"role": "user"})] { assert!(decode_anthropic_message(&invalid).is_err()); } @@ -1599,3 +1613,27 @@ fn anthropic_streaming_codec_keeps_partial_json_when_unparseable() { assert_eq!(block["id"], json!("toolu_p")); assert_eq!(block["input"], json!("{\"q\": \"trun")); } + +#[test] +fn anthropic_helpers_cover_invalid_and_provider_native_values() { + let codec = AnthropicMessagesCodec; + for invalid_request in [ + json!({}), + json!({"messages": false}), + json!({"messages": [], "tools": false}), + ] { + assert!(codec.decode(&make_request(invalid_request)).is_err()); + } + + assert!(decode_anthropic_tool_choice(&Json::Null).is_none()); + assert!(decode_anthropic_tool_choice(&json!({"type": "provider_choice"})).is_none()); + assert!(decode_parallel_tool_calls(&json!(false)).unwrap().is_none()); + + let mut object = serde_json::Map::from_iter([("remove".into(), json!(true))]); + set_or_remove_json(&mut object, "remove", None); + assert!(!object.contains_key("remove")); + set_or_remove_json(&mut object, "insert", Some(json!(42))); + assert_eq!(object["insert"], json!(42)); + + assert!(AnthropicMessagesStreamingCodec::default().finalizer()().is_object()); +} diff --git a/crates/core/tests/unit/codec/openai_chat_tests.rs b/crates/core/tests/unit/codec/openai_chat_tests.rs index 872373d8f..cb3fdb3b8 100644 --- a/crates/core/tests/unit/codec/openai_chat_tests.rs +++ b/crates/core/tests/unit/codec/openai_chat_tests.rs @@ -939,6 +939,12 @@ fn test_helper_and_error_paths_cover_remaining_chat_branches() { #[test] fn chat_request_component_branch_matrix() { + assert_chat_content_and_tool_call_branches(); + assert_chat_message_decode_branches(); + assert_chat_message_encoding_tool_and_choice_branches(); +} + +fn assert_chat_content_and_tool_call_branches() { assert!(decode_chat_content(&json!({})).is_err()); for invalid in [ json!(42), @@ -1023,7 +1029,9 @@ fn chat_request_component_branch_matrix() { .unwrap() .is_some() ); +} +fn assert_chat_message_decode_branches() { for invalid in [ json!(42), json!({"content": "x"}), @@ -1059,7 +1067,9 @@ fn chat_request_component_branch_matrix() { Message::ProviderNative { .. } )); } +} +fn assert_chat_message_encoding_tool_and_choice_branches() { let tool_call = ToolCall { id: "call".into(), call_type: "function".into(), @@ -1754,3 +1764,48 @@ fn openai_chat_streaming_codec_skips_null_usage_chunks() { assert_eq!(assembled["usage"]["prompt_tokens"], json!(1)); assert_eq!(assembled["usage"]["total_tokens"], json!(2)); } + +#[test] +fn openai_chat_helpers_cover_provider_edge_values() { + let refusal = decode_chat_content_part(&json!({ + "type": "refusal", + "refusal": "cannot comply", + "provider_field": true + })) + .unwrap(); + assert!(matches!( + refusal, + ContentPart::Refusal { refusal, extra } + if refusal == "cannot comply" && extra["provider_field"] == json!(true) + )); + + assert!(decode_chat_message(&json!({"role": "tool"})).is_err()); + assert!(matches!( + decode_chat_tool_choice(&json!("none")), + ToolChoice::None + )); + assert!(matches!( + decode_chat_tool_choice(&json!("required")), + ToolChoice::Required + )); + assert_eq!( + encode_chat_tool_choice(&ToolChoice::Auto).unwrap(), + json!("auto") + ); + + let mut object = serde_json::Map::from_iter([("remove".into(), json!(true))]); + set_or_remove_json(&mut object, "remove", None); + assert!(!object.contains_key("remove")); + set_or_remove_json(&mut object, "insert", Some(json!(42))); + assert_eq!(object["insert"], json!(42)); + + let codec = OpenAIChatCodec; + for invalid_request in [ + json!({"messages": [], "stop": false}), + json!({"messages": [], "functions": {}}), + json!({"messages": [], "modalities": false}), + ] { + assert!(codec.decode(&make_request(invalid_request)).is_err()); + } + assert!(OpenAIChatStreamingCodec::default().finalizer()().is_object()); +} diff --git a/crates/core/tests/unit/codec/openai_responses_tests.rs b/crates/core/tests/unit/codec/openai_responses_tests.rs index 8dadc2680..0d978811e 100644 --- a/crates/core/tests/unit/codec/openai_responses_tests.rs +++ b/crates/core/tests/unit/codec/openai_responses_tests.rs @@ -104,8 +104,11 @@ fn test_decode_full_response() { assert_eq!(usage.cache_read_tokens, Some(10)); assert_eq!(usage.cache_write_tokens, None); - // API specific fields - match resp.api_specific.unwrap() { + assert_full_response_api_specific(resp.api_specific.unwrap()); +} + +fn assert_full_response_api_specific(api_specific: ApiSpecificResponse) { + match api_specific { ApiSpecificResponse::OpenAIResponses { output_items, status, @@ -1150,6 +1153,12 @@ fn test_helper_and_error_paths_cover_remaining_responses_branches() { #[test] fn responses_request_component_branch_matrix() { + assert_responses_decode_component_branches(); + assert_responses_encode_content_and_item_branches(); + assert_responses_tool_and_choice_branches(); +} + +fn assert_responses_decode_component_branches() { assert!(decode_responses_content(&json!({})).is_err()); for invalid in [ json!(42), @@ -1205,7 +1214,9 @@ fn responses_request_component_branch_matrix() { Message::ProviderNative { .. } )); } +} +fn assert_responses_encode_content_and_item_branches() { let content = MessageContent::Parts(vec![ ContentPart::Text { text: "hello".into(), @@ -1326,7 +1337,9 @@ fn responses_request_component_branch_matrix() { }) .is_err() ); +} +fn assert_responses_tool_and_choice_branches() { for invalid in [ json!(42), json!({"type": "function"}), @@ -1606,3 +1619,35 @@ fn openai_responses_streaming_codec_ignores_per_token_deltas() { Some(MessageContent::Text("Hello".to_string())) ); } + +#[test] +fn responses_helpers_cover_invalid_and_provider_edge_values() { + let codec = OpenAIResponsesCodec; + for invalid_request in [ + json!({}), + json!({"input": false}), + json!({"input": [], "instructions": false}), + ] { + assert!(codec.decode(&make_request(invalid_request)).is_err()); + } + + assert!(matches!( + decode_openai_or_anthropic_tool_choice(&json!({ + "type": "function", + "function": {"name": "lookup"} + })), + ToolChoice::Specific(_) + )); + assert_eq!( + encode_responses_tool_choice(&ToolChoice::Required).unwrap(), + json!("required") + ); + + let mut object = serde_json::Map::from_iter([("remove".into(), json!(true))]); + set_or_remove_json(&mut object, "remove", None); + assert!(!object.contains_key("remove")); + set_or_remove_json(&mut object, "insert", Some(json!(42))); + assert_eq!(object["insert"], json!(42)); + + assert!(OpenAIResponsesStreamingCodec::default().finalizer()().is_object()); +} diff --git a/crates/core/tests/unit/dynamic_worker_tests.rs b/crates/core/tests/unit/dynamic_worker_tests.rs index e4812eab5..b2b439792 100644 --- a/crates/core/tests/unit/dynamic_worker_tests.rs +++ b/crates/core/tests/unit/dynamic_worker_tests.rs @@ -3,6 +3,9 @@ use std::sync::{Arc, Mutex}; +#[cfg(unix)] +use std::os::unix::fs::PermissionsExt; + use crate::api::event::{BaseEvent, MarkEvent}; use crate::api::optimization::{ LlmOptimizationRecorder, record_llm_optimization_contribution, scope_llm_optimization_recorder, @@ -119,11 +122,78 @@ fn python_worker_launch_clears_host_python_environment() { } } +#[cfg(unix)] +#[test] +fn python_worker_process_launch_uses_the_managed_interpreter_and_endpoint_file() { + let plugin_id = "acme.python.launch"; + let digest = Sha256::digest(plugin_id.as_bytes()) + .iter() + .map(|byte| format!("{byte:02x}")) + .collect::(); + let temp = tempfile::tempdir().unwrap(); + let managed = temp.path().join(MANAGED_ENVIRONMENTS_DIR).join(digest); + let interpreter = managed.join("bin/python"); + std::fs::create_dir_all(interpreter.parent().unwrap()).unwrap(); + let probe = temp.path().join("launch-probe"); + std::fs::write( + &interpreter, + format!( + "#!/bin/sh\nprintf '%s\\n%s\\n' \"$0\" \"$NEMO_RELAY_WORKER_ENDPOINT_FILE\" > '{}'\nexit 0\n", + probe.display() + ), + ) + .unwrap(); + let mut permissions = std::fs::metadata(&interpreter).unwrap().permissions(); + permissions.set_mode(0o755); + std::fs::set_permissions(&interpreter, permissions).unwrap(); + + let endpoint_file = temp.path().join("worker-endpoint"); + let mut child = spawn_worker_process(WorkerProcessLaunch { + runtime: WorkerRuntime::Python, + manifest_path: &temp.path().join("plugin.toml"), + environment_ref: managed.to_str(), + plugin_id, + entrypoint: "acme_worker:create_plugin", + activation_id: "activation", + auth_token: "token", + host_endpoint: "http://127.0.0.1:1", + worker_endpoint: "http://127.0.0.1:2", + worker_endpoint_file: Some(&endpoint_file), + }) + .unwrap(); + assert!(child.wait().unwrap().success()); + let recorded = std::fs::read_to_string(&probe).unwrap(); + let mut lines = recorded.lines(); + assert_eq!(lines.next(), interpreter.to_str()); + assert_eq!(lines.next(), endpoint_file.to_str()); + assert_eq!(lines.next(), None); +} + #[cfg(not(unix))] #[test] fn empty_worker_endpoint_announcement_is_retried() { let temp = tempfile::tempdir().unwrap(); let announcement = temp.path().join("worker-endpoint"); + let missing = temp.path().join("missing-endpoint"); + let directory = temp.path().join("endpoint-directory"); + std::fs::create_dir(&directory).unwrap(); + + assert_eq!( + normalize_worker_tcp_endpoint(" tcp://127.0.0.1:50051 ").unwrap(), + "http://127.0.0.1:50051" + ); + assert_eq!( + normalize_worker_tcp_endpoint("http://127.0.0.1:50051").unwrap(), + "http://127.0.0.1:50051" + ); + assert!(normalize_worker_tcp_endpoint("tcp://").is_err()); + assert!(normalize_worker_tcp_endpoint("https://127.0.0.1").is_err()); + assert!( + resolve_worker_connect_endpoint(&WorkerConnectEndpoint::Announced(missing)) + .unwrap() + .is_none() + ); + assert!(resolve_worker_connect_endpoint(&WorkerConnectEndpoint::Announced(directory)).is_err()); std::fs::write(&announcement, "").unwrap(); let endpoint = WorkerConnectEndpoint::Announced(announcement); @@ -135,14 +205,18 @@ fn empty_worker_endpoint_announcement_is_retried() { ); } -#[test] -fn response_helpers_cover_error_and_unexpected_shapes() { - enable_operational_logs(); - let worker_error = WorkerError { +fn test_worker_error() -> WorkerError { + WorkerError { code: "worker.failed".into(), message: "boom".into(), retryable: false, - }; + } +} + +#[test] +fn response_helpers_cover_json_and_guardrail_error_shapes() { + enable_operational_logs(); + let worker_error = test_worker_error(); let error = json_from_invoke_response(InvokeResponse { result: Some(InvokeResult::Json(JsonResult { @@ -211,6 +285,12 @@ fn response_helpers_cover_error_and_unexpected_shapes() { .to_string() .contains("guardrail returned unexpected") ); +} + +#[test] +fn response_helpers_cover_stream_optional_and_cancellation_errors() { + enable_operational_logs(); + let worker_error = test_worker_error(); assert!( json_from_stream_chunk(StreamChunk { @@ -237,6 +317,74 @@ fn response_helpers_cover_error_and_unexpected_shapes() { .to_string() .contains("stream chunk was empty") ); + + assert!( + optional_json_from_invoke_response(InvokeResponse { + result: Some(InvokeResult::Json(JsonResult { + value: None, + error: Some(worker_error.clone()), + })), + }) + .unwrap_err() + .to_string() + .contains("worker.failed") + ); + assert!( + optional_json_from_invoke_response(InvokeResponse { + result: Some(InvokeResult::Json(JsonResult { + value: Some(JsonEnvelope { + schema: JSON_SCHEMA.into(), + json: b"{".to_vec(), + }), + error: None, + })), + }) + .unwrap_err() + .to_string() + .contains("invalid JSON result") + ); + assert_eq!( + optional_json_from_invoke_response(InvokeResponse { + result: Some(InvokeResult::Empty(EmptyResult {})), + }) + .expect("empty result must map to no replacement JSON"), + None + ); + assert!( + optional_json_from_invoke_response(InvokeResponse { + result: Some(InvokeResult::Guardrail(GuardrailResult { + block_reason: String::new(), + })), + }) + .unwrap_err() + .to_string() + .contains("unexpected LLM sanitizer result") + ); + + let typed_error = typed_json_result::( + JSON_SCHEMA, + Err(FlowError::Internal("typed result failed".into())), + ); + assert!(typed_error.value.is_none()); + assert!( + typed_error + .error + .is_some_and(|error| error.message.contains("typed result failed")) + ); + assert!( + worker_error_to_flow(WorkerError { + code: "worker.cancelled".into(), + message: "cancelled by caller".into(), + retryable: false, + }) + .to_string() + .contains("worker invocation cancelled") + ); + assert!( + worker_status_to_flow("unused", Status::cancelled("cancelled by transport")) + .to_string() + .contains("worker invocation cancelled") + ); } #[test] @@ -421,6 +569,16 @@ fn worker_endpoints_fail_when_host_socket_cannot_bind() { let _ = std::fs::remove_dir_all(&activation_dir); } +#[cfg(unix)] +#[tokio::test] +async fn worker_unix_connection_reports_a_missing_socket() { + let missing = std::env::temp_dir().join(format!("nmrw-missing-{}", Uuid::now_v7())); + let error = connect_worker(&WorkerConnectEndpoint::Unix(missing.clone())) + .await + .expect_err("a missing worker socket should fail to connect"); + assert!(error.to_string().contains(&missing.display().to_string())); +} + #[tokio::test(flavor = "multi_thread")] async fn callback_helpers_cover_worker_response_edges() { enable_operational_logs(); @@ -1852,6 +2010,18 @@ async fn host_runtime_codec_capabilities_are_directional_authorized_and_ephemera .expect_err("a capability cannot be reused by another invocation"); assert_eq!(wrong_invocation.code(), tonic::Code::PermissionDenied); + let wrong_invocation = state + .response_codec(&response_capability, "another-invocation") + .err() + .expect("response capability must remain invocation-bound"); + assert_eq!(wrong_invocation.code(), tonic::Code::PermissionDenied); + + let wrong_direction = state + .response_codec(&request_capability, invocation_id) + .err() + .expect("request capability cannot decode responses"); + assert_eq!(wrong_direction.code(), tonic::Code::InvalidArgument); + let decoded = service .decode_llm_codec_request(Request::new(LlmCodecDecodeRequest { activation_id: ACTIVATION_ID.into(), diff --git a/crates/core/tests/unit/llm_api_tests.rs b/crates/core/tests/unit/llm_api_tests.rs index 62615e9e2..ee989abc9 100644 --- a/crates/core/tests/unit/llm_api_tests.rs +++ b/crates/core/tests/unit/llm_api_tests.rs @@ -146,7 +146,7 @@ impl LlmResponseCodec for RuntimeIdentityCodec { } #[test] -fn sanitizer_context_preserves_all_codec_identity_states() { +fn request_sanitizer_context_preserves_all_codec_identity_states() { let identity_only_request = crate::api::runtime::LlmSanitizeRequestContext::with_identity( LlmCodecIdentity::Runtime("identity-only.request.v1".into()), ); @@ -155,15 +155,7 @@ fn sanitizer_context_preserves_all_codec_identity_states() { &LlmCodecIdentity::Runtime("identity-only.request.v1".into()) ); assert!(identity_only_request.resolve_codec().is_none()); - - let identity_only_response = crate::api::runtime::LlmSanitizeResponseContext::with_identity( - LlmCodecIdentity::Runtime("identity-only.response.v1".into()), - ); - assert_eq!( - identity_only_response.codec(), - &LlmCodecIdentity::Runtime("identity-only.response.v1".into()) - ); - assert!(identity_only_response.resolve_codec().is_none()); + assert!(format!("{identity_only_request:?}").contains("identity-only.request.v1")); assert_eq!( sanitize_context_for_request_codec(None).codec(), @@ -192,6 +184,20 @@ fn sanitizer_context_preserves_all_codec_identity_states() { .codec(), &LlmCodecIdentity::Opaque ); +} + +#[test] +fn response_sanitizer_context_preserves_all_codec_identity_states() { + let identity_only_response = crate::api::runtime::LlmSanitizeResponseContext::with_identity( + LlmCodecIdentity::Runtime("identity-only.response.v1".into()), + ); + assert_eq!( + identity_only_response.codec(), + &LlmCodecIdentity::Runtime("identity-only.response.v1".into()) + ); + assert!(identity_only_response.resolve_codec().is_none()); + assert!(format!("{identity_only_response:?}").contains("identity-only.response.v1")); + assert_eq!( sanitize_context_for_response_codec(Some(&OpenAIChatCodec as &dyn LlmResponseCodec)) .codec(), diff --git a/crates/core/tests/unit/native_plugin_tests.rs b/crates/core/tests/unit/native_plugin_tests.rs index e180909bd..7a7e203ce 100644 --- a/crates/core/tests/unit/native_plugin_tests.rs +++ b/crates/core/tests/unit/native_plugin_tests.rs @@ -15,6 +15,8 @@ use nemo_relay_plugin::{ NemoRelayNativeLlmSanitizeResponseContext, NemoRelayNativeLlmStreamNextFn, NemoRelayNativeToolNextFn, }; +#[cfg(unix)] +use nemo_relay_plugin::{NemoRelayNativePluginRegisterFn, NemoRelayNativePluginValidateFn}; use serde_json::json; use crate::api::optimization::{ @@ -185,592 +187,1318 @@ unsafe extern "C" fn stop_after_first_native_stream_item( false } -struct InvokeNativeNextThenReturnState { - callback_state: u32, - invoke_status: AtomicUsize, - started: Mutex>, -} +#[test] +fn native_async_entrypoints_reject_null_handles() { + unsafe { + assert_eq!( + native_async_completion_resolve_json(ptr::null(), ptr::null()), + NemoRelayStatus::NullPointer + ); + assert_eq!( + native_async_completion_reject(ptr::null(), ptr::null()), + NemoRelayStatus::NullPointer + ); + assert!(native_async_completion_is_cancelled(ptr::null())); + native_async_completion_release(ptr::null()); + native_async_next_release(ptr::null()); -unsafe extern "C" fn invoke_native_next_then_return_state( - user_data: *mut c_void, - invocation_json: *const NemoRelayNativeString, - next: *const NemoRelayNativeAsyncNext, - completion: *const NemoRelayNativeAsyncCompletion, -) -> u32 { - let state = unsafe { &*user_data.cast::() }; - let status = unsafe { native_async_next_invoke(next, invocation_json, completion) }; - state - .invoke_status - .store(status as usize, Ordering::Release); - if status == NemoRelayStatus::Ok { - let _ = state - .started - .lock() - .unwrap_or_else(|error| error.into_inner()) - .recv_timeout(Duration::from_secs(1)); - } - unsafe { native_async_next_release(next) }; - state.callback_state -} + assert_eq!( + native_async_stream_push_json(ptr::null(), ptr::null()), + NemoRelayStatus::NullPointer + ); + assert_eq!( + native_async_stream_finish(ptr::null()), + NemoRelayStatus::NullPointer + ); + assert_eq!( + native_async_stream_reject(ptr::null(), ptr::null()), + NemoRelayStatus::NullPointer + ); + assert!(native_async_stream_is_cancelled(ptr::null())); + native_async_stream_release(ptr::null()); -unsafe extern "C" fn invoke_native_stream_next_then_return_state( - user_data: *mut c_void, - invocation_json: *const NemoRelayNativeString, - next: *const NemoRelayNativeAsyncNext, - stream: *const NemoRelayNativeAsyncStream, -) -> u32 { - let state = unsafe { &*user_data.cast::() }; - let request = read_native_string(invocation_json) - .ok() - .and_then(|invocation| serde_json::from_str::(&invocation).ok()) - .and_then(|invocation| invocation.get("request").cloned()) - .and_then(|request| native_string_from_json(&request)); - let status = if let Some(request) = request { - let status = unsafe { + assert_eq!( + native_async_next_invoke(ptr::null(), ptr::null(), ptr::null()), + NemoRelayStatus::NullPointer + ); + assert_eq!( + native_async_next_invoke_result( + ptr::null(), + ptr::null(), + complete_native_next_result, + ptr::null_mut(), + ), + NemoRelayStatus::NullPointer + ); + assert_eq!( native_async_next_invoke_stream( - next, - request, - stream, + ptr::null(), + ptr::null(), + ptr::null(), accept_native_stream_item, ptr::null_mut(), - ) - }; - unsafe { native_string_free(request) }; - status - } else { - NemoRelayStatus::InvalidJson - }; - state - .invoke_status - .store(status as usize, Ordering::Release); - if status == NemoRelayStatus::Ok { - let _ = state - .started - .lock() - .unwrap_or_else(|error| error.into_inner()) - .recv_timeout(Duration::from_secs(1)); - } - unsafe { - native_async_next_release(next); - native_async_stream_release(stream); + ), + NemoRelayStatus::NullPointer + ); } - state.callback_state } -unsafe extern "C" fn invoke_detached_next_and_finish_replacement_stream( - user_data: *mut c_void, - invocation_json: *const NemoRelayNativeString, - next: *const NemoRelayNativeAsyncNext, - stream: *const NemoRelayNativeAsyncStream, -) -> u32 { - let invoke_status = unsafe { &*user_data.cast::() }; - let request = read_native_string(invocation_json) - .ok() - .and_then(|invocation| serde_json::from_str::(&invocation).ok()) - .and_then(|invocation| invocation.get("request").cloned()) - .and_then(|request| native_string_from_json(&request)); - let status = if let Some(request) = request { - let status = unsafe { - native_async_next_invoke_stream( - next, - request, - stream, - accept_native_stream_item, - ptr::null_mut(), - ) - }; - unsafe { native_string_free(request) }; - status - } else { - NemoRelayStatus::InvalidJson - }; - invoke_status.store(status as usize, Ordering::Release); - let replacement = native_string_from_json(&json!({"source": "replacement"})).unwrap(); - unsafe { - let _ = native_async_stream_push_json(stream, replacement); - native_string_free(replacement); - let _ = native_async_stream_finish(stream); - native_async_next_release(next); - native_async_stream_release(stream); - } - NemoRelayNativeAsyncCallbackState::Complete as u32 +#[cfg(unix)] +unsafe extern "C" fn native_test_validate_invalid_json( + _user_data: *mut c_void, + _config: *const NemoRelayNativeString, + out: *mut *mut NemoRelayNativeString, +) -> NemoRelayStatus { + unsafe { *out = native_string("not-json") }; + NemoRelayStatus::Ok } -struct FailingNativeCodec; +#[cfg(unix)] +unsafe extern "C" fn native_test_validate_error( + _user_data: *mut c_void, + _config: *const NemoRelayNativeString, + out: *mut *mut NemoRelayNativeString, +) -> NemoRelayStatus { + unsafe { *out = native_string("unused") }; + set_native_last_error("validation callback failed"); + NemoRelayStatus::InvalidArg +} -impl LlmCodec for FailingNativeCodec { - fn decode(&self, _request: &LlmRequest) -> FlowResult { - Err(FlowError::Internal("request decode rejected".into())) - } +#[cfg(unix)] +unsafe extern "C" fn native_test_validate_empty( + _user_data: *mut c_void, + _config: *const NemoRelayNativeString, + _out: *mut *mut NemoRelayNativeString, +) -> NemoRelayStatus { + NemoRelayStatus::Ok +} - fn encode( - &self, - _annotated: &AnnotatedLlmRequest, - _original: &LlmRequest, - ) -> FlowResult { - Err(FlowError::Internal("request encode rejected".into())) - } +unsafe extern "C" fn native_test_register_ok( + _user_data: *mut c_void, + _config: *const NemoRelayNativeString, + _ctx: *mut NemoRelayNativePluginContext, +) -> NemoRelayStatus { + NemoRelayStatus::Ok } -impl LlmResponseCodec for FailingNativeCodec { - fn decode_response(&self, _response: &Json) -> FlowResult { - Err(FlowError::Internal("response decode rejected".into())) +#[cfg(unix)] +unsafe extern "C" fn native_test_register_error( + _user_data: *mut c_void, + _config: *const NemoRelayNativeString, + _ctx: *mut NemoRelayNativePluginContext, +) -> NemoRelayStatus { + set_native_last_error("registration callback failed"); + NemoRelayStatus::InvalidArg +} + +#[cfg(unix)] +fn native_test_adapter( + validate: Option, + register: Option, +) -> NativePluginAdapter { + let plugin = NemoRelayNativePluginV1 { + validate, + register, + ..Default::default() + }; + NativePluginAdapter { + plugin_kind: "test.native.adapter".into(), + allows_multiple_components: false, + instance: Arc::new(NativePluginInstance { + plugin_kind: "test.native.adapter".into(), + relay_compat: "^0.7".into(), + allows_multiple_components: false, + plugin: Mutex::new(plugin), + _library: libloading::os::unix::Library::this().into(), + }), } } -struct PanickingNativeCodec; +#[cfg(unix)] +#[tokio::test] +async fn native_plugin_adapter_covers_validation_and_registration_results() { + let no_validate = native_test_adapter(None, Some(native_test_register_ok)); + assert_eq!(no_validate.plugin_kind(), "test.native.adapter"); + assert!(!no_validate.allows_multiple_components()); + assert!(no_validate.validate(&Map::new()).is_empty()); + + let empty = native_test_adapter( + Some(native_test_validate_empty), + Some(native_test_register_ok), + ); + assert!(empty.validate(&Map::new()).is_empty()); + + let invalid = native_test_adapter( + Some(native_test_validate_invalid_json), + Some(native_test_register_ok), + ); + let diagnostics = invalid.validate(&Map::new()); + assert_eq!(diagnostics.len(), 1); + assert!(diagnostics[0].message.contains("invalid diagnostics JSON")); + + let failing = native_test_adapter( + Some(native_test_validate_error), + Some(native_test_register_error), + ); + let diagnostics = failing.validate(&Map::new()); + assert_eq!(diagnostics.len(), 1); + assert_eq!(diagnostics[0].message, "validation callback failed"); + + let mut context = PluginRegistrationContext::new(); + assert!(empty.register(&Map::new(), &mut context).await.is_ok()); + let error = failing + .register(&Map::new(), &mut context) + .await + .expect_err("registration callback should fail"); + assert!(error.to_string().contains("registration callback failed")); + + let missing = native_test_adapter(Some(native_test_validate_empty), None); + let error = missing + .register(&Map::new(), &mut context) + .await + .expect_err("missing registration callback should fail"); + assert!( + error + .to_string() + .contains("did not return a register callback") + ); +} -impl LlmCodec for PanickingNativeCodec { - fn decode(&self, _request: &LlmRequest) -> FlowResult { - panic!("request decode panic") - } +#[test] +fn native_loader_helpers_cover_compatibility_descriptor_and_digest_edges() { + assert_native_compatibility_edges(); + assert_native_descriptor_edges(); + assert_native_digest_edges(); + assert_native_host_api_versions(); +} - fn encode( - &self, - _annotated: &AnnotatedLlmRequest, - _original: &LlmRequest, - ) -> FlowResult { - panic!("request encode panic") - } +fn assert_native_compatibility_edges() { + assert!(validate_relay_compatibility(None).is_err()); + assert!(validate_relay_compatibility(Some(" ")).is_err()); + assert!(validate_relay_compatibility(Some("not a requirement")).is_err()); + assert!(validate_relay_compatibility(Some(">=999.0.0")).is_err()); + assert!(validate_relay_compatibility(Some("^0.7")).is_ok()); } -impl LlmResponseCodec for PanickingNativeCodec { - fn decode_response(&self, _response: &Json) -> FlowResult { - panic!("response decode panic") - } +fn assert_native_descriptor_edges() { + let mut descriptor = NemoRelayNativePluginV1 { + struct_size: 0, + ..Default::default() + }; + assert!(validate_plugin_descriptor("test", &descriptor).is_err()); + descriptor.struct_size = std::mem::size_of::(); + assert!(validate_plugin_descriptor("test", &descriptor).is_err()); + descriptor.plugin_kind = native_string("test"); + assert!(validate_plugin_descriptor("test", &descriptor).is_err()); + descriptor.register = Some(native_test_register_ok); + assert!(validate_plugin_descriptor("test", &descriptor).is_ok()); + drop_native_plugin_descriptor(&mut descriptor); } -#[test] -fn native_string_and_json_helpers_cover_abi_boundaries() { - clear_native_last_error(); +fn assert_native_digest_edges() { + let temp = tempfile::tempdir().unwrap(); + let manifest = temp.path().join("plugin.toml"); + let library = temp.path().join("plugin.bin"); + std::fs::write(&library, b"native plugin bytes").unwrap(); assert_eq!( - unsafe { native_string_new(ptr::null(), 0, ptr::null_mut()) }, - NemoRelayStatus::NullPointer + resolve_manifest_relative_path(&manifest, "plugin.bin"), + library ); - assert_last_error_contains("out string pointer is null"); + assert_eq!( + resolve_manifest_relative_path(&manifest, temp.path().to_str().unwrap()), + temp.path() + ); + assert_eq!(hex_digest([0x00, 0xab, 0xff]), "00abff"); + let digest = hex_digest(Sha256::digest(b"native plugin bytes")); + assert!(verify_sha256(&library, &format!("sha256:{}", digest.to_uppercase())).is_ok()); + assert!(verify_sha256(&library, "sha256:00").is_err()); + assert!(verify_sha256(&temp.path().join("missing"), "00").is_err()); +} - let mut out = ptr::null_mut(); +fn assert_native_host_api_versions() { + let current = native_host_api(); + let legacy = native_host_api_legacy(); + assert!(!current.is_null()); + assert!(!legacy.is_null()); assert_eq!( - unsafe { native_string_new(ptr::null(), 1, &mut out) }, - NemoRelayStatus::NullPointer + unsafe { (*legacy).abi_version }, + NEMO_RELAY_NATIVE_ABI_VERSION_LEGACY ); - assert!(out.is_null()); - assert_last_error_contains("string data pointer is null"); +} - let invalid_utf8 = [0xff]; +#[tokio::test] +async fn native_async_wait_and_rejection_cover_dropped_and_aborted_continuations() { + let (sender, receiver) = tokio::sync::oneshot::channel::>(); + drop(sender); + let mut wait = NativeAsyncWait { + completion: Arc::new(NativeAsyncCompletion { + sender: Mutex::new(None), + cancelled: AtomicBool::new(false), + next_invoked: AtomicBool::new(false), + next_abort: Mutex::new(None), + before_settlement_lock: None, + _callback_user_data: None, + }), + receiver, + completed: false, + }; + let error = wait + .receive() + .await + .expect_err("a dropped callback must fail the wait"); + assert!(error.to_string().contains("dropped without settling")); + + let task = tokio::spawn(std::future::pending::<()>()); + let abort = task.abort_handle(); + let (sender, receiver) = tokio::sync::oneshot::channel(); + let completion = Arc::new(NativeAsyncCompletion { + sender: Mutex::new(Some(sender)), + cancelled: AtomicBool::new(false), + next_invoked: AtomicBool::new(true), + next_abort: Mutex::new(Some(abort)), + before_settlement_lock: None, + _callback_user_data: None, + }); + let completion_ref = + Arc::into_raw(Arc::clone(&completion)) as *const NemoRelayNativeAsyncCompletion; + let invalid_message = Box::into_raw(Box::new(NativeHostString(vec![0xff]))).cast(); assert_eq!( - unsafe { native_string_new(invalid_utf8.as_ptr(), invalid_utf8.len(), &mut out) }, - NemoRelayStatus::InvalidUtf8 + unsafe { native_async_completion_reject(completion_ref, invalid_message) }, + NemoRelayStatus::InvalidArg ); - assert!(out.is_null()); - assert_last_error_contains("not valid UTF-8"); + unsafe { native_string_free(invalid_message) }; - let text = native_string("hello"); - assert_eq!(unsafe { native_string_len(text) }, 5); + let message = native_string("continuation rejected"); assert_eq!( - unsafe { std::slice::from_raw_parts(native_string_data(text), 5) }, - b"hello" + unsafe { native_async_completion_reject(completion_ref, message) }, + NemoRelayStatus::Ok ); - assert!(unsafe { native_string_data(ptr::null()) }.is_null()); - assert_eq!(unsafe { native_string_len(ptr::null()) }, 0); - assert_eq!(take_native_string(text).unwrap(), "hello"); - unsafe { native_string_free(ptr::null_mut()) }; - - let empty = native_string(""); - assert_eq!(read_native_string(empty).unwrap(), ""); - unsafe { native_string_free(empty) }; - assert_eq!(read_native_string(ptr::null()).unwrap(), ""); + unsafe { native_string_free(message) }; + let error = receiver.await.unwrap().unwrap_err(); + assert!(error.to_string().contains("continuation rejected")); + assert!(task.await.unwrap_err().is_cancelled()); + unsafe { native_async_completion_release(completion_ref) }; +} - let bad = Box::into_raw(Box::new(NativeHostString(vec![0xff]))) as *mut NemoRelayNativeString; - assert!(read_native_string(bad).is_err()); - assert_eq!( - optional_json_from_native_string(bad, "bad json"), - Err(NemoRelayStatus::InvalidUtf8) - ); - unsafe { native_last_error_set(bad) }; - assert_last_error_contains("not valid UTF-8"); - unsafe { native_string_free(bad) }; +#[test] +fn native_stream_callback_guard_covers_terminal_drop_modes() { + let (sender, _receiver) = tokio::sync::mpsc::channel(1); + let stream = Arc::new(NativeAsyncStream { + sender: Mutex::new(Some(sender)), + cancelled: AtomicBool::new(false), + settled: AtomicBool::new(false), + downstream_aborts: Mutex::new(HashMap::new()), + settlement: Mutex::new(()), + before_settlement_lock: None, + _callback_user_data: None, + }); + let callback_state = NativeStreamCallbackState::default(); + let user_data: *mut c_void = ptr::from_ref(&callback_state).cast_mut().cast(); - let message = native_string("explicit native error"); - unsafe { native_last_error_set(message) }; - assert_eq!( - native_last_error_message().as_deref(), - Some("explicit native error") - ); - unsafe { native_string_free(message) }; - unsafe { native_last_error_clear() }; - assert!(native_last_error_message().is_none()); + let mut inactive = NativeAsyncStreamCallbackGuard { + cb: record_native_stream_result, + user_data: user_data as usize, + stream: Arc::clone(&stream), + _library_guard: None, + active: false, + }; + inactive.fail("ignored"); - set_native_last_error("specific fallback"); + stream.cancelled.store(true, Ordering::Release); + let mut cancelled = NativeAsyncStreamCallbackGuard { + cb: record_native_stream_result, + user_data: user_data as usize, + stream: Arc::clone(&stream), + _library_guard: None, + active: true, + }; + cancelled.fail("cancellation owns settlement"); + drop(cancelled); + + stream.cancelled.store(false, Ordering::Release); + stream.settled.store(true, Ordering::Release); + drop(NativeAsyncStreamCallbackGuard { + cb: record_native_stream_result, + user_data: user_data as usize, + stream: Arc::clone(&stream), + _library_guard: None, + active: true, + }); + + stream.settled.store(false, Ordering::Release); + drop(NativeAsyncStreamCallbackGuard { + cb: record_native_stream_result, + user_data: user_data as usize, + stream, + _library_guard: None, + active: true, + }); + assert_eq!(callback_state.callbacks.load(Ordering::Acquire), 3); + assert!(callback_state.done.load(Ordering::Acquire)); +} + +#[tokio::test] +async fn native_async_stream_forwarding_reports_conversion_and_stream_errors() { + let make_stream_state = || { + let (sender, _receiver) = tokio::sync::mpsc::channel(1); + Arc::new(NativeAsyncStream { + sender: Mutex::new(Some(sender)), + cancelled: AtomicBool::new(false), + settled: AtomicBool::new(false), + downstream_aborts: Mutex::new(HashMap::new()), + settlement: Mutex::new(()), + before_settlement_lock: None, + _callback_user_data: None, + }) + }; + + let conversion_state = NativeStreamCallbackState::default(); + let mut conversion_guard = NativeAsyncStreamCallbackGuard { + cb: record_native_stream_result, + user_data: ptr::from_ref(&conversion_state) as usize, + stream: make_stream_state(), + _library_guard: None, + active: true, + }; + forward_native_async_next_stream_with( + LlmJsonStream::new(tokio_stream::iter([Ok(json!({"chunk": true}))])), + &mut conversion_guard, + |_| None, + ) + .await; assert!( - json_from_native_string(ptr::null_mut(), "generic fallback") - .unwrap_err() - .to_string() - .contains("specific fallback") + conversion_state + .error + .lock() + .unwrap() + .as_deref() + .is_some_and(|error| error.contains("failed to serialize or allocate")) ); - clear_native_last_error(); + assert!(!conversion_state.done.load(Ordering::Acquire)); + + let stream_error_state = NativeStreamCallbackState::default(); + let mut stream_error_guard = NativeAsyncStreamCallbackGuard { + cb: record_native_stream_result, + user_data: ptr::from_ref(&stream_error_state) as usize, + stream: make_stream_state(), + _library_guard: None, + active: true, + }; + forward_native_async_next_stream( + LlmJsonStream::new(tokio_stream::iter([Err(FlowError::Internal( + "provider stream failed".into(), + ))])), + &mut stream_error_guard, + ) + .await; assert!( - json_from_native_string(ptr::null_mut(), "generic fallback") - .unwrap_err() - .to_string() - .contains("generic fallback") + stream_error_state + .error + .lock() + .unwrap() + .as_deref() + .is_some_and(|error| error.contains("provider stream failed")) ); +} - let invalid_json = native_string("{"); - assert!( - take_json_from_native_string(invalid_json, "unused") - .unwrap_err() - .to_string() - .contains("invalid JSON") +#[tokio::test] +async fn native_async_result_entrypoint_covers_llm_and_stream_continuations() { + let llm_next = Arc::new(NativeAsyncNext::new( + NativeAsyncNextInner::Llm(Arc::new(|request| { + Box::pin(async move { Ok(request.content) }) + })), + tokio::runtime::Handle::current(), + None, + )); + let llm_ref = Arc::into_raw(llm_next) as *const NemoRelayNativeAsyncNext; + let invalid = native_string("{}"); + assert_eq!( + unsafe { + native_async_next_invoke_result( + llm_ref, + invalid, + complete_native_next_result, + ptr::null_mut(), + ) + }, + NemoRelayStatus::InvalidJson ); + unsafe { native_string_free(invalid) }; + + let request = native_string_from_json( + &serde_json::to_value(LlmRequest { + headers: Map::new(), + content: json!({"message": "hello"}), + }) + .unwrap(), + ) + .unwrap(); + let (sender, receiver) = tokio::sync::oneshot::channel::>(); assert_eq!( - optional_json_from_native_string(ptr::null(), "optional"), - Ok(None) + unsafe { + native_async_next_invoke_result( + llm_ref, + request, + complete_native_next_result, + Box::into_raw(Box::new(sender)).cast(), + ) + }, + NemoRelayStatus::Ok ); - let valid_json = native_string(r#"{"value":1}"#); assert_eq!( - optional_json_from_native_string(valid_json, "optional").unwrap(), - Some(json!({"value": 1})) + receiver.await.unwrap().unwrap(), + json!({"message": "hello"}) ); - unsafe { native_string_free(valid_json) }; - let invalid_json = native_string("not-json"); + unsafe { + native_string_free(request); + native_async_next_release(llm_ref); + } + + let stream_next = Arc::new(NativeAsyncNext::new( + NativeAsyncNextInner::LlmStream(Arc::new(|_request| { + Box::pin(async { Ok(LlmJsonStream::new(tokio_stream::empty())) }) + })), + tokio::runtime::Handle::current(), + None, + )); + let stream_ref = Arc::into_raw(stream_next) as *const NemoRelayNativeAsyncNext; + let invocation = native_string("null"); assert_eq!( - optional_json_from_native_string(invalid_json, "optional"), - Err(NemoRelayStatus::InvalidJson) + unsafe { + native_async_next_invoke_result( + stream_ref, + invocation, + complete_native_next_result, + ptr::null_mut(), + ) + }, + NemoRelayStatus::InvalidArg ); - assert_last_error_contains("optional is not valid JSON"); - unsafe { native_string_free(invalid_json) }; + unsafe { + native_string_free(invocation); + native_async_next_release(stream_ref); + } +} +#[test] +fn native_async_stream_entrypoints_cover_closed_full_and_settled_channels() { + let chunk = native_string("null"); + + let no_sender = Arc::new(NativeAsyncStream { + sender: Mutex::new(None), + cancelled: AtomicBool::new(false), + settled: AtomicBool::new(false), + downstream_aborts: Mutex::new(HashMap::new()), + settlement: Mutex::new(()), + before_settlement_lock: None, + _callback_user_data: None, + }); + let no_sender_ref = Arc::into_raw(no_sender) as *const NemoRelayNativeAsyncStream; assert_eq!( - parse_json_arg(ptr::null(), "null JSON").unwrap_err(), - NemoRelayStatus::InvalidJson + unsafe { native_async_stream_push_json(no_sender_ref, chunk) }, + NemoRelayStatus::InvalidArg ); - let request = LlmRequest { - headers: Map::new(), - content: json!({"model": "test"}), - }; - let request_json = native_string_from_json(&serde_json::to_value(&request).unwrap()).unwrap(); assert_eq!( - parse_llm_request_arg(request_json, "request").unwrap(), - request + unsafe { native_async_stream_finish(no_sender_ref) }, + NemoRelayStatus::InvalidArg ); - unsafe { native_string_free(request_json) }; - let wrong_shape = native_string(r#"{"headers":[]}"#); assert_eq!( - parse_llm_request_arg(wrong_shape, "request").unwrap_err(), - NemoRelayStatus::InvalidJson + unsafe { native_async_stream_reject(no_sender_ref, chunk) }, + NemoRelayStatus::InvalidArg ); - assert_last_error_contains("was not an LLM request"); - unsafe { native_string_free(wrong_shape) }; + unsafe { native_async_stream_release(no_sender_ref) }; + let (full_sender, _full_receiver) = tokio::sync::mpsc::channel::>(1); + full_sender.try_send(Ok(Json::Null)).unwrap(); + let full = Arc::new(NativeAsyncStream { + sender: Mutex::new(Some(full_sender)), + cancelled: AtomicBool::new(false), + settled: AtomicBool::new(false), + downstream_aborts: Mutex::new(HashMap::new()), + settlement: Mutex::new(()), + before_settlement_lock: None, + _callback_user_data: None, + }); + let full_ref = Arc::into_raw(full) as *const NemoRelayNativeAsyncStream; assert_eq!( - write_native_json(&json!({"ok": true}), ptr::null_mut()), - NemoRelayStatus::NullPointer + unsafe { native_async_stream_push_json(full_ref, chunk) }, + NemoRelayStatus::Internal ); - let mut json_out = ptr::null_mut(); assert_eq!( - write_native_json(&json!({"ok": true}), &mut json_out), - NemoRelayStatus::Ok + unsafe { native_async_stream_reject(full_ref, chunk) }, + NemoRelayStatus::Internal ); + unsafe { native_async_stream_release(full_ref) }; + + let (closed_sender, closed_receiver) = tokio::sync::mpsc::channel::>(1); + drop(closed_receiver); + let closed = Arc::new(NativeAsyncStream { + sender: Mutex::new(Some(closed_sender)), + cancelled: AtomicBool::new(false), + settled: AtomicBool::new(false), + downstream_aborts: Mutex::new(HashMap::new()), + settlement: Mutex::new(()), + before_settlement_lock: None, + _callback_user_data: None, + }); + let closed_ref = Arc::into_raw(closed) as *const NemoRelayNativeAsyncStream; assert_eq!( - take_json_from_native_string(json_out, "unused").unwrap(), - json!({"ok": true}) + unsafe { native_async_stream_push_json(closed_ref, chunk) }, + NemoRelayStatus::InvalidArg ); + assert_eq!( + unsafe { native_async_stream_reject(closed_ref, chunk) }, + NemoRelayStatus::InvalidArg + ); + unsafe { native_async_stream_release(closed_ref) }; - let host_api = unsafe { &*native_host_api() }; - assert_eq!(host_api.abi_version, NEMO_RELAY_NATIVE_ABI_VERSION); + let (settled_sender, _settled_receiver) = tokio::sync::mpsc::channel(1); + let settled = Arc::new(NativeAsyncStream { + sender: Mutex::new(Some(settled_sender)), + cancelled: AtomicBool::new(false), + settled: AtomicBool::new(true), + downstream_aborts: Mutex::new(HashMap::new()), + settlement: Mutex::new(()), + before_settlement_lock: None, + _callback_user_data: None, + }); + let settled_ref = Arc::into_raw(settled) as *const NemoRelayNativeAsyncStream; assert_eq!( - host_api.struct_size, - std::mem::size_of::() + unsafe { native_async_stream_push_json(settled_ref, chunk) }, + NemoRelayStatus::InvalidArg + ); + assert_eq!( + unsafe { native_async_stream_finish(settled_ref) }, + NemoRelayStatus::InvalidArg + ); + assert_eq!( + unsafe { native_async_stream_reject(settled_ref, chunk) }, + NemoRelayStatus::InvalidArg ); + unsafe { + native_async_stream_release(settled_ref); + native_string_free(chunk); + } } -#[test] -fn native_async_next_abi_runs_tool_llm_and_stream_continuations() { - let runtime = tokio::runtime::Builder::new_current_thread() - .enable_all() - .build() - .unwrap(); - let cases: Vec<(NativeAsyncNextInner, Json, Json)> = vec![ - ( - NativeAsyncNextInner::Tool(Arc::new(|value| Box::pin(async move { Ok(value) }))), - json!({"tool": true}), - json!({"result": {"tool": true}, "pending_marks": []}), - ), - ( - NativeAsyncNextInner::Llm(Arc::new(|request| { - Box::pin(async move { Ok(request.content) }) - })), - serde_json::to_value(LlmRequest { - headers: Map::new(), - content: json!({"llm": true}), - }) - .unwrap(), - json!({"llm": true}), - ), - ]; - - for (inner, invocation, expected) in cases { - let next = Arc::new(NativeAsyncNext::new(inner, runtime.handle().clone(), None)); +#[tokio::test] +async fn native_async_result_entrypoint_reports_provider_errors_and_panics() { + for next_fn in [ + Arc::new(|_value| { + Box::pin(async { Err(FlowError::Internal("provider failed".into())) }) + as Pin> + Send>> + }) as ToolExecutionNextFn, + Arc::new(|_value| { + Box::pin(async { + panic!("provider panicked"); + #[allow(unreachable_code)] + Ok(Json::Null) + }) as Pin> + Send>> + }) as ToolExecutionNextFn, + ] { + let next = Arc::new(NativeAsyncNext::new( + NativeAsyncNextInner::Tool(next_fn), + tokio::runtime::Handle::current(), + None, + )); let next_ref = Arc::into_raw(next) as *const NemoRelayNativeAsyncNext; - let (sender, receiver) = tokio::sync::oneshot::channel(); - let completion = Arc::new(NativeAsyncCompletion { - sender: Mutex::new(Some(sender)), - cancelled: AtomicBool::new(false), - next_invoked: AtomicBool::new(false), - next_abort: Mutex::new(None), - before_settlement_lock: None, - _callback_user_data: None, - }); - let completion_ref = - Arc::into_raw(Arc::clone(&completion)) as *const NemoRelayNativeAsyncCompletion; - let invocation = native_string_from_json(&invocation).unwrap(); + let invocation = native_string("null"); + let (sender, receiver) = + tokio::sync::oneshot::channel::>(); assert_eq!( - unsafe { native_async_next_invoke(next_ref, invocation, completion_ref) }, + unsafe { + native_async_next_invoke_result( + next_ref, + invocation, + complete_native_next_result, + Box::into_raw(Box::new(sender)).cast(), + ) + }, NemoRelayStatus::Ok ); - assert_eq!(runtime.block_on(receiver).unwrap().unwrap(), expected); + let error = receiver.await.unwrap().unwrap_err(); + assert!(error.contains("provider")); unsafe { native_string_free(invocation); native_async_next_release(next_ref); - native_async_completion_release(completion_ref); } } +} - let next = Arc::new(NativeAsyncNext::new( - NativeAsyncNextInner::LlmStream(Arc::new(|_request| { - Box::pin(async { - Ok(LlmJsonStream::new(tokio_stream::iter(vec![ - Ok(json!({"chunk": 1})), - Ok(json!({"chunk": 2})), - ]))) - }) - })), - runtime.handle().clone(), - None, - )); - let next_ref = Arc::into_raw(next) as *const NemoRelayNativeAsyncNext; - let (sender, _receiver) = tokio::sync::oneshot::channel(); - let completion = Arc::new(NativeAsyncCompletion { +#[tokio::test] +async fn native_async_stream_next_entrypoint_validates_handle_kind_and_request() { + let (sender, _receiver) = tokio::sync::mpsc::channel(1); + let stream = Arc::new(NativeAsyncStream { sender: Mutex::new(Some(sender)), cancelled: AtomicBool::new(false), - next_invoked: AtomicBool::new(false), - next_abort: Mutex::new(None), + settled: AtomicBool::new(false), + downstream_aborts: Mutex::new(HashMap::new()), + settlement: Mutex::new(()), before_settlement_lock: None, _callback_user_data: None, }); - let completion_ref = - Arc::into_raw(Arc::clone(&completion)) as *const NemoRelayNativeAsyncCompletion; - let invocation = native_string_from_json( - &serde_json::to_value(LlmRequest { - headers: Map::new(), - content: json!({"stream": true}), - }) - .unwrap(), - ) - .unwrap(); - assert_eq!( - unsafe { native_async_next_invoke(next_ref, invocation, completion_ref) }, - NemoRelayStatus::InvalidArg - ); - assert_last_error_contains("async_next_invoke_stream"); - unsafe { - native_string_free(invocation); - native_async_next_release(next_ref); - native_async_completion_release(completion_ref); - } -} + let stream_ref = Arc::into_raw(stream) as *const NemoRelayNativeAsyncStream; + let invocation = native_string("null"); -#[test] -fn native_async_next_reports_a_revoked_continuation_without_calling_the_provider() { - let runtime = tokio::runtime::Builder::new_current_thread() - .enable_all() - .build() - .unwrap(); - let provider_calls = Arc::new(AtomicUsize::new(0)); - let (lease, guard) = MiddlewareContinuationLease::capture(); - let next = Arc::new(NativeAsyncNext::new( - NativeAsyncNextInner::Tool({ - let provider_calls = provider_calls.clone(); - Arc::new(move |value| { - let provider_calls = provider_calls.clone(); - let invocation = lease.begin(); - Box::pin(async move { - invocation? - .invoke(|| async move { - provider_calls.fetch_add(1, Ordering::SeqCst); - Ok(value) - }) - .await - }) - }) - }), - runtime.handle().clone(), + let tool_next = Arc::new(NativeAsyncNext::new( + NativeAsyncNextInner::Tool(Arc::new(|value| Box::pin(async move { Ok(value) }))), + tokio::runtime::Handle::current(), None, )); - let next_ref = Arc::into_raw(next) as *const NemoRelayNativeAsyncNext; - let (sender, receiver) = tokio::sync::oneshot::channel(); - let completion = Arc::new(NativeAsyncCompletion { - sender: Mutex::new(Some(sender)), - cancelled: AtomicBool::new(false), - next_invoked: AtomicBool::new(false), - next_abort: Mutex::new(None), - before_settlement_lock: None, - _callback_user_data: None, - }); - let completion_ref = - Arc::into_raw(Arc::clone(&completion)) as *const NemoRelayNativeAsyncCompletion; - let invocation = native_string_from_json(&json!({"tool": true})).unwrap(); - - drop(guard); + let tool_ref = Arc::into_raw(tool_next) as *const NemoRelayNativeAsyncNext; assert_eq!( - unsafe { native_async_next_invoke(next_ref, invocation, completion_ref) }, - NemoRelayStatus::Ok - ); - let error = runtime - .block_on(receiver) - .expect("native completion should settle") - .expect_err("revoked continuation should reject"); - assert!( - error - .to_string() - .contains("execution continuation is no longer active") + unsafe { + native_async_next_invoke_stream( + tool_ref, + invocation, + stream_ref, + accept_native_stream_item, + ptr::null_mut(), + ) + }, + NemoRelayStatus::InvalidArg ); - assert_eq!(provider_calls.load(Ordering::SeqCst), 0); - - unsafe { - native_string_free(invocation); - native_async_next_release(next_ref); - native_async_completion_release(completion_ref); - } -} -#[test] -fn native_async_next_result_supports_repeated_concurrent_calls() { - let runtime = tokio::runtime::Builder::new_current_thread() - .enable_all() - .build() - .unwrap(); - let provider_calls = Arc::new(AtomicUsize::new(0)); - let next = Arc::new(NativeAsyncNext::new( - NativeAsyncNextInner::Tool({ - let provider_calls = provider_calls.clone(); - Arc::new(move |value| { - provider_calls.fetch_add(1, Ordering::SeqCst); - Box::pin(async move { - tokio::task::yield_now().await; - Ok(json!({ - "value": value, - "scope": crate::api::runtime::task_scope_top().uuid.to_string(), - })) - }) - }) - }), - runtime.handle().clone(), + let stream_next = Arc::new(NativeAsyncNext::new( + NativeAsyncNextInner::LlmStream(Arc::new(|_request| { + Box::pin(async { Ok(LlmJsonStream::new(tokio_stream::empty())) }) + })), + tokio::runtime::Handle::current(), None, )); - let next_ref = Arc::into_raw(next) as *const NemoRelayNativeAsyncNext; - let first = native_string_from_json(&json!({"branch": "first"})).unwrap(); - let second = native_string_from_json(&json!({"branch": "second"})).unwrap(); - let (first_tx, first_rx) = tokio::sync::oneshot::channel::>(); - let (second_tx, second_rx) = - tokio::sync::oneshot::channel::>(); - let first_stack = create_scope_stack(); - let first_scope = first_stack - .read() - .unwrap_or_else(|error| error.into_inner()) - .top() - .uuid - .to_string(); - let second_stack = create_scope_stack(); - let second_scope = second_stack - .read() - .unwrap_or_else(|error| error.into_inner()) - .top() - .uuid - .to_string(); - + let next_ref = Arc::into_raw(stream_next) as *const NemoRelayNativeAsyncNext; assert_eq!( - with_scope_stack(first_stack, || unsafe { - native_async_next_invoke_result( + unsafe { + native_async_next_invoke_stream( next_ref, - first, - complete_native_next_result, - Box::into_raw(Box::new(first_tx)).cast(), + invocation, + ptr::null(), + accept_native_stream_item, + ptr::null_mut(), ) - }), - NemoRelayStatus::Ok + }, + NemoRelayStatus::NullPointer ); assert_eq!( - with_scope_stack(second_stack, || unsafe { - native_async_next_invoke_result( + unsafe { + native_async_next_invoke_stream( next_ref, - second, - complete_native_next_result, - Box::into_raw(Box::new(second_tx)).cast(), + invocation, + stream_ref, + accept_native_stream_item, + ptr::null_mut(), ) - }), - NemoRelayStatus::Ok - ); - let (first_result, second_result) = - runtime.block_on(async { tokio::join!(first_rx, second_rx) }); - - assert_eq!( - first_result.unwrap().unwrap(), - json!({ - "value": {"branch": "first"}, - "scope": first_scope, - }) - ); - assert_eq!( - second_result.unwrap().unwrap(), - json!({ - "value": {"branch": "second"}, - "scope": second_scope, - }) + }, + NemoRelayStatus::InvalidJson ); - assert_eq!(provider_calls.load(Ordering::SeqCst), 2); unsafe { - native_string_free(first); - native_string_free(second); + native_string_free(invocation); + native_async_next_release(tool_ref); native_async_next_release(next_ref); + native_async_stream_release(stream_ref); } } -#[test] -fn native_async_next_result_uses_captured_scope_on_an_unbound_plugin_thread() { - let runtime = tokio::runtime::Builder::new_current_thread() - .enable_all() - .build() - .unwrap(); - let captured_stack = create_scope_stack(); - let captured_scope = captured_stack - .read() - .unwrap_or_else(|error| error.into_inner()) - .top() - .uuid - .to_string(); - let next = with_scope_stack(captured_stack, || { - Arc::new(NativeAsyncNext::new( - NativeAsyncNextInner::Tool(Arc::new(|value| { - Box::pin(async move { - Ok(json!({ - "value": value, - "scope": crate::api::runtime::task_scope_top().uuid.to_string(), - })) - }) - })), - runtime.handle().clone(), - None, - )) - }); - let next_ref = Arc::into_raw(next) as *const NemoRelayNativeAsyncNext; - let invocation = native_string_from_json(&json!({"thread": "plugin"})).unwrap(); - let (sender, receiver) = tokio::sync::oneshot::channel::>(); - let next_address = next_ref as usize; - let invocation_address = invocation as usize; - let sender_address = Box::into_raw(Box::new(sender)) as usize; +struct InvokeNativeNextThenReturnState { + callback_state: u32, + invoke_status: AtomicUsize, + started: Mutex>, +} + +unsafe extern "C" fn invoke_native_next_then_return_state( + user_data: *mut c_void, + invocation_json: *const NemoRelayNativeString, + next: *const NemoRelayNativeAsyncNext, + completion: *const NemoRelayNativeAsyncCompletion, +) -> u32 { + let state = unsafe { &*user_data.cast::() }; + let status = unsafe { native_async_next_invoke(next, invocation_json, completion) }; + state + .invoke_status + .store(status as usize, Ordering::Release); + if status == NemoRelayStatus::Ok { + let _ = state + .started + .lock() + .unwrap_or_else(|error| error.into_inner()) + .recv_timeout(Duration::from_secs(1)); + } + unsafe { native_async_next_release(next) }; + state.callback_state +} + +unsafe extern "C" fn invoke_native_stream_next_then_return_state( + user_data: *mut c_void, + invocation_json: *const NemoRelayNativeString, + next: *const NemoRelayNativeAsyncNext, + stream: *const NemoRelayNativeAsyncStream, +) -> u32 { + let state = unsafe { &*user_data.cast::() }; + let request = read_native_string(invocation_json) + .ok() + .and_then(|invocation| serde_json::from_str::(&invocation).ok()) + .and_then(|invocation| invocation.get("request").cloned()) + .and_then(|request| native_string_from_json(&request)); + let status = if let Some(request) = request { + let status = unsafe { + native_async_next_invoke_stream( + next, + request, + stream, + accept_native_stream_item, + ptr::null_mut(), + ) + }; + unsafe { native_string_free(request) }; + status + } else { + NemoRelayStatus::InvalidJson + }; + state + .invoke_status + .store(status as usize, Ordering::Release); + if status == NemoRelayStatus::Ok { + let _ = state + .started + .lock() + .unwrap_or_else(|error| error.into_inner()) + .recv_timeout(Duration::from_secs(1)); + } + unsafe { + native_async_next_release(next); + native_async_stream_release(stream); + } + state.callback_state +} + +unsafe extern "C" fn invoke_detached_next_and_finish_replacement_stream( + user_data: *mut c_void, + invocation_json: *const NemoRelayNativeString, + next: *const NemoRelayNativeAsyncNext, + stream: *const NemoRelayNativeAsyncStream, +) -> u32 { + let invoke_status = unsafe { &*user_data.cast::() }; + let request = read_native_string(invocation_json) + .ok() + .and_then(|invocation| serde_json::from_str::(&invocation).ok()) + .and_then(|invocation| invocation.get("request").cloned()) + .and_then(|request| native_string_from_json(&request)); + let status = if let Some(request) = request { + let status = unsafe { + native_async_next_invoke_stream( + next, + request, + stream, + accept_native_stream_item, + ptr::null_mut(), + ) + }; + unsafe { native_string_free(request) }; + status + } else { + NemoRelayStatus::InvalidJson + }; + invoke_status.store(status as usize, Ordering::Release); + let replacement = native_string_from_json(&json!({"source": "replacement"})).unwrap(); + unsafe { + let _ = native_async_stream_push_json(stream, replacement); + native_string_free(replacement); + let _ = native_async_stream_finish(stream); + native_async_next_release(next); + native_async_stream_release(stream); + } + NemoRelayNativeAsyncCallbackState::Complete as u32 +} + +struct FailingNativeCodec; + +impl LlmCodec for FailingNativeCodec { + fn decode(&self, _request: &LlmRequest) -> FlowResult { + Err(FlowError::Internal("request decode rejected".into())) + } + + fn encode( + &self, + _annotated: &AnnotatedLlmRequest, + _original: &LlmRequest, + ) -> FlowResult { + Err(FlowError::Internal("request encode rejected".into())) + } +} + +impl LlmResponseCodec for FailingNativeCodec { + fn decode_response(&self, _response: &Json) -> FlowResult { + Err(FlowError::Internal("response decode rejected".into())) + } +} + +struct PanickingNativeCodec; + +impl LlmCodec for PanickingNativeCodec { + fn decode(&self, _request: &LlmRequest) -> FlowResult { + panic!("request decode panic") + } + + fn encode( + &self, + _annotated: &AnnotatedLlmRequest, + _original: &LlmRequest, + ) -> FlowResult { + panic!("request encode panic") + } +} + +impl LlmResponseCodec for PanickingNativeCodec { + fn decode_response(&self, _response: &Json) -> FlowResult { + panic!("response decode panic") + } +} + +#[test] +fn native_string_and_json_helpers_cover_abi_boundaries() { + assert_native_string_allocation_boundaries(); + assert_native_error_string_boundaries(); + assert_native_json_parsing_boundaries(); + assert_native_json_output_and_host_api(); +} + +fn assert_native_string_allocation_boundaries() { + clear_native_last_error(); + assert_eq!( + unsafe { native_string_new(ptr::null(), 0, ptr::null_mut()) }, + NemoRelayStatus::NullPointer + ); + assert_last_error_contains("out string pointer is null"); + + let mut out = ptr::null_mut(); + assert_eq!( + unsafe { native_string_new(ptr::null(), 1, &mut out) }, + NemoRelayStatus::NullPointer + ); + assert!(out.is_null()); + assert_last_error_contains("string data pointer is null"); + + let invalid_utf8 = [0xff]; + assert_eq!( + unsafe { native_string_new(invalid_utf8.as_ptr(), invalid_utf8.len(), &mut out) }, + NemoRelayStatus::InvalidUtf8 + ); + assert!(out.is_null()); + assert_last_error_contains("not valid UTF-8"); + + let text = native_string("hello"); + assert_eq!(unsafe { native_string_len(text) }, 5); + assert_eq!( + unsafe { std::slice::from_raw_parts(native_string_data(text), 5) }, + b"hello" + ); + assert!(unsafe { native_string_data(ptr::null()) }.is_null()); + assert_eq!(unsafe { native_string_len(ptr::null()) }, 0); + assert_eq!(take_native_string(text).unwrap(), "hello"); + unsafe { native_string_free(ptr::null_mut()) }; + + let empty = native_string(""); + assert_eq!(read_native_string(empty).unwrap(), ""); + unsafe { native_string_free(empty) }; + assert_eq!(read_native_string(ptr::null()).unwrap(), ""); +} + +fn assert_native_error_string_boundaries() { + let bad = Box::into_raw(Box::new(NativeHostString(vec![0xff]))) as *mut NemoRelayNativeString; + assert!(read_native_string(bad).is_err()); + assert_eq!( + optional_json_from_native_string(bad, "bad json"), + Err(NemoRelayStatus::InvalidUtf8) + ); + unsafe { native_last_error_set(bad) }; + assert_last_error_contains("not valid UTF-8"); + unsafe { native_string_free(bad) }; + + let message = native_string("explicit native error"); + unsafe { native_last_error_set(message) }; + assert_eq!( + native_last_error_message().as_deref(), + Some("explicit native error") + ); + unsafe { native_string_free(message) }; + unsafe { native_last_error_clear() }; + assert!(native_last_error_message().is_none()); + + set_native_last_error("specific fallback"); + assert!( + json_from_native_string(ptr::null_mut(), "generic fallback") + .unwrap_err() + .to_string() + .contains("specific fallback") + ); + clear_native_last_error(); + assert!( + json_from_native_string(ptr::null_mut(), "generic fallback") + .unwrap_err() + .to_string() + .contains("generic fallback") + ); +} + +fn assert_native_json_parsing_boundaries() { + let invalid_json = native_string("{"); + assert!( + take_json_from_native_string(invalid_json, "unused") + .unwrap_err() + .to_string() + .contains("invalid JSON") + ); + assert_eq!( + optional_json_from_native_string(ptr::null(), "optional"), + Ok(None) + ); + let valid_json = native_string(r#"{"value":1}"#); + assert_eq!( + optional_json_from_native_string(valid_json, "optional").unwrap(), + Some(json!({"value": 1})) + ); + unsafe { native_string_free(valid_json) }; + let invalid_json = native_string("not-json"); + assert_eq!( + optional_json_from_native_string(invalid_json, "optional"), + Err(NemoRelayStatus::InvalidJson) + ); + assert_last_error_contains("optional is not valid JSON"); + unsafe { native_string_free(invalid_json) }; + + assert_eq!( + parse_json_arg(ptr::null(), "null JSON").unwrap_err(), + NemoRelayStatus::InvalidJson + ); + let request = LlmRequest { + headers: Map::new(), + content: json!({"model": "test"}), + }; + let request_json = native_string_from_json(&serde_json::to_value(&request).unwrap()).unwrap(); + assert_eq!( + parse_llm_request_arg(request_json, "request").unwrap(), + request + ); + unsafe { native_string_free(request_json) }; + let wrong_shape = native_string(r#"{"headers":[]}"#); + assert_eq!( + parse_llm_request_arg(wrong_shape, "request").unwrap_err(), + NemoRelayStatus::InvalidJson + ); + assert_last_error_contains("was not an LLM request"); + unsafe { native_string_free(wrong_shape) }; +} + +fn assert_native_json_output_and_host_api() { + assert_eq!( + write_native_json(&json!({"ok": true}), ptr::null_mut()), + NemoRelayStatus::NullPointer + ); + let mut json_out = ptr::null_mut(); + assert_eq!( + write_native_json(&json!({"ok": true}), &mut json_out), + NemoRelayStatus::Ok + ); + assert_eq!( + take_json_from_native_string(json_out, "unused").unwrap(), + json!({"ok": true}) + ); + + let host_api = unsafe { &*native_host_api() }; + assert_eq!(host_api.abi_version, NEMO_RELAY_NATIVE_ABI_VERSION); + assert_eq!( + host_api.struct_size, + std::mem::size_of::() + ); +} + +#[test] +fn native_async_next_abi_runs_tool_llm_and_stream_continuations() { + let runtime = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .unwrap(); + let cases: Vec<(NativeAsyncNextInner, Json, Json)> = vec![ + ( + NativeAsyncNextInner::Tool(Arc::new(|value| Box::pin(async move { Ok(value) }))), + json!({"tool": true}), + json!({"result": {"tool": true}, "pending_marks": []}), + ), + ( + NativeAsyncNextInner::Llm(Arc::new(|request| { + Box::pin(async move { Ok(request.content) }) + })), + serde_json::to_value(LlmRequest { + headers: Map::new(), + content: json!({"llm": true}), + }) + .unwrap(), + json!({"llm": true}), + ), + ]; + + for (inner, invocation, expected) in cases { + let next = Arc::new(NativeAsyncNext::new(inner, runtime.handle().clone(), None)); + let next_ref = Arc::into_raw(next) as *const NemoRelayNativeAsyncNext; + let (sender, receiver) = tokio::sync::oneshot::channel(); + let completion = Arc::new(NativeAsyncCompletion { + sender: Mutex::new(Some(sender)), + cancelled: AtomicBool::new(false), + next_invoked: AtomicBool::new(false), + next_abort: Mutex::new(None), + before_settlement_lock: None, + _callback_user_data: None, + }); + let completion_ref = + Arc::into_raw(Arc::clone(&completion)) as *const NemoRelayNativeAsyncCompletion; + let invocation = native_string_from_json(&invocation).unwrap(); + assert_eq!( + unsafe { native_async_next_invoke(next_ref, invocation, completion_ref) }, + NemoRelayStatus::Ok + ); + assert_eq!(runtime.block_on(receiver).unwrap().unwrap(), expected); + unsafe { + native_string_free(invocation); + native_async_next_release(next_ref); + native_async_completion_release(completion_ref); + } + } + + let next = Arc::new(NativeAsyncNext::new( + NativeAsyncNextInner::LlmStream(Arc::new(|_request| { + Box::pin(async { + Ok(LlmJsonStream::new(tokio_stream::iter(vec![ + Ok(json!({"chunk": 1})), + Ok(json!({"chunk": 2})), + ]))) + }) + })), + runtime.handle().clone(), + None, + )); + let next_ref = Arc::into_raw(next) as *const NemoRelayNativeAsyncNext; + let (sender, _receiver) = tokio::sync::oneshot::channel(); + let completion = Arc::new(NativeAsyncCompletion { + sender: Mutex::new(Some(sender)), + cancelled: AtomicBool::new(false), + next_invoked: AtomicBool::new(false), + next_abort: Mutex::new(None), + before_settlement_lock: None, + _callback_user_data: None, + }); + let completion_ref = + Arc::into_raw(Arc::clone(&completion)) as *const NemoRelayNativeAsyncCompletion; + let invocation = native_string_from_json( + &serde_json::to_value(LlmRequest { + headers: Map::new(), + content: json!({"stream": true}), + }) + .unwrap(), + ) + .unwrap(); + assert_eq!( + unsafe { native_async_next_invoke(next_ref, invocation, completion_ref) }, + NemoRelayStatus::InvalidArg + ); + assert_last_error_contains("async_next_invoke_stream"); + unsafe { + native_string_free(invocation); + native_async_next_release(next_ref); + native_async_completion_release(completion_ref); + } +} + +#[test] +fn native_async_next_reports_a_revoked_continuation_without_calling_the_provider() { + let runtime = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .unwrap(); + let provider_calls = Arc::new(AtomicUsize::new(0)); + let (lease, guard) = MiddlewareContinuationLease::capture(); + let next = Arc::new(NativeAsyncNext::new( + NativeAsyncNextInner::Tool({ + let provider_calls = provider_calls.clone(); + Arc::new(move |value| { + let provider_calls = provider_calls.clone(); + let invocation = lease.begin(); + Box::pin(async move { + invocation? + .invoke(|| async move { + provider_calls.fetch_add(1, Ordering::SeqCst); + Ok(value) + }) + .await + }) + }) + }), + runtime.handle().clone(), + None, + )); + let next_ref = Arc::into_raw(next) as *const NemoRelayNativeAsyncNext; + let (sender, receiver) = tokio::sync::oneshot::channel(); + let completion = Arc::new(NativeAsyncCompletion { + sender: Mutex::new(Some(sender)), + cancelled: AtomicBool::new(false), + next_invoked: AtomicBool::new(false), + next_abort: Mutex::new(None), + before_settlement_lock: None, + _callback_user_data: None, + }); + let completion_ref = + Arc::into_raw(Arc::clone(&completion)) as *const NemoRelayNativeAsyncCompletion; + let invocation = native_string_from_json(&json!({"tool": true})).unwrap(); + + drop(guard); + assert_eq!( + unsafe { native_async_next_invoke(next_ref, invocation, completion_ref) }, + NemoRelayStatus::Ok + ); + let error = runtime + .block_on(receiver) + .expect("native completion should settle") + .expect_err("revoked continuation should reject"); + assert!( + error + .to_string() + .contains("execution continuation is no longer active") + ); + assert_eq!(provider_calls.load(Ordering::SeqCst), 0); + + unsafe { + native_string_free(invocation); + native_async_next_release(next_ref); + native_async_completion_release(completion_ref); + } +} + +#[test] +fn native_async_next_result_supports_repeated_concurrent_calls() { + let runtime = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .unwrap(); + let provider_calls = Arc::new(AtomicUsize::new(0)); + let next = Arc::new(NativeAsyncNext::new( + NativeAsyncNextInner::Tool({ + let provider_calls = provider_calls.clone(); + Arc::new(move |value| { + provider_calls.fetch_add(1, Ordering::SeqCst); + Box::pin(async move { + tokio::task::yield_now().await; + Ok(json!({ + "value": value, + "scope": crate::api::runtime::task_scope_top().uuid.to_string(), + })) + }) + }) + }), + runtime.handle().clone(), + None, + )); + let next_ref = Arc::into_raw(next) as *const NemoRelayNativeAsyncNext; + let first = native_string_from_json(&json!({"branch": "first"})).unwrap(); + let second = native_string_from_json(&json!({"branch": "second"})).unwrap(); + let (first_tx, first_rx) = tokio::sync::oneshot::channel::>(); + let (second_tx, second_rx) = + tokio::sync::oneshot::channel::>(); + let first_stack = create_scope_stack(); + let first_scope = first_stack + .read() + .unwrap_or_else(|error| error.into_inner()) + .top() + .uuid + .to_string(); + let second_stack = create_scope_stack(); + let second_scope = second_stack + .read() + .unwrap_or_else(|error| error.into_inner()) + .top() + .uuid + .to_string(); + + assert_eq!( + with_scope_stack(first_stack, || unsafe { + native_async_next_invoke_result( + next_ref, + first, + complete_native_next_result, + Box::into_raw(Box::new(first_tx)).cast(), + ) + }), + NemoRelayStatus::Ok + ); + assert_eq!( + with_scope_stack(second_stack, || unsafe { + native_async_next_invoke_result( + next_ref, + second, + complete_native_next_result, + Box::into_raw(Box::new(second_tx)).cast(), + ) + }), + NemoRelayStatus::Ok + ); + let (first_result, second_result) = + runtime.block_on(async { tokio::join!(first_rx, second_rx) }); + + assert_eq!( + first_result.unwrap().unwrap(), + json!({ + "value": {"branch": "first"}, + "scope": first_scope, + }) + ); + assert_eq!( + second_result.unwrap().unwrap(), + json!({ + "value": {"branch": "second"}, + "scope": second_scope, + }) + ); + assert_eq!(provider_calls.load(Ordering::SeqCst), 2); + + unsafe { + native_string_free(first); + native_string_free(second); + native_async_next_release(next_ref); + } +} + +#[test] +fn native_async_next_result_uses_captured_scope_on_an_unbound_plugin_thread() { + let runtime = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .unwrap(); + let captured_stack = create_scope_stack(); + let captured_scope = captured_stack + .read() + .unwrap_or_else(|error| error.into_inner()) + .top() + .uuid + .to_string(); + let next = with_scope_stack(captured_stack, || { + Arc::new(NativeAsyncNext::new( + NativeAsyncNextInner::Tool(Arc::new(|value| { + Box::pin(async move { + Ok(json!({ + "value": value, + "scope": crate::api::runtime::task_scope_top().uuid.to_string(), + })) + }) + })), + runtime.handle().clone(), + None, + )) + }); + let next_ref = Arc::into_raw(next) as *const NemoRelayNativeAsyncNext; + let invocation = native_string_from_json(&json!({"thread": "plugin"})).unwrap(); + let (sender, receiver) = tokio::sync::oneshot::channel::>(); + let next_address = next_ref as usize; + let invocation_address = invocation as usize; + let sender_address = Box::into_raw(Box::new(sender)) as usize; assert_eq!( std::thread::spawn(move || unsafe { @@ -2581,6 +3309,10 @@ unsafe extern "C" fn count_scope_callback(user_data: *mut c_void) -> NemoRelaySt NemoRelayStatus::Ok } +unsafe extern "C" fn fail_scope_callback(_user_data: *mut c_void) -> NemoRelayStatus { + NemoRelayStatus::InvalidArg +} + #[test] fn native_scope_stack_abi_covers_lifecycle_and_validation() { let _runtime_guard = crate::shared_runtime::runtime_owner_test_mutex() @@ -2589,6 +3321,18 @@ fn native_scope_stack_abi_covers_lifecycle_and_validation() { crate::shared_runtime::reset_runtime_owner_for_tests(); let _global_context_restore = GlobalContextRestore::replace_with_empty(); let _restore = ThreadScopeStackRestore::capture(); + + assert_native_scope_stack_null_validation(); + let stack = create_active_native_scope_stack(); + let strings = NativeScopeTestStrings::new(); + let scope = assert_native_scope_lifecycle(&strings); + assert_native_scope_push_validation(&strings); + assert_native_scope_pop_and_mark_validation(scope, &strings); + assert_native_scope_stack_binding_lifecycle(); + free_native_scope_test_resources(stack, scope, strings); +} + +fn assert_native_scope_stack_null_validation() { assert_eq!( unsafe { native_scope_stack_create(ptr::null_mut()) }, NemoRelayStatus::NullPointer @@ -2615,7 +3359,25 @@ fn native_scope_stack_abi_covers_lifecycle_and_validation() { unsafe { native_scope_get_current(ptr::null_mut()) }, NemoRelayStatus::NullPointer ); + assert_eq!( + unsafe { + native_scope_push( + ptr::null(), + NemoRelayNativeScopeType::Custom, + ptr::null(), + 0, + ptr::null(), + ptr::null(), + ptr::null(), + ptr::null(), + ptr::null_mut(), + ) + }, + NemoRelayStatus::NullPointer + ); +} +fn create_active_native_scope_stack() -> *mut NemoRelayNativeScopeStack { let mut stack = ptr::null_mut(); assert_eq!( unsafe { native_scope_stack_create(&mut stack) }, @@ -2637,78 +3399,246 @@ fn native_scope_stack_abi_covers_lifecycle_and_validation() { (&calls as *const AtomicUsize).cast_mut().cast(), ) }, - NemoRelayStatus::Ok + NemoRelayStatus::Ok + ); + assert_eq!(calls.load(Ordering::SeqCst), 1); + assert_eq!( + unsafe { native_scope_stack_with_current(stack, fail_scope_callback, ptr::null_mut()) }, + NemoRelayStatus::InvalidArg + ); + assert_last_error_contains("scope-stack callback returned InvalidArg"); + + stack +} + +struct NativeScopeTestStrings { + name: *mut NemoRelayNativeString, + data: *mut NemoRelayNativeString, + metadata: *mut NemoRelayNativeString, + input: *mut NemoRelayNativeString, + mark_name: *mut NemoRelayNativeString, + output: *mut NemoRelayNativeString, + invalid: *mut NemoRelayNativeString, +} + +impl NativeScopeTestStrings { + fn new() -> Self { + Self { + name: native_string("native-scope"), + data: native_string(r#"{"source":"native"}"#), + metadata: native_string(r#"{"test":true}"#), + input: native_string(r#"{"input":1}"#), + mark_name: native_string("native-mark"), + output: native_string(r#"{"output":1}"#), + invalid: native_string("not-json"), + } + } +} + +fn assert_native_scope_lifecycle( + strings: &NativeScopeTestStrings, +) -> *mut NemoRelayNativeScopeHandle { + let timestamp = 0_i64; + let mut scope = ptr::null_mut(); + assert_eq!( + unsafe { + native_scope_push( + strings.name, + NemoRelayNativeScopeType::Custom, + ptr::null(), + 0, + strings.data, + strings.metadata, + strings.input, + ×tamp, + &mut scope, + ) + }, + NemoRelayStatus::Ok + ); + assert!(!scope.is_null()); + assert_eq!(native_scope_ref(scope).unwrap().name, "native-scope"); + + let mut current = ptr::null_mut(); + assert_eq!( + unsafe { native_scope_get_current(&mut current) }, + NemoRelayStatus::Ok + ); + assert_eq!(native_scope_ref(current).unwrap().name, "native-scope"); + unsafe { native_scope_handle_free(current) }; + + assert_eq!( + unsafe { + native_emit_mark( + strings.mark_name, + scope, + strings.data, + strings.metadata, + ×tamp, + ) + }, + NemoRelayStatus::Ok + ); + assert_eq!( + unsafe { native_scope_pop(scope, strings.output, strings.metadata, ×tamp) }, + NemoRelayStatus::Ok + ); + + scope +} + +fn assert_native_scope_push_validation(strings: &NativeScopeTestStrings) { + let mut invalid_scope = ptr::null_mut(); + assert_eq!( + unsafe { + native_scope_push( + strings.name, + NemoRelayNativeScopeType::Custom, + ptr::null(), + 0, + strings.invalid, + ptr::null(), + ptr::null(), + ptr::null(), + &mut invalid_scope, + ) + }, + NemoRelayStatus::InvalidJson + ); + assert!(invalid_scope.is_null()); + let invalid_const = strings.invalid.cast_const(); + for (data_arg, metadata_arg, input_arg) in [ + (ptr::null(), invalid_const, ptr::null()), + (ptr::null(), ptr::null(), invalid_const), + ] { + assert_eq!( + unsafe { + native_scope_push( + strings.name, + NemoRelayNativeScopeType::Custom, + ptr::null(), + 0, + data_arg, + metadata_arg, + input_arg, + ptr::null(), + &mut invalid_scope, + ) + }, + NemoRelayStatus::InvalidJson + ); + assert!(invalid_scope.is_null()); + } + let invalid_timestamp = i64::MAX; + assert_eq!( + unsafe { + native_scope_push( + strings.name, + NemoRelayNativeScopeType::Custom, + ptr::null(), + 0, + ptr::null(), + ptr::null(), + ptr::null(), + &invalid_timestamp, + &mut invalid_scope, + ) + }, + NemoRelayStatus::InvalidArg ); - assert_eq!(calls.load(Ordering::SeqCst), 1); - let name = native_string("native-scope"); - let data = native_string(r#"{"source":"native"}"#); - let metadata = native_string(r#"{"test":true}"#); - let input = native_string(r#"{"input":1}"#); - let timestamp = 0_i64; - let mut scope = ptr::null_mut(); + let invalid_name = Box::into_raw(Box::new(NativeHostString(vec![0xff]))).cast(); assert_eq!( unsafe { native_scope_push( - name, + invalid_name, NemoRelayNativeScopeType::Custom, ptr::null(), 0, - data, - metadata, - input, - ×tamp, - &mut scope, + ptr::null(), + ptr::null(), + ptr::null(), + ptr::null(), + &mut invalid_scope, ) }, - NemoRelayStatus::Ok + NemoRelayStatus::InvalidUtf8 ); - assert!(!scope.is_null()); - assert_eq!(native_scope_ref(scope).unwrap().name, "native-scope"); - - let mut current = ptr::null_mut(); assert_eq!( - unsafe { native_scope_get_current(&mut current) }, - NemoRelayStatus::Ok + unsafe { + native_emit_mark( + invalid_name, + ptr::null(), + ptr::null(), + ptr::null(), + ptr::null(), + ) + }, + NemoRelayStatus::InvalidUtf8 ); - assert_eq!(native_scope_ref(current).unwrap().name, "native-scope"); - unsafe { native_scope_handle_free(current) }; + unsafe { native_string_free(invalid_name) }; +} - let mark_name = native_string("native-mark"); +fn assert_native_scope_pop_and_mark_validation( + scope: *mut NemoRelayNativeScopeHandle, + strings: &NativeScopeTestStrings, +) { + let invalid_timestamp = i64::MAX; assert_eq!( - unsafe { native_emit_mark(mark_name, scope, data, metadata, ×tamp) }, - NemoRelayStatus::Ok + unsafe { native_scope_pop(scope, strings.invalid, ptr::null(), ptr::null()) }, + NemoRelayStatus::InvalidJson ); - let output = native_string(r#"{"output":1}"#); assert_eq!( - unsafe { native_scope_pop(scope, output, metadata, ×tamp) }, - NemoRelayStatus::Ok + unsafe { native_scope_pop(scope, ptr::null(), strings.invalid, ptr::null()) }, + NemoRelayStatus::InvalidJson + ); + assert_eq!( + unsafe { native_scope_pop(scope, ptr::null(), ptr::null(), &invalid_timestamp) }, + NemoRelayStatus::InvalidArg ); - - let invalid = native_string("not-json"); - let mut invalid_scope = ptr::null_mut(); assert_eq!( unsafe { - native_scope_push( - name, - NemoRelayNativeScopeType::Custom, + native_emit_mark( + strings.mark_name, + scope, + strings.invalid, ptr::null(), - 0, - invalid, ptr::null(), + ) + }, + NemoRelayStatus::InvalidJson + ); + assert_eq!( + unsafe { + native_emit_mark( + strings.mark_name, + scope, ptr::null(), + strings.invalid, ptr::null(), - &mut invalid_scope, ) }, NemoRelayStatus::InvalidJson ); - assert!(invalid_scope.is_null()); + assert_eq!( + unsafe { + native_emit_mark( + strings.mark_name, + scope, + ptr::null(), + ptr::null(), + &invalid_timestamp, + ) + }, + NemoRelayStatus::InvalidArg + ); assert_eq!( unsafe { native_scope_pop(ptr::null(), ptr::null(), ptr::null(), ptr::null()) }, NemoRelayStatus::NullPointer ); +} +fn assert_native_scope_stack_binding_lifecycle() { let mut binding = ptr::null_mut(); assert_eq!( unsafe { native_scope_stack_capture_thread(&mut binding) }, @@ -2724,8 +3654,22 @@ fn native_scope_stack_abi_covers_lifecycle_and_validation() { NemoRelayStatus::Ok ); unsafe { native_scope_stack_binding_free(disposable_binding) }; +} - for value in [name, data, metadata, input, mark_name, output, invalid] { +fn free_native_scope_test_resources( + stack: *mut NemoRelayNativeScopeStack, + scope: *mut NemoRelayNativeScopeHandle, + strings: NativeScopeTestStrings, +) { + for value in [ + strings.name, + strings.data, + strings.metadata, + strings.input, + strings.mark_name, + strings.output, + strings.invalid, + ] { unsafe { native_string_free(value) }; } unsafe { @@ -2791,184 +3735,841 @@ unsafe extern "C" fn noop_json( NemoRelayStatus::Ok } -unsafe extern "C" fn noop_llm_conditional( - _user_data: *mut c_void, - _request_json: *const NemoRelayNativeString, - _out_reason: *mut *mut NemoRelayNativeString, -) -> NemoRelayStatus { - NemoRelayStatus::Ok -} +unsafe extern "C" fn noop_llm_conditional( + _user_data: *mut c_void, + _request_json: *const NemoRelayNativeString, + _out_reason: *mut *mut NemoRelayNativeString, +) -> NemoRelayStatus { + NemoRelayStatus::Ok +} + +unsafe extern "C" fn noop_llm_request_intercept( + _user_data: *mut c_void, + _name: *const NemoRelayNativeString, + _request_json: *const NemoRelayNativeString, + _annotated_json: *const NemoRelayNativeString, + _out_outcome_json: *mut *mut NemoRelayNativeString, +) -> NemoRelayStatus { + NemoRelayStatus::Ok +} + +unsafe extern "C" fn noop_llm_execution( + _user_data: *mut c_void, + _name: *const NemoRelayNativeString, + _request_json: *const NemoRelayNativeString, + _next_fn: NemoRelayNativeLlmNextFn, + _next_ctx: *mut c_void, + _out_json: *mut *mut NemoRelayNativeString, +) -> NemoRelayStatus { + NemoRelayStatus::Ok +} + +unsafe extern "C" fn noop_llm_stream_execution( + _user_data: *mut c_void, + _name: *const NemoRelayNativeString, + _request_json: *const NemoRelayNativeString, + _next_fn: NemoRelayNativeLlmStreamNextFn, + _next_ctx: *mut c_void, + _out_stream: *mut NemoRelayNativeLlmStreamV1, +) -> NemoRelayStatus { + NemoRelayStatus::Ok +} + +#[test] +fn native_registration_entrypoints_reject_null_contexts() { + unsafe { + assert_eq!( + native_plugin_context_register_subscriber( + ptr::null_mut(), + ptr::null(), + noop_subscriber, + ptr::null_mut(), + None, + ), + NemoRelayStatus::NullPointer + ); + assert_eq!( + native_plugin_context_register_tool_sanitize_request_guardrail( + ptr::null_mut(), + ptr::null(), + 0, + noop_tool_json, + ptr::null_mut(), + None, + ), + NemoRelayStatus::NullPointer + ); + assert_eq!( + native_plugin_context_register_tool_sanitize_response_guardrail( + ptr::null_mut(), + ptr::null(), + 0, + noop_tool_json, + ptr::null_mut(), + None, + ), + NemoRelayStatus::NullPointer + ); + assert_eq!( + native_plugin_context_register_tool_conditional_execution_guardrail( + ptr::null_mut(), + ptr::null(), + 0, + noop_tool_conditional, + ptr::null_mut(), + None, + ), + NemoRelayStatus::NullPointer + ); + assert_eq!( + native_plugin_context_register_tool_request_intercept( + ptr::null_mut(), + ptr::null(), + 0, + false, + noop_tool_json, + ptr::null_mut(), + None, + ), + NemoRelayStatus::NullPointer + ); + assert_eq!( + native_plugin_context_register_tool_execution_intercept( + ptr::null_mut(), + ptr::null(), + 0, + noop_tool_execution, + ptr::null_mut(), + None, + ), + NemoRelayStatus::NullPointer + ); + assert_eq!( + native_plugin_context_register_llm_sanitize_request_guardrail( + ptr::null_mut(), + ptr::null(), + 0, + noop_llm_request, + ptr::null_mut(), + None, + ), + NemoRelayStatus::NullPointer + ); + assert_eq!( + native_plugin_context_register_llm_sanitize_response_guardrail( + ptr::null_mut(), + ptr::null(), + 0, + noop_json, + ptr::null_mut(), + None, + ), + NemoRelayStatus::NullPointer + ); + assert_eq!( + native_plugin_context_register_llm_conditional_execution_guardrail( + ptr::null_mut(), + ptr::null(), + 0, + noop_llm_conditional, + ptr::null_mut(), + None, + ), + NemoRelayStatus::NullPointer + ); + assert_eq!( + native_plugin_context_register_llm_request_intercept( + ptr::null_mut(), + ptr::null(), + 0, + false, + noop_llm_request_intercept, + ptr::null_mut(), + None, + ), + NemoRelayStatus::NullPointer + ); + assert_eq!( + native_plugin_context_register_llm_execution_intercept( + ptr::null_mut(), + ptr::null(), + 0, + noop_llm_execution, + ptr::null_mut(), + None, + ), + NemoRelayStatus::NullPointer + ); + assert_eq!( + native_plugin_context_register_llm_stream_execution_intercept( + ptr::null_mut(), + ptr::null(), + 0, + noop_llm_stream_execution, + ptr::null_mut(), + None, + ), + NemoRelayStatus::NullPointer + ); + } + assert_last_error_contains("plugin context is null"); +} + +#[cfg(unix)] +#[test] +fn native_registration_entrypoints_reject_invalid_host_contexts_and_names() { + let instance = Arc::new(NativePluginInstance { + plugin_kind: "test.native".into(), + relay_compat: "^0.7".into(), + allows_multiple_components: false, + plugin: Mutex::new(NemoRelayNativePluginV1::default()), + _library: libloading::os::unix::Library::this().into(), + }); + let mut invalid_host = NativeHostPluginContext { + ctx: ptr::null_mut(), + instance: Arc::clone(&instance), + }; + assert_eq!( + unsafe { + native_plugin_context_register_subscriber( + ptr::from_mut(&mut invalid_host).cast(), + ptr::null(), + noop_subscriber, + ptr::null_mut(), + None, + ) + }, + NemoRelayStatus::NullPointer + ); -unsafe extern "C" fn noop_llm_request_intercept( - _user_data: *mut c_void, - _name: *const NemoRelayNativeString, - _request_json: *const NemoRelayNativeString, - _annotated_json: *const NemoRelayNativeString, - _out_outcome_json: *mut *mut NemoRelayNativeString, -) -> NemoRelayStatus { - NemoRelayStatus::Ok -} + let mut registration = PluginRegistrationContext::new(); + let mut host = NativeHostPluginContext { + ctx: ptr::from_mut(&mut registration), + instance, + }; + let ctx = ptr::from_mut(&mut host).cast(); -unsafe extern "C" fn noop_llm_execution( - _user_data: *mut c_void, - _name: *const NemoRelayNativeString, - _request_json: *const NemoRelayNativeString, - _next_fn: NemoRelayNativeLlmNextFn, - _next_ctx: *mut c_void, - _out_json: *mut *mut NemoRelayNativeString, -) -> NemoRelayStatus { - NemoRelayStatus::Ok + assert_registration_entrypoints_reject_invalid_names(ctx); + assert_registration_entrypoints_accept_valid_names(ctx); + assert_async_registration_entrypoints_validate_contracts(ctx); + assert_async_request_registration_rejects_legacy_relay_contract(); } -unsafe extern "C" fn noop_llm_stream_execution( - _user_data: *mut c_void, - _name: *const NemoRelayNativeString, - _request_json: *const NemoRelayNativeString, - _next_fn: NemoRelayNativeLlmStreamNextFn, - _next_ctx: *mut c_void, - _out_stream: *mut NemoRelayNativeLlmStreamV1, -) -> NemoRelayStatus { - NemoRelayStatus::Ok +#[cfg(unix)] +fn assert_registration_entrypoints_reject_invalid_names(ctx: *mut NemoRelayNativePluginContext) { + let invalid_name = Box::into_raw(Box::new(NativeHostString(vec![0xff]))).cast(); + unsafe { + assert_eq!( + native_plugin_context_register_subscriber( + ctx, + invalid_name, + noop_subscriber, + ptr::null_mut(), + None, + ), + NemoRelayStatus::InvalidUtf8 + ); + assert_eq!( + native_plugin_context_register_tool_sanitize_request_guardrail( + ctx, + invalid_name, + 0, + noop_tool_json, + ptr::null_mut(), + None, + ), + NemoRelayStatus::InvalidUtf8 + ); + assert_eq!( + native_plugin_context_register_tool_sanitize_response_guardrail( + ctx, + invalid_name, + 0, + noop_tool_json, + ptr::null_mut(), + None, + ), + NemoRelayStatus::InvalidUtf8 + ); + assert_eq!( + native_plugin_context_register_tool_conditional_execution_guardrail( + ctx, + invalid_name, + 0, + noop_tool_conditional, + ptr::null_mut(), + None, + ), + NemoRelayStatus::InvalidUtf8 + ); + assert_eq!( + native_plugin_context_register_tool_request_intercept( + ctx, + invalid_name, + 0, + false, + noop_tool_json, + ptr::null_mut(), + None, + ), + NemoRelayStatus::InvalidUtf8 + ); + assert_eq!( + native_plugin_context_register_tool_execution_intercept( + ctx, + invalid_name, + 0, + noop_tool_execution, + ptr::null_mut(), + None, + ), + NemoRelayStatus::InvalidUtf8 + ); + assert_eq!( + native_plugin_context_register_llm_sanitize_request_guardrail( + ctx, + invalid_name, + 0, + noop_llm_request, + ptr::null_mut(), + None, + ), + NemoRelayStatus::InvalidUtf8 + ); + assert_eq!( + native_plugin_context_register_llm_sanitize_response_guardrail( + ctx, + invalid_name, + 0, + noop_json, + ptr::null_mut(), + None, + ), + NemoRelayStatus::InvalidUtf8 + ); + assert_eq!( + native_plugin_context_register_llm_conditional_execution_guardrail( + ctx, + invalid_name, + 0, + noop_llm_conditional, + ptr::null_mut(), + None, + ), + NemoRelayStatus::InvalidUtf8 + ); + assert_eq!( + native_plugin_context_register_llm_request_intercept( + ctx, + invalid_name, + 0, + false, + noop_llm_request_intercept, + ptr::null_mut(), + None, + ), + NemoRelayStatus::InvalidUtf8 + ); + assert_eq!( + native_plugin_context_register_llm_execution_intercept( + ctx, + invalid_name, + 0, + noop_llm_execution, + ptr::null_mut(), + None, + ), + NemoRelayStatus::InvalidUtf8 + ); + assert_eq!( + native_plugin_context_register_llm_stream_execution_intercept( + ctx, + invalid_name, + 0, + noop_llm_stream_execution, + ptr::null_mut(), + None, + ), + NemoRelayStatus::InvalidUtf8 + ); + assert_eq!( + native_plugin_context_register_async_stream_middleware( + ctx, + invalid_name, + 0, + invoke_native_stream_next_then_return_state, + ptr::null_mut(), + None, + ), + NemoRelayStatus::InvalidUtf8 + ); + assert_eq!( + native_plugin_context_register_async_middleware( + ctx, + NemoRelayNativeAsyncMiddlewareKind::ToolSanitizeRequest as u32, + invalid_name, + 0, + false, + resolve_async_static_json, + ptr::null_mut(), + None, + ), + NemoRelayStatus::InvalidUtf8 + ); + drop(Box::from_raw(invalid_name as *mut NativeHostString)); + } } -#[test] -fn native_registration_entrypoints_reject_null_contexts() { +#[cfg(unix)] +fn assert_registration_entrypoints_accept_valid_names(ctx: *mut NemoRelayNativePluginContext) { unsafe { + let name = native_string("registered"); assert_eq!( native_plugin_context_register_subscriber( - ptr::null_mut(), - ptr::null(), + ctx, + name, noop_subscriber, ptr::null_mut(), None, ), - NemoRelayStatus::NullPointer + NemoRelayStatus::Ok ); assert_eq!( native_plugin_context_register_tool_sanitize_request_guardrail( - ptr::null_mut(), - ptr::null(), + ctx, + name, 0, noop_tool_json, ptr::null_mut(), None, ), - NemoRelayStatus::NullPointer + NemoRelayStatus::Ok ); assert_eq!( native_plugin_context_register_tool_sanitize_response_guardrail( + ctx, + name, + 0, + noop_tool_json, ptr::null_mut(), - ptr::null(), + None, + ), + NemoRelayStatus::Ok + ); + assert_eq!( + native_plugin_context_register_tool_conditional_execution_guardrail( + ctx, + name, + 0, + noop_tool_conditional, + ptr::null_mut(), + None, + ), + NemoRelayStatus::Ok + ); + assert_eq!( + native_plugin_context_register_tool_request_intercept( + ctx, + name, 0, + false, noop_tool_json, ptr::null_mut(), None, ), - NemoRelayStatus::NullPointer + NemoRelayStatus::Ok + ); + assert_eq!( + native_plugin_context_register_tool_execution_intercept( + ctx, + name, + 0, + noop_tool_execution, + ptr::null_mut(), + None, + ), + NemoRelayStatus::Ok + ); + assert_eq!( + native_plugin_context_register_llm_sanitize_request_guardrail( + ctx, + name, + 0, + noop_llm_request, + ptr::null_mut(), + None, + ), + NemoRelayStatus::Ok + ); + assert_eq!( + native_plugin_context_register_llm_sanitize_response_guardrail( + ctx, + name, + 0, + noop_json, + ptr::null_mut(), + None, + ), + NemoRelayStatus::Ok + ); + assert_eq!( + native_plugin_context_register_llm_conditional_execution_guardrail( + ctx, + name, + 0, + noop_llm_conditional, + ptr::null_mut(), + None, + ), + NemoRelayStatus::Ok + ); + assert_eq!( + native_plugin_context_register_llm_request_intercept( + ctx, + name, + 0, + false, + noop_llm_request_intercept, + ptr::null_mut(), + None, + ), + NemoRelayStatus::Ok + ); + assert_eq!( + native_plugin_context_register_llm_execution_intercept( + ctx, + name, + 0, + noop_llm_execution, + ptr::null_mut(), + None, + ), + NemoRelayStatus::Ok ); assert_eq!( - native_plugin_context_register_tool_conditional_execution_guardrail( - ptr::null_mut(), - ptr::null(), + native_plugin_context_register_llm_stream_execution_intercept( + ctx, + name, 0, - noop_tool_conditional, + noop_llm_stream_execution, ptr::null_mut(), None, ), - NemoRelayStatus::NullPointer + NemoRelayStatus::Ok ); + native_string_free(name); + } +} + +#[cfg(unix)] +fn assert_async_registration_entrypoints_validate_contracts( + ctx: *mut NemoRelayNativePluginContext, +) { + unsafe { assert_eq!( - native_plugin_context_register_tool_request_intercept( + native_plugin_context_register_async_stream_middleware( ptr::null_mut(), ptr::null(), 0, - false, - noop_tool_json, + invoke_native_stream_next_then_return_state, ptr::null_mut(), None, ), NemoRelayStatus::NullPointer ); assert_eq!( - native_plugin_context_register_tool_execution_intercept( + native_plugin_context_register_async_middleware( ptr::null_mut(), + 0, ptr::null(), 0, - noop_tool_execution, + false, + resolve_async_static_json, ptr::null_mut(), None, ), NemoRelayStatus::NullPointer ); + + let name = native_string("async-registered"); assert_eq!( - native_plugin_context_register_llm_sanitize_request_guardrail( - ptr::null_mut(), - ptr::null(), + native_plugin_context_register_async_middleware( + ctx, + u32::MAX, + name, 0, - noop_llm_request, + false, + resolve_async_static_json, ptr::null_mut(), None, ), - NemoRelayStatus::NullPointer + NemoRelayStatus::InvalidArg ); assert_eq!( - native_plugin_context_register_llm_sanitize_response_guardrail( - ptr::null_mut(), - ptr::null(), + native_plugin_context_register_async_middleware( + ctx, + NemoRelayNativeAsyncMiddlewareKind::LlmStreamExecutionIntercept as u32, + name, 0, - noop_json, + false, + resolve_async_static_json, ptr::null_mut(), None, ), - NemoRelayStatus::NullPointer + NemoRelayStatus::InvalidArg ); assert_eq!( - native_plugin_context_register_llm_conditional_execution_guardrail( - ptr::null_mut(), - ptr::null(), + native_plugin_context_register_async_middleware( + ctx, + NemoRelayNativeAsyncMiddlewareKind::ToolSanitizeRequest as u32, + name, 0, - noop_llm_conditional, + false, + resolve_async_static_json, ptr::null_mut(), None, ), - NemoRelayStatus::NullPointer + NemoRelayStatus::Ok ); assert_eq!( - native_plugin_context_register_llm_request_intercept( - ptr::null_mut(), - ptr::null(), + native_plugin_context_register_async_middleware( + ctx, + NemoRelayNativeAsyncMiddlewareKind::ToolSanitizeRequest as u32, + name, 0, false, - noop_llm_request_intercept, + resolve_async_static_json, ptr::null_mut(), None, ), - NemoRelayStatus::NullPointer + NemoRelayStatus::Internal ); + + let stream_name = native_string("async-stream-registered"); assert_eq!( - native_plugin_context_register_llm_execution_intercept( - ptr::null_mut(), - ptr::null(), + native_plugin_context_register_async_stream_middleware( + ctx, + stream_name, 0, - noop_llm_execution, + invoke_native_stream_next_then_return_state, ptr::null_mut(), None, ), - NemoRelayStatus::NullPointer + NemoRelayStatus::Ok ); assert_eq!( - native_plugin_context_register_llm_stream_execution_intercept( - ptr::null_mut(), - ptr::null(), + native_plugin_context_register_async_stream_middleware( + ctx, + stream_name, 0, - noop_llm_stream_execution, + invoke_native_stream_next_then_return_state, ptr::null_mut(), None, ), - NemoRelayStatus::NullPointer + NemoRelayStatus::Internal ); + native_string_free(name); + native_string_free(stream_name); + } +} + +#[cfg(unix)] +fn assert_async_request_registration_rejects_legacy_relay_contract() { + let instance = Arc::new(NativePluginInstance { + plugin_kind: "test.native.legacy".into(), + relay_compat: "^0.5".into(), + allows_multiple_components: false, + plugin: Mutex::new(NemoRelayNativePluginV1::default()), + _library: libloading::os::unix::Library::this().into(), + }); + let mut registration = PluginRegistrationContext::new(); + let mut host = NativeHostPluginContext { + ctx: ptr::from_mut(&mut registration), + instance, + }; + let name = native_string("legacy-request"); + assert_eq!( + unsafe { + native_plugin_context_register_async_middleware( + ptr::from_mut(&mut host).cast(), + NemoRelayNativeAsyncMiddlewareKind::LlmRequestIntercept as u32, + name, + 0, + false, + resolve_async_static_json, + ptr::null_mut(), + None, + ) + }, + NemoRelayStatus::InvalidArg + ); + assert_last_error_contains("excludes Relay 0.5"); + unsafe { native_string_free(name) }; +} + +#[cfg(unix)] +unsafe extern "C" fn resolve_async_static_json( + user_data: *mut c_void, + _invocation_json: *const NemoRelayNativeString, + _next: *const NemoRelayNativeAsyncNext, + completion: *const NemoRelayNativeAsyncCompletion, +) -> u32 { + assert_eq!( + unsafe { native_async_completion_resolve_json(completion, user_data.cast()) }, + NemoRelayStatus::Ok + ); + NemoRelayNativeAsyncCallbackState::Complete as u32 +} + +#[cfg(unix)] +#[tokio::test] +async fn native_async_wrappers_validate_callback_result_shapes() { + let instance = Arc::new(NativePluginInstance { + plugin_kind: "test.native.async".into(), + relay_compat: "^0.7".into(), + allows_multiple_components: false, + plugin: Mutex::new(NemoRelayNativePluginV1::default()), + _library: libloading::os::unix::Library::this().into(), + }); + let result = native_string("true"); + let user_data = result.cast(); + + let tool_json = wrap_native_async_tool_json( + Arc::clone(&instance), + resolve_async_static_json, + user_data, + None, + ); + assert_eq!( + tool_json("tool".into(), json!({})).await.unwrap(), + json!(true) + ); + + let tool_conditional = wrap_native_async_tool_conditional( + Arc::clone(&instance), + resolve_async_static_json, + user_data, + None, + ); + assert!( + tool_conditional("tool".into(), json!({})) + .await + .unwrap_err() + .to_string() + .contains("expected string or null") + ); + + let request = LlmRequest { + headers: Map::new(), + content: json!({"messages": []}), + }; + let llm_conditional = wrap_native_async_llm_conditional( + Arc::clone(&instance), + resolve_async_static_json, + user_data, + None, + ); + assert!( + llm_conditional(request.clone()) + .await + .unwrap_err() + .to_string() + .contains("expected string or null") + ); + + let sanitize_request = wrap_native_async_llm_sanitize_request( + Arc::clone(&instance), + resolve_async_static_json, + user_data, + None, + ); + assert!( + sanitize_request( + request.clone(), + LlmSanitizeRequestContext::for_request_codec(None), + ) + .await + .is_err() + ); + + let sanitize_response = wrap_native_async_llm_sanitize_response( + Arc::clone(&instance), + resolve_async_static_json, + user_data, + None, + ); + assert_eq!( + sanitize_response( + json!({"response": true}), + LlmSanitizeResponseContext::for_response_codec(None), + ) + .await + .unwrap(), + Some(json!(true)) + ); + + let request_intercept = wrap_native_async_llm_request_intercept( + Arc::clone(&instance), + resolve_async_static_json, + user_data, + None, + ); + assert!( + request_intercept("model".into(), request, None) + .await + .unwrap_err() + .to_string() + .contains("invalid native async LLM intercept outcome") + ); + + let tool_execution = wrap_native_async_tool_execution( + Arc::clone(&instance), + resolve_async_static_json, + user_data, + None, + ); + assert!( + tool_execution("tool", json!({}), tool_next(Ok(Json::Null))) + .await + .unwrap_err() + .to_string() + .contains("invalid native async tool outcome") + ); + + let fields = EventSanitizeFields::default(); + let fields_result = native_string_from_json(&serde_json::to_value(&fields).unwrap()).unwrap(); + let event_sanitize = wrap_native_async_event_sanitize( + Arc::clone(&instance), + resolve_async_static_json, + fields_result.cast(), + None, + ); + let event = Event::Mark(crate::api::event::MarkEvent::new( + crate::api::event::BaseEvent::builder() + .name("native-async-event") + .build(), + None, + None, + )); + assert_eq!( + event_sanitize(Arc::new(event), fields.clone()) + .await + .unwrap(), + fields + ); + + drop(event_sanitize); + drop(tool_execution); + drop(request_intercept); + drop(sanitize_response); + drop(sanitize_request); + drop(llm_conditional); + drop(tool_conditional); + drop(tool_json); + unsafe { + native_string_free(fields_result); + native_string_free(result); } - assert_last_error_contains("plugin context is null"); } #[test] @@ -3100,52 +4701,125 @@ fn native_codec_operations_contain_codec_panics() { &mut output, ) }, - NemoRelayStatus::Internal + NemoRelayStatus::Internal + ); + assert!(output.is_null()); + assert_last_error_contains("request codec decode panicked"); + + assert_eq!( + unsafe { + native_llm_request_codec_encode( + ptr::from_ref(&request_codec).cast(), + annotated_json, + request_json, + &mut output, + ) + }, + NemoRelayStatus::Internal + ); + assert!(output.is_null()); + assert_last_error_contains("request codec encode panicked"); + + assert_eq!( + unsafe { + native_llm_response_codec_decode( + ptr::from_ref(&response_codec).cast(), + response_json, + &mut output, + ) + }, + NemoRelayStatus::Internal + ); + assert!(output.is_null()); + assert_last_error_contains("response codec decode panicked"); + + unsafe { + native_string_free(request_json); + native_string_free(annotated_json); + native_string_free(response_json); + } +} + +#[test] +fn native_codec_operations_clear_output_slots_on_null_arguments() { + let request_codec = NativeHostLlmRequestCodec(Arc::new(OpenAIChatCodec) as Arc); + let response_codec = + NativeHostLlmResponseCodec(Arc::new(OpenAIChatCodec) as Arc); + let request_json = native_string( + r#"{"headers":{},"content":{"model":"gpt-test","messages":[{"role":"user","content":"hello"}]}}"#, + ); + let request = LlmRequest { + headers: Map::new(), + content: json!({ + "model": "gpt-test", + "messages": [{"role": "user", "content": "hello"}] + }), + }; + let annotated = OpenAIChatCodec.decode(&request).unwrap(); + let annotated_json = native_string(&serde_json::to_string(&annotated).unwrap()); + + assert_eq!( + unsafe { native_llm_request_codec_decode(ptr::null(), request_json, &mut ptr::null_mut()) }, + NemoRelayStatus::NullPointer + ); + assert_eq!( + unsafe { + native_llm_request_codec_decode( + ptr::from_ref(&request_codec).cast(), + request_json, + ptr::null_mut(), + ) + }, + NemoRelayStatus::NullPointer + ); + assert_eq!( + unsafe { + native_llm_request_codec_encode( + ptr::null(), + annotated_json, + request_json, + &mut ptr::null_mut(), + ) + }, + NemoRelayStatus::NullPointer + ); + assert_eq!( + unsafe { + native_llm_request_codec_encode( + ptr::from_ref(&request_codec).cast(), + annotated_json, + ptr::null(), + &mut ptr::null_mut(), + ) + }, + NemoRelayStatus::NullPointer ); - assert!(output.is_null()); - assert_last_error_contains("request codec decode panicked"); - assert_eq!( unsafe { native_llm_request_codec_encode( ptr::from_ref(&request_codec).cast(), annotated_json, request_json, - &mut output, + ptr::null_mut(), ) }, - NemoRelayStatus::Internal + NemoRelayStatus::NullPointer + ); + assert_eq!( + unsafe { + native_llm_response_codec_decode(ptr::null(), request_json, &mut ptr::null_mut()) + }, + NemoRelayStatus::NullPointer ); - assert!(output.is_null()); - assert_last_error_contains("request codec encode panicked"); - assert_eq!( unsafe { native_llm_response_codec_decode( ptr::from_ref(&response_codec).cast(), - response_json, - &mut output, + request_json, + ptr::null_mut(), ) }, - NemoRelayStatus::Internal - ); - assert!(output.is_null()); - assert_last_error_contains("response codec decode panicked"); - - unsafe { - native_string_free(request_json); - native_string_free(annotated_json); - native_string_free(response_json); - } -} - -#[test] -fn native_codec_operations_clear_output_slots_on_null_arguments() { - let request_codec = NativeHostLlmRequestCodec(Arc::new(OpenAIChatCodec) as Arc); - let response_codec = - NativeHostLlmResponseCodec(Arc::new(OpenAIChatCodec) as Arc); - let request_json = native_string( - r#"{"headers":{},"content":{"model":"gpt-test","messages":[{"role":"user","content":"hello"}]}}"#, + NemoRelayStatus::NullPointer ); let request_decode_sentinel = native_string("request-decode-sentinel"); @@ -3200,6 +4874,7 @@ fn native_codec_operations_clear_output_slots_on_null_arguments() { assert_last_error_contains("response codec decode response is null"); unsafe { native_string_free(response_decode_sentinel); + native_string_free(annotated_json); native_string_free(request_json); } } @@ -3226,6 +4901,127 @@ unsafe extern "C" fn tool_json_error( NemoRelayStatus::InvalidArg } +#[cfg(unix)] +unsafe extern "C" fn tool_conditional_error( + _user_data: *mut c_void, + _name: *const NemoRelayNativeString, + _args_json: *const NemoRelayNativeString, + out_reason: *mut *mut NemoRelayNativeString, +) -> NemoRelayStatus { + unsafe { *out_reason = native_string("discarded reason") }; + set_native_last_error("tool conditional failed"); + NemoRelayStatus::InvalidArg +} + +#[cfg(unix)] +unsafe extern "C" fn tool_conditional_reason( + _user_data: *mut c_void, + _name: *const NemoRelayNativeString, + _args_json: *const NemoRelayNativeString, + out_reason: *mut *mut NemoRelayNativeString, +) -> NemoRelayStatus { + unsafe { *out_reason = native_string("blocked by tool policy") }; + NemoRelayStatus::Ok +} + +#[cfg(unix)] +unsafe extern "C" fn llm_conditional_error( + _user_data: *mut c_void, + _request_json: *const NemoRelayNativeString, + out_reason: *mut *mut NemoRelayNativeString, +) -> NemoRelayStatus { + unsafe { *out_reason = native_string("discarded reason") }; + set_native_last_error("LLM conditional failed"); + NemoRelayStatus::InvalidArg +} + +#[cfg(unix)] +unsafe extern "C" fn llm_conditional_reason( + _user_data: *mut c_void, + _request_json: *const NemoRelayNativeString, + out_reason: *mut *mut NemoRelayNativeString, +) -> NemoRelayStatus { + unsafe { *out_reason = native_string("blocked by LLM policy") }; + NemoRelayStatus::Ok +} + +unsafe extern "C" fn llm_request_error( + _user_data: *mut c_void, + _request_json: *const NemoRelayNativeString, + _context: NemoRelayNativeLlmSanitizeRequestContext, + out_request_json: *mut *mut NemoRelayNativeString, +) -> NemoRelayStatus { + unsafe { *out_request_json = native_string(r#"{"discarded":true}"#) }; + set_native_last_error("LLM request sanitizer failed"); + NemoRelayStatus::InvalidArg +} + +unsafe extern "C" fn llm_response_error( + _user_data: *mut c_void, + _response_json: *const NemoRelayNativeString, + _context: NemoRelayNativeLlmSanitizeResponseContext, + out_response_json: *mut *mut NemoRelayNativeString, +) -> NemoRelayStatus { + unsafe { *out_response_json = native_string(r#"{"discarded":true}"#) }; + set_native_last_error("LLM response sanitizer failed"); + NemoRelayStatus::InvalidArg +} + +#[cfg(unix)] +unsafe extern "C" fn tool_execution_error( + _user_data: *mut c_void, + _name: *const NemoRelayNativeString, + _args_json: *const NemoRelayNativeString, + _next_fn: NemoRelayNativeToolNextFn, + _next_ctx: *mut c_void, + out_outcome_json: *mut *mut NemoRelayNativeString, +) -> NemoRelayStatus { + unsafe { *out_outcome_json = native_string(r#"{"discarded":true}"#) }; + set_native_last_error("tool execution failed"); + NemoRelayStatus::InvalidArg +} + +#[cfg(unix)] +unsafe extern "C" fn llm_request_intercept_error( + _user_data: *mut c_void, + _name: *const NemoRelayNativeString, + _request_json: *const NemoRelayNativeString, + _annotated_json: *const NemoRelayNativeString, + out_outcome_json: *mut *mut NemoRelayNativeString, +) -> NemoRelayStatus { + unsafe { *out_outcome_json = native_string(r#"{"discarded":true}"#) }; + set_native_last_error("LLM request intercept failed"); + NemoRelayStatus::InvalidArg +} + +#[cfg(unix)] +unsafe extern "C" fn llm_execution_error( + _user_data: *mut c_void, + _name: *const NemoRelayNativeString, + _request_json: *const NemoRelayNativeString, + _next_fn: NemoRelayNativeLlmNextFn, + _next_ctx: *mut c_void, + out_json: *mut *mut NemoRelayNativeString, +) -> NemoRelayStatus { + unsafe { *out_json = native_string(r#"{"discarded":true}"#) }; + set_native_last_error("LLM execution failed"); + NemoRelayStatus::InvalidArg +} + +#[cfg(unix)] +unsafe extern "C" fn llm_stream_execution_error( + _user_data: *mut c_void, + _name: *const NemoRelayNativeString, + _request_json: *const NemoRelayNativeString, + _next_fn: NemoRelayNativeLlmStreamNextFn, + _next_ctx: *mut c_void, + out_stream: *mut NemoRelayNativeLlmStreamV1, +) -> NemoRelayStatus { + unsafe { *out_stream = NemoRelayNativeLlmStreamV1::default() }; + set_native_last_error("LLM stream execution failed"); + NemoRelayStatus::InvalidArg +} + unsafe extern "C" fn llm_request_echo( _user_data: *mut c_void, request_json: *const NemoRelayNativeString, @@ -3321,7 +5117,7 @@ fn native_callback_helpers_cover_success_error_and_invalid_output() { LlmSanitizeRequestContext::default(), ) .unwrap(), - Some(request) + Some(request.clone()) ); let request = LlmRequest { @@ -3336,7 +5132,7 @@ fn native_callback_helpers_cover_success_error_and_invalid_output() { LlmSanitizeRequestContext::default(), ) .unwrap(), - Some(request) + Some(request.clone()) ); let response = json!({"message": "alias"}); @@ -3348,7 +5144,169 @@ fn native_callback_helpers_cover_success_error_and_invalid_output() { LlmSanitizeResponseContext::default(), ) .unwrap(), - Some(response) + Some(response.clone()) + ); + + assert!( + call_llm_sanitize_request_callback( + llm_request_error, + ptr::null_mut(), + &request, + LlmSanitizeRequestContext::default(), + ) + .unwrap_err() + .to_string() + .contains("LLM request sanitizer failed") + ); + assert_eq!( + call_llm_sanitize_request_callback( + noop_llm_request, + ptr::null_mut(), + &request, + LlmSanitizeRequestContext::default(), + ) + .unwrap(), + None + ); + assert!( + call_llm_sanitize_response_callback( + llm_response_error, + ptr::null_mut(), + &response, + LlmSanitizeResponseContext::default(), + ) + .unwrap_err() + .to_string() + .contains("LLM response sanitizer failed") + ); + assert_eq!( + call_llm_sanitize_response_callback( + noop_json, + ptr::null_mut(), + &response, + LlmSanitizeResponseContext::default(), + ) + .unwrap(), + None + ); +} + +#[cfg(unix)] +#[tokio::test] +async fn native_callback_wrappers_release_error_outputs_and_preserve_reasons() { + let instance = Arc::new(NativePluginInstance { + plugin_kind: "test.native.callback-errors".into(), + relay_compat: "^0.7".into(), + allows_multiple_components: false, + plugin: Mutex::new(NemoRelayNativePluginV1::default()), + _library: libloading::os::unix::Library::this().into(), + }); + let request = LlmRequest { + headers: Map::new(), + content: json!({"model": "test"}), + }; + + let tool_conditional = wrap_tool_conditional_fn( + Arc::clone(&instance), + tool_conditional_error, + ptr::null_mut(), + None, + ); + assert!( + tool_conditional("tool".into(), json!({})) + .await + .unwrap_err() + .to_string() + .contains("tool conditional failed") + ); + let tool_conditional = wrap_tool_conditional_fn( + Arc::clone(&instance), + tool_conditional_reason, + ptr::null_mut(), + None, + ); + assert_eq!( + tool_conditional("tool".into(), json!({})).await.unwrap(), + Some("blocked by tool policy".into()) + ); + + let tool_execution = wrap_tool_execution_fn( + Arc::clone(&instance), + tool_execution_error, + ptr::null_mut(), + None, + ); + assert!( + tool_execution("tool", json!({}), tool_next(Ok(Json::Null))) + .await + .unwrap_err() + .to_string() + .contains("tool execution failed") + ); + + let llm_conditional = wrap_llm_conditional_fn( + Arc::clone(&instance), + llm_conditional_error, + ptr::null_mut(), + None, + ); + assert!( + llm_conditional(request.clone()) + .await + .unwrap_err() + .to_string() + .contains("LLM conditional failed") + ); + let llm_conditional = wrap_llm_conditional_fn( + Arc::clone(&instance), + llm_conditional_reason, + ptr::null_mut(), + None, + ); + assert_eq!( + llm_conditional(request.clone()).await.unwrap(), + Some("blocked by LLM policy".into()) + ); + + let request_intercept = wrap_llm_request_intercept_fn( + Arc::clone(&instance), + llm_request_intercept_error, + ptr::null_mut(), + None, + ); + assert!( + request_intercept("model".into(), request.clone(), None) + .await + .unwrap_err() + .to_string() + .contains("LLM request intercept failed") + ); + + let llm_execution = wrap_llm_execution_fn( + Arc::clone(&instance), + llm_execution_error, + ptr::null_mut(), + None, + ); + assert!( + llm_execution("model", request.clone(), llm_next(Ok(Json::Null))) + .await + .unwrap_err() + .to_string() + .contains("LLM execution failed") + ); + + let stream_next: LlmStreamExecutionNextFn = + Arc::new(|_| Box::pin(async { Ok(LlmJsonStream::new(tokio_stream::empty())) })); + let llm_stream_execution = + wrap_llm_stream_execution_fn(instance, llm_stream_execution_error, ptr::null_mut(), None); + assert!( + llm_stream_execution("model", request, stream_next) + .await + .err() + .expect("native stream callback should fail") + .to_string() + .contains("LLM stream execution failed") ); } @@ -3572,6 +5530,16 @@ fn native_non_streaming_continuations_cover_success_and_error_paths() { NemoRelayStatus::NotFound ); unsafe { drop(Box::from_raw(next as *mut ToolExecutionNextFn)) }; + + let panicking_next: ToolExecutionNextFn = + Arc::new(|_| Box::pin(async { panic!("tool next panic") })); + let next = Box::into_raw(Box::new(panicking_next)) as *mut c_void; + assert_eq!( + unsafe { native_tool_next(args, next, &mut out) }, + NemoRelayStatus::Internal + ); + assert_last_error_contains("native tool next panicked"); + unsafe { drop(Box::from_raw(next as *mut ToolExecutionNextFn)) }; unsafe { native_string_free(args) }; let request = LlmRequest { @@ -3599,6 +5567,16 @@ fn native_non_streaming_continuations_cover_success_and_error_paths() { unsafe { native_llm_next(request_json, next, &mut out) }, NemoRelayStatus::GuardrailRejected ); + unsafe { drop(Box::from_raw(next as *mut LlmExecutionNextFn)) }; + + let panicking_next: LlmExecutionNextFn = + Arc::new(|_| Box::pin(async { panic!("LLM next panic") })); + let next = Box::into_raw(Box::new(panicking_next)) as *mut c_void; + assert_eq!( + unsafe { native_llm_next(request_json, next, &mut out) }, + NemoRelayStatus::Internal + ); + assert_last_error_contains("native LLM next panicked"); unsafe { drop(Box::from_raw(next as *mut LlmExecutionNextFn)); native_string_free(request_json); @@ -3611,7 +5589,9 @@ enum NativeStreamItem { InvalidJson, Null, Error(NemoRelayStatus), + ErrorWithJson(NemoRelayStatus), End, + EndWithJson, } struct TestNativeStream { @@ -3634,7 +5614,15 @@ unsafe extern "C" fn test_native_stream_poll( } NativeStreamItem::Null => NemoRelayStatus::Ok, NativeStreamItem::Error(status) => status, + NativeStreamItem::ErrorWithJson(status) => { + unsafe { *out_json = native_string(r#"{"discarded":true}"#) }; + status + } NativeStreamItem::End => NemoRelayStatus::StreamEnd, + NativeStreamItem::EndWithJson => { + unsafe { *out_json = native_string(r#"{"discarded":true}"#) }; + NemoRelayStatus::StreamEnd + } } } @@ -3700,6 +5688,7 @@ async fn native_stream_adapter_covers_chunks_end_errors_and_cancellation() { NativeStreamItem::InvalidJson, NativeStreamItem::Null, NativeStreamItem::Error(NemoRelayStatus::InvalidArg), + NativeStreamItem::ErrorWithJson(NemoRelayStatus::InvalidArg), ] { let (raw, _, drop_count) = test_native_stream([item]); let mut stream = native_stream_to_relay_stream(raw, None, None).unwrap(); @@ -3717,6 +5706,20 @@ async fn native_stream_adapter_covers_chunks_end_errors_and_cancellation() { raw.next = None; assert!(NativeRelayLlmStream::from_raw(raw, None, None).is_err()); assert_eq!(drop_count.load(Ordering::SeqCst), 1); + + let (raw, _, drop_count) = test_native_stream([NativeStreamItem::EndWithJson]); + let mut stream = native_stream_to_relay_stream(raw, None, None).unwrap(); + assert!(stream.next().await.is_none()); + assert_eq!(drop_count.load(Ordering::SeqCst), 1); + + let mut invalid = NativeRelayLlmStream { + raw: NemoRelayNativeLlmStreamV1::default(), + finished: false, + _next_ctx: None, + _callback_user_data: None, + }; + assert!(invalid.next().await.unwrap().is_err()); + assert!(invalid.next().await.is_none()); } #[tokio::test] @@ -3751,6 +5754,10 @@ async fn relay_stream_adapter_covers_poll_end_error_and_cancel() { unsafe { poll(raw.user_data, &mut out) }, NemoRelayStatus::StreamEnd ); + assert_eq!( + unsafe { poll(raw.user_data, &mut out) }, + NemoRelayStatus::StreamEnd + ); assert_eq!( unsafe { cancel_relay_llm_stream(raw.user_data) }, NemoRelayStatus::Ok @@ -3775,6 +5782,10 @@ async fn relay_stream_adapter_covers_poll_end_error_and_cancel() { let _guard = mutex.lock().unwrap(); panic!("poison native stream lock"); })); + assert_eq!( + unsafe { raw.next.unwrap()(raw.user_data, &mut out) }, + NemoRelayStatus::Internal + ); assert_eq!( unsafe { cancel_relay_llm_stream(raw.user_data) }, NemoRelayStatus::Internal @@ -3823,6 +5834,16 @@ fn native_stream_continuation_covers_success_and_error() { unsafe { native_llm_stream_next(request_json, next_ctx, &mut raw) }, NemoRelayStatus::NotFound ); + unsafe { drop(Box::from_raw(next_ctx as *mut LlmStreamExecutionNextFn)) }; + + let next: LlmStreamExecutionNextFn = + Arc::new(|_| Box::pin(async { panic!("stream next panic") })); + let next_ctx = Box::into_raw(Box::new(next)) as *mut c_void; + assert_eq!( + unsafe { native_llm_stream_next(request_json, next_ctx, &mut raw) }, + NemoRelayStatus::Internal + ); + assert_last_error_contains("native LLM stream next panicked"); unsafe { drop(Box::from_raw(next_ctx as *mut LlmStreamExecutionNextFn)); native_string_free(request_json); diff --git a/crates/core/tests/unit/observability/atof_tests.rs b/crates/core/tests/unit/observability/atof_tests.rs index 5a897cc35..1e2d8e3e7 100644 --- a/crates/core/tests/unit/observability/atof_tests.rs +++ b/crates/core/tests/unit/observability/atof_tests.rs @@ -713,11 +713,7 @@ fn subscriber_preserves_wire_format_llm_lifecycle_payloads_as_raw_jsonl() { "anthropic.messages", "claude-sonnet-4", "/v1/messages", - json!({ - "model": "claude-sonnet-4", - "messages": [{"role": "user", "content": "Find the file."}], - "tools": [{"name": "search", "input_schema": {"type": "object"}}] - }), + anthropic_wire_start_payload(), ), wire_format_llm_event( anthropic_uuid, @@ -726,21 +722,7 @@ fn subscriber_preserves_wire_format_llm_lifecycle_payloads_as_raw_jsonl() { "anthropic.messages", "claude-sonnet-4", "/v1/messages", - json!({ - "id": "msg_01", - "type": "message", - "content": [ - {"type": "text", "text": "I will search."}, - {"type": "tool_use", "id": "toolu_01", "name": "search", "input": {"query": "file"}} - ], - "usage": { - "input_tokens": 11, - "output_tokens": 7, - "cache_read_input_tokens": 3, - "cache_creation_input_tokens": 5, - "cost": {"total": 0.0042} - } - }), + anthropic_wire_end_payload(), ), wire_format_llm_event( responses_uuid, @@ -749,11 +731,7 @@ fn subscriber_preserves_wire_format_llm_lifecycle_payloads_as_raw_jsonl() { "openai.responses", "gpt-4o", "/v1/responses", - json!({ - "model": "gpt-4o", - "input": "Find the weather.", - "tools": [{"type": "function", "name": "get_weather"}] - }), + responses_wire_start_payload(), ), wire_format_llm_event( responses_uuid, @@ -762,20 +740,7 @@ fn subscriber_preserves_wire_format_llm_lifecycle_payloads_as_raw_jsonl() { "openai.responses", "gpt-4o", "/v1/responses", - json!({ - "id": "resp_1", - "output": [ - {"type": "message", "content": [{"type": "output_text", "text": "I will check."}]}, - {"type": "function_call", "call_id": "call_weather_1", "name": "get_weather", "arguments": "{\"city\":\"SF\"}"} - ], - "usage": { - "input_tokens": 75, - "output_tokens": 20, - "total_tokens": 95, - "input_tokens_details": {"cached_tokens": 10}, - "cost_usd": 0.005 - } - }), + responses_wire_end_payload(), ), wire_format_llm_event( chat_uuid, @@ -784,11 +749,7 @@ fn subscriber_preserves_wire_format_llm_lifecycle_payloads_as_raw_jsonl() { "openai.chat_completions", "gpt-4o", "/v1/chat/completions", - json!({ - "model": "gpt-4o", - "messages": [{"role": "user", "content": "Inspect the files."}], - "tools": [{"type": "function", "function": {"name": "read"}}] - }), + chat_wire_start_payload(), ), wire_format_llm_event( chat_uuid, @@ -797,22 +758,7 @@ fn subscriber_preserves_wire_format_llm_lifecycle_payloads_as_raw_jsonl() { "openai.chat_completions", "gpt-4o", "/v1/chat/completions", - json!({ - "choices": [{ - "message": { - "role": "assistant", - "content": "I will inspect.", - "tool_calls": [{"id": "call_read_1", "function": {"name": "read", "arguments": "{\"path\":\"api.py\"}"}}] - } - }], - "usage": { - "prompt_tokens": 3, - "completion_tokens": 4, - "total_tokens": 7, - "prompt_tokens_details": {"cached_tokens": 2}, - "cost_usd": 0.001 - } - }), + chat_wire_end_payload(), ), ]; @@ -823,8 +769,21 @@ fn subscriber_preserves_wire_format_llm_lifecycle_payloads_as_raw_jsonl() { let lines = read_jsonl(exporter.path().expect("file sink path")); assert_eq!(lines.len(), events.len()); + assert_wire_lines_match_events(&lines, &events); + assert_wire_lines_common_shape(&lines, parent_uuid); + assert_anthropic_wire_lines(&lines); + assert_responses_wire_lines(&lines); + assert_chat_wire_lines(&lines); +} + +fn assert_wire_lines_match_events(lines: &[Json], events: &[Event]) { for (line, event) in lines.iter().zip(events.iter()) { assert_eq!(line, &event.try_to_json_value().unwrap()); + } +} + +fn assert_wire_lines_common_shape(lines: &[Json], parent_uuid: Uuid) { + for line in lines { assert_eq!(line["kind"], "scope"); assert_eq!(line["atof_version"], "0.1"); assert_eq!(line["parent_uuid"], parent_uuid.to_string()); @@ -834,7 +793,9 @@ fn subscriber_preserves_wire_format_llm_lifecycle_payloads_as_raw_jsonl() { assert_eq!(line["metadata"]["source"], "openclaw.public_plugin"); assert_eq!(line["metadata"]["provider_payload_exact"], true); } +} +fn assert_anthropic_wire_lines(lines: &[Json]) { assert_eq!(lines[0]["name"], "anthropic.messages"); assert_eq!(lines[0]["scope_category"], "start"); assert_eq!(lines[0]["metadata"]["gateway_path"], "/v1/messages"); @@ -847,7 +808,9 @@ fn subscriber_preserves_wire_format_llm_lifecycle_payloads_as_raw_jsonl() { assert_eq!(lines[1]["data"]["content"][1]["type"], "tool_use"); assert_eq!(lines[1]["data"]["usage"]["cache_creation_input_tokens"], 5); assert_eq!(lines[1]["data"]["usage"]["cost"]["total"], 0.0042); +} +fn assert_responses_wire_lines(lines: &[Json]) { assert_eq!(lines[2]["metadata"]["gateway_path"], "/v1/responses"); assert_eq!(lines[2]["data"]["input"], "Find the weather."); assert_eq!(lines[3]["data"]["output"][1]["type"], "function_call"); @@ -856,7 +819,9 @@ fn subscriber_preserves_wire_format_llm_lifecycle_payloads_as_raw_jsonl() { 10 ); assert_eq!(lines[3]["data"]["usage"]["cost_usd"], 0.005); +} +fn assert_chat_wire_lines(lines: &[Json]) { assert_eq!(lines[4]["metadata"]["gateway_path"], "/v1/chat/completions"); assert_eq!( lines[4]["data"]["messages"][0]["content"], @@ -873,6 +838,84 @@ fn subscriber_preserves_wire_format_llm_lifecycle_payloads_as_raw_jsonl() { assert_eq!(lines[5]["data"]["usage"]["cost_usd"], 0.001); } +fn anthropic_wire_start_payload() -> Json { + json!({ + "model": "claude-sonnet-4", + "messages": [{"role": "user", "content": "Find the file."}], + "tools": [{"name": "search", "input_schema": {"type": "object"}}] + }) +} + +fn anthropic_wire_end_payload() -> Json { + json!({ + "id": "msg_01", + "type": "message", + "content": [ + {"type": "text", "text": "I will search."}, + {"type": "tool_use", "id": "toolu_01", "name": "search", "input": {"query": "file"}} + ], + "usage": { + "input_tokens": 11, + "output_tokens": 7, + "cache_read_input_tokens": 3, + "cache_creation_input_tokens": 5, + "cost": {"total": 0.0042} + } + }) +} + +fn responses_wire_start_payload() -> Json { + json!({ + "model": "gpt-4o", + "input": "Find the weather.", + "tools": [{"type": "function", "name": "get_weather"}] + }) +} + +fn responses_wire_end_payload() -> Json { + json!({ + "id": "resp_1", + "output": [ + {"type": "message", "content": [{"type": "output_text", "text": "I will check."}]}, + {"type": "function_call", "call_id": "call_weather_1", "name": "get_weather", "arguments": "{\"city\":\"SF\"}"} + ], + "usage": { + "input_tokens": 75, + "output_tokens": 20, + "total_tokens": 95, + "input_tokens_details": {"cached_tokens": 10}, + "cost_usd": 0.005 + } + }) +} + +fn chat_wire_start_payload() -> Json { + json!({ + "model": "gpt-4o", + "messages": [{"role": "user", "content": "Inspect the files."}], + "tools": [{"type": "function", "function": {"name": "read"}}] + }) +} + +fn chat_wire_end_payload() -> Json { + json!({ + "choices": [{ + "message": { + "role": "assistant", + "content": "I will inspect.", + "tool_calls": [{"id": "call_read_1", "function": {"name": "read", "arguments": "{\"path\":\"api.py\"}"}}] + } + }], + "usage": { + "prompt_tokens": 3, + "completion_tokens": 4, + "total_tokens": 7, + "prompt_tokens_details": {"cached_tokens": 2}, + "cost_usd": 0.001 + } + }) +} + #[test] fn openclaw_subagent_events_preserve_nested_and_fallback_parent_uuid() { let dir = temp_dir("atof-openclaw-subagent-parentage"); @@ -1629,6 +1672,73 @@ fn http_endpoint_worker_reports_request_transport_failure() { worker.join().unwrap(); } +#[test] +#[cfg(feature = "atof-streaming")] +fn ndjson_endpoint_worker_streams_events_flushes_and_closes() { + enable_operational_logs(); + let listener = TcpListener::bind("127.0.0.1:0").unwrap(); + let url = format!("http://{}", listener.local_addr().unwrap()); + let server = std::thread::spawn(move || { + let (mut stream, _) = listener.accept().unwrap(); + stream + .set_read_timeout(Some(std::time::Duration::from_secs(5))) + .unwrap(); + let body = read_http_request(&mut stream); + assert_eq!(body, "{\"kind\":\"mark\"}\n"); + stream + .write_all(b"HTTP/1.1 200 OK\r\nContent-Length: 0\r\nConnection: close\r\n\r\n") + .unwrap(); + }); + + let endpoint = validate_endpoint_config( + AtofEndpointConfig::new(url, AtofEndpointTransport::Ndjson).with_timeout_millis(5_000), + ) + .unwrap(); + let (tx, rx) = tokio::sync::mpsc::unbounded_channel(); + let worker = std::thread::spawn(move || run_endpoint_worker(0, endpoint, rx)); + + tx.send(EndpointMessage::Event("{\"kind\":\"mark\"}".into())) + .unwrap(); + let (flush_tx, flush_rx) = std::sync::mpsc::channel(); + tx.send(EndpointMessage::Flush(flush_tx)).unwrap(); + flush_rx + .recv_timeout(std::time::Duration::from_secs(5)) + .unwrap(); + let (close_tx, close_rx) = std::sync::mpsc::channel(); + tx.send(EndpointMessage::Close(close_tx)).unwrap(); + close_rx + .recv_timeout(std::time::Duration::from_secs(5)) + .unwrap(); + worker.join().unwrap(); + server.join().unwrap(); +} + +#[test] +fn atof_config_helpers_cover_file_path_and_replace_dots_policy() { + let config = AtofExporterConfig::new() + .with_output_directory("output") + .with_filename("events.jsonl"); + assert_eq!(config.path(), Some(PathBuf::from("output/events.jsonl"))); + assert_eq!( + AtofEndpointFieldNamePolicy::ReplaceDots.as_str(), + "replace_dots" + ); +} + +#[cfg(target_os = "linux")] +#[test] +fn atof_file_sink_reports_deferred_dev_full_write_failures() { + let exporter = AtofExporter::new( + AtofExporterConfig::new() + .with_output_directory("/dev") + .with_filename("full"), + ) + .unwrap(); + exporter.subscriber()(&make_mark_event("write-failure")); + assert!(exporter.force_flush().is_err()); + assert!(exporter.shutdown().is_err()); +} + #[test] #[cfg(feature = "atof-streaming")] fn websocket_helpers_cover_invalid_headers_and_timeout_reconnect_path() { @@ -1656,6 +1766,48 @@ fn websocket_helpers_cover_invalid_headers_and_timeout_reconnect_path() { }); } +#[test] +#[cfg(feature = "atof-streaming")] +fn endpoint_validation_helpers_cover_header_retry_and_collision_edges() { + enable_operational_logs(); + + let invalid_name = std::collections::HashMap::from([("bad header".into(), "value".into())]); + assert!(resolved_header_map(&invalid_name, &std::collections::HashMap::new()).is_err()); + + let invalid_value = std::collections::HashMap::from([("x-test".into(), "bad\nvalue".into())]); + assert!(resolved_header_map(&invalid_value, &std::collections::HashMap::new()).is_err()); + + let headers = std::collections::HashMap::from([("x-test".into(), "value".into())]); + let header_env = std::collections::HashMap::from([("x-test".into(), "UNUSED_ENV".into())]); + assert!(resolved_header_map(&headers, &header_env).is_err()); + + let mut retry = WebSocketRetryState::default(); + retry.record_failure(7); + retry.record_failure(7); + assert_eq!(retry.attempts, 2); + assert!(retry.warning_emitted); + retry.record_recovered(7); + assert_eq!(retry.attempts, 0); + assert!(!retry.warning_emitted); + assert!(retry.access_validated); + retry.record_recovered(7); + + let object = + serde_json::Map::from_iter([("key".into(), Json::Null), ("key_2".into(), Json::Null)]); + assert_eq!(collision_free_key(&object, "new".into()), "new"); + assert_eq!(collision_free_key(&object, "key".into()), "key_3"); + + for (transport, name) in [ + (AtofEndpointTransport::HttpPost, "http_post"), + (AtofEndpointTransport::Websocket, "websocket"), + (AtofEndpointTransport::Ndjson, "ndjson"), + ] { + assert_eq!(AtofEndpointTransport::parse(name), Some(transport)); + assert_eq!(transport.as_str(), name); + } + assert_eq!(AtofEndpointTransport::parse("invalid"), None); +} + #[test] #[cfg(feature = "atof-streaming")] fn ndjson_upload_close_timeout_acknowledges_close() { diff --git a/crates/core/tests/unit/observability/openinference_tests.rs b/crates/core/tests/unit/observability/openinference_tests.rs index 652311287..fed092236 100644 --- a/crates/core/tests/unit/observability/openinference_tests.rs +++ b/crates/core/tests/unit/observability/openinference_tests.rs @@ -79,6 +79,16 @@ fn attr_map(attributes: &[KeyValue]) -> HashMap { .collect() } +fn finished_span_named<'a>( + spans: &'a [opentelemetry_sdk::trace::SpanData], + name: &str, +) -> &'a opentelemetry_sdk::trace::SpanData { + spans + .iter() + .find(|span| span.name.as_ref() == name) + .unwrap_or_else(|| panic!("missing span {name}")) +} + fn assert_attr(attributes: &HashMap, key: &str, value: &str) { assert_eq!(attributes.get(key).map(String::as_str), Some(value)); } @@ -918,22 +928,17 @@ fn session_identity_is_projected_on_trace_roots_and_marks_only() { subscriber.force_flush().unwrap(); let spans = exporter.get_finished_spans().unwrap(); - let root = spans - .iter() - .find(|span| span.name.as_ref() == "identity-root") - .unwrap(); - let child = spans - .iter() - .find(|span| span.name.as_ref() == "identity-child") - .unwrap(); - let second_root = spans - .iter() - .find(|span| span.name.as_ref() == "identity-second-root") - .unwrap(); - let orphan_mark = spans - .iter() - .find(|span| span.name.as_ref() == "mark:session.start") - .unwrap(); + assert_session_root_and_child_identity(&spans, &instance_id); + assert_session_mark_identity(&spans); +} + +fn assert_session_root_and_child_identity( + spans: &[opentelemetry_sdk::trace::SpanData], + instance_id: &str, +) { + let root = finished_span_named(spans, "identity-root"); + let child = finished_span_named(spans, "identity-child"); + let second_root = finished_span_named(spans, "identity-second-root"); let root_attributes = attr_map(&root.attributes); assert_eq!(root_attributes["session.id"], "logical-session"); @@ -971,7 +976,12 @@ fn session_identity_is_projected_on_trace_roots_and_marks_only() { root.span_context.trace_id(), second_root.span_context.trace_id() ); +} +fn assert_session_mark_identity(spans: &[opentelemetry_sdk::trace::SpanData]) { + let root = finished_span_named(spans, "identity-root"); + let orphan_mark = finished_span_named(spans, "mark:session.start"); + let root_attributes = attr_map(&root.attributes); let mark_attributes = attr_map(&root.events.events[0].attributes); assert_eq!(mark_attributes["session.id"], "logical-session"); assert_eq!(mark_attributes["user.id"], "alice"); @@ -1083,12 +1093,18 @@ fn registered_subscriber_emits_spans_for_scope_push_pop_and_marks() { let spans = exporter.get_finished_spans().unwrap(); assert_eq!(spans.len(), 1); + assert_registered_scope_span(&spans[0]); +} - let span = &spans[0]; +fn assert_registered_scope_span(span: &opentelemetry_sdk::trace::SpanData) { assert_eq!(span.name.as_ref(), "otel_scope"); assert_eq!(span.events.events.len(), 1); assert_eq!(span.events.events[0].name.as_ref(), "otel_mark"); + assert_registered_scope_attributes(span); + assert_registered_scope_event_attributes(span); +} +fn assert_registered_scope_attributes(span: &opentelemetry_sdk::trace::SpanData) { let attributes = attr_map(&span.attributes); assert_eq!( attributes.get("openinference.span.kind"), @@ -1116,7 +1132,9 @@ fn registered_subscriber_emits_spans_for_scope_push_pop_and_marks() { attributes.get("metadata"), Some(&"{\"phase\":\"start\"}".to_string()) ); +} +fn assert_registered_scope_event_attributes(span: &opentelemetry_sdk::trace::SpanData) { let event_attributes = attr_map(&span.events.events[0].attributes); assert_eq!( event_attributes.get("nemo_relay.mark.data.step"), @@ -2846,18 +2864,15 @@ fn tool_projection_emits_generic_mark_as_parented_openinference_tool_span() { let spans = exporter.get_finished_spans().unwrap(); assert_eq!(spans.len(), 2); - let parent = spans - .iter() - .find(|span| span.name.as_ref() == "agent-turn") - .unwrap(); - let projected = spans - .iter() - .find(|span| span.name.as_ref() == "mark:plugin.output_compacted") - .unwrap(); + let parent = finished_span_named(&spans, "agent-turn"); + let projected = finished_span_named(&spans, "mark:plugin.output_compacted"); assert!(parent.events.events.is_empty()); assert_eq!(projected.parent_span_id, parent.span_context.span_id()); assert_eq!(projected.start_time, projected.end_time); + assert_tool_projection_attributes(projected); +} +fn assert_tool_projection_attributes(projected: &opentelemetry_sdk::trace::SpanData) { let attributes = attr_map(&projected.attributes); assert_eq!( attributes.get("openinference.span.kind"), @@ -3513,6 +3528,13 @@ fn pre_epoch_timestamps_round_trip_through_system_time() { #[test] fn helper_functions_cover_additional_openinference_branches() { + assert_openinference_scope_type_branches(); + assert_openinference_llm_common_attribute_branches(); + assert_openinference_tool_and_mark_attribute_branches(); + assert_openinference_input_usage_and_time_branches(); +} + +fn assert_openinference_scope_type_branches() { let function_end = make_end_event(Uuid::now_v7(), None, "fn-scope", ScopeType::Function, None); assert_eq!(span_name(&function_end), "fn-scope"); assert_eq!( @@ -3548,7 +3570,9 @@ fn helper_functions_cover_additional_openinference_branches() { assert_eq!(openinference_span_kind(Some(ScopeType::Custom)), "CHAIN"); assert_eq!(openinference_span_kind(Some(ScopeType::Unknown)), "CHAIN"); assert_eq!(openinference_span_kind(None), "CHAIN"); +} +fn assert_openinference_llm_common_attribute_branches() { let llm_end = Event::Scope(ScopeEvent::new( BaseEvent::builder() .name("chat") @@ -3615,7 +3639,9 @@ fn helper_functions_cover_additional_openinference_branches() { llm_attributes.get("metadata"), Some(&"{\"phase\":\"done\"}".to_string()) ); +} +fn assert_openinference_tool_and_mark_attribute_branches() { let tool_start = Event::Scope(ScopeEvent::new( BaseEvent::builder() .name("lookup") @@ -3689,7 +3715,9 @@ fn helper_functions_cover_additional_openinference_branches() { mark_attributes.get("nemo_relay.mark.metadata.source"), Some(&"unit".to_string()) ); +} +fn assert_openinference_input_usage_and_time_branches() { let llm_with_scalar_input = make_start_event( Uuid::now_v7(), None, diff --git a/crates/core/tests/unit/observability/otel_tests.rs b/crates/core/tests/unit/observability/otel_tests.rs index b703ecf8e..2399cf66d 100644 --- a/crates/core/tests/unit/observability/otel_tests.rs +++ b/crates/core/tests/unit/observability/otel_tests.rs @@ -329,6 +329,16 @@ fn attr_map(attributes: &[KeyValue]) -> HashMap { .collect() } +fn finished_span_named<'a>( + spans: &'a [opentelemetry_sdk::trace::SpanData], + name: &str, +) -> &'a opentelemetry_sdk::trace::SpanData { + spans + .iter() + .find(|span| span.name.as_ref() == name) + .unwrap_or_else(|| panic!("missing span {name}")) +} + fn make_start_event( uuid: Uuid, parent_uuid: Option, @@ -646,7 +656,11 @@ fn config_defaults_and_builder_overrides_are_applied() { .with_mark_exclude_names(["notification"]) .with_attribute_mapping("nemo_relay.model_name", "model.alias") .with_timeout(Duration::from_millis(1250)); + assert_config_builder_overrides(&config); + assert_config_defaults(&OpenTelemetryConfig::default()); +} +fn assert_config_builder_overrides(config: &OpenTelemetryConfig) { assert_eq!(config.transport, OtlpTransport::HttpBinary); assert_eq!(config.endpoint, "http://localhost:4318/v1/traces"); assert_eq!( @@ -665,8 +679,9 @@ fn config_defaults_and_builder_overrides_are_applied() { assert_eq!(config.mark_exclude_names, vec!["notification"]); assert_eq!(config.attribute_mappings.len(), 1); assert_eq!(config.timeout, Duration::from_millis(1250)); +} - let defaults = OpenTelemetryConfig::default(); +fn assert_config_defaults(defaults: &OpenTelemetryConfig) { assert_eq!(defaults.transport, OtlpTransport::HttpBinary); assert_eq!(defaults.service_name, "unknown_service"); assert_eq!(defaults.instrumentation_scope, "opentelemetry"); @@ -940,22 +955,17 @@ fn session_identity_is_projected_on_trace_roots_and_marks_only() { subscriber.force_flush().unwrap(); let spans = exporter.get_finished_spans().unwrap(); - let root = spans - .iter() - .find(|span| span.name.as_ref() == "identity-root") - .unwrap(); - let child = spans - .iter() - .find(|span| span.name.as_ref() == "identity-child") - .unwrap(); - let second_root = spans - .iter() - .find(|span| span.name.as_ref() == "identity-second-root") - .unwrap(); - let orphan_mark = spans - .iter() - .find(|span| span.name.as_ref() == "mark:session.start") - .unwrap(); + assert_session_root_and_child_identity(&spans, &instance_id); + assert_session_mark_identity(&spans); +} + +fn assert_session_root_and_child_identity( + spans: &[opentelemetry_sdk::trace::SpanData], + instance_id: &str, +) { + let root = finished_span_named(spans, "identity-root"); + let child = finished_span_named(spans, "identity-child"); + let second_root = finished_span_named(spans, "identity-second-root"); let root_attributes = attr_map(&root.attributes); assert_eq!(root_attributes["session.id"], "logical-session"); @@ -996,7 +1006,12 @@ fn session_identity_is_projected_on_trace_roots_and_marks_only() { root.span_context.trace_id(), second_root.span_context.trace_id() ); +} +fn assert_session_mark_identity(spans: &[opentelemetry_sdk::trace::SpanData]) { + let root = finished_span_named(spans, "identity-root"); + let orphan_mark = finished_span_named(spans, "mark:session.start"); + let root_attributes = attr_map(&root.attributes); let mark_attributes = attr_map(&root.events.events[0].attributes); assert_eq!(mark_attributes["session.id"], "logical-session"); assert_eq!(mark_attributes["user.id"], "alice"); @@ -1613,6 +1628,106 @@ fn gen_ai_projection_prefers_standard_names_and_normalized_provider_details() { } } +#[test] +fn gen_ai_projection_covers_optional_request_controls_and_finish_reasons() { + let request = Event::Scope(ScopeEvent::new( + BaseEvent::builder() + .uuid(Uuid::now_v7()) + .name("chat") + .data(json!({ + "headers": {}, + "content": { + "model": "gpt-5", + "messages": [{"role": "user", "content": "hello"}], + "max_tokens": 512, + "temperature": 0.4, + "top_p": 0.8, + "stop": ["done"], + "n": 1 + } + })) + .build(), + ScopeCategory::Start, + Vec::new(), + EventCategory::llm(), + None, + )); + let attributes = attr_map(&crate::observability::otel_genai::start_attributes( + &request, + )); + for (key, expected) in [ + ("gen_ai.provider.name", "openai"), + ("gen_ai.request.max_tokens", "512"), + ("gen_ai.request.temperature", "0.4"), + ("gen_ai.request.top_p", "0.8"), + ("gen_ai.request.stop_sequences", "[\"done\"]"), + ] { + assert_eq!(attributes.get(key), Some(&expected.to_string())); + } + assert!(!attributes.contains_key("gen_ai.request.choice.count")); + + for (reason, expected) in [ + (FinishReason::Complete, "stop"), + (FinishReason::Length, "length"), + (FinishReason::ContentFilter, "content_filter"), + ( + FinishReason::Unknown("provider_reason".to_string()), + "provider_reason", + ), + ] { + let event = make_scope_event_with_profile( + ScopeCategory::End, + Uuid::now_v7(), + None, + "chat", + ScopeType::Llm, + None, + Some( + CategoryProfile::builder() + .annotated_response(std::sync::Arc::new(AnnotatedLlmResponse { + finish_reason: Some(reason), + ..empty_annotated_response() + })) + .build(), + ), + ); + let attributes = attr_map(&crate::observability::otel_genai::end_attributes(&event)); + assert_eq!( + attributes.get("gen_ai.response.finish_reasons"), + Some(&format!("[\"{expected}\"]")) + ); + } +} + +#[test] +fn gen_ai_projection_reads_nested_scalar_fallbacks() { + let event = Event::Scope(ScopeEvent::new( + BaseEvent::builder() + .uuid(Uuid::now_v7()) + .name("embed") + .metadata(json!({ + "request": {"provider": true, "server_port": 8443}, + "response": {"response_model": 17}, + "usage": {"input_tokens": u64::MAX} + })) + .data(json!({"usage": {"prompt_tokens": 23}})) + .build(), + ScopeCategory::End, + Vec::new(), + EventCategory::from(ScopeType::Embedder), + None, + )); + let attributes = attr_map(&crate::observability::otel_genai::end_attributes(&event)); + assert_eq!( + attributes.get("gen_ai.response.model"), + Some(&"17".to_string()) + ); + assert_eq!( + attributes.get("gen_ai.usage.input_tokens"), + Some(&"23".to_string()) + ); +} + #[test] fn http_config_exports_scope_push_pop_and_marks_without_tokio_runtime() { let _guard = crate::observability::test_mutex().lock().unwrap(); @@ -2462,6 +2577,14 @@ fn llm_end_with_unannotated_openai_response_without_usage_omits_cost() { #[test] fn helper_functions_cover_additional_otel_branches() { + assert_otel_scope_and_model_attribute_branches(); + assert_otel_tool_attribute_branches(); + assert_otel_catalog_cost_branches(); + assert_otel_normalized_cost_branches(); + assert_otel_manual_cost_branches(); +} + +fn assert_otel_scope_and_model_attribute_branches() { let function_end = make_end_event(Uuid::now_v7(), None, "fn-scope", ScopeType::Function, None); assert_eq!(span_name(&function_end), "fn-scope"); assert_eq!( @@ -2531,7 +2654,9 @@ fn helper_functions_cover_additional_otel_branches() { response_model_attributes.get("nemo_relay.model_name"), Some(&"response-model".to_string()) ); +} +fn assert_otel_tool_attribute_branches() { let tool_event = Event::Scope(ScopeEvent::new( BaseEvent::builder() .name("lookup") @@ -2574,7 +2699,9 @@ fn helper_functions_cover_additional_otel_branches() { tool_end_attributes.get("nemo_relay.end.data.result"), Some(&"true".to_string()) ); +} +fn assert_otel_catalog_cost_branches() { { let _pricing_guard = pricing_test_mutex().lock().unwrap(); install_test_pricing("priced-model"); @@ -2653,7 +2780,9 @@ fn helper_functions_cover_additional_otel_branches() { Some(&"USD".to_string()) ); } +} +fn assert_otel_normalized_cost_branches() { let normalized_cost_event = make_scope_event_with_profile( ScopeCategory::End, Uuid::now_v7(), @@ -2747,7 +2876,9 @@ fn helper_functions_cover_additional_otel_branches() { Some(&"EUR".to_string()) ); } +} +fn assert_otel_manual_cost_branches() { { let _pricing_guard = pricing_test_mutex().lock().unwrap(); install_test_pricing("priced-model"); diff --git a/crates/core/tests/unit/observability/plugin_component_tests.rs b/crates/core/tests/unit/observability/plugin_component_tests.rs index 7a299edc0..4e5b039f6 100644 --- a/crates/core/tests/unit/observability/plugin_component_tests.rs +++ b/crates/core/tests/unit/observability/plugin_component_tests.rs @@ -83,7 +83,7 @@ fn start_http_capture_server(expected_requests: usize) -> (String, Arc (String, std::thread::JoinHandle>) { @@ -135,7 +135,7 @@ fn start_http_status_server( stream.read_exact(&mut body)?; write!( stream, - "HTTP/1.1 {status}\r\nContent-Length: 0\r\nConnection: close\r\n\r\n" + "HTTP/1.1 {status}\r\nContent-Length: 0\r\nETag: \"test-etag\"\r\nConnection: close\r\n\r\n" )?; stream.flush() }); @@ -408,16 +408,7 @@ fn default_config_and_component_conversion_cover_public_shape() { assert!(!atof.enabled); assert!(atof.sinks.is_empty()); - let parsed_atof: AtofSectionConfig = serde_json::from_value(json!({ - "sinks": [{"type": "stream", "name": "switchyard", "url": "http://localhost/events"}] - })) - .unwrap(); - let AtofSinkSectionConfig::Stream(stream) = &parsed_atof.sinks[0] else { - panic!("expected stream sink"); - }; - assert_eq!(stream.name.as_deref(), Some("switchyard")); - assert_eq!(stream.transport, "http_post"); - assert_eq!(stream.field_name_policy, "preserve"); + assert_default_stream_sink_shape(); let atif = AtifSectionConfig::default(); assert!(!atif.enabled); @@ -459,6 +450,19 @@ fn default_config_and_component_conversion_cover_public_shape() { assert_eq!(generic.config["atif"]["agent_name"], json!("NeMo Relay")); } +fn assert_default_stream_sink_shape() { + let parsed_atof: AtofSectionConfig = serde_json::from_value(json!({ + "sinks": [{"type": "stream", "name": "switchyard", "url": "http://localhost/events"}] + })) + .unwrap(); + let AtofSinkSectionConfig::Stream(stream) = &parsed_atof.sinks[0] else { + panic!("expected stream sink"); + }; + assert_eq!(stream.name.as_deref(), Some("switchyard")); + assert_eq!(stream.transport, "http_post"); + assert_eq!(stream.field_name_policy, "preserve"); +} + #[test] fn version_three_rejects_removed_otlp_controls() { let report = validate_plugin_config(&plugin_config(json!({ @@ -630,6 +634,96 @@ fn validate_opentelemetry_section_reports_empty_and_malformed_endpoints() { } } +#[test] +fn opentelemetry_registration_rejects_an_empty_endpoint_list() { + let mut context = PluginRegistrationContext::new(); + let error = register_opentelemetry( + OpenTelemetrySectionConfig { + enabled: true, + endpoints: Vec::new(), + }, + &mut context, + ) + .unwrap_err(); + assert!(error.to_string().contains("at least one endpoint")); +} + +#[test] +fn atof_stream_header_validation_reports_invalid_values_and_environment_names() { + let _guard = crate::observability::test_mutex() + .lock() + .unwrap_or_else(|error| error.into_inner()); + let policy = ConfigPolicy::default(); + let mut diagnostics = Vec::new(); + + validate_atof_stream_header_env(&mut diagnostics, &policy, "header_env.empty", ""); + validate_atof_stream_header_env( + &mut diagnostics, + &policy, + "header_env.padded", + " PADDED_ENV ", + ); + + let blank = "NEMO_RELAY_TEST_BLANK_ATOF_STREAM_HEADER_ENV"; + // SAFETY: the observability mutex serializes access to this test-only variable. + unsafe { std::env::set_var(blank, " ") }; + validate_atof_stream_header_env(&mut diagnostics, &policy, "header_env.blank", blank); + // SAFETY: cleanup of the test-only environment variable. + unsafe { std::env::remove_var(blank) }; + + assert!(diagnostics.iter().any(|diagnostic| { + diagnostic + .message + .contains("non-empty environment variable") + })); + assert!( + diagnostics + .iter() + .any(|diagnostic| diagnostic.message.contains("surrounding whitespace")) + ); + assert!(diagnostics.iter().any(|diagnostic| { + diagnostic + .message + .contains("environment variable that is blank") + })); + + #[cfg(feature = "atof-streaming")] + { + diagnostics.clear(); + validate_atof_stream_header( + &mut diagnostics, + &policy, + "headers.invalid", + "x-test", + "invalid\nvalue", + ); + assert!( + diagnostics + .iter() + .any(|diagnostic| diagnostic.message.contains("value is invalid")) + ); + } +} + +#[test] +fn atif_dispatcher_surfaces_fatal_and_disabled_local_sink_states() { + let mut dispatcher = AtifDispatcher::new(AtifSectionConfig::default()); + dispatcher.fatal_error = Some("fatal export failure".into()); + assert!( + dispatcher + .last_error_result() + .unwrap_err() + .to_string() + .contains("fatal export failure") + ); + + dispatcher.fatal_error = None; + dispatcher + .sink_errors + .insert(SinkLabel::Local, "local write failed".into()); + assert!(dispatcher.sink_targets().is_empty()); +} + #[test] fn opentelemetry_endpoint_header_env_rejects_missing_and_duplicate_headers() { let _guard = crate::observability::test_mutex().lock().unwrap(); @@ -3286,3 +3380,136 @@ fn atif_storage_private_helpers_resolve_env_and_key_prefix_branches() { std::env::remove_var(token); } } + +#[cfg(feature = "object-store")] +fn http_storage_config(endpoint: impl Into) -> HttpStorageConfig { + HttpStorageConfig { + endpoint: endpoint.into(), + headers: std::collections::HashMap::new(), + header_env: std::collections::HashMap::new(), + timeout_millis: 1_000, + } +} + +#[test] +#[cfg(feature = "object-store")] +fn http_upload_config_rejects_endpoint_timeout_and_header_errors() { + let _guard = crate::observability::test_mutex() + .lock() + .unwrap_or_else(|error| error.into_inner()); + for endpoint in [" http://example.com", "://", "ftp://example.com"] { + assert!(HttpUploadConfig::resolve(2, &http_storage_config(endpoint)).is_err()); + } + + let mut config = http_storage_config("https://example.com/atif"); + config.timeout_millis = 0; + assert!(HttpUploadConfig::resolve(2, &config).is_err()); + + config.timeout_millis = 1_000; + config.headers.insert("bad header".into(), "value".into()); + assert!(HttpUploadConfig::resolve(2, &config).is_err()); + + config.headers.clear(); + config.headers.insert("x-bad".into(), "bad\nvalue".into()); + assert!(HttpUploadConfig::resolve(2, &config).is_err()); + + let variable = "NEMO_RELAY_TEST_ATIF_HTTP_RESOLVE_ZZZZ"; + // SAFETY: this uniquely named environment variable is serialized by the observability mutex. + unsafe { std::env::set_var(variable, "Bearer resolved") }; + config.headers.clear(); + config + .header_env + .insert("authorization".into(), variable.into()); + let resolved = HttpUploadConfig::resolve(2, &config).unwrap(); + assert_eq!( + resolved.headers.get("authorization").map(String::as_str), + Some("Bearer resolved") + ); + // SAFETY: cleanup of the test-only environment variable. + unsafe { std::env::remove_var(variable) }; +} + +#[tokio::test] +#[cfg(feature = "object-store")] +async fn post_atif_http_reports_transport_failure() { + let config = HttpUploadConfig::resolve(0, &http_storage_config("http://127.0.0.1:1")).unwrap(); + let client = reqwest::Client::builder() + .timeout(config.timeout) + .build() + .unwrap(); + assert!( + post_atif_http( + &client, + &config, + "trajectory.json".into(), + "session".into(), + b"{}".to_vec(), + ) + .await + .unwrap_err() + .to_string() + .contains("HTTP ATIF upload") + ); +} + +#[test] +#[cfg(feature = "object-store")] +fn s3_remote_storage_uploads_to_a_custom_http_endpoint() { + let _guard = crate::observability::test_mutex().lock().unwrap(); + let variable = "NEMO_RELAY_TEST_S3_UPLOAD_SECRET_ZZZZ"; + // SAFETY: the observability mutex serializes access to this test-only variable. + unsafe { std::env::set_var(variable, "secret") }; + let (endpoint, server) = start_http_status_server("200 OK"); + let storage = AtifRemoteStorage::from_config( + 7, + &AtifStorageConfig::S3(S3StorageConfig { + bucket: "test-bucket".into(), + key_prefix: Some("trajectories".into()), + access_key_id: Some("access".into()), + secret_access_key_var: Some(variable.into()), + session_token_var: None, + region: Some("us-east-1".into()), + endpoint_url: Some(endpoint), + allow_http: Some(true), + }), + ) + .unwrap(); + let result = storage.put("trajectory.json", "session", b"{}"); + server.join().unwrap().unwrap(); + // SAFETY: cleanup of the test-only environment variable. + unsafe { std::env::remove_var(variable) }; + result.unwrap(); +} + +#[test] +fn observability_private_editor_and_validation_helpers_cover_edge_configs() { + assert_eq!( + default_atof_file_sink_editor_value(), + json!({ + "type": "file", + "mode": "append" + }) + ); + assert_eq!( + default_atof_stream_sink_editor_value()["transport"], + json!("http_post") + ); + assert_eq!( + default_opentelemetry_endpoint_editor_value()["service_name"], + json!("unknown_service") + ); + let field = otel_editor_field("optional", EditorFieldKind::String, &[], true); + assert_eq!(field.name, "optional"); + assert!(field.optional); + + let plugin = ObservabilityPlugin; + assert!(!plugin.allows_multiple_components()); + for value in [ + json!({"atof": {"enabled": true, "filename": "removed.jsonl"}}), + json!({"atof": {"enabled": true, "sinks": []}}), + json!({"opentelemetry": {"enabled": true, "endpoints": []}}), + ] { + let config = value.as_object().unwrap(); + assert!(!plugin.validate(config).is_empty()); + } +} diff --git a/crates/core/tests/unit/plugin_tests.rs b/crates/core/tests/unit/plugin_tests.rs index 8601e18da..4e79c9e45 100644 --- a/crates/core/tests/unit/plugin_tests.rs +++ b/crates/core/tests/unit/plugin_tests.rs @@ -707,6 +707,33 @@ fn test_layer_config_preserves_multi_instance_kinds() { assert_eq!(components[1]["config"]["tag"], json!("second")); } +#[test] +fn config_layering_helpers_cover_malformed_and_scalar_shapes() { + let mut scalar = json!("lower"); + layer_config(&mut scalar, json!("higher")); + assert_eq!(scalar, json!("higher")); + + let mut malformed_left = json!({"not": "components"}); + merge_plugin_components(&mut malformed_left, json!([])); + assert_eq!(malformed_left, json!([])); + + let mut malformed_right = json!([]); + merge_plugin_components(&mut malformed_right, json!({"not": "components"})); + assert_eq!(malformed_right, json!({"not": "components"})); + + let mut components = json!([]); + merge_plugin_components(&mut components, json!([{"enabled": true}])); + assert_eq!(components, json!([{"enabled": true}])); + + let mut malformed_component = json!("lower"); + merge_plugin_component(&mut malformed_component, json!({"kind": "higher"})); + assert_eq!(malformed_component, json!({"kind": "higher"})); + + let mut object = json!({"existing": true}); + merge_json_value(&mut object, json!({"added": true})); + assert_eq!(object, json!({"existing": true, "added": true})); +} + #[test] fn test_config_report_has_errors() { let report = ConfigReport { @@ -2257,6 +2284,37 @@ fn test_load_plugin_config_files_rejects_version_before_layering() { } } +#[test] +fn test_plugin_config_loading_reports_read_parse_and_version_type_errors() { + let dir = tempfile::tempdir().unwrap(); + let unreadable = dir.path().join("directory.toml"); + std::fs::create_dir(&unreadable).unwrap(); + assert!( + load_plugin_config_files([unreadable]) + .unwrap_err() + .to_string() + .contains("failed to read") + ); + + let malformed = dir.path().join("malformed.toml"); + std::fs::write(&malformed, "version = [").unwrap(); + assert!( + load_plugin_config_files([malformed]) + .unwrap_err() + .to_string() + .contains("invalid plugin TOML") + ); + + assert!( + merge_plugin_config_documents([ + (dir.path().join("typed.toml"), json!({"version": "one"}),) + ]) + .unwrap_err() + .to_string() + .contains("invalid plugin config version") + ); +} + #[test] fn test_default_plugin_config_paths_order_user_project_system() { let dir = tempfile::tempdir().unwrap(); diff --git a/crates/core/tests/unit/plugins/nemo_guardrails/component_tests.rs b/crates/core/tests/unit/plugins/nemo_guardrails/component_tests.rs index b1318e7f7..153d88b4d 100644 --- a/crates/core/tests/unit/plugins/nemo_guardrails/component_tests.rs +++ b/crates/core/tests/unit/plugins/nemo_guardrails/component_tests.rs @@ -499,7 +499,15 @@ fn invalid_shapes_and_values_are_reported() { .lock() .unwrap_or_else(|err| err.into_inner()); reset_runtime(); + assert_invalid_shape_and_mode(); + assert_invalid_local_config(); + assert_invalid_remote_identity_and_codec(); + assert_remote_tool_surface_validation(); + assert_empty_and_mixed_config_values(); + assert_request_defaults_validation(); +} +fn assert_invalid_shape_and_mode() { let invalid_shape = validate_plugin_config(&plugin_config(json!({ "version": "one", }))); @@ -534,7 +542,9 @@ fn invalid_shapes_and_values_are_reported() { .any(|diag| diag.field.as_deref() == Some("mode") && diag.message.contains("mode must be 'remote' or 'local'")) ); +} +fn assert_invalid_local_config() { let local_missing_source = validate_plugin_config(&plugin_config(json!({ "mode": "local", "codec": "openai_chat", @@ -576,7 +586,9 @@ fn invalid_shapes_and_values_are_reported() { .any(|diag| diag.field.as_deref() == Some("remote") && diag.message.contains("cannot be used when mode is 'local'")) ); +} +fn assert_invalid_remote_identity_and_codec() { let remote_missing_identity = validate_plugin_config(&plugin_config(json!({ "mode": "remote", "codec": "openai_chat", @@ -664,7 +676,9 @@ fn invalid_shapes_and_values_are_reported() { .contains("remote mode currently supports only codec = 'openai_chat'") }) ); +} +fn assert_remote_tool_surface_validation() { let unsupported_remote_tool_input = validate_plugin_config(&plugin_config(json!({ "mode": "remote", "codec": "openai_chat", @@ -697,7 +711,9 @@ fn invalid_shapes_and_values_are_reported() { } }))); assert!(!supported_remote_tool_output.has_errors()); +} +fn assert_empty_and_mixed_config_values() { let remote_empty_fields = validate_plugin_config(&plugin_config(json!({ "mode": "remote", "codec": "openai_chat", @@ -810,7 +826,9 @@ fn invalid_shapes_and_values_are_reported() { .iter() .any(|diag| diag.field.as_deref() == Some("local.python_path")) ); +} +fn assert_request_defaults_validation() { let local_request_defaults = validate_plugin_config(&plugin_config(json!({ "mode": "local", "codec": "openai_chat", diff --git a/crates/core/tests/unit/plugins/nemo_guardrails/local_python_tests.rs b/crates/core/tests/unit/plugins/nemo_guardrails/local_python_tests.rs index 5fa7418e2..3353112bd 100644 --- a/crates/core/tests/unit/plugins/nemo_guardrails/local_python_tests.rs +++ b/crates/core/tests/unit/plugins/nemo_guardrails/local_python_tests.rs @@ -112,6 +112,26 @@ fn python_executable_uses_python_environment_before_default() { ); } +#[test] +fn local_worker_start_reports_an_unavailable_python_executable() { + let config = NeMoGuardrailsConfig { + local: Some(LocalBackendConfig { + python_executable: Some("nemo-relay-python-that-does-not-exist".to_string()), + ..LocalBackendConfig::default() + }), + ..NeMoGuardrailsConfig::default() + }; + + let error = LocalGuardrailsWorker::start(&config) + .err() + .expect("an unavailable executable should fail worker startup"); + assert!( + error + .to_string() + .contains("failed to start NeMo Guardrails local Python worker") + ); +} + #[test] fn worker_python_path_prepends_configured_path_to_inherited_pythonpath() { let configured = std::path::PathBuf::from("fake-guardrails"); @@ -514,6 +534,200 @@ async fn guarded_provider_stream_reports_block_after_forwarded_chunks() { assert!(chunk_rx.recv().await.is_none()); } +#[tokio::test] +async fn stream_monitor_errors_are_forwarded_to_the_provider_stream() { + async fn panicking_monitor() -> FlowResult<()> { + panic!("monitor panicked"); + } + + let blocked = Arc::new(Mutex::new(None)); + + for monitor in [ + tokio::spawn(async { Err(FlowError::Internal("monitor failed".into())) }), + tokio::spawn(panicking_monitor()), + ] { + let (chunk_tx, mut chunk_rx) = mpsc::channel(1); + assert!(send_stream_monitor_error(monitor, &chunk_tx, &blocked).await); + assert!(chunk_rx.recv().await.unwrap().is_err()); + } + + let (chunk_tx, mut chunk_rx) = mpsc::channel(1); + *blocked.lock().unwrap() = Some("blocked output".into()); + assert!(send_stream_monitor_error(tokio::spawn(async { Ok(()) }), &chunk_tx, &blocked).await); + assert!( + chunk_rx + .recv() + .await + .unwrap() + .unwrap_err() + .to_string() + .contains("blocked output") + ); + + *blocked.lock().unwrap() = None; + assert!(!send_stream_monitor_error(tokio::spawn(async { Ok(()) }), &chunk_tx, &blocked).await); +} + +#[tokio::test] +async fn guarded_provider_stream_forwards_provider_and_channel_failures() { + let provider_error = FlowError::Internal("provider failed".into()); + let provider_stream = LlmJsonStream::new(tokio_stream::iter(vec![Err(provider_error)])); + let (text_tx, mut text_rx) = mpsc::channel(2); + let (chunk_tx, mut chunk_rx) = mpsc::channel(1); + let (_cancel_tx, cancel_rx) = watch::channel(false); + let (closed_tx, closed_rx) = watch::channel(None); + forward_guarded_provider_stream( + provider_stream, + LocalGuardrailsCodec::OpenAIChat, + text_tx, + chunk_tx, + tokio::spawn(async { Ok(()) }), + Arc::new(Mutex::new(None)), + cancel_rx, + closed_tx, + ) + .await; + assert!(chunk_rx.recv().await.unwrap().is_err()); + assert_eq!(text_rx.recv().await, Some(None)); + assert!(closed_rx.borrow().as_ref().unwrap().is_ok()); + + let provider_stream = LlmJsonStream::new(tokio_stream::iter(vec![Ok(json!({ + "choices": [{"delta": {"content": "hello"}}] + }))])); + let (text_tx, text_rx) = mpsc::channel(1); + drop(text_rx); + let (chunk_tx, mut chunk_rx) = mpsc::channel(1); + let (_cancel_tx, cancel_rx) = watch::channel(false); + let (closed_tx, _closed_rx) = watch::channel(None); + forward_guarded_provider_stream( + provider_stream, + LocalGuardrailsCodec::OpenAIChat, + text_tx, + chunk_tx, + tokio::spawn(async { Err(FlowError::Internal("monitor closed".into())) }), + Arc::new(Mutex::new(None)), + cancel_rx, + closed_tx, + ) + .await; + assert!(chunk_rx.recv().await.unwrap().is_err()); +} + +#[tokio::test] +async fn guarded_provider_stream_handles_preblocked_dropped_and_cancelled_consumers() { + let stream_chunk = || { + LlmJsonStream::new(tokio_stream::iter(vec![Ok(json!({ + "choices": [{"delta": {"content": "hello"}}] + }))])) + }; + + let (text_tx, _text_rx) = mpsc::channel(3); + let (chunk_tx, mut chunk_rx) = mpsc::channel(2); + let (_cancel_tx, cancel_rx) = watch::channel(false); + let (closed_tx, _closed_rx) = watch::channel(None); + forward_guarded_provider_stream( + stream_chunk(), + LocalGuardrailsCodec::OpenAIChat, + text_tx, + chunk_tx, + tokio::spawn(async { Ok(()) }), + Arc::new(Mutex::new(Some("already blocked".into()))), + cancel_rx, + closed_tx, + ) + .await; + assert!(chunk_rx.recv().await.unwrap().is_err()); + + let (text_tx, _text_rx) = mpsc::channel(3); + let (chunk_tx, chunk_rx) = mpsc::channel(1); + drop(chunk_rx); + let (_cancel_tx, cancel_rx) = watch::channel(false); + let (closed_tx, _closed_rx) = watch::channel(None); + forward_guarded_provider_stream( + stream_chunk(), + LocalGuardrailsCodec::OpenAIChat, + text_tx, + chunk_tx, + tokio::spawn(async { Ok(()) }), + Arc::new(Mutex::new(None)), + cancel_rx, + closed_tx, + ) + .await; + + let (text_tx, mut text_rx) = mpsc::channel(1); + let (chunk_tx, _chunk_rx) = mpsc::channel(1); + let (cancel_tx, cancel_rx) = watch::channel(false); + cancel_tx.send_replace(true); + let (closed_tx, _closed_rx) = watch::channel(None); + forward_guarded_provider_stream( + stream_chunk(), + LocalGuardrailsCodec::OpenAIChat, + text_tx, + chunk_tx, + tokio::spawn(async { std::future::pending::>().await }), + Arc::new(Mutex::new(None)), + cancel_rx, + closed_tx, + ) + .await; + assert_eq!(text_rx.recv().await, Some(None)); +} + +#[test] +fn local_codec_and_rewrite_helpers_cover_all_provider_surfaces() { + for (codec, surface) in [ + ( + LocalGuardrailsCodec::OpenAIChat, + ProviderSurface::OpenAIChat, + ), + ( + LocalGuardrailsCodec::OpenAIResponses, + ProviderSurface::OpenAIResponses, + ), + ( + LocalGuardrailsCodec::AnthropicMessages, + ProviderSurface::AnthropicMessages, + ), + ] { + assert_eq!(codec.provider_surface(), surface); + assert_eq!( + LocalGuardrailsCodec::from_provider_surface(surface).provider_surface(), + surface + ); + } + + let mut config = NeMoGuardrailsConfig { + input: false, + output: false, + ..Default::default() + }; + assert!(resolve_codec(&config).unwrap().is_none()); + config.input = true; + assert!(resolve_codec(&config).is_err()); + config.codec = Some("unsupported".into()); + assert!(resolve_codec(&config).is_err()); + + let mut annotated = AnnotatedLlmRequest { + messages: vec![Message::Assistant { + content: None, + tool_calls: None, + name: None, + }], + ..Default::default() + }; + replace_last_role_content(&mut annotated, "assistant", "rewritten".into()).unwrap(); + assert!(matches!( + &annotated.messages[0], + Message::Assistant { + content: Some(MessageContent::Text(content)), + .. + } if content == "rewritten" + )); + assert!(replace_last_role_content(&mut annotated, "user", "missing".into()).is_err()); + assert!(modified_tool_payload("[]", "arguments").is_err()); +} + #[test] fn parse_check_result_rejects_unknown_status() { assert!(matches!( @@ -531,6 +745,293 @@ fn parse_check_result_rejects_unknown_status() { .contains("unexpected worker check status: surprising"), "unexpected error: {error}" ); + assert!(parse_check_result(json!({"status": 7})).is_err()); +} + +#[test] +fn worker_envelope_helpers_cover_delivery_shutdown_and_default_results() { + assert!(set_request_id(&mut Json::Null, "1").is_err()); + let mut payload = json!({"command": "check"}); + set_request_id(&mut payload, "request-1").unwrap(); + assert_eq!(payload["id"], json!("request-1")); + + let waiters = Arc::new(Mutex::new(HashMap::new())); + let stream_events = Arc::new(Mutex::new(HashMap::new())); + let (waiter_tx, waiter_rx) = std_mpsc::channel(); + waiters.lock().unwrap().insert("unary".into(), waiter_tx); + dispatch_worker_envelope( + &waiters, + &stream_events, + WorkerEnvelope { + id: "unary".into(), + ok: true, + result: Some(json!({"ok": true})), + error: None, + event: None, + message: None, + }, + ); + assert!(waiter_rx.recv().unwrap().ok); + + let (stream_tx, mut stream_rx) = mpsc::unbounded_channel(); + stream_events + .lock() + .unwrap() + .insert("stream".into(), stream_tx); + dispatch_worker_envelope( + &waiters, + &stream_events, + WorkerEnvelope { + id: "stream".into(), + ok: true, + result: None, + error: None, + event: Some("done".into()), + message: None, + }, + ); + assert_eq!(stream_rx.try_recv().unwrap().event.as_deref(), Some("done")); + + let (waiter_tx, waiter_rx) = std_mpsc::channel(); + waiters.lock().unwrap().insert("closed".into(), waiter_tx); + let (stream_tx, mut stream_rx) = mpsc::unbounded_channel(); + stream_events + .lock() + .unwrap() + .insert("closed-stream".into(), stream_tx); + notify_worker_closed(&waiters, &stream_events, "worker gone".into()); + assert_eq!( + waiter_rx.recv().unwrap().error.as_deref(), + Some("worker gone") + ); + assert_eq!( + stream_rx.try_recv().unwrap().error.as_deref(), + Some("worker gone") + ); + + assert_eq!( + worker_result(WorkerEnvelope { + id: "ok".into(), + ok: true, + result: None, + error: None, + event: None, + message: None, + }) + .unwrap(), + Json::Null + ); + assert!( + worker_result(WorkerEnvelope { + id: "error".into(), + ok: false, + result: None, + error: None, + event: None, + message: None, + }) + .unwrap_err() + .to_string() + .contains("worker failed") + ); +} + +#[test] +fn worker_command_writer_reports_stored_and_closed_channel_errors() { + let (sender, receiver) = std_mpsc::channel(); + let writer = WorkerCommandWriter { + sender, + error: Arc::new(Mutex::new(Some("broken pipe".into()))), + handle: None, + }; + assert!( + writer + .send("ignored".into()) + .unwrap_err() + .to_string() + .contains("broken pipe") + ); + drop(receiver); + + let (sender, receiver) = std_mpsc::channel(); + drop(receiver); + let writer = WorkerCommandWriter { + sender, + error: Arc::new(Mutex::new(None)), + handle: None, + }; + assert!( + writer + .send("ignored".into()) + .unwrap_err() + .to_string() + .contains("channel closed") + ); +} + +#[cfg(unix)] +#[test] +fn worker_reader_handles_blank_valid_invalid_and_eof_lines() { + { + let worker = monitor_test_worker(); + let (sender, receiver) = std_mpsc::channel(); + worker.waiters.lock().unwrap().insert("ok".into(), sender); + let mut valid_source = Command::new("sh") + .arg("-c") + .arg("printf '\n{\"id\":\"ok\",\"ok\":true}\n'") + .stdout(Stdio::piped()) + .spawn() + .unwrap(); + worker.spawn_reader(valid_source.stdout.take().unwrap()); + assert!(receiver.recv_timeout(Duration::from_secs(1)).unwrap().ok); + valid_source.wait().unwrap(); + } + + let worker = monitor_test_worker(); + let (sender, receiver) = std_mpsc::channel(); + worker + .waiters + .lock() + .unwrap() + .insert("invalid".into(), sender); + let mut invalid_source = Command::new("sh") + .arg("-c") + .arg("printf 'not-json\n'") + .stdout(Stdio::piped()) + .spawn() + .unwrap(); + worker.spawn_reader(invalid_source.stdout.take().unwrap()); + assert!( + receiver + .recv_timeout(Duration::from_secs(1)) + .unwrap() + .error + .unwrap() + .contains("invalid worker response") + ); + invalid_source.wait().unwrap(); +} + +#[cfg(unix)] +#[test] +fn closed_worker_writer_cleans_up_unary_and_stream_registrations() { + let worker = monitor_test_worker(); + let mut request = json!({"command": "check"}); + assert!(worker.send_request(&mut request).is_err()); + assert!(worker.waiters.lock().unwrap().is_empty()); + + assert!(worker.start_stream(vec![json!({"role": "user"})]).is_err()); + assert!(worker.stream_events.lock().unwrap().is_empty()); +} + +#[tokio::test] +async fn guarded_stream_close_reports_an_early_cleanup_exit() { + let (_chunk_tx, chunk_rx) = mpsc::channel(1); + let (cancel, _cancel_rx) = watch::channel(false); + let (closed_tx, closed) = watch::channel(None); + drop(closed_tx); + let mut stream = GuardedProviderStream { + receiver: ReceiverStream::new(chunk_rx), + cancel, + closed, + }; + assert!( + Pin::new(&mut stream) + .close() + .await + .unwrap_err() + .to_string() + .contains("cleanup task ended early") + ); +} + +#[cfg(unix)] +fn monitor_test_worker() -> Arc { + Arc::new(LocalGuardrailsWorker { + writer: Mutex::new(None), + child: Mutex::new(Command::new("sleep").arg("60").spawn().unwrap()), + waiters: Arc::new(Mutex::new(HashMap::new())), + stream_events: Arc::new(Mutex::new(HashMap::new())), + next_id: AtomicU64::new(0), + shutdown_started: AtomicBool::new(false), + }) +} + +#[cfg(unix)] +async fn run_monitor_event(event: Option) -> (FlowResult<()>, Option) { + let worker = monitor_test_worker(); + let (text_tx, text_rx) = mpsc::channel(1); + let (event_tx, event_rx) = mpsc::unbounded_channel(); + let blocked = Arc::new(Mutex::new(None)); + if let Some(event) = event { + event_tx.send(event).unwrap(); + } + drop(event_tx); + let result = monitor_guardrails_stream( + worker, + "stream-id".into(), + text_rx, + event_rx, + Arc::clone(&blocked), + ) + .await; + drop(text_tx); + let message = blocked.lock().unwrap().clone(); + (result, message) +} + +#[cfg(unix)] +#[tokio::test] +async fn stream_monitor_handles_terminal_worker_event_variants() { + let envelope = |ok, event: &str, error: Option<&str>, message: Option<&str>| WorkerEnvelope { + id: "stream-id".into(), + ok, + result: None, + error: error.map(str::to_string), + event: Some(event.into()), + message: message.map(str::to_string), + }; + + let (result, blocked) = run_monitor_event(Some(envelope( + true, + "blocked", + None, + Some("policy blocked output"), + ))) + .await; + assert!(result.is_ok()); + assert_eq!(blocked.as_deref(), Some("policy blocked output")); + + assert!( + run_monitor_event(Some(envelope(true, "done", None, None))) + .await + .0 + .is_ok() + ); + assert!( + run_monitor_event(Some(envelope(false, "error", None, None))) + .await + .0 + .unwrap_err() + .to_string() + .contains("worker stream failed") + ); + assert!( + run_monitor_event(Some(envelope(true, "unexpected", None, None))) + .await + .0 + .unwrap_err() + .to_string() + .contains("unknown stream event") + ); + assert!( + run_monitor_event(None) + .await + .0 + .unwrap_err() + .to_string() + .contains("closed unexpectedly") + ); } #[test] @@ -573,6 +1074,23 @@ fn stream_text_extraction_handles_supported_codecs() { ), Some("hello".to_string()) ); + for (codec, chunk) in [ + (LocalGuardrailsCodec::OpenAIChat, Json::Null), + ( + LocalGuardrailsCodec::OpenAIResponses, + json!({"type": "response.completed", "delta": "ignored"}), + ), + ( + LocalGuardrailsCodec::AnthropicMessages, + json!({"type": "message_delta", "delta": {"type": "text_delta", "text": "ignored"}}), + ), + ( + LocalGuardrailsCodec::AnthropicMessages, + json!({"type": "content_block_delta", "delta": {"type": "input_json_delta"}}), + ), + ] { + assert_eq!(extract_stream_text(codec, &chunk), None); + } } #[cfg(unix)] diff --git a/crates/core/tests/unit/plugins/nemo_guardrails/remote_coverage_tests.rs b/crates/core/tests/unit/plugins/nemo_guardrails/remote_coverage_tests.rs index 2e16c7bc6..2ae2aa296 100644 --- a/crates/core/tests/unit/plugins/nemo_guardrails/remote_coverage_tests.rs +++ b/crates/core/tests/unit/plugins/nemo_guardrails/remote_coverage_tests.rs @@ -91,6 +91,24 @@ fn spawn_http_response( format!("http://{address}") } +fn spawn_truncated_http_response(status: &'static str, content_type: &'static str) -> String { + let listener = TcpListener::bind("127.0.0.1:0").unwrap(); + let address = listener.local_addr().unwrap(); + std::thread::spawn(move || { + let (mut stream, _) = listener.accept().unwrap(); + stream + .set_read_timeout(Some(Duration::from_secs(2))) + .unwrap(); + read_http_request(&mut stream); + write!( + stream, + "HTTP/1.1 {status}\r\nContent-Type: {content_type}\r\nContent-Length: 64\r\nConnection: close\r\n\r\npartial" + ) + .unwrap(); + }); + format!("http://{address}") +} + fn read_http_request(stream: &mut std::net::TcpStream) { let mut request = Vec::new(); let mut buffer = [0; 1024]; @@ -205,6 +223,22 @@ fn request_body_and_guardrails_config_helpers_cover_defaults() { runtime.build_request_body(&invalid_request, false), "request content is not an object", ); + assert_flow_error_contains( + runtime.build_request_body( + &LlmRequest { + headers: Map::new(), + content: json!({ + "model": "gpt-4o-mini", + "messages": [], + "tools": [{"type": "function", "function": {"name": "lookup"}}] + }), + }, + false, + ), + "does not support OpenAI tool definitions", + ); + runtime.record_access_status(reqwest::StatusCode::OK); + runtime.record_access_status(reqwest::StatusCode::OK); let defaults = RequestDefaultsConfig { context: Some(json!({"tenant": "test"})), @@ -772,6 +806,53 @@ async fn remote_execute_transport_and_stream_status_errors_are_reported() { ); } +#[tokio::test] +#[allow(clippy::await_holding_lock)] +async fn truncated_remote_bodies_report_buffered_streaming_and_tool_errors() { + let _guard = crate::plugins::nemo_guardrails::test_mutex() + .lock() + .unwrap_or_else(|err| err.into_inner()); + crate::shared_runtime::reset_runtime_owner_for_tests(); + crate::api::runtime::set_thread_scope_stack(crate::api::runtime::create_scope_stack()); + + let buffered = + runtime_with_endpoint(spawn_truncated_http_response("200 OK", "application/json")); + assert_flow_error_contains( + buffered.execute(simple_chat_request(), false).await, + "failed to read remote response body", + ); + + let status = runtime_with_endpoint(spawn_truncated_http_response( + "503 Service Unavailable", + "text/plain", + )); + assert_flow_error_contains( + status.execute_stream(simple_chat_request()).await, + "failed to read remote stream error body", + ); + + let streaming = + runtime_with_endpoint(spawn_truncated_http_response("200 OK", "text/event-stream")); + let mut stream = streaming + .execute_stream(simple_chat_request()) + .await + .expect("stream opens after headers"); + assert_flow_error_contains( + stream + .next() + .await + .expect("truncated body reports an error"), + "failed to read remote stream chunk", + ); + + let tool = runtime_with_endpoint(spawn_truncated_http_response("200 OK", "application/json")); + assert_flow_error_contains( + tool.check_tool_input("weather_lookup", &json!({"city": "Phoenix"})) + .await, + "failed to read remote response body", + ); +} + #[tokio::test] #[allow(clippy::await_holding_lock)] async fn remote_execute_stream_yields_completed_events_and_reports_malformed_final_event() { @@ -848,6 +929,25 @@ async fn remote_execute_stream_reports_malformed_named_final_event() { ); } +#[tokio::test] +#[allow(clippy::await_holding_lock)] +async fn remote_stream_decoder_flushes_an_unterminated_final_event() { + let _guard = crate::plugins::nemo_guardrails::test_mutex() + .lock() + .unwrap_or_else(|err| err.into_inner()); + crate::shared_runtime::reset_runtime_owner_for_tests(); + crate::api::runtime::set_thread_scope_stack(crate::api::runtime::create_scope_stack()); + + let runtime = runtime_with_endpoint(spawn_http_response( + "200 OK", + "text/event-stream", + "data: {\"final\":true}", + )); + let mut stream = runtime.execute_stream(simple_chat_request()).await.unwrap(); + assert_eq!(stream.next().await.unwrap().unwrap()["final"], json!(true)); + assert!(stream.next().await.is_none()); +} + #[tokio::test] #[allow(clippy::await_holding_lock)] async fn remote_tool_input_checks_cover_rewrite_block_noop_and_invalid_json() { diff --git a/crates/core/tests/unit/subscriber_dispatcher_tests.rs b/crates/core/tests/unit/subscriber_dispatcher_tests.rs index 3a8f2c883..00bb3c9aa 100644 --- a/crates/core/tests/unit/subscriber_dispatcher_tests.rs +++ b/crates/core/tests/unit/subscriber_dispatcher_tests.rs @@ -3,8 +3,9 @@ use super::native::{ DispatcherLoopState, DispatcherMessage, PendingFlush, PublicationLineage, PublicationPermit, dispatcher_sender, enqueue_dispatch_message, flush_queued_subscribers, flush_subscribers, - register_async_publication, register_pending_publication, sanitize_event_snapshot, - set_sanitizer_runtime_failure_for_test, spawn_background_publication, + prepare_for_fork, register_async_publication, register_pending_publication, + resume_after_fork_parent, sanitize_event_snapshot, set_sanitizer_runtime_failure_for_test, + spawn_background_publication, }; use super::{EventSubscriberFn, publication_context}; use crate::api::registry::RegistryRecord; @@ -12,6 +13,17 @@ use crate::api::runtime::EventSanitizeFn; use crate::api::runtime::scope_stack::current_scope_stack; use std::sync::{Arc, Mutex, mpsc}; +#[test] +fn subscriber_dispatcher_parent_fork_hooks_validate_balanced_calls() { + let _lock = crate::shared_runtime::runtime_owner_test_mutex() + .lock() + .unwrap_or_else(|error| error.into_inner()); + prepare_for_fork(); + assert!(std::panic::catch_unwind(prepare_for_fork).is_err()); + resume_after_fork_parent(); + assert!(std::panic::catch_unwind(resume_after_fork_parent).is_err()); +} + #[test] fn flush_waits_for_active_but_not_later_publication_barriers() { let _lock = crate::shared_runtime::runtime_owner_test_mutex() diff --git a/crates/core/tests/unit/types_tests.rs b/crates/core/tests/unit/types_tests.rs index 720651f60..2dfdace5c 100644 --- a/crates/core/tests/unit/types_tests.rs +++ b/crates/core/tests/unit/types_tests.rs @@ -207,11 +207,14 @@ fn llm_request_serializes_explicit_headers_and_content() { #[test] fn event_accessors_cover_scope_tool_llm_and_mark_variants() { let parent_uuid = Some(Uuid::now_v7()); - let scope_uuid = Uuid::now_v7(); - let tool_uuid = Uuid::now_v7(); - let llm_uuid = Uuid::now_v7(); - let mark_uuid = Uuid::now_v7(); + assert_scope_event_accessors(parent_uuid); + assert_tool_event_accessors(parent_uuid); + assert_llm_event_accessors(parent_uuid); + assert_mark_event_accessors(parent_uuid); +} +fn assert_scope_event_accessors(parent_uuid: Option) { + let scope_uuid = Uuid::now_v7(); let scope_event = Event::Scope(ScopeEvent::new( BaseEvent::builder() .parent_uuid_opt(parent_uuid) @@ -239,7 +242,10 @@ fn event_accessors_cover_scope_tool_llm_and_mark_variants() { assert_eq!(scope_event.scope_type(), Some(ScopeType::Function)); assert_eq!(scope_event.input(), Some(&json!({"task": "classify"}))); assert!(scope_event.timestamp().timestamp() > 0); +} +fn assert_tool_event_accessors(parent_uuid: Option) { + let tool_uuid = Uuid::now_v7(); let tool_event = Event::Scope(ScopeEvent::new( BaseEvent::builder() .parent_uuid_opt(parent_uuid) @@ -266,7 +272,10 @@ fn event_accessors_cover_scope_tool_llm_and_mark_variants() { assert_eq!(tool_event.tool_call_id(), Some("tool-call-1")); assert_eq!(tool_event.scope_type(), Some(ScopeType::Tool)); assert_eq!(tool_event.model_name(), None); +} +fn assert_llm_event_accessors(parent_uuid: Option) { + let llm_uuid = Uuid::now_v7(); let llm_event = Event::Scope(ScopeEvent::new( BaseEvent::builder() .parent_uuid_opt(parent_uuid) @@ -288,7 +297,10 @@ fn event_accessors_cover_scope_tool_llm_and_mark_variants() { assert_eq!(llm_event.model_name(), Some("gpt-test")); assert_eq!(llm_event.scope_type(), Some(ScopeType::Llm)); assert_eq!(llm_event.output(), None); +} +fn assert_mark_event_accessors(parent_uuid: Option) { + let mark_uuid = Uuid::now_v7(); let mark_event = Event::Mark(MarkEvent::new( BaseEvent::builder() .parent_uuid_opt(parent_uuid) @@ -818,6 +830,11 @@ fn atof_event_builders_construct_concrete_events() { #[test] fn base_event_and_flattened_specialized_builders_work() { + assert_base_and_tool_builders(); + assert_llm_and_mark_builders(); +} + +fn assert_base_and_tool_builders() { let base = BaseEvent::builder() .parent_uuid(Uuid::nil()) .name("base-name") @@ -869,7 +886,9 @@ fn base_event_and_flattened_specialized_builders_work() { assert_eq!(tool_end.base.data, None); assert_eq!(tool_end.base.metadata, None); assert_eq!(tool_end.category_profile, None); +} +fn assert_llm_and_mark_builders() { let llm_start = ScopeEvent::new( BaseEvent::builder().name("llm-start").build(), ScopeCategory::Start, diff --git a/crates/ffi/tests/integration/api/coverage_sweeps_tests.rs b/crates/ffi/tests/integration/api/coverage_sweeps_tests.rs index 8708edb9e..e65e06406 100644 --- a/crates/ffi/tests/integration/api/coverage_sweeps_tests.rs +++ b/crates/ffi/tests/integration/api/coverage_sweeps_tests.rs @@ -14,7 +14,7 @@ fn test_ffi_scope_and_event_remaining_error_paths() { unsafe { let stack = fresh_scope_stack(); let mut parent = ptr::null_mut(); - assert_eq!(nemo_relay_get_handle(&mut parent), NemoRelayStatus::Ok); + assert_status!(nemo_relay_get_handle(&mut parent), NemoRelayStatus::Ok); let scope_name = cstring("ffi_child_scope_with_parent"); let data = cstring(r#"{"scope":"child"}"#); @@ -24,7 +24,7 @@ fn test_ffi_scope_and_event_remaining_error_paths() { let invalid = invalid_utf8.as_ptr() as *const c_char; let mut child = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_push_scope( scope_name.as_ptr(), NemoRelayScopeType::Function, @@ -38,7 +38,7 @@ fn test_ffi_scope_and_event_remaining_error_paths() { NemoRelayStatus::Ok ); assert!(take_string(nemo_relay_scope_handle_parent_uuid(child)).is_some()); - assert_eq!( + assert_status!( nemo_relay_push_scope( invalid, NemoRelayScopeType::Function, @@ -51,7 +51,7 @@ fn test_ffi_scope_and_event_remaining_error_paths() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_push_scope( scope_name.as_ptr(), NemoRelayScopeType::Function, @@ -64,7 +64,7 @@ fn test_ffi_scope_and_event_remaining_error_paths() { ), NemoRelayStatus::InvalidJson ); - assert_eq!( + assert_status!( nemo_relay_push_scope( scope_name.as_ptr(), NemoRelayScopeType::Function, @@ -79,7 +79,7 @@ fn test_ffi_scope_and_event_remaining_error_paths() { ); let event_name = cstring("ffi_event_with_parent"); - assert_eq!( + assert_status!( nemo_relay_event( event_name.as_ptr(), parent, @@ -88,11 +88,11 @@ fn test_ffi_scope_and_event_remaining_error_paths() { ), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_event(invalid, parent, ptr::null(), ptr::null()), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_event( event_name.as_ptr(), parent, @@ -101,7 +101,7 @@ fn test_ffi_scope_and_event_remaining_error_paths() { ), NemoRelayStatus::InvalidJson ); - assert_eq!( + assert_status!( nemo_relay_event( event_name.as_ptr(), parent, @@ -111,15 +111,15 @@ fn test_ffi_scope_and_event_remaining_error_paths() { NemoRelayStatus::InvalidJson ); - assert_eq!( + assert_status!( nemo_relay_pop_scope(ptr::null(), ptr::null()), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_pop_scope(child, ptr::null()), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_pop_scope(child, ptr::null()), NemoRelayStatus::NotFound ); @@ -138,7 +138,7 @@ fn test_ffi_tool_and_llm_parent_utf8_and_shape_paths() { unsafe { let stack = fresh_scope_stack(); let mut parent = ptr::null_mut(); - assert_eq!(nemo_relay_get_handle(&mut parent), NemoRelayStatus::Ok); + assert_status!(nemo_relay_get_handle(&mut parent), NemoRelayStatus::Ok); let tool_name = cstring("ffi_tool_call_utf8"); let tool_args = cstring(r#"{"value":1}"#); @@ -151,7 +151,7 @@ fn test_ffi_tool_and_llm_parent_utf8_and_shape_paths() { let invalid = invalid_utf8.as_ptr() as *const c_char; let mut tool_handle = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_tool_call( tool_name.as_ptr(), tool_args.as_ptr(), @@ -165,7 +165,7 @@ fn test_ffi_tool_and_llm_parent_utf8_and_shape_paths() { NemoRelayStatus::Ok ); assert!(take_string(nemo_relay_tool_handle_parent_uuid(tool_handle)).is_some()); - assert_eq!( + assert_status!( nemo_relay_tool_call( invalid, tool_args.as_ptr(), @@ -178,7 +178,7 @@ fn test_ffi_tool_and_llm_parent_utf8_and_shape_paths() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_tool_call( tool_name.as_ptr(), tool_args.as_ptr(), @@ -191,7 +191,7 @@ fn test_ffi_tool_and_llm_parent_utf8_and_shape_paths() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_tool_call_end( tool_handle, tool_result.as_ptr(), @@ -200,7 +200,7 @@ fn test_ffi_tool_and_llm_parent_utf8_and_shape_paths() { ), NemoRelayStatus::InvalidJson ); - assert_eq!( + assert_status!( nemo_relay_tool_call_end( tool_handle, tool_result.as_ptr(), @@ -221,7 +221,7 @@ fn test_ffi_tool_and_llm_parent_utf8_and_shape_paths() { let model_name = cstring("ffi-model-override"); let mut llm_handle = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_llm_call( llm_name.as_ptr(), request.as_ptr(), @@ -235,7 +235,7 @@ fn test_ffi_tool_and_llm_parent_utf8_and_shape_paths() { NemoRelayStatus::Ok ); assert!(take_string(nemo_relay_llm_handle_parent_uuid(llm_handle)).is_some()); - assert_eq!( + assert_status!( nemo_relay_llm_call( invalid, request.as_ptr(), @@ -248,7 +248,7 @@ fn test_ffi_tool_and_llm_parent_utf8_and_shape_paths() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_llm_call( llm_name.as_ptr(), request.as_ptr(), @@ -261,7 +261,7 @@ fn test_ffi_tool_and_llm_parent_utf8_and_shape_paths() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_llm_call_end( llm_handle, response.as_ptr(), @@ -270,7 +270,7 @@ fn test_ffi_tool_and_llm_parent_utf8_and_shape_paths() { ), NemoRelayStatus::InvalidJson ); - assert_eq!( + assert_status!( nemo_relay_llm_call_end( llm_handle, response.as_ptr(), @@ -281,7 +281,7 @@ fn test_ffi_tool_and_llm_parent_utf8_and_shape_paths() { ); let mut out = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_llm_call_execute( llm_name.as_ptr(), invalid_shape.as_ptr(), @@ -309,7 +309,7 @@ fn test_ffi_tool_and_llm_parent_utf8_and_shape_paths() { ); let mut stream = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_llm_stream_call_execute( llm_name.as_ptr(), invalid_shape.as_ptr(), @@ -354,7 +354,7 @@ fn test_ffi_global_registry_invalid_utf8_name_sweep() { let invalid = invalid_utf8.as_ptr() as *const c_char; unsafe { - assert_eq!( + assert_status!( nemo_relay_register_tool_sanitize_request_guardrail( invalid, 1, @@ -364,11 +364,11 @@ fn test_ffi_global_registry_invalid_utf8_name_sweep() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_deregister_tool_sanitize_request_guardrail(invalid), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_register_tool_sanitize_response_guardrail( invalid, 1, @@ -378,11 +378,11 @@ fn test_ffi_global_registry_invalid_utf8_name_sweep() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_deregister_tool_sanitize_response_guardrail(invalid), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_register_tool_conditional_execution_guardrail( invalid, 1, @@ -392,11 +392,11 @@ fn test_ffi_global_registry_invalid_utf8_name_sweep() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_deregister_tool_conditional_execution_guardrail(invalid), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_register_tool_request_intercept( invalid, 1, @@ -407,11 +407,11 @@ fn test_ffi_global_registry_invalid_utf8_name_sweep() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_deregister_tool_request_intercept(invalid), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_register_tool_execution_intercept( invalid, 1, @@ -421,12 +421,12 @@ fn test_ffi_global_registry_invalid_utf8_name_sweep() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_deregister_tool_execution_intercept(invalid), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_register_llm_sanitize_request_guardrail( invalid, 1, @@ -436,11 +436,11 @@ fn test_ffi_global_registry_invalid_utf8_name_sweep() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_deregister_llm_sanitize_request_guardrail(invalid), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_register_llm_sanitize_response_guardrail( invalid, 1, @@ -450,11 +450,11 @@ fn test_ffi_global_registry_invalid_utf8_name_sweep() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_deregister_llm_sanitize_response_guardrail(invalid), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_register_llm_conditional_execution_guardrail( invalid, 1, @@ -464,11 +464,11 @@ fn test_ffi_global_registry_invalid_utf8_name_sweep() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_deregister_llm_conditional_execution_guardrail(invalid), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_register_llm_request_intercept( invalid, 1, @@ -479,11 +479,11 @@ fn test_ffi_global_registry_invalid_utf8_name_sweep() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_deregister_llm_request_intercept(invalid), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_register_llm_execution_intercept( invalid, 1, @@ -493,11 +493,11 @@ fn test_ffi_global_registry_invalid_utf8_name_sweep() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_deregister_llm_execution_intercept(invalid), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_register_llm_stream_execution_intercept( invalid, 1, @@ -507,15 +507,15 @@ fn test_ffi_global_registry_invalid_utf8_name_sweep() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_deregister_llm_stream_execution_intercept(invalid), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_register_subscriber(invalid, subscriber_cb, ptr::null_mut(), None), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_deregister_subscriber(invalid), NemoRelayStatus::InvalidUtf8 ); @@ -534,7 +534,7 @@ fn test_ffi_scope_registry_invalid_utf8_scope_and_name_sweeps() { let stack = fresh_scope_stack(); let scope_name = cstring("scope-registry-invalid"); let mut scope = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_push_scope( scope_name.as_ptr(), NemoRelayScopeType::Function, @@ -551,7 +551,7 @@ fn test_ffi_scope_registry_invalid_utf8_scope_and_name_sweeps() { let invalid_name = invalid_utf8.as_ptr() as *const c_char; let valid_name = cstring("scope-registry-valid-name"); - assert_eq!( + assert_status!( nemo_relay_scope_register_tool_sanitize_request_guardrail( invalid_scope, valid_name.as_ptr(), @@ -562,14 +562,14 @@ fn test_ffi_scope_registry_invalid_utf8_scope_and_name_sweeps() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_tool_sanitize_request_guardrail( invalid_scope, valid_name.as_ptr(), ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_register_tool_sanitize_request_guardrail( scope_uuid.as_ptr(), invalid_name, @@ -580,7 +580,7 @@ fn test_ffi_scope_registry_invalid_utf8_scope_and_name_sweeps() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_tool_sanitize_request_guardrail( scope_uuid.as_ptr(), invalid_name @@ -588,7 +588,7 @@ fn test_ffi_scope_registry_invalid_utf8_scope_and_name_sweeps() { NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_register_tool_execution_intercept( invalid_scope, valid_name.as_ptr(), @@ -599,14 +599,14 @@ fn test_ffi_scope_registry_invalid_utf8_scope_and_name_sweeps() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_tool_execution_intercept( invalid_scope, valid_name.as_ptr() ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_register_tool_execution_intercept( scope_uuid.as_ptr(), invalid_name, @@ -617,12 +617,12 @@ fn test_ffi_scope_registry_invalid_utf8_scope_and_name_sweeps() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_tool_execution_intercept(scope_uuid.as_ptr(), invalid_name), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_register_llm_sanitize_request_guardrail( invalid_scope, valid_name.as_ptr(), @@ -633,14 +633,14 @@ fn test_ffi_scope_registry_invalid_utf8_scope_and_name_sweeps() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_llm_sanitize_request_guardrail( invalid_scope, valid_name.as_ptr(), ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_register_llm_sanitize_request_guardrail( scope_uuid.as_ptr(), invalid_name, @@ -651,7 +651,7 @@ fn test_ffi_scope_registry_invalid_utf8_scope_and_name_sweeps() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_llm_sanitize_request_guardrail( scope_uuid.as_ptr(), invalid_name @@ -659,7 +659,7 @@ fn test_ffi_scope_registry_invalid_utf8_scope_and_name_sweeps() { NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_register_llm_execution_intercept( invalid_scope, valid_name.as_ptr(), @@ -670,11 +670,11 @@ fn test_ffi_scope_registry_invalid_utf8_scope_and_name_sweeps() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_llm_execution_intercept(invalid_scope, valid_name.as_ptr()), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_register_llm_execution_intercept( scope_uuid.as_ptr(), invalid_name, @@ -685,12 +685,12 @@ fn test_ffi_scope_registry_invalid_utf8_scope_and_name_sweeps() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_llm_execution_intercept(scope_uuid.as_ptr(), invalid_name), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_register_subscriber( invalid_scope, valid_name.as_ptr(), @@ -700,11 +700,11 @@ fn test_ffi_scope_registry_invalid_utf8_scope_and_name_sweeps() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_subscriber(invalid_scope, valid_name.as_ptr()), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_register_subscriber( scope_uuid.as_ptr(), invalid_name, @@ -714,12 +714,12 @@ fn test_ffi_scope_registry_invalid_utf8_scope_and_name_sweeps() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_subscriber(scope_uuid.as_ptr(), invalid_name), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_pop_scope(scope, ptr::null()), NemoRelayStatus::Ok ); @@ -753,55 +753,55 @@ fn test_ffi_adaptive_and_observability_entry_points_from_integration_binary() { let invalid_json = cstring("{"); let mut out_json = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_adaptive_validate_config(config.as_ptr(), &mut out_json), NemoRelayStatus::Ok ); assert_eq!(returned_json(out_json)["diagnostics"], json!([])); - assert_eq!( + assert_status!( nemo_relay_adaptive_validate_config(config.as_ptr(), ptr::null_mut()), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_adaptive_validate_config(invalid_json.as_ptr(), &mut out_json), NemoRelayStatus::InvalidJson ); let mut runtime = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_adaptive_runtime_create(config.as_ptr(), ptr::null_mut()), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_adaptive_runtime_create(invalid_json.as_ptr(), &mut runtime), NemoRelayStatus::InvalidJson ); - assert_eq!( + assert_status!( nemo_relay_adaptive_runtime_create(config.as_ptr(), &mut runtime), NemoRelayStatus::Ok ); assert!(!runtime.is_null()); - assert_eq!( + assert_status!( nemo_relay_adaptive_runtime_register(ptr::null_mut()), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_adaptive_runtime_register(runtime), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_adaptive_runtime_wait_for_idle(ptr::null_mut()), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_adaptive_runtime_wait_for_idle(runtime), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_adaptive_runtime_report_json(runtime, ptr::null_mut()), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_adaptive_runtime_report_json(runtime, &mut out_json), NemoRelayStatus::Ok ); @@ -810,7 +810,7 @@ fn test_ffi_adaptive_and_observability_entry_points_from_integration_binary() { let stack = fresh_scope_stack(); let scope_name = cstring("ffi_adaptive_integration_scope"); let mut scope = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_push_scope( scope_name.as_ptr(), NemoRelayScopeType::Agent, @@ -823,11 +823,11 @@ fn test_ffi_adaptive_and_observability_entry_points_from_integration_binary() { ), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_adaptive_runtime_bind_scope(runtime, scope), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_adaptive_runtime_bind_scope(runtime, ptr::null()), NemoRelayStatus::NullPointer ); @@ -849,7 +849,7 @@ fn test_ffi_adaptive_and_observability_entry_points_from_integration_binary() { }) .to_string(), ); - assert_eq!( + assert_status!( nemo_relay_adaptive_runtime_build_cache_request_facts( runtime, cache_options.as_ptr(), @@ -859,7 +859,7 @@ fn test_ffi_adaptive_and_observability_entry_points_from_integration_binary() { ); let facts = returned_json(out_json); assert_eq!(facts["provider"], json!("openai")); - assert_eq!( + assert_status!( nemo_relay_adaptive_runtime_build_cache_request_facts( runtime, cache_options.as_ptr(), @@ -922,7 +922,7 @@ fn test_ffi_adaptive_and_observability_entry_points_from_integration_binary() { }) .to_string(), ); - assert_eq!( + assert_status!( nemo_relay_adaptive_build_cache_telemetry_event( telemetry_options.as_ptr(), &mut out_json, @@ -932,7 +932,7 @@ fn test_ffi_adaptive_and_observability_entry_points_from_integration_binary() { let event = returned_json(out_json); assert_eq!(event["cache_read_tokens"], json!(10)); assert_eq!(event["hit_rate"], json!(10.0 / 60.0)); - assert_eq!( + assert_status!( nemo_relay_adaptive_build_cache_telemetry_event( telemetry_options.as_ptr(), ptr::null_mut(), @@ -951,7 +951,7 @@ fn test_ffi_adaptive_and_observability_entry_points_from_integration_binary() { }) .to_string(), ); - assert_eq!( + assert_status!( nemo_relay_adaptive_build_cache_telemetry_event( no_usage_options.as_ptr(), &mut out_json, @@ -1034,11 +1034,11 @@ fn test_ffi_adaptive_and_observability_entry_points_from_integration_binary() { expected ); } - assert_eq!( + assert_status!( nemo_relay_adaptive_build_cache_telemetry_event(invalid_json.as_ptr(), &mut out_json,), NemoRelayStatus::InvalidJson ); - assert_eq!( + assert_status!( nemo_relay_adaptive_runtime_build_cache_request_facts( ptr::null_mut(), cache_options.as_ptr(), @@ -1046,32 +1046,32 @@ fn test_ffi_adaptive_and_observability_entry_points_from_integration_binary() { ), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_adaptive_set_latency_sensitivity(0), NemoRelayStatus::InvalidArg ); - assert_eq!( + assert_status!( nemo_relay_adaptive_runtime_deregister(runtime), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_adaptive_runtime_deregister(ptr::null_mut()), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_adaptive_runtime_shutdown(runtime), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_adaptive_runtime_shutdown(ptr::null_mut()), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_adaptive_runtime_report_json(runtime, &mut out_json), NemoRelayStatus::InvalidArg ); - assert_eq!( + assert_status!( nemo_relay_pop_scope(scope, ptr::null()), NemoRelayStatus::Ok ); @@ -1079,21 +1079,21 @@ fn test_ffi_adaptive_and_observability_entry_points_from_integration_binary() { nemo_relay_scope_stack_free(stack); types::nemo_relay_adaptive_runtime_free(runtime); - assert_eq!( + assert_status!( nemo_relay_observability_default_config_json(&mut out_json), NemoRelayStatus::Ok ); assert_eq!(returned_json(out_json)["version"], json!(3)); - assert_eq!( + assert_status!( nemo_relay_observability_default_config_json(ptr::null_mut()), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_observability_component_spec_json(ptr::null(), true, &mut out_json), NemoRelayStatus::Ok ); assert_eq!(returned_json(out_json)["kind"], json!("observability")); - assert_eq!( + assert_status!( nemo_relay_observability_component_spec_json( invalid_json.as_ptr(), true, @@ -1111,7 +1111,7 @@ fn test_ffi_adaptive_and_observability_entry_points_from_integration_binary() { let happy_filename = cstring("happy-events.jsonl"); let happy_name = cstring(&unique_name("ffi_atof_happy")); let mut happy_atof = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_atof_exporter_create( happy_dir.as_ptr(), append.as_ptr(), @@ -1122,7 +1122,7 @@ fn test_ffi_adaptive_and_observability_entry_points_from_integration_binary() { ); assert!(!happy_atof.is_null()); let mut path_ptr = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_atof_exporter_path(happy_atof, &mut path_ptr), NemoRelayStatus::Ok ); @@ -1131,19 +1131,19 @@ fn test_ffi_adaptive_and_observability_entry_points_from_integration_binary() { path.ends_with("happy-events.jsonl"), "unexpected ATOF exporter path: {path}" ); - assert_eq!( + assert_status!( nemo_relay_atof_exporter_register(happy_atof, happy_name.as_ptr()), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_atof_exporter_force_flush(happy_atof), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_atof_exporter_shutdown(happy_atof), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_atof_exporter_deregister(happy_name.as_ptr()), NemoRelayStatus::Ok ); @@ -1152,7 +1152,7 @@ fn test_ffi_adaptive_and_observability_entry_points_from_integration_binary() { let invalid_utf8 = [0xffu8, 0]; let invalid = invalid_utf8.as_ptr() as *const c_char; let mut atof = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_atof_exporter_create( ptr::null(), append.as_ptr(), @@ -1161,7 +1161,7 @@ fn test_ffi_adaptive_and_observability_entry_points_from_integration_binary() { ), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_atof_exporter_create( ptr::null(), bad_mode.as_ptr(), @@ -1170,31 +1170,31 @@ fn test_ffi_adaptive_and_observability_entry_points_from_integration_binary() { ), NemoRelayStatus::InvalidArg ); - assert_eq!( + assert_status!( nemo_relay_atof_exporter_create(invalid, append.as_ptr(), filename.as_ptr(), &mut atof), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_atof_exporter_create(ptr::null(), invalid, filename.as_ptr(), &mut atof), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_atof_exporter_create(ptr::null(), append.as_ptr(), invalid, &mut atof), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_atof_exporter_register(ptr::null(), filename.as_ptr()), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_atof_exporter_force_flush(ptr::null()), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_atof_exporter_shutdown(ptr::null()), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_atof_exporter_path(ptr::null(), &mut out_json), NemoRelayStatus::NullPointer ); diff --git a/crates/ffi/tests/integration/api_tests.rs b/crates/ffi/tests/integration/api_tests.rs index 140f089ce..4637d6778 100644 --- a/crates/ffi/tests/integration/api_tests.rs +++ b/crates/ffi/tests/integration/api_tests.rs @@ -45,6 +45,17 @@ static COLLECTED_CHUNKS: OnceLock>> = OnceLock::new(); static FINALIZER_CALLS: OnceLock> = OnceLock::new(); static PLUGIN_FREES: OnceLock> = OnceLock::new(); +#[track_caller] +fn assert_native_status(actual: NemoRelayStatus, expected: NemoRelayStatus) { + assert_eq!(actual, expected); +} + +macro_rules! assert_status { + ($actual:expr, $expected:expr $(,)?) => { + assert_native_status($actual, $expected) + }; +} + fn event_log() -> &'static Mutex> { EVENT_LOG.get_or_init(|| Mutex::new(Vec::new())) } diff --git a/crates/ffi/tests/integration/plugin_activation_tests.rs b/crates/ffi/tests/integration/plugin_activation_tests.rs index bab91f2d9..b4c412962 100644 --- a/crates/ffi/tests/integration/plugin_activation_tests.rs +++ b/crates/ffi/tests/integration/plugin_activation_tests.rs @@ -98,26 +98,7 @@ fn run_discovered_config_activation_test() { // below proves this attempt did not retain the process-wide host claim. let config = cstring(r#"{"version":1,"components":[]}"#); let empty_specs = cstring("[]"); - let mut empty_activation = ptr::null_mut(); - let mut empty_report = ptr::null_mut(); - assert_eq!( - unsafe { - api::nemo_relay_initialize_with_dynamic_plugins( - config.as_ptr(), - empty_specs.as_ptr(), - &mut empty_activation, - &mut empty_report, - ) - }, - NemoRelayStatus::InvalidArg - ); - assert!(empty_activation.is_null()); - assert!(empty_report.is_null()); - assert!( - unsafe { read_last_error() } - .unwrap_or_default() - .contains("at least one dynamic plugin") - ); + assert_empty_dynamic_specs_rejected(&config, &empty_specs); std::fs::write( &plugins_toml, @@ -158,22 +139,7 @@ source = "project-file" "config": {} }])); - // The file-only component and its config must survive the merge. - assert_eq!(report["diagnostics"], json!([])); - assert_eq!(DISCOVERED_STATIC_REGISTRATIONS.load(Ordering::SeqCst), 1); - assert_eq!( - DISCOVERED_STATIC_CONFIG.lock().unwrap().as_ref(), - Some(&json!({"source": "project-file"})) - ); - assert!(plugin_kinds().iter().any(|kind| kind == "fixture_native")); - - // Mutating the file after startup has no effect: discovery is one-shot. - std::fs::write(&plugins_toml, "invalid = [").expect("mutate plugin config after startup"); - let intercepted = tool_request_intercepts("ffi-layered-tool", json!({"input": true})); - assert_eq!(intercepted["file_static"], true); - assert_eq!(intercepted["static_saw_dynamic"], false); - assert_eq!(intercepted["native_plugin"], true); - assert_eq!(DISCOVERED_STATIC_CALLBACKS.load(Ordering::SeqCst), 1); + assert_discovered_activation(&report, &plugins_toml); unsafe { assert_eq!( @@ -194,6 +160,48 @@ source = "project-file" ); } +fn assert_empty_dynamic_specs_rejected(config: &CString, empty_specs: &CString) { + let mut empty_activation = ptr::null_mut(); + let mut empty_report = ptr::null_mut(); + assert_eq!( + unsafe { + api::nemo_relay_initialize_with_dynamic_plugins( + config.as_ptr(), + empty_specs.as_ptr(), + &mut empty_activation, + &mut empty_report, + ) + }, + NemoRelayStatus::InvalidArg + ); + assert!(empty_activation.is_null()); + assert!(empty_report.is_null()); + assert!( + unsafe { read_last_error() } + .unwrap_or_default() + .contains("at least one dynamic plugin") + ); +} + +fn assert_discovered_activation(report: &Json, plugins_toml: &Path) { + // The file-only component and its config must survive the merge. + assert_eq!(report["diagnostics"], json!([])); + assert_eq!(DISCOVERED_STATIC_REGISTRATIONS.load(Ordering::SeqCst), 1); + assert_eq!( + DISCOVERED_STATIC_CONFIG.lock().unwrap().as_ref(), + Some(&json!({"source": "project-file"})) + ); + assert!(plugin_kinds().iter().any(|kind| kind == "fixture_native")); + + // Mutating the file after startup has no effect: discovery is one-shot. + std::fs::write(plugins_toml, "invalid = [").expect("mutate plugin config after startup"); + let intercepted = tool_request_intercepts("ffi-layered-tool", json!({"input": true})); + assert_eq!(intercepted["file_static"], true); + assert_eq!(intercepted["static_saw_dynamic"], false); + assert_eq!(intercepted["native_plugin"], true); + assert_eq!(DISCOVERED_STATIC_CALLBACKS.load(Ordering::SeqCst), 1); +} + unsafe extern "C" fn discovered_static_register( _user_data: *mut libc::c_void, plugin_config_json: *const c_char, diff --git a/crates/ffi/tests/unit/api/core_tests.rs b/crates/ffi/tests/unit/api/core_tests.rs index ad1a769a3..d55a0294f 100644 --- a/crates/ffi/tests/unit/api/core_tests.rs +++ b/crates/ffi/tests/unit/api/core_tests.rs @@ -14,7 +14,7 @@ fn test_ffi_llm_request_intercept_outcome_json_allocation_and_validation() { let marks = cstring(r#"[{"name":"first"},{"name":"second","data":{"order":2}}]"#); let mut outcome_json = ptr::null_mut(); - assert_eq!( + assert_status!( unsafe { api::nemo_relay_llm_request_intercept_outcome_json_new( request, @@ -33,7 +33,7 @@ fn test_ffi_llm_request_intercept_outcome_json_allocation_and_validation() { assert_eq!(outcome["optimization_contributions"], json!([])); let malformed_contributions = cstring(r#"[{"producer":1}]"#); - assert_eq!( + assert_status!( unsafe { api::nemo_relay_llm_request_intercept_outcome_json_new_v2( request, @@ -51,7 +51,7 @@ fn test_ffi_llm_request_intercept_outcome_json_allocation_and_validation() { "[{}]", include_str!("../../../../types/tests/fixtures/llm_optimization_contribution_v1.json") )); - assert_eq!( + assert_status!( unsafe { api::nemo_relay_llm_request_intercept_outcome_json_new_v2( request, @@ -71,7 +71,7 @@ fn test_ffi_llm_request_intercept_outcome_json_allocation_and_validation() { assert_eq!(outcome["optimization_contributions"], json!([expected])); let malformed_marks = cstring(r#"{"name":"not-an-array"}"#); - assert_eq!( + assert_status!( unsafe { api::nemo_relay_llm_request_intercept_outcome_json_new( request, @@ -84,7 +84,7 @@ fn test_ffi_llm_request_intercept_outcome_json_allocation_and_validation() { ); assert!(outcome_json.is_null()); outcome_json = std::ptr::dangling_mut(); - assert_eq!( + assert_status!( unsafe { api::nemo_relay_llm_request_intercept_outcome_json_new( ptr::null(), @@ -134,7 +134,7 @@ fn test_ffi_plugin_config_validate_initialize_and_clear() { ); let mut report_json = ptr::null_mut(); - assert_eq!( + assert_status!( unsafe { nemo_relay_validate_plugin_config(config.as_ptr(), &mut report_json) }, NemoRelayStatus::Ok ); @@ -142,7 +142,7 @@ fn test_ffi_plugin_config_validate_initialize_and_clear() { assert_eq!(report["diagnostics"], json!([])); let mut kinds_json = ptr::null_mut(); - assert_eq!( + assert_status!( unsafe { nemo_relay_list_plugin_kinds_json(&mut kinds_json) }, NemoRelayStatus::Ok ); @@ -159,7 +159,7 @@ fn test_ffi_plugin_config_validate_initialize_and_clear() { ); let mut configured_json = ptr::null_mut(); - assert_eq!( + assert_status!( unsafe { nemo_relay_initialize_plugins(config.as_ptr(), &mut configured_json) }, NemoRelayStatus::Ok ); @@ -167,17 +167,17 @@ fn test_ffi_plugin_config_validate_initialize_and_clear() { assert_eq!(configured_report["diagnostics"], json!([])); let mut active_json = ptr::null_mut(); - assert_eq!( + assert_status!( unsafe { nemo_relay_active_plugin_report_json(&mut active_json) }, NemoRelayStatus::Ok ); let active_report = unsafe { returned_json(active_json) }; assert_eq!(active_report["diagnostics"], json!([])); - assert_eq!(nemo_relay_clear_plugin_configuration(), NemoRelayStatus::Ok); + assert_status!(nemo_relay_clear_plugin_configuration(), NemoRelayStatus::Ok); let mut cleared_json = ptr::null_mut(); - assert_eq!( + assert_status!( unsafe { nemo_relay_active_plugin_report_json(&mut cleared_json) }, NemoRelayStatus::Ok ); @@ -234,13 +234,13 @@ fn test_ffi_observability_plugin_file_sinks() { "observability" ); let mut default_config_json = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_observability_default_config_json(&mut default_config_json), NemoRelayStatus::Ok ); assert_eq!(returned_json(default_config_json)["version"], json!(3)); let mut component_json = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_observability_component_spec_json(ptr::null(), true, &mut component_json), NemoRelayStatus::Ok ); @@ -249,14 +249,14 @@ fn test_ffi_observability_plugin_file_sinks() { assert_eq!(component["enabled"], true); let mut report_json = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_validate_plugin_config(config.as_ptr(), &mut report_json), NemoRelayStatus::Ok ); assert_eq!(returned_json(report_json)["diagnostics"], json!([])); let mut initialized_json = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_initialize_plugins(config.as_ptr(), &mut initialized_json), NemoRelayStatus::Ok ); @@ -266,7 +266,7 @@ fn test_ffi_observability_plugin_file_sinks() { let scope_name = cstring("ffi-observability-agent"); let input = cstring(r#"{"agent":true}"#); let mut scope = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_push_scope( scope_name.as_ptr(), NemoRelayScopeType::Agent, @@ -283,17 +283,17 @@ fn test_ffi_observability_plugin_file_sinks() { let mark_name = cstring("ffi-observability-mark"); let mark_data = cstring(r#"{"step":1}"#); - assert_eq!( + assert_status!( nemo_relay_event(mark_name.as_ptr(), scope, mark_data.as_ptr(), ptr::null()), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_pop_scope(scope, ptr::null()), NemoRelayStatus::Ok ); nemo_relay_scope_handle_free(scope); nemo_relay_scope_stack_free(stack); - assert_eq!(nemo_relay_clear_plugin_configuration(), NemoRelayStatus::Ok); + assert_status!(nemo_relay_clear_plugin_configuration(), NemoRelayStatus::Ok); let jsonl = std::fs::read_to_string(dir.join("events.jsonl")).unwrap(); assert_eq!(jsonl.trim().lines().count(), 3); @@ -344,7 +344,7 @@ fn test_ffi_observability_plugin_atif_splits_multiple_top_level_agents() { unsafe { let mut initialized_json = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_initialize_plugins(config.as_ptr(), &mut initialized_json), NemoRelayStatus::Ok ); @@ -355,7 +355,7 @@ fn test_ffi_observability_plugin_atif_splits_multiple_top_level_agents() { let first_name = cstring("ffi-first-agent"); let first_input = cstring(r#"{"agent":"first"}"#); let mut first = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_push_scope( first_name.as_ptr(), NemoRelayScopeType::Agent, @@ -372,7 +372,7 @@ fn test_ffi_observability_plugin_atif_splits_multiple_top_level_agents() { let first_mark = cstring("ffi-first-mark"); let first_mark_data = cstring(r#"{"agent":"first"}"#); - assert_eq!( + assert_status!( nemo_relay_event( first_mark.as_ptr(), first, @@ -385,7 +385,7 @@ fn test_ffi_observability_plugin_atif_splits_multiple_top_level_agents() { let nested_name = cstring("ffi-nested-agent"); let nested_input = cstring(r#"{"agent":"nested"}"#); let mut nested = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_push_scope( nested_name.as_ptr(), NemoRelayScopeType::Agent, @@ -400,7 +400,7 @@ fn test_ffi_observability_plugin_atif_splits_multiple_top_level_agents() { ); let nested_mark = cstring("ffi-nested-mark"); let nested_mark_data = cstring(r#"{"agent":"nested"}"#); - assert_eq!( + assert_status!( nemo_relay_event( nested_mark.as_ptr(), nested, @@ -409,12 +409,12 @@ fn test_ffi_observability_plugin_atif_splits_multiple_top_level_agents() { ), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_pop_scope(nested, ptr::null()), NemoRelayStatus::Ok ); nemo_relay_scope_handle_free(nested); - assert_eq!( + assert_status!( nemo_relay_pop_scope(first, ptr::null()), NemoRelayStatus::Ok ); @@ -423,7 +423,7 @@ fn test_ffi_observability_plugin_atif_splits_multiple_top_level_agents() { let second_name = cstring("ffi-second-agent"); let second_input = cstring(r#"{"agent":"second"}"#); let mut second = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_push_scope( second_name.as_ptr(), NemoRelayScopeType::Agent, @@ -439,7 +439,7 @@ fn test_ffi_observability_plugin_atif_splits_multiple_top_level_agents() { let second_uuid = take_string(nemo_relay_scope_handle_uuid(second)).unwrap(); let second_mark = cstring("ffi-second-mark"); let second_mark_data = cstring(r#"{"agent":"second"}"#); - assert_eq!( + assert_status!( nemo_relay_event( second_mark.as_ptr(), second, @@ -448,13 +448,13 @@ fn test_ffi_observability_plugin_atif_splits_multiple_top_level_agents() { ), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_pop_scope(second, ptr::null()), NemoRelayStatus::Ok ); nemo_relay_scope_handle_free(second); nemo_relay_scope_stack_free(stack); - assert_eq!(nemo_relay_clear_plugin_configuration(), NemoRelayStatus::Ok); + assert_status!(nemo_relay_clear_plugin_configuration(), NemoRelayStatus::Ok); let files = std::fs::read_dir(&dir) .unwrap() @@ -498,7 +498,7 @@ fn test_ffi_plugin_top_level_null_and_invalid_paths() { let invalid_shape = cstring(r#"{"version":"bad","components":"nope"}"#); unsafe { - assert_eq!( + assert_status!( nemo_relay_validate_plugin_config(valid_config.as_ptr(), ptr::null_mut()), NemoRelayStatus::NullPointer ); @@ -509,11 +509,11 @@ fn test_ffi_plugin_top_level_null_and_invalid_paths() { ); let mut out_json = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_validate_plugin_config(invalid_json.as_ptr(), &mut out_json), NemoRelayStatus::InvalidJson ); - assert_eq!( + assert_status!( nemo_relay_validate_plugin_config(invalid_shape.as_ptr(), &mut out_json), NemoRelayStatus::InvalidJson ); @@ -523,28 +523,28 @@ fn test_ffi_plugin_top_level_null_and_invalid_paths() { .contains("invalid type") ); - assert_eq!( + assert_status!( nemo_relay_initialize_plugins(valid_config.as_ptr(), ptr::null_mut()), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_initialize_plugins(invalid_json.as_ptr(), &mut out_json), NemoRelayStatus::InvalidJson ); - assert_eq!( + assert_status!( nemo_relay_initialize_plugins(invalid_shape.as_ptr(), &mut out_json), NemoRelayStatus::InvalidJson ); - assert_eq!( + assert_status!( nemo_relay_active_plugin_report_json(ptr::null_mut()), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_list_plugin_kinds_json(ptr::null_mut()), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_register_plugin( ptr::null(), None, @@ -554,7 +554,7 @@ fn test_ffi_plugin_top_level_null_and_invalid_paths() { ), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_deregister_plugin(ptr::null()), NemoRelayStatus::NullPointer ); @@ -567,7 +567,7 @@ fn test_ffi_error_paths_and_scope_stack() { reset_globals(); unsafe { - assert_eq!( + assert_status!( nemo_relay_get_handle(ptr::null_mut()), NemoRelayStatus::NullPointer ); @@ -576,7 +576,7 @@ fn test_ffi_error_paths_and_scope_stack() { let name = cstring("ffi_invalid_scope"); let invalid_json = cstring("{"); let mut handle = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_push_scope( name.as_ptr(), NemoRelayScopeType::Agent, @@ -594,7 +594,7 @@ fn test_ffi_error_paths_and_scope_stack() { assert!(nemo_relay_scope_stack_active()); let mut root = ptr::null_mut(); - assert_eq!(nemo_relay_get_handle(&mut root), NemoRelayStatus::Ok); + assert_status!(nemo_relay_get_handle(&mut root), NemoRelayStatus::Ok); let root_uuid = take_string(nemo_relay_scope_handle_uuid(root)).unwrap(); assert!(!root_uuid.is_empty()); assert_eq!( @@ -608,7 +608,7 @@ fn test_ffi_error_paths_and_scope_stack() { let scope_data = cstring(r#"{"scope":true}"#); let scope_metadata = cstring(r#"{"meta":"ok"}"#); let mut scope = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_push_scope( scope_name.as_ptr(), NemoRelayScopeType::Function, @@ -645,7 +645,7 @@ fn test_ffi_error_paths_and_scope_stack() { .unwrap(), json!({"meta": "ok"}) ); - assert_eq!( + assert_status!( nemo_relay_pop_scope(scope, ptr::null()), NemoRelayStatus::Ok ); @@ -664,7 +664,7 @@ fn test_ffi_pop_scope_merges_scope_metadata() { let stack = fresh_scope_stack(); let subscriber_name = unique_name("ffi_scope_end_metadata_subscriber"); let subscriber_name_c = cstring(&subscriber_name); - assert_eq!( + assert_status!( nemo_relay_register_subscriber( subscriber_name_c.as_ptr(), subscriber_cb, @@ -679,7 +679,7 @@ fn test_ffi_pop_scope_merges_scope_metadata() { let end_metadata = cstring(r#"{"c":3.5,"d":4}"#); let invalid_end_metadata = cstring("{"); let mut scope = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_push_scope( scope_name.as_ptr(), NemoRelayScopeType::Function, @@ -692,7 +692,7 @@ fn test_ffi_pop_scope_merges_scope_metadata() { ), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( api::nemo_relay_pop_scope( scope, ptr::null(), @@ -701,11 +701,11 @@ fn test_ffi_pop_scope_merges_scope_metadata() { ), NemoRelayStatus::InvalidJson ); - assert_eq!( + assert_status!( api::nemo_relay_pop_scope(scope, ptr::null(), end_metadata.as_ptr(), ptr::null(),), NemoRelayStatus::Ok ); - assert_eq!(nemo_relay_flush_subscribers(), NemoRelayStatus::Ok); + assert_status!(nemo_relay_flush_subscribers(), NemoRelayStatus::Ok); let events = lock_unpoisoned(event_log()).clone(); let end_event = events @@ -721,7 +721,7 @@ fn test_ffi_pop_scope_merges_scope_metadata() { json!({"a": 1, "b": 2, "c": 3.5, "d": 4}) ); - assert_eq!( + assert_status!( nemo_relay_deregister_subscriber(subscriber_name_c.as_ptr()), NemoRelayStatus::Ok ); @@ -746,7 +746,7 @@ fn test_ffi_tool_lifecycle_execute_and_helpers() { let stack = fresh_scope_stack(); let subscriber_name = unique_name("ffi_subscriber"); let subscriber_name_c = cstring(&subscriber_name); - assert_eq!( + assert_status!( nemo_relay_register_subscriber( subscriber_name_c.as_ptr(), subscriber_cb, @@ -758,7 +758,7 @@ fn test_ffi_tool_lifecycle_execute_and_helpers() { let intercept_name = unique_name("ffi_tool_intercept"); let intercept_name_c = cstring(&intercept_name); - assert_eq!( + assert_status!( nemo_relay_register_tool_request_intercept( intercept_name_c.as_ptr(), 1, @@ -772,7 +772,7 @@ fn test_ffi_tool_lifecycle_execute_and_helpers() { let conditional_name = unique_name("ffi_tool_conditional"); let conditional_name_c = cstring(&conditional_name); - assert_eq!( + assert_status!( nemo_relay_register_tool_conditional_execution_guardrail( conditional_name_c.as_ptr(), 1, @@ -786,7 +786,7 @@ fn test_ffi_tool_lifecycle_execute_and_helpers() { let tool_name = cstring("ffi_tool"); let args = cstring(r#"{"value": 1}"#); let mut intercepted_out = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_tool_request_intercepts( tool_name.as_ptr(), args.as_ptr(), @@ -797,7 +797,7 @@ fn test_ffi_tool_lifecycle_execute_and_helpers() { let intercepted_json = returned_json(intercepted_out); assert_eq!(intercepted_json["intercepted"], json!(true)); - assert_eq!( + assert_status!( nemo_relay_tool_conditional_execution(tool_name.as_ptr(), args.as_ptr()), NemoRelayStatus::Ok ); @@ -805,7 +805,7 @@ fn test_ffi_tool_lifecycle_execute_and_helpers() { let tool_call_id = cstring("call_ffi_123"); let metadata = cstring(r#"{"source":"ffi-test"}"#); let mut handle: *mut FfiToolHandle = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_tool_call( tool_name.as_ptr(), args.as_ptr(), @@ -827,14 +827,14 @@ fn test_ffi_tool_lifecycle_execute_and_helpers() { assert!(take_string(nemo_relay_tool_handle_parent_uuid(handle)).is_some()); let result = cstring(r#"{"ok": true}"#); - assert_eq!( + assert_status!( nemo_relay_tool_call_end(handle, result.as_ptr(), ptr::null(), ptr::null()), NemoRelayStatus::Ok ); nemo_relay_tool_handle_free(handle); let mut execute_out = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_tool_call_execute( tool_name.as_ptr(), args.as_ptr(), @@ -853,7 +853,7 @@ fn test_ffi_tool_lifecycle_execute_and_helpers() { assert_eq!(executed_json["intercepted"], json!(true)); assert_eq!(executed_json["executed"], json!(true)); - assert_eq!(nemo_relay_flush_subscribers(), NemoRelayStatus::Ok); + assert_status!(nemo_relay_flush_subscribers(), NemoRelayStatus::Ok); let events = lock_unpoisoned(event_log()).clone(); assert!(events.iter().any(|event| event["name"] == "ffi_tool")); assert!(events.iter().any(|event| { @@ -875,7 +875,7 @@ fn test_ffi_tool_lifecycle_execute_and_helpers() { let mark_name = cstring("ffi_mark"); let mark_data = cstring(r#"{"mark":true}"#); let mark_metadata = cstring(r#"{"origin":"ffi"}"#); - assert_eq!( + assert_status!( nemo_relay_event( mark_name.as_ptr(), ptr::null(), @@ -884,7 +884,7 @@ fn test_ffi_tool_lifecycle_execute_and_helpers() { ), NemoRelayStatus::Ok ); - assert_eq!(nemo_relay_flush_subscribers(), NemoRelayStatus::Ok); + assert_status!(nemo_relay_flush_subscribers(), NemoRelayStatus::Ok); let events = lock_unpoisoned(event_log()).clone(); assert!(events.iter().any(|event| { event["name"] == "ffi_mark" @@ -895,15 +895,15 @@ fn test_ffi_tool_lifecycle_execute_and_helpers() { && event["metadata"] == json!({"origin": "ffi"}) })); - assert_eq!( + assert_status!( nemo_relay_deregister_tool_request_intercept(intercept_name_c.as_ptr()), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_deregister_tool_conditional_execution_guardrail(conditional_name_c.as_ptr()), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_deregister_subscriber(subscriber_name_c.as_ptr()), NemoRelayStatus::Ok ); @@ -918,13 +918,13 @@ fn synchronous_ffi_middleware_helper_works_inside_tokio_with_scope_local_visibil unsafe { let mut scope = ptr::null_mut(); - assert_eq!(api::nemo_relay_get_handle(&mut scope), NemoRelayStatus::Ok); + assert_status!(api::nemo_relay_get_handle(&mut scope), NemoRelayStatus::Ok); let scope_uuid = cstring( &take_string(nemo_relay_scope_handle_uuid(scope)) .expect("current scope should have a UUID"), ); let intercept_name = cstring(&unique_name("ffi_tokio_scope_intercept")); - assert_eq!( + assert_status!( nemo_relay_scope_register_tool_request_intercept( scope_uuid.as_ptr(), intercept_name.as_ptr(), @@ -947,10 +947,10 @@ fn synchronous_ffi_middleware_helper_works_inside_tokio_with_scope_local_visibil let status = runtime.block_on(async { nemo_relay_tool_request_intercepts(tool_name.as_ptr(), args.as_ptr(), &mut output) }); - assert_eq!(status, NemoRelayStatus::Ok); + assert_status!(status, NemoRelayStatus::Ok); assert_eq!(returned_json(output)["intercepted"], json!(true)); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_tool_request_intercept( scope_uuid.as_ptr(), intercept_name.as_ptr(), @@ -982,7 +982,7 @@ fn test_ffi_manual_lifecycle_timestamps_accept_unix_micros() { let stack = fresh_scope_stack(); let subscriber_name = unique_name("ffi_timestamp_subscriber"); let subscriber_name_c = cstring(&subscriber_name); - assert_eq!( + assert_status!( nemo_relay_register_subscriber( subscriber_name_c.as_ptr(), subscriber_cb, @@ -1004,7 +1004,7 @@ fn test_ffi_manual_lifecycle_timestamps_accept_unix_micros() { let scope_name = cstring("ffi_ts_scope"); let mut scope: *mut FfiScopeHandle = ptr::null_mut(); - assert_eq!( + assert_status!( api::nemo_relay_push_scope( scope_name.as_ptr(), NemoRelayScopeType::Agent, @@ -1020,7 +1020,7 @@ fn test_ffi_manual_lifecycle_timestamps_accept_unix_micros() { ); let mark_name = cstring("ffi_ts_mark"); - assert_eq!( + assert_status!( api::nemo_relay_event( mark_name.as_ptr(), scope, @@ -1034,7 +1034,7 @@ fn test_ffi_manual_lifecycle_timestamps_accept_unix_micros() { let tool_name = cstring("ffi_ts_tool"); let tool_args = cstring(r#"{"x":1}"#); let mut tool: *mut FfiToolHandle = ptr::null_mut(); - assert_eq!( + assert_status!( api::nemo_relay_tool_call( tool_name.as_ptr(), tool_args.as_ptr(), @@ -1049,7 +1049,7 @@ fn test_ffi_manual_lifecycle_timestamps_accept_unix_micros() { NemoRelayStatus::Ok ); let tool_result = cstring(r#"{"ok":true}"#); - assert_eq!( + assert_status!( api::nemo_relay_tool_call_end( tool, tool_result.as_ptr(), @@ -1064,7 +1064,7 @@ fn test_ffi_manual_lifecycle_timestamps_accept_unix_micros() { let llm_request = cstring(r#"{"headers":{},"content":{"messages":[],"model":"test-model"}}"#); let mut llm: *mut FfiLLMHandle = ptr::null_mut(); - assert_eq!( + assert_status!( api::nemo_relay_llm_call( llm_name.as_ptr(), llm_request.as_ptr(), @@ -1079,7 +1079,7 @@ fn test_ffi_manual_lifecycle_timestamps_accept_unix_micros() { NemoRelayStatus::Ok ); let llm_response = cstring(r#"{"ok":true}"#); - assert_eq!( + assert_status!( api::nemo_relay_llm_call_end( llm, llm_response.as_ptr(), @@ -1090,12 +1090,12 @@ fn test_ffi_manual_lifecycle_timestamps_accept_unix_micros() { NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( api::nemo_relay_pop_scope(scope, ptr::null(), ptr::null(), ×tamps[6]), NemoRelayStatus::Ok ); - assert_eq!(nemo_relay_flush_subscribers(), NemoRelayStatus::Ok); + assert_status!(nemo_relay_flush_subscribers(), NemoRelayStatus::Ok); let events = lock_unpoisoned(event_log()).clone(); let observed: Vec<_> = events .iter() @@ -1124,7 +1124,7 @@ fn test_ffi_manual_lifecycle_timestamps_accept_unix_micros() { ] ); - assert_eq!( + assert_status!( nemo_relay_deregister_subscriber(subscriber_name_c.as_ptr()), NemoRelayStatus::Ok ); @@ -1141,7 +1141,7 @@ fn test_ffi_manual_lifecycle_timestamps_reject_out_of_range_unix_micros() { reset_globals(); fn assert_invalid_timestamp(status: NemoRelayStatus) { - assert_eq!(status, NemoRelayStatus::InvalidArg); + assert_status!(status, NemoRelayStatus::InvalidArg); assert!( unsafe { read_last_error() } .unwrap_or_default() @@ -1170,7 +1170,7 @@ fn test_ffi_manual_lifecycle_timestamps_reject_out_of_range_unix_micros() { let scope_name = cstring("ffi_valid_ts_scope"); let mut scope: *mut FfiScopeHandle = ptr::null_mut(); - assert_eq!( + assert_status!( api::nemo_relay_push_scope( scope_name.as_ptr(), NemoRelayScopeType::Agent, @@ -1212,7 +1212,7 @@ fn test_ffi_manual_lifecycle_timestamps_reject_out_of_range_unix_micros() { let tool_name = cstring("ffi_valid_ts_tool"); let mut tool: *mut FfiToolHandle = ptr::null_mut(); - assert_eq!( + assert_status!( api::nemo_relay_tool_call( tool_name.as_ptr(), tool_args.as_ptr(), @@ -1234,7 +1234,7 @@ fn test_ffi_manual_lifecycle_timestamps_reject_out_of_range_unix_micros() { ptr::null(), &invalid_timestamp, )); - assert_eq!( + assert_status!( api::nemo_relay_tool_call_end( tool, tool_result.as_ptr(), @@ -1264,7 +1264,7 @@ fn test_ffi_manual_lifecycle_timestamps_reject_out_of_range_unix_micros() { let llm_name = cstring("ffi_valid_ts_llm"); let mut llm: *mut FfiLLMHandle = ptr::null_mut(); - assert_eq!( + assert_status!( api::nemo_relay_llm_call( llm_name.as_ptr(), llm_request.as_ptr(), @@ -1286,7 +1286,7 @@ fn test_ffi_manual_lifecycle_timestamps_reject_out_of_range_unix_micros() { ptr::null(), &invalid_timestamp, )); - assert_eq!( + assert_status!( api::nemo_relay_llm_call_end( llm, llm_response.as_ptr(), @@ -1303,7 +1303,7 @@ fn test_ffi_manual_lifecycle_timestamps_reject_out_of_range_unix_micros() { ptr::null(), &invalid_timestamp, )); - assert_eq!( + assert_status!( api::nemo_relay_pop_scope(scope, ptr::null(), ptr::null(), ptr::null()), NemoRelayStatus::Ok ); @@ -1332,7 +1332,7 @@ fn test_ffi_additional_null_and_invalid_json_paths() { let mut out_json: *mut c_char = ptr::null_mut(); let mut stream: *mut FfiStream = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_tool_call( name.as_ptr(), args.as_ptr(), @@ -1345,7 +1345,7 @@ fn test_ffi_additional_null_and_invalid_json_paths() { ), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_tool_call( name.as_ptr(), invalid_json.as_ptr(), @@ -1358,7 +1358,7 @@ fn test_ffi_additional_null_and_invalid_json_paths() { ), NemoRelayStatus::InvalidJson ); - assert_eq!( + assert_status!( nemo_relay_tool_call( name.as_ptr(), args.as_ptr(), @@ -1372,7 +1372,7 @@ fn test_ffi_additional_null_and_invalid_json_paths() { NemoRelayStatus::InvalidJson ); - assert_eq!( + assert_status!( nemo_relay_tool_call( name.as_ptr(), args.as_ptr(), @@ -1385,25 +1385,25 @@ fn test_ffi_additional_null_and_invalid_json_paths() { ), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_tool_call_end(ptr::null(), args.as_ptr(), ptr::null(), ptr::null()), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_tool_call_end(handle, invalid_json.as_ptr(), ptr::null(), ptr::null()), NemoRelayStatus::InvalidJson ); - assert_eq!( + assert_status!( nemo_relay_tool_call_end(handle, args.as_ptr(), invalid_json.as_ptr(), ptr::null(),), NemoRelayStatus::InvalidJson ); - assert_eq!( + assert_status!( nemo_relay_tool_call_end(handle, args.as_ptr(), ptr::null(), ptr::null()), NemoRelayStatus::Ok ); nemo_relay_tool_handle_free(handle); - assert_eq!( + assert_status!( nemo_relay_tool_call_execute( name.as_ptr(), args.as_ptr(), @@ -1418,7 +1418,7 @@ fn test_ffi_additional_null_and_invalid_json_paths() { ), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_tool_call_execute( name.as_ptr(), invalid_json.as_ptr(), @@ -1434,7 +1434,7 @@ fn test_ffi_additional_null_and_invalid_json_paths() { NemoRelayStatus::InvalidJson ); - assert_eq!( + assert_status!( nemo_relay_llm_call( name.as_ptr(), request.as_ptr(), @@ -1447,7 +1447,7 @@ fn test_ffi_additional_null_and_invalid_json_paths() { ), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_llm_call( name.as_ptr(), invalid_json.as_ptr(), @@ -1460,7 +1460,7 @@ fn test_ffi_additional_null_and_invalid_json_paths() { ), NemoRelayStatus::InvalidJson ); - assert_eq!( + assert_status!( nemo_relay_llm_call( name.as_ptr(), invalid_request_shape.as_ptr(), @@ -1479,7 +1479,7 @@ fn test_ffi_additional_null_and_invalid_json_paths() { .contains("failed to parse native_json as LlmRequest") ); - assert_eq!( + assert_status!( nemo_relay_llm_call( name.as_ptr(), request.as_ptr(), @@ -1492,21 +1492,21 @@ fn test_ffi_additional_null_and_invalid_json_paths() { ), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_llm_call_end(ptr::null(), args.as_ptr(), ptr::null(), ptr::null()), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_llm_call_end(llm_handle, invalid_json.as_ptr(), ptr::null(), ptr::null(),), NemoRelayStatus::InvalidJson ); - assert_eq!( + assert_status!( nemo_relay_llm_call_end(llm_handle, args.as_ptr(), ptr::null(), ptr::null()), NemoRelayStatus::Ok ); nemo_relay_llm_handle_free(llm_handle); - assert_eq!( + assert_status!( nemo_relay_llm_call_execute( name.as_ptr(), request.as_ptr(), @@ -1527,7 +1527,7 @@ fn test_ffi_additional_null_and_invalid_json_paths() { ), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_llm_call_execute( name.as_ptr(), invalid_request_shape.as_ptr(), @@ -1549,7 +1549,7 @@ fn test_ffi_additional_null_and_invalid_json_paths() { NemoRelayStatus::InvalidJson ); - assert_eq!( + assert_status!( nemo_relay_llm_stream_call_execute( name.as_ptr(), request.as_ptr(), @@ -1572,7 +1572,7 @@ fn test_ffi_additional_null_and_invalid_json_paths() { ), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_llm_stream_call_execute( name.as_ptr(), invalid_request_shape.as_ptr(), diff --git a/crates/ffi/tests/unit/api/coverage_sweeps_tests.rs b/crates/ffi/tests/unit/api/coverage_sweeps_tests.rs index b778b9a2c..6fd0c8a06 100644 --- a/crates/ffi/tests/unit/api/coverage_sweeps_tests.rs +++ b/crates/ffi/tests/unit/api/coverage_sweeps_tests.rs @@ -43,13 +43,13 @@ fn test_ffi_additional_duplicate_registration_sweeps_for_missing_global_wrappers macro_rules! assert_already_exists { ($expr:expr) => { - assert_eq!($expr, NemoRelayStatus::AlreadyExists); + assert_status!($expr, NemoRelayStatus::AlreadyExists); }; } unsafe { let tool_san_req = cstring(&unique_name("dup_tool_san_req_extra")); - assert_eq!( + assert_status!( nemo_relay_register_tool_sanitize_request_guardrail( tool_san_req.as_ptr(), 1, @@ -66,13 +66,13 @@ fn test_ffi_additional_duplicate_registration_sweeps_for_missing_global_wrappers ptr::null_mut(), None, )); - assert_eq!( + assert_status!( nemo_relay_deregister_tool_sanitize_request_guardrail(tool_san_req.as_ptr()), NemoRelayStatus::Ok ); let tool_san_resp = cstring(&unique_name("dup_tool_san_resp_extra")); - assert_eq!( + assert_status!( nemo_relay_register_tool_sanitize_response_guardrail( tool_san_resp.as_ptr(), 1, @@ -89,13 +89,13 @@ fn test_ffi_additional_duplicate_registration_sweeps_for_missing_global_wrappers ptr::null_mut(), None, )); - assert_eq!( + assert_status!( nemo_relay_deregister_tool_sanitize_response_guardrail(tool_san_resp.as_ptr()), NemoRelayStatus::Ok ); let tool_exec = cstring(&unique_name("dup_tool_exec_extra")); - assert_eq!( + assert_status!( nemo_relay_register_tool_execution_intercept( tool_exec.as_ptr(), 1, @@ -112,13 +112,13 @@ fn test_ffi_additional_duplicate_registration_sweeps_for_missing_global_wrappers ptr::null_mut(), None, )); - assert_eq!( + assert_status!( nemo_relay_deregister_tool_execution_intercept(tool_exec.as_ptr()), NemoRelayStatus::Ok ); let llm_san_req = cstring(&unique_name("dup_llm_san_req_extra")); - assert_eq!( + assert_status!( nemo_relay_register_llm_sanitize_request_guardrail( llm_san_req.as_ptr(), 1, @@ -135,13 +135,13 @@ fn test_ffi_additional_duplicate_registration_sweeps_for_missing_global_wrappers ptr::null_mut(), None, )); - assert_eq!( + assert_status!( nemo_relay_deregister_llm_sanitize_request_guardrail(llm_san_req.as_ptr()), NemoRelayStatus::Ok ); let llm_exec = cstring(&unique_name("dup_llm_exec_extra")); - assert_eq!( + assert_status!( nemo_relay_register_llm_execution_intercept( llm_exec.as_ptr(), 1, @@ -158,13 +158,13 @@ fn test_ffi_additional_duplicate_registration_sweeps_for_missing_global_wrappers ptr::null_mut(), None, )); - assert_eq!( + assert_status!( nemo_relay_deregister_llm_execution_intercept(llm_exec.as_ptr()), NemoRelayStatus::Ok ); let llm_stream_exec = cstring(&unique_name("dup_llm_stream_exec_extra")); - assert_eq!( + assert_status!( nemo_relay_register_llm_stream_execution_intercept( llm_stream_exec.as_ptr(), 1, @@ -181,7 +181,7 @@ fn test_ffi_additional_duplicate_registration_sweeps_for_missing_global_wrappers ptr::null_mut(), None, )); - assert_eq!( + assert_status!( nemo_relay_deregister_llm_stream_execution_intercept(llm_stream_exec.as_ptr()), NemoRelayStatus::Ok ); @@ -196,7 +196,7 @@ fn test_ffi_runtime_owner_conflict_and_llm_shape_error_sweeps() { unsafe { let stack = fresh_scope_stack(); let mut parent = ptr::null_mut(); - assert_eq!(nemo_relay_get_handle(&mut parent), NemoRelayStatus::Ok); + assert_status!(nemo_relay_get_handle(&mut parent), NemoRelayStatus::Ok); let scope_uuid = cstring( &take_string(nemo_relay_scope_handle_uuid(parent)).expect("scope uuid should exist"), ); @@ -210,7 +210,7 @@ fn test_ffi_runtime_owner_conflict_and_llm_shape_error_sweeps() { let llm_response = cstring(r#"{"content":"ok","role":"assistant","tool_calls":[]}"#); let mut tool_handle = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_tool_call( tool_name.as_ptr(), tool_args.as_ptr(), @@ -225,7 +225,7 @@ fn test_ffi_runtime_owner_conflict_and_llm_shape_error_sweeps() { ); let mut llm_handle = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_llm_call( llm_name.as_ptr(), llm_request.as_ptr(), @@ -241,7 +241,7 @@ fn test_ffi_runtime_owner_conflict_and_llm_shape_error_sweeps() { let malformed_request = cstring(r#"{"headers":[],"content":"bad"}"#); let mut transformed_out = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_llm_request_intercepts( llm_name.as_ptr(), malformed_request.as_ptr(), @@ -254,7 +254,7 @@ fn test_ffi_runtime_owner_conflict_and_llm_shape_error_sweeps() { .unwrap_or_default() .contains("failed to parse native_json as LlmRequest") ); - assert_eq!( + assert_status!( nemo_relay_llm_conditional_execution(malformed_request.as_ptr()), NemoRelayStatus::InvalidJson ); @@ -275,7 +275,7 @@ fn test_ffi_runtime_owner_conflict_and_llm_shape_error_sweeps() { let conflict_fragment = "multiple bindings in one process"; let mut out_json = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_tool_request_intercepts( tool_name.as_ptr(), tool_args.as_ptr(), @@ -289,7 +289,7 @@ fn test_ffi_runtime_owner_conflict_and_llm_shape_error_sweeps() { .contains(conflict_fragment) ); - assert_eq!( + assert_status!( nemo_relay_llm_request_intercepts( llm_name.as_ptr(), llm_request.as_ptr(), @@ -303,7 +303,7 @@ fn test_ffi_runtime_owner_conflict_and_llm_shape_error_sweeps() { .contains(conflict_fragment) ); - assert_eq!( + assert_status!( nemo_relay_llm_conditional_execution(llm_request.as_ptr()), NemoRelayStatus::InvalidArg ); @@ -314,7 +314,7 @@ fn test_ffi_runtime_owner_conflict_and_llm_shape_error_sweeps() { ); let mut conflict_scope = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_get_handle(&mut conflict_scope), NemoRelayStatus::InvalidArg ); @@ -323,7 +323,7 @@ fn test_ffi_runtime_owner_conflict_and_llm_shape_error_sweeps() { .unwrap_or_default() .contains(conflict_fragment) ); - assert_eq!( + assert_status!( nemo_relay_push_scope( tool_name.as_ptr(), NemoRelayScopeType::Function, @@ -341,7 +341,7 @@ fn test_ffi_runtime_owner_conflict_and_llm_shape_error_sweeps() { .unwrap_or_default() .contains(conflict_fragment) ); - assert_eq!( + assert_status!( nemo_relay_event(tool_name.as_ptr(), parent, ptr::null(), ptr::null()), NemoRelayStatus::InvalidArg ); @@ -352,7 +352,7 @@ fn test_ffi_runtime_owner_conflict_and_llm_shape_error_sweeps() { ); let mut conflict_tool_handle = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_tool_call( tool_name.as_ptr(), tool_args.as_ptr(), @@ -370,7 +370,7 @@ fn test_ffi_runtime_owner_conflict_and_llm_shape_error_sweeps() { .unwrap_or_default() .contains(conflict_fragment) ); - assert_eq!( + assert_status!( nemo_relay_tool_call_end(tool_handle, tool_result.as_ptr(), ptr::null(), ptr::null()), NemoRelayStatus::InvalidArg ); @@ -380,7 +380,7 @@ fn test_ffi_runtime_owner_conflict_and_llm_shape_error_sweeps() { .contains(conflict_fragment) ); - assert_eq!( + assert_status!( nemo_relay_llm_call( llm_name.as_ptr(), llm_request.as_ptr(), @@ -398,7 +398,7 @@ fn test_ffi_runtime_owner_conflict_and_llm_shape_error_sweeps() { .unwrap_or_default() .contains(conflict_fragment) ); - assert_eq!( + assert_status!( nemo_relay_llm_call_end(llm_handle, llm_response.as_ptr(), ptr::null(), ptr::null()), NemoRelayStatus::InvalidArg ); @@ -409,123 +409,123 @@ fn test_ffi_runtime_owner_conflict_and_llm_shape_error_sweeps() { ); let global_name = cstring("conflict-global"); - assert_eq!( + assert_status!( nemo_relay_deregister_tool_sanitize_request_guardrail(global_name.as_ptr()), NemoRelayStatus::InvalidArg ); - assert_eq!( + assert_status!( nemo_relay_deregister_tool_conditional_execution_guardrail(global_name.as_ptr()), NemoRelayStatus::InvalidArg ); - assert_eq!( + assert_status!( nemo_relay_deregister_tool_request_intercept(global_name.as_ptr()), NemoRelayStatus::InvalidArg ); - assert_eq!( + assert_status!( nemo_relay_deregister_tool_execution_intercept(global_name.as_ptr()), NemoRelayStatus::InvalidArg ); - assert_eq!( + assert_status!( nemo_relay_deregister_llm_sanitize_request_guardrail(global_name.as_ptr()), NemoRelayStatus::InvalidArg ); - assert_eq!( + assert_status!( nemo_relay_deregister_llm_sanitize_response_guardrail(global_name.as_ptr()), NemoRelayStatus::InvalidArg ); - assert_eq!( + assert_status!( nemo_relay_deregister_llm_conditional_execution_guardrail(global_name.as_ptr()), NemoRelayStatus::InvalidArg ); - assert_eq!( + assert_status!( nemo_relay_deregister_llm_request_intercept(global_name.as_ptr()), NemoRelayStatus::InvalidArg ); - assert_eq!( + assert_status!( nemo_relay_deregister_llm_execution_intercept(global_name.as_ptr()), NemoRelayStatus::InvalidArg ); - assert_eq!( + assert_status!( nemo_relay_deregister_llm_stream_execution_intercept(global_name.as_ptr()), NemoRelayStatus::InvalidArg ); - assert_eq!( + assert_status!( nemo_relay_deregister_subscriber(global_name.as_ptr()), NemoRelayStatus::InvalidArg ); let scope_name = cstring("conflict-scope"); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_tool_sanitize_request_guardrail( scope_uuid.as_ptr(), scope_name.as_ptr(), ), NemoRelayStatus::InvalidArg ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_tool_conditional_execution_guardrail( scope_uuid.as_ptr(), scope_name.as_ptr(), ), NemoRelayStatus::InvalidArg ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_tool_request_intercept( scope_uuid.as_ptr(), scope_name.as_ptr() ), NemoRelayStatus::InvalidArg ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_tool_execution_intercept( scope_uuid.as_ptr(), scope_name.as_ptr(), ), NemoRelayStatus::InvalidArg ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_llm_sanitize_request_guardrail( scope_uuid.as_ptr(), scope_name.as_ptr(), ), NemoRelayStatus::InvalidArg ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_llm_sanitize_response_guardrail( scope_uuid.as_ptr(), scope_name.as_ptr(), ), NemoRelayStatus::InvalidArg ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_llm_conditional_execution_guardrail( scope_uuid.as_ptr(), scope_name.as_ptr(), ), NemoRelayStatus::InvalidArg ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_llm_request_intercept( scope_uuid.as_ptr(), scope_name.as_ptr(), ), NemoRelayStatus::InvalidArg ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_llm_execution_intercept( scope_uuid.as_ptr(), scope_name.as_ptr(), ), NemoRelayStatus::InvalidArg ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_llm_stream_execution_intercept( scope_uuid.as_ptr(), scope_name.as_ptr(), ), NemoRelayStatus::InvalidArg ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_subscriber(scope_uuid.as_ptr(), scope_name.as_ptr()), NemoRelayStatus::InvalidArg ); @@ -544,7 +544,7 @@ fn test_ffi_additional_duplicate_registration_sweeps_for_missing_scope_wrappers( macro_rules! assert_already_exists { ($expr:expr) => { - assert_eq!($expr, NemoRelayStatus::AlreadyExists); + assert_status!($expr, NemoRelayStatus::AlreadyExists); }; } @@ -552,7 +552,7 @@ fn test_ffi_additional_duplicate_registration_sweeps_for_missing_scope_wrappers( let stack = fresh_scope_stack(); let scope_name = cstring("dup_scope_extra"); let mut scope = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_push_scope( scope_name.as_ptr(), NemoRelayScopeType::Function, @@ -568,7 +568,7 @@ fn test_ffi_additional_duplicate_registration_sweeps_for_missing_scope_wrappers( let scope_uuid = cstring(&take_string(nemo_relay_scope_handle_uuid(scope)).unwrap()); let tool_san_req = cstring(&unique_name("dup_scope_tool_san_req_extra")); - assert_eq!( + assert_status!( nemo_relay_scope_register_tool_sanitize_request_guardrail( scope_uuid.as_ptr(), tool_san_req.as_ptr(), @@ -587,7 +587,7 @@ fn test_ffi_additional_duplicate_registration_sweeps_for_missing_scope_wrappers( ptr::null_mut(), None, )); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_tool_sanitize_request_guardrail( scope_uuid.as_ptr(), tool_san_req.as_ptr(), @@ -596,7 +596,7 @@ fn test_ffi_additional_duplicate_registration_sweeps_for_missing_scope_wrappers( ); let tool_san_resp = cstring(&unique_name("dup_scope_tool_san_resp_extra")); - assert_eq!( + assert_status!( nemo_relay_scope_register_tool_sanitize_response_guardrail( scope_uuid.as_ptr(), tool_san_resp.as_ptr(), @@ -615,7 +615,7 @@ fn test_ffi_additional_duplicate_registration_sweeps_for_missing_scope_wrappers( ptr::null_mut(), None, )); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_tool_sanitize_response_guardrail( scope_uuid.as_ptr(), tool_san_resp.as_ptr(), @@ -624,7 +624,7 @@ fn test_ffi_additional_duplicate_registration_sweeps_for_missing_scope_wrappers( ); let tool_exec = cstring(&unique_name("dup_scope_tool_exec_extra")); - assert_eq!( + assert_status!( nemo_relay_scope_register_tool_execution_intercept( scope_uuid.as_ptr(), tool_exec.as_ptr(), @@ -643,7 +643,7 @@ fn test_ffi_additional_duplicate_registration_sweeps_for_missing_scope_wrappers( ptr::null_mut(), None, )); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_tool_execution_intercept( scope_uuid.as_ptr(), tool_exec.as_ptr(), @@ -652,7 +652,7 @@ fn test_ffi_additional_duplicate_registration_sweeps_for_missing_scope_wrappers( ); let llm_san_req = cstring(&unique_name("dup_scope_llm_san_req_extra")); - assert_eq!( + assert_status!( nemo_relay_scope_register_llm_sanitize_request_guardrail( scope_uuid.as_ptr(), llm_san_req.as_ptr(), @@ -671,7 +671,7 @@ fn test_ffi_additional_duplicate_registration_sweeps_for_missing_scope_wrappers( ptr::null_mut(), None, )); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_llm_sanitize_request_guardrail( scope_uuid.as_ptr(), llm_san_req.as_ptr(), @@ -680,7 +680,7 @@ fn test_ffi_additional_duplicate_registration_sweeps_for_missing_scope_wrappers( ); let llm_san_resp = cstring(&unique_name("dup_scope_llm_san_resp_extra")); - assert_eq!( + assert_status!( nemo_relay_scope_register_llm_sanitize_response_guardrail( scope_uuid.as_ptr(), llm_san_resp.as_ptr(), @@ -699,7 +699,7 @@ fn test_ffi_additional_duplicate_registration_sweeps_for_missing_scope_wrappers( ptr::null_mut(), None, )); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_llm_sanitize_response_guardrail( scope_uuid.as_ptr(), llm_san_resp.as_ptr(), @@ -708,7 +708,7 @@ fn test_ffi_additional_duplicate_registration_sweeps_for_missing_scope_wrappers( ); let llm_exec = cstring(&unique_name("dup_scope_llm_exec_extra")); - assert_eq!( + assert_status!( nemo_relay_scope_register_llm_execution_intercept( scope_uuid.as_ptr(), llm_exec.as_ptr(), @@ -727,7 +727,7 @@ fn test_ffi_additional_duplicate_registration_sweeps_for_missing_scope_wrappers( ptr::null_mut(), None, )); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_llm_execution_intercept( scope_uuid.as_ptr(), llm_exec.as_ptr(), @@ -736,7 +736,7 @@ fn test_ffi_additional_duplicate_registration_sweeps_for_missing_scope_wrappers( ); let llm_stream_exec = cstring(&unique_name("dup_scope_llm_stream_exec_extra")); - assert_eq!( + assert_status!( nemo_relay_scope_register_llm_stream_execution_intercept( scope_uuid.as_ptr(), llm_stream_exec.as_ptr(), @@ -755,7 +755,7 @@ fn test_ffi_additional_duplicate_registration_sweeps_for_missing_scope_wrappers( ptr::null_mut(), None, )); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_llm_stream_execution_intercept( scope_uuid.as_ptr(), llm_stream_exec.as_ptr(), @@ -763,7 +763,7 @@ fn test_ffi_additional_duplicate_registration_sweeps_for_missing_scope_wrappers( NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_pop_scope(scope, ptr::null()), NemoRelayStatus::Ok ); @@ -781,7 +781,7 @@ fn test_ffi_global_tool_registration_invalid_utf8_name_sweep() { let invalid = invalid_utf8.as_ptr() as *const c_char; unsafe { - assert_eq!( + assert_status!( nemo_relay_register_tool_sanitize_request_guardrail( invalid, 1, @@ -791,11 +791,11 @@ fn test_ffi_global_tool_registration_invalid_utf8_name_sweep() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_deregister_tool_sanitize_request_guardrail(invalid), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_register_tool_sanitize_response_guardrail( invalid, 1, @@ -805,11 +805,11 @@ fn test_ffi_global_tool_registration_invalid_utf8_name_sweep() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_deregister_tool_sanitize_response_guardrail(invalid), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_register_tool_conditional_execution_guardrail( invalid, 1, @@ -819,11 +819,11 @@ fn test_ffi_global_tool_registration_invalid_utf8_name_sweep() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_deregister_tool_conditional_execution_guardrail(invalid), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_register_tool_request_intercept( invalid, 1, @@ -834,11 +834,11 @@ fn test_ffi_global_tool_registration_invalid_utf8_name_sweep() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_deregister_tool_request_intercept(invalid), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_register_tool_execution_intercept( invalid, 1, @@ -848,7 +848,7 @@ fn test_ffi_global_tool_registration_invalid_utf8_name_sweep() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_deregister_tool_execution_intercept(invalid), NemoRelayStatus::InvalidUtf8 ); @@ -864,7 +864,7 @@ fn test_ffi_global_llm_and_subscriber_registration_invalid_utf8_name_sweep() { let invalid = invalid_utf8.as_ptr() as *const c_char; unsafe { - assert_eq!( + assert_status!( nemo_relay_register_llm_sanitize_request_guardrail( invalid, 1, @@ -874,11 +874,11 @@ fn test_ffi_global_llm_and_subscriber_registration_invalid_utf8_name_sweep() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_deregister_llm_sanitize_request_guardrail(invalid), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_register_llm_sanitize_response_guardrail( invalid, 1, @@ -888,11 +888,11 @@ fn test_ffi_global_llm_and_subscriber_registration_invalid_utf8_name_sweep() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_deregister_llm_sanitize_response_guardrail(invalid), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_register_llm_conditional_execution_guardrail( invalid, 1, @@ -902,11 +902,11 @@ fn test_ffi_global_llm_and_subscriber_registration_invalid_utf8_name_sweep() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_deregister_llm_conditional_execution_guardrail(invalid), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_register_llm_request_intercept( invalid, 1, @@ -917,11 +917,11 @@ fn test_ffi_global_llm_and_subscriber_registration_invalid_utf8_name_sweep() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_deregister_llm_request_intercept(invalid), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_register_llm_execution_intercept( invalid, 1, @@ -931,11 +931,11 @@ fn test_ffi_global_llm_and_subscriber_registration_invalid_utf8_name_sweep() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_deregister_llm_execution_intercept(invalid), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_register_llm_stream_execution_intercept( invalid, 1, @@ -945,15 +945,15 @@ fn test_ffi_global_llm_and_subscriber_registration_invalid_utf8_name_sweep() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_deregister_llm_stream_execution_intercept(invalid), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_register_subscriber(invalid, subscriber_cb, ptr::null_mut(), None), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_deregister_subscriber(invalid), NemoRelayStatus::InvalidUtf8 ); @@ -970,7 +970,7 @@ fn test_ffi_scope_tool_registration_invalid_utf8_scope_uuid_sweep() { let name = cstring("scope-tool-invalid-scope"); unsafe { - assert_eq!( + assert_status!( nemo_relay_scope_register_tool_sanitize_request_guardrail( invalid_scope, name.as_ptr(), @@ -981,14 +981,14 @@ fn test_ffi_scope_tool_registration_invalid_utf8_scope_uuid_sweep() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_tool_sanitize_request_guardrail( invalid_scope, name.as_ptr(), ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_register_tool_sanitize_response_guardrail( invalid_scope, name.as_ptr(), @@ -999,14 +999,14 @@ fn test_ffi_scope_tool_registration_invalid_utf8_scope_uuid_sweep() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_tool_sanitize_response_guardrail( invalid_scope, name.as_ptr(), ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_register_tool_conditional_execution_guardrail( invalid_scope, name.as_ptr(), @@ -1017,14 +1017,14 @@ fn test_ffi_scope_tool_registration_invalid_utf8_scope_uuid_sweep() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_tool_conditional_execution_guardrail( invalid_scope, name.as_ptr(), ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_register_tool_request_intercept( invalid_scope, name.as_ptr(), @@ -1036,11 +1036,11 @@ fn test_ffi_scope_tool_registration_invalid_utf8_scope_uuid_sweep() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_tool_request_intercept(invalid_scope, name.as_ptr()), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_register_tool_execution_intercept( invalid_scope, name.as_ptr(), @@ -1051,7 +1051,7 @@ fn test_ffi_scope_tool_registration_invalid_utf8_scope_uuid_sweep() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_tool_execution_intercept(invalid_scope, name.as_ptr()), NemoRelayStatus::InvalidUtf8 ); @@ -1068,7 +1068,7 @@ fn test_ffi_scope_llm_and_subscriber_registration_invalid_utf8_scope_uuid_sweep( let name = cstring("scope-llm-invalid-scope"); unsafe { - assert_eq!( + assert_status!( nemo_relay_scope_register_llm_sanitize_request_guardrail( invalid_scope, name.as_ptr(), @@ -1079,14 +1079,14 @@ fn test_ffi_scope_llm_and_subscriber_registration_invalid_utf8_scope_uuid_sweep( ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_llm_sanitize_request_guardrail( invalid_scope, name.as_ptr(), ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_register_llm_sanitize_response_guardrail( invalid_scope, name.as_ptr(), @@ -1097,14 +1097,14 @@ fn test_ffi_scope_llm_and_subscriber_registration_invalid_utf8_scope_uuid_sweep( ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_llm_sanitize_response_guardrail( invalid_scope, name.as_ptr(), ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_register_llm_conditional_execution_guardrail( invalid_scope, name.as_ptr(), @@ -1115,14 +1115,14 @@ fn test_ffi_scope_llm_and_subscriber_registration_invalid_utf8_scope_uuid_sweep( ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_llm_conditional_execution_guardrail( invalid_scope, name.as_ptr(), ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_register_llm_request_intercept( invalid_scope, name.as_ptr(), @@ -1134,11 +1134,11 @@ fn test_ffi_scope_llm_and_subscriber_registration_invalid_utf8_scope_uuid_sweep( ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_llm_request_intercept(invalid_scope, name.as_ptr()), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_register_llm_execution_intercept( invalid_scope, name.as_ptr(), @@ -1149,11 +1149,11 @@ fn test_ffi_scope_llm_and_subscriber_registration_invalid_utf8_scope_uuid_sweep( ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_llm_execution_intercept(invalid_scope, name.as_ptr()), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_register_llm_stream_execution_intercept( invalid_scope, name.as_ptr(), @@ -1164,14 +1164,14 @@ fn test_ffi_scope_llm_and_subscriber_registration_invalid_utf8_scope_uuid_sweep( ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_llm_stream_execution_intercept( invalid_scope, name.as_ptr() ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_register_subscriber( invalid_scope, name.as_ptr(), @@ -1181,7 +1181,7 @@ fn test_ffi_scope_llm_and_subscriber_registration_invalid_utf8_scope_uuid_sweep( ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_subscriber(invalid_scope, name.as_ptr()), NemoRelayStatus::InvalidUtf8 ); @@ -1197,7 +1197,7 @@ fn test_ffi_scope_tool_registration_invalid_utf8_name_sweep() { let stack = fresh_scope_stack(); let scope_name = cstring("scope-tool-invalid-name"); let mut scope = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_push_scope( scope_name.as_ptr(), NemoRelayScopeType::Function, @@ -1214,7 +1214,7 @@ fn test_ffi_scope_tool_registration_invalid_utf8_name_sweep() { let invalid_utf8 = [0xffu8, 0]; let invalid = invalid_utf8.as_ptr() as *const c_char; - assert_eq!( + assert_status!( nemo_relay_scope_register_tool_sanitize_request_guardrail( scope_uuid.as_ptr(), invalid, @@ -1225,14 +1225,14 @@ fn test_ffi_scope_tool_registration_invalid_utf8_name_sweep() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_tool_sanitize_request_guardrail( scope_uuid.as_ptr(), invalid ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_register_tool_sanitize_response_guardrail( scope_uuid.as_ptr(), invalid, @@ -1243,14 +1243,14 @@ fn test_ffi_scope_tool_registration_invalid_utf8_name_sweep() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_tool_sanitize_response_guardrail( scope_uuid.as_ptr(), invalid ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_register_tool_conditional_execution_guardrail( scope_uuid.as_ptr(), invalid, @@ -1261,14 +1261,14 @@ fn test_ffi_scope_tool_registration_invalid_utf8_name_sweep() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_tool_conditional_execution_guardrail( scope_uuid.as_ptr(), invalid, ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_register_tool_request_intercept( scope_uuid.as_ptr(), invalid, @@ -1280,11 +1280,11 @@ fn test_ffi_scope_tool_registration_invalid_utf8_name_sweep() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_tool_request_intercept(scope_uuid.as_ptr(), invalid), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_register_tool_execution_intercept( scope_uuid.as_ptr(), invalid, @@ -1295,12 +1295,12 @@ fn test_ffi_scope_tool_registration_invalid_utf8_name_sweep() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_tool_execution_intercept(scope_uuid.as_ptr(), invalid), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_pop_scope(scope, ptr::null()), NemoRelayStatus::Ok ); @@ -1318,7 +1318,7 @@ fn test_ffi_scope_llm_and_subscriber_registration_invalid_utf8_name_sweep() { let stack = fresh_scope_stack(); let scope_name = cstring("scope-llm-invalid-name"); let mut scope = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_push_scope( scope_name.as_ptr(), NemoRelayScopeType::Function, @@ -1335,7 +1335,7 @@ fn test_ffi_scope_llm_and_subscriber_registration_invalid_utf8_name_sweep() { let invalid_utf8 = [0xffu8, 0]; let invalid = invalid_utf8.as_ptr() as *const c_char; - assert_eq!( + assert_status!( nemo_relay_scope_register_llm_sanitize_request_guardrail( scope_uuid.as_ptr(), invalid, @@ -1346,14 +1346,14 @@ fn test_ffi_scope_llm_and_subscriber_registration_invalid_utf8_name_sweep() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_llm_sanitize_request_guardrail( scope_uuid.as_ptr(), invalid ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_register_llm_sanitize_response_guardrail( scope_uuid.as_ptr(), invalid, @@ -1364,14 +1364,14 @@ fn test_ffi_scope_llm_and_subscriber_registration_invalid_utf8_name_sweep() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_llm_sanitize_response_guardrail( scope_uuid.as_ptr(), invalid ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_register_llm_conditional_execution_guardrail( scope_uuid.as_ptr(), invalid, @@ -1382,14 +1382,14 @@ fn test_ffi_scope_llm_and_subscriber_registration_invalid_utf8_name_sweep() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_llm_conditional_execution_guardrail( scope_uuid.as_ptr(), invalid, ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_register_llm_request_intercept( scope_uuid.as_ptr(), invalid, @@ -1401,11 +1401,11 @@ fn test_ffi_scope_llm_and_subscriber_registration_invalid_utf8_name_sweep() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_llm_request_intercept(scope_uuid.as_ptr(), invalid), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_register_llm_execution_intercept( scope_uuid.as_ptr(), invalid, @@ -1416,11 +1416,11 @@ fn test_ffi_scope_llm_and_subscriber_registration_invalid_utf8_name_sweep() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_llm_execution_intercept(scope_uuid.as_ptr(), invalid), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_register_llm_stream_execution_intercept( scope_uuid.as_ptr(), invalid, @@ -1431,14 +1431,14 @@ fn test_ffi_scope_llm_and_subscriber_registration_invalid_utf8_name_sweep() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_llm_stream_execution_intercept( scope_uuid.as_ptr(), invalid ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_register_subscriber( scope_uuid.as_ptr(), invalid, @@ -1448,12 +1448,12 @@ fn test_ffi_scope_llm_and_subscriber_registration_invalid_utf8_name_sweep() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_subscriber(scope_uuid.as_ptr(), invalid), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_pop_scope(scope, ptr::null()), NemoRelayStatus::Ok ); @@ -1470,7 +1470,7 @@ fn test_ffi_scope_and_event_parent_and_utf8_paths() { unsafe { let stack = fresh_scope_stack(); let mut parent = ptr::null_mut(); - assert_eq!(nemo_relay_get_handle(&mut parent), NemoRelayStatus::Ok); + assert_status!(nemo_relay_get_handle(&mut parent), NemoRelayStatus::Ok); let scope_name = cstring("ffi_child_scope_with_parent"); let data = cstring(r#"{"scope":"child"}"#); @@ -1480,7 +1480,7 @@ fn test_ffi_scope_and_event_parent_and_utf8_paths() { let invalid = invalid_utf8.as_ptr() as *const c_char; let mut child = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_push_scope( scope_name.as_ptr(), NemoRelayScopeType::Function, @@ -1494,7 +1494,7 @@ fn test_ffi_scope_and_event_parent_and_utf8_paths() { NemoRelayStatus::Ok ); assert!(take_string(nemo_relay_scope_handle_parent_uuid(child)).is_some()); - assert_eq!( + assert_status!( nemo_relay_push_scope( invalid, NemoRelayScopeType::Function, @@ -1507,7 +1507,7 @@ fn test_ffi_scope_and_event_parent_and_utf8_paths() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_push_scope( scope_name.as_ptr(), NemoRelayScopeType::Function, @@ -1520,7 +1520,7 @@ fn test_ffi_scope_and_event_parent_and_utf8_paths() { ), NemoRelayStatus::InvalidJson ); - assert_eq!( + assert_status!( nemo_relay_push_scope( scope_name.as_ptr(), NemoRelayScopeType::Function, @@ -1535,7 +1535,7 @@ fn test_ffi_scope_and_event_parent_and_utf8_paths() { ); let event_name = cstring("ffi_event_with_parent"); - assert_eq!( + assert_status!( nemo_relay_event( event_name.as_ptr(), parent, @@ -1544,11 +1544,11 @@ fn test_ffi_scope_and_event_parent_and_utf8_paths() { ), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_event(invalid, parent, ptr::null(), ptr::null()), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_event( event_name.as_ptr(), parent, @@ -1557,7 +1557,7 @@ fn test_ffi_scope_and_event_parent_and_utf8_paths() { ), NemoRelayStatus::InvalidJson ); - assert_eq!( + assert_status!( nemo_relay_event( event_name.as_ptr(), parent, @@ -1567,7 +1567,7 @@ fn test_ffi_scope_and_event_parent_and_utf8_paths() { NemoRelayStatus::InvalidJson ); - assert_eq!( + assert_status!( nemo_relay_pop_scope(child, ptr::null()), NemoRelayStatus::Ok ); @@ -1585,7 +1585,7 @@ fn test_ffi_tool_call_parent_tool_call_id_and_utf8_paths() { unsafe { let stack = fresh_scope_stack(); let mut parent = ptr::null_mut(); - assert_eq!(nemo_relay_get_handle(&mut parent), NemoRelayStatus::Ok); + assert_status!(nemo_relay_get_handle(&mut parent), NemoRelayStatus::Ok); let name = cstring("ffi_tool_call_utf8"); let args = cstring(r#"{"value":1}"#); @@ -1598,7 +1598,7 @@ fn test_ffi_tool_call_parent_tool_call_id_and_utf8_paths() { let invalid = invalid_utf8.as_ptr() as *const c_char; let mut handle = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_tool_call( name.as_ptr(), args.as_ptr(), @@ -1612,7 +1612,7 @@ fn test_ffi_tool_call_parent_tool_call_id_and_utf8_paths() { NemoRelayStatus::Ok ); assert!(take_string(nemo_relay_tool_handle_parent_uuid(handle)).is_some()); - assert_eq!( + assert_status!( nemo_relay_tool_call( invalid, args.as_ptr(), @@ -1625,7 +1625,7 @@ fn test_ffi_tool_call_parent_tool_call_id_and_utf8_paths() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_tool_call( name.as_ptr(), args.as_ptr(), @@ -1638,11 +1638,11 @@ fn test_ffi_tool_call_parent_tool_call_id_and_utf8_paths() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_tool_call_end(handle, result.as_ptr(), ptr::null(), invalid_json.as_ptr()), NemoRelayStatus::InvalidJson ); - assert_eq!( + assert_status!( nemo_relay_tool_call_end(handle, result.as_ptr(), data.as_ptr(), metadata.as_ptr()), NemoRelayStatus::Ok ); @@ -1661,7 +1661,7 @@ fn test_ffi_llm_call_parent_model_and_utf8_paths() { unsafe { let stack = fresh_scope_stack(); let mut parent = ptr::null_mut(); - assert_eq!(nemo_relay_get_handle(&mut parent), NemoRelayStatus::Ok); + assert_status!(nemo_relay_get_handle(&mut parent), NemoRelayStatus::Ok); let name = cstring("ffi_llm_call_utf8"); let request = cstring( @@ -1676,7 +1676,7 @@ fn test_ffi_llm_call_parent_model_and_utf8_paths() { let invalid = invalid_utf8.as_ptr() as *const c_char; let mut handle = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_llm_call( name.as_ptr(), request.as_ptr(), @@ -1690,7 +1690,7 @@ fn test_ffi_llm_call_parent_model_and_utf8_paths() { NemoRelayStatus::Ok ); assert!(take_string(nemo_relay_llm_handle_parent_uuid(handle)).is_some()); - assert_eq!( + assert_status!( nemo_relay_llm_call( invalid, request.as_ptr(), @@ -1703,7 +1703,7 @@ fn test_ffi_llm_call_parent_model_and_utf8_paths() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_llm_call( name.as_ptr(), request.as_ptr(), @@ -1716,7 +1716,7 @@ fn test_ffi_llm_call_parent_model_and_utf8_paths() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_llm_call_end( handle, response.as_ptr(), @@ -1725,7 +1725,7 @@ fn test_ffi_llm_call_parent_model_and_utf8_paths() { ), NemoRelayStatus::InvalidJson ); - assert_eq!( + assert_status!( nemo_relay_llm_call_end(handle, response.as_ptr(), data.as_ptr(), metadata.as_ptr()), NemoRelayStatus::Ok ); @@ -1748,7 +1748,7 @@ fn test_ffi_llm_execute_and_stream_shape_and_out_error_paths() { r#"{"headers":{},"content":{"messages":[{"role":"user","content":"hi"}],"model":"ffi-model"}}"#, ); - assert_eq!( + assert_status!( nemo_relay_llm_call_execute( name.as_ptr(), request.as_ptr(), @@ -1770,7 +1770,7 @@ fn test_ffi_llm_execute_and_stream_shape_and_out_error_paths() { NemoRelayStatus::NullPointer ); let mut out = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_llm_call_execute( name.as_ptr(), invalid_shape.as_ptr(), @@ -1797,7 +1797,7 @@ fn test_ffi_llm_execute_and_stream_shape_and_out_error_paths() { .contains("failed to parse native_json as LlmRequest") ); - assert_eq!( + assert_status!( nemo_relay_llm_stream_call_execute( name.as_ptr(), request.as_ptr(), @@ -1821,7 +1821,7 @@ fn test_ffi_llm_execute_and_stream_shape_and_out_error_paths() { NemoRelayStatus::NullPointer ); let mut stream = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_llm_stream_call_execute( name.as_ptr(), invalid_shape.as_ptr(), @@ -1896,7 +1896,7 @@ fn test_ffi_llm_helper_invalid_shape_and_intercept_failure_paths() { let invalid_shape = cstring(r#"{"headers":[],"content":1}"#); let mut out = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_llm_request_intercepts(name.as_ptr(), invalid_shape.as_ptr(), &mut out), NemoRelayStatus::InvalidJson ); @@ -1906,7 +1906,7 @@ fn test_ffi_llm_helper_invalid_shape_and_intercept_failure_paths() { .contains("failed to parse native_json as LlmRequest") ); - assert_eq!( + assert_status!( nemo_relay_llm_conditional_execution(invalid_shape.as_ptr()), NemoRelayStatus::InvalidJson ); @@ -1917,7 +1917,7 @@ fn test_ffi_llm_helper_invalid_shape_and_intercept_failure_paths() { ); let intercept_name = cstring(&unique_name("ffi_llm_request_intercept_fail")); - assert_eq!( + assert_status!( nemo_relay_register_llm_request_intercept( intercept_name.as_ptr(), 1, @@ -1928,7 +1928,7 @@ fn test_ffi_llm_helper_invalid_shape_and_intercept_failure_paths() { ), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_llm_request_intercepts(name.as_ptr(), valid_request.as_ptr(), &mut out), NemoRelayStatus::Internal ); @@ -1937,7 +1937,7 @@ fn test_ffi_llm_helper_invalid_shape_and_intercept_failure_paths() { .unwrap_or_default() .contains("llm request intercept callback failed") ); - assert_eq!( + assert_status!( nemo_relay_deregister_llm_request_intercept(intercept_name.as_ptr()), NemoRelayStatus::Ok ); @@ -1954,7 +1954,7 @@ fn test_ffi_helper_and_lifecycle_callback_failure_paths() { unsafe { let stack = fresh_scope_stack(); let mut parent = ptr::null_mut(); - assert_eq!(nemo_relay_get_handle(&mut parent), NemoRelayStatus::Ok); + assert_status!(nemo_relay_get_handle(&mut parent), NemoRelayStatus::Ok); let tool_name = cstring("ffi_tool_failure_sweep"); let tool_args = cstring(r#"{"value":9}"#); @@ -1964,7 +1964,7 @@ fn test_ffi_helper_and_lifecycle_callback_failure_paths() { let llm_response = cstring(r#"{"content":"ok","role":"assistant","tool_calls":[]}"#); let tool_intercept = cstring(&unique_name("ffi_tool_helper_fail")); - assert_eq!( + assert_status!( nemo_relay_register_tool_request_intercept( tool_intercept.as_ptr(), 1, @@ -1976,7 +1976,7 @@ fn test_ffi_helper_and_lifecycle_callback_failure_paths() { NemoRelayStatus::Ok ); let mut tool_out = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_tool_request_intercepts( tool_name.as_ptr(), tool_args.as_ptr(), @@ -1989,13 +1989,13 @@ fn test_ffi_helper_and_lifecycle_callback_failure_paths() { .unwrap_or_default() .contains("tool sanitize callback failed") ); - assert_eq!( + assert_status!( nemo_relay_deregister_tool_request_intercept(tool_intercept.as_ptr()), NemoRelayStatus::Ok ); let llm_intercept = cstring(&unique_name("ffi_llm_helper_fail")); - assert_eq!( + assert_status!( nemo_relay_register_llm_request_intercept( llm_intercept.as_ptr(), 1, @@ -2007,7 +2007,7 @@ fn test_ffi_helper_and_lifecycle_callback_failure_paths() { NemoRelayStatus::Ok ); let mut llm_out = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_llm_request_intercepts( llm_name.as_ptr(), llm_request.as_ptr(), @@ -2020,13 +2020,13 @@ fn test_ffi_helper_and_lifecycle_callback_failure_paths() { .unwrap_or_default() .contains("llm request intercept callback failed") ); - assert_eq!( + assert_status!( nemo_relay_deregister_llm_request_intercept(llm_intercept.as_ptr()), NemoRelayStatus::Ok ); let mut llm_handle = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_llm_call( llm_name.as_ptr(), llm_request.as_ptr(), @@ -2039,14 +2039,14 @@ fn test_ffi_helper_and_lifecycle_callback_failure_paths() { ), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_llm_call_end(llm_handle, llm_response.as_ptr(), ptr::null(), ptr::null()), NemoRelayStatus::Ok ); nemo_relay_llm_handle_free(llm_handle); let mut tool_handle = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_tool_call( tool_name.as_ptr(), tool_args.as_ptr(), @@ -2061,7 +2061,7 @@ fn test_ffi_helper_and_lifecycle_callback_failure_paths() { ); let tool_result = cstring(r#"{"done":true}"#); - assert_eq!( + assert_status!( nemo_relay_tool_call_end(tool_handle, tool_result.as_ptr(), ptr::null(), ptr::null()), NemoRelayStatus::Ok ); @@ -2071,7 +2071,7 @@ fn test_ffi_helper_and_lifecycle_callback_failure_paths() { let invalid_name = invalid_utf8.as_ptr() as *const c_char; let invalid_json = cstring("{"); let mut exec_out = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_tool_call_execute( invalid_name, tool_args.as_ptr(), @@ -2086,7 +2086,7 @@ fn test_ffi_helper_and_lifecycle_callback_failure_paths() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_tool_call_execute( tool_name.as_ptr(), tool_args.as_ptr(), @@ -2120,7 +2120,7 @@ fn test_ffi_scope_registry_missing_scope_and_null_out_sweeps() { let invalid_utf8 = [0xffu8, 0]; let invalid_name = invalid_utf8.as_ptr() as *const c_char; - assert_eq!( + assert_status!( nemo_relay_push_scope( scope_name.as_ptr(), NemoRelayScopeType::Function, @@ -2141,7 +2141,7 @@ fn test_ffi_scope_registry_missing_scope_and_null_out_sweeps() { macro_rules! assert_missing_scope { ($expr:expr) => { - assert_eq!($expr, NemoRelayStatus::NotFound); + assert_status!($expr, NemoRelayStatus::NotFound); }; } @@ -2303,7 +2303,7 @@ fn test_ffi_scope_registry_missing_scope_and_null_out_sweeps() { )); let mut scope = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_push_scope( scope_name.as_ptr(), NemoRelayScopeType::Function, @@ -2318,19 +2318,19 @@ fn test_ffi_scope_registry_missing_scope_and_null_out_sweeps() { ); let scope_uuid = cstring(&take_string(nemo_relay_scope_handle_uuid(scope)).unwrap()); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_llm_stream_execution_intercept( scope_uuid.as_ptr(), invalid_name, ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_subscriber(scope_uuid.as_ptr(), invalid_name), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_pop_scope(scope, ptr::null()), NemoRelayStatus::Ok ); @@ -2347,7 +2347,7 @@ fn test_ffi_llm_lifecycle_additional_error_paths() { unsafe { let stack = fresh_scope_stack(); let mut parent = ptr::null_mut(); - assert_eq!(nemo_relay_get_handle(&mut parent), NemoRelayStatus::Ok); + assert_status!(nemo_relay_get_handle(&mut parent), NemoRelayStatus::Ok); let name = cstring("ffi_llm_lifecycle_extra"); let request = cstring( @@ -2357,7 +2357,7 @@ fn test_ffi_llm_lifecycle_additional_error_paths() { let invalid_json = cstring("{"); let mut handle = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_llm_call( name.as_ptr(), request.as_ptr(), @@ -2370,7 +2370,7 @@ fn test_ffi_llm_lifecycle_additional_error_paths() { ), NemoRelayStatus::InvalidJson ); - assert_eq!( + assert_status!( nemo_relay_llm_call( name.as_ptr(), request.as_ptr(), @@ -2384,7 +2384,7 @@ fn test_ffi_llm_lifecycle_additional_error_paths() { NemoRelayStatus::InvalidJson ); - assert_eq!( + assert_status!( nemo_relay_llm_call( name.as_ptr(), request.as_ptr(), @@ -2397,7 +2397,7 @@ fn test_ffi_llm_lifecycle_additional_error_paths() { ), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_llm_call_end( handle, response.as_ptr(), @@ -2406,7 +2406,7 @@ fn test_ffi_llm_lifecycle_additional_error_paths() { ), NemoRelayStatus::InvalidJson ); - assert_eq!( + assert_status!( nemo_relay_llm_call_end(handle, response.as_ptr(), ptr::null(), ptr::null()), NemoRelayStatus::Ok ); @@ -2425,7 +2425,7 @@ fn test_ffi_llm_execute_and_stream_additional_input_paths() { unsafe { let stack = fresh_scope_stack(); let mut parent = ptr::null_mut(); - assert_eq!(nemo_relay_get_handle(&mut parent), NemoRelayStatus::Ok); + assert_status!(nemo_relay_get_handle(&mut parent), NemoRelayStatus::Ok); let name = cstring("ffi_llm_execute_extra"); let request = cstring( @@ -2442,7 +2442,7 @@ fn test_ffi_llm_execute_and_stream_additional_input_paths() { let mut stream = ptr::null_mut(); let mut chunk = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_llm_call_execute( invalid_name, request.as_ptr(), @@ -2463,7 +2463,7 @@ fn test_ffi_llm_execute_and_stream_additional_input_paths() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_llm_call_execute( name.as_ptr(), invalid_json.as_ptr(), @@ -2484,7 +2484,7 @@ fn test_ffi_llm_execute_and_stream_additional_input_paths() { ), NemoRelayStatus::InvalidJson ); - assert_eq!( + assert_status!( nemo_relay_llm_call_execute( name.as_ptr(), request.as_ptr(), @@ -2505,7 +2505,7 @@ fn test_ffi_llm_execute_and_stream_additional_input_paths() { ), NemoRelayStatus::InvalidJson ); - assert_eq!( + assert_status!( nemo_relay_llm_call_execute( name.as_ptr(), request.as_ptr(), @@ -2530,7 +2530,7 @@ fn test_ffi_llm_execute_and_stream_additional_input_paths() { assert_eq!(decoded["id"], json!("chatcmpl-ffi")); assert_eq!(decoded["model"], json!("codec-model")); - assert_eq!( + assert_status!( nemo_relay_llm_stream_call_execute( invalid_name, request.as_ptr(), @@ -2553,7 +2553,7 @@ fn test_ffi_llm_execute_and_stream_additional_input_paths() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_llm_stream_call_execute( name.as_ptr(), invalid_json.as_ptr(), @@ -2576,7 +2576,7 @@ fn test_ffi_llm_execute_and_stream_additional_input_paths() { ), NemoRelayStatus::InvalidJson ); - assert_eq!( + assert_status!( nemo_relay_llm_stream_call_execute( name.as_ptr(), request.as_ptr(), @@ -2599,7 +2599,7 @@ fn test_ffi_llm_execute_and_stream_additional_input_paths() { ), NemoRelayStatus::InvalidJson ); - assert_eq!( + assert_status!( nemo_relay_llm_stream_call_execute( name.as_ptr(), request.as_ptr(), @@ -2622,7 +2622,7 @@ fn test_ffi_llm_execute_and_stream_additional_input_paths() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_llm_stream_call_execute( name.as_ptr(), request.as_ptr(), @@ -2681,53 +2681,53 @@ fn test_ffi_adaptive_runtime_and_cache_helper_paths() { let invalid_json = cstring("{"); let mut out_json = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_adaptive_validate_config(ptr::null(), &mut out_json), NemoRelayStatus::InvalidJson ); - assert_eq!( + assert_status!( nemo_relay_adaptive_validate_config(config.as_ptr(), ptr::null_mut()), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_adaptive_validate_config(invalid_json.as_ptr(), &mut out_json), NemoRelayStatus::InvalidJson ); - assert_eq!( + assert_status!( nemo_relay_adaptive_validate_config(config.as_ptr(), &mut out_json), NemoRelayStatus::Ok ); assert_eq!(returned_json(out_json)["diagnostics"], json!([])); let mut runtime = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_adaptive_runtime_create(config.as_ptr(), ptr::null_mut()), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_adaptive_runtime_create(invalid_json.as_ptr(), &mut runtime), NemoRelayStatus::InvalidJson ); - assert_eq!( + assert_status!( nemo_relay_adaptive_runtime_create(config.as_ptr(), &mut runtime), NemoRelayStatus::Ok ); assert!(!runtime.is_null()); - assert_eq!( + assert_status!( nemo_relay_adaptive_runtime_report_json(runtime, ptr::null_mut()), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_adaptive_runtime_report_json(runtime, &mut out_json), NemoRelayStatus::Ok ); assert_eq!(returned_json(out_json)["diagnostics"], json!([])); - assert_eq!( + assert_status!( nemo_relay_adaptive_runtime_register(runtime), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_adaptive_runtime_wait_for_idle(runtime), NemoRelayStatus::Ok ); @@ -2735,7 +2735,7 @@ fn test_ffi_adaptive_runtime_and_cache_helper_paths() { let stack = fresh_scope_stack(); let scope_name = cstring("ffi_adaptive_bound_scope"); let mut scope = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_push_scope( scope_name.as_ptr(), NemoRelayScopeType::Agent, @@ -2748,11 +2748,11 @@ fn test_ffi_adaptive_runtime_and_cache_helper_paths() { ), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_adaptive_runtime_bind_scope(runtime, ptr::null()), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_adaptive_runtime_bind_scope(runtime, scope), NemoRelayStatus::Ok ); @@ -2775,7 +2775,7 @@ fn test_ffi_adaptive_runtime_and_cache_helper_paths() { }) .to_string(), ); - assert_eq!( + assert_status!( nemo_relay_adaptive_runtime_build_cache_request_facts( runtime, cache_facts_options.as_ptr(), @@ -2827,7 +2827,7 @@ fn test_ffi_adaptive_runtime_and_cache_helper_paths() { expected ); } - assert_eq!( + assert_status!( nemo_relay_adaptive_runtime_build_cache_request_facts( runtime, invalid_json.as_ptr(), @@ -2835,7 +2835,7 @@ fn test_ffi_adaptive_runtime_and_cache_helper_paths() { ), NemoRelayStatus::InvalidJson ); - assert_eq!( + assert_status!( nemo_relay_adaptive_runtime_build_cache_request_facts( runtime, cache_facts_options.as_ptr(), @@ -2863,7 +2863,7 @@ fn test_ffi_adaptive_runtime_and_cache_helper_paths() { }) .to_string(), ); - assert_eq!( + assert_status!( nemo_relay_adaptive_build_cache_telemetry_event( telemetry_options.as_ptr(), &mut out_json, @@ -2886,7 +2886,7 @@ fn test_ffi_adaptive_runtime_and_cache_helper_paths() { }) .to_string(), ); - assert_eq!( + assert_status!( nemo_relay_adaptive_build_cache_telemetry_event( no_usage_options.as_ptr(), &mut out_json, @@ -2970,40 +2970,40 @@ fn test_ffi_adaptive_runtime_and_cache_helper_paths() { expected ); } - assert_eq!( + assert_status!( nemo_relay_adaptive_build_cache_telemetry_event( telemetry_options.as_ptr(), ptr::null_mut(), ), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_adaptive_build_cache_telemetry_event(invalid_json.as_ptr(), &mut out_json), NemoRelayStatus::InvalidJson ); - assert_eq!( + assert_status!( nemo_relay_adaptive_set_latency_sensitivity(0), NemoRelayStatus::InvalidArg ); - assert_eq!( + assert_status!( nemo_relay_adaptive_runtime_deregister(runtime), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_adaptive_runtime_shutdown(runtime), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_adaptive_runtime_shutdown(runtime), NemoRelayStatus::InvalidArg ); - assert_eq!( + assert_status!( nemo_relay_adaptive_runtime_register(runtime), NemoRelayStatus::InvalidArg ); - assert_eq!( + assert_status!( nemo_relay_pop_scope(scope, ptr::null()), NemoRelayStatus::Ok ); @@ -3020,11 +3020,11 @@ fn test_ffi_scope_stack_propagation_and_thread_binding_entry_points() { unsafe { let mut context_json = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_capture_propagation_context_json(ptr::null_mut()), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_capture_propagation_context_json(&mut context_json), NemoRelayStatus::Ok ); @@ -3032,7 +3032,7 @@ fn test_ffi_scope_stack_propagation_and_thread_binding_entry_points() { assert_eq!(context["version"], json!(1)); let root_uuid = cstring("018f13f0-7c1a-7a80-8000-000000000701"); - assert_eq!( + assert_status!( nemo_relay_capture_propagation_context_with_root_json( root_uuid.as_ptr(), ptr::null_mut(), @@ -3040,14 +3040,14 @@ fn test_ffi_scope_stack_propagation_and_thread_binding_entry_points() { NemoRelayStatus::NullPointer ); let invalid_root = cstring("not-a-uuid"); - assert_eq!( + assert_status!( nemo_relay_capture_propagation_context_with_root_json( invalid_root.as_ptr(), &mut context_json, ), NemoRelayStatus::InvalidArg ); - assert_eq!( + assert_status!( nemo_relay_capture_propagation_context_with_root_json( root_uuid.as_ptr(), &mut context_json, @@ -3058,7 +3058,7 @@ fn test_ffi_scope_stack_propagation_and_thread_binding_entry_points() { returned_json(context_json)["root_uuid"], root_uuid.to_str().unwrap() ); - assert_eq!( + assert_status!( nemo_relay_capture_propagation_context_with_root_json(ptr::null(), &mut context_json), NemoRelayStatus::Ok ); @@ -3067,49 +3067,49 @@ fn test_ffi_scope_stack_propagation_and_thread_binding_entry_points() { let payload = cstring(&context.to_string()); let invalid_json = cstring("{"); let mut stack = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_scope_stack_create_from_propagation_json(payload.as_ptr(), ptr::null_mut()), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_scope_stack_create_from_propagation_json(ptr::null(), &mut stack), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_scope_stack_create_from_propagation_json(invalid_json.as_ptr(), &mut stack), NemoRelayStatus::InvalidJson ); - assert_eq!( + assert_status!( nemo_relay_scope_stack_create_from_propagation_json(payload.as_ptr(), &mut stack), NemoRelayStatus::Ok ); assert!(!stack.is_null()); - assert_eq!( + assert_status!( nemo_relay_scope_stack_set_thread(ptr::null()), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_scope_stack_set_thread(stack), NemoRelayStatus::Ok ); assert!(nemo_relay_scope_stack_active()); let mut binding = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_scope_stack_capture_thread(ptr::null_mut()), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_scope_stack_capture_thread(&mut binding), NemoRelayStatus::Ok ); assert!(!binding.is_null()); - assert_eq!( + assert_status!( nemo_relay_scope_stack_restore_thread(ptr::null_mut()), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_scope_stack_restore_thread(binding), NemoRelayStatus::Ok ); @@ -3124,16 +3124,16 @@ fn test_ffi_observability_exporter_error_lifecycles() { unsafe { let mut out_json = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_observability_default_config_json(ptr::null_mut()), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_observability_component_spec_json(ptr::null(), true, ptr::null_mut()), NemoRelayStatus::NullPointer ); let invalid_config = cstring("{"); - assert_eq!( + assert_status!( nemo_relay_observability_component_spec_json( invalid_config.as_ptr(), true, @@ -3151,7 +3151,7 @@ fn test_ffi_observability_exporter_error_lifecycles() { let mode = cstring("overwrite"); let filename = cstring("events.jsonl"); let mut exporter = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_atof_exporter_create( directory.as_ptr(), mode.as_ptr(), @@ -3160,36 +3160,36 @@ fn test_ffi_observability_exporter_error_lifecycles() { ), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_atof_exporter_register(ptr::null(), filename.as_ptr()), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_atof_exporter_path(exporter, ptr::null_mut()), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_atof_exporter_force_flush(ptr::null()), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_atof_exporter_shutdown(ptr::null()), NemoRelayStatus::NullPointer ); let subscriber_name = cstring(&unique_name("ffi_atof_subscriber")); - assert_eq!( + assert_status!( nemo_relay_atof_exporter_register(exporter, subscriber_name.as_ptr()), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_atof_exporter_register(exporter, subscriber_name.as_ptr()), NemoRelayStatus::AlreadyExists ); - assert_eq!( + assert_status!( nemo_relay_atof_exporter_deregister(subscriber_name.as_ptr()), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_atof_exporter_deregister(subscriber_name.as_ptr()), NemoRelayStatus::Ok ); @@ -3198,7 +3198,7 @@ fn test_ffi_observability_exporter_error_lifecycles() { let otel_type = cstring("full"); let endpoint = cstring("http://127.0.0.1:4318/v1/traces"); let mut subscriber = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_create( otel_type.as_ptr(), ptr::null(), @@ -3214,20 +3214,20 @@ fn test_ffi_observability_exporter_error_lifecycles() { ), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_shutdown(subscriber), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_force_flush(subscriber), NemoRelayStatus::Internal ); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_shutdown(subscriber), NemoRelayStatus::Internal ); let missing_name = cstring(&unique_name("missing_otel_subscriber")); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_deregister(missing_name.as_ptr()), NemoRelayStatus::Ok ); @@ -3254,15 +3254,15 @@ fn test_ffi_observability_component_and_constructor_error_paths() { .to_string(), ); - assert_eq!( + assert_status!( nemo_relay_observability_default_config_json(ptr::null_mut()), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_observability_component_spec_json(ptr::null(), true, ptr::null_mut(),), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_observability_component_spec_json( invalid_json.as_ptr(), true, @@ -3270,7 +3270,7 @@ fn test_ffi_observability_component_and_constructor_error_paths() { ), NemoRelayStatus::InvalidJson ); - assert_eq!( + assert_status!( nemo_relay_observability_component_spec_json( invalid_config.as_ptr(), true, @@ -3278,7 +3278,7 @@ fn test_ffi_observability_component_and_constructor_error_paths() { ), NemoRelayStatus::InvalidJson ); - assert_eq!( + assert_status!( nemo_relay_observability_component_spec_json( explicit_config.as_ptr(), false, @@ -3297,7 +3297,7 @@ fn test_ffi_observability_component_and_constructor_error_paths() { let bad_mode = cstring("bad-mode"); let mut atof = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_atof_exporter_create( ptr::null(), append.as_ptr(), @@ -3306,7 +3306,7 @@ fn test_ffi_observability_component_and_constructor_error_paths() { ), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_atof_exporter_create( ptr::null(), bad_mode.as_ptr(), @@ -3315,15 +3315,15 @@ fn test_ffi_observability_component_and_constructor_error_paths() { ), NemoRelayStatus::InvalidArg ); - assert_eq!( + assert_status!( nemo_relay_atof_exporter_create(invalid, append.as_ptr(), filename.as_ptr(), &mut atof,), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_atof_exporter_create(ptr::null(), invalid, filename.as_ptr(), &mut atof,), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_atof_exporter_create(ptr::null(), append.as_ptr(), invalid, &mut atof,), NemoRelayStatus::InvalidUtf8 ); @@ -3331,7 +3331,7 @@ fn test_ffi_observability_component_and_constructor_error_paths() { let invalid_map_shape = cstring(r#"["not-an-object"]"#); let endpoint = cstring("http://localhost:4318/v1/traces"); let mut otel = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_create( c"full".as_ptr(), ptr::null(), @@ -3348,7 +3348,7 @@ fn test_ffi_observability_component_and_constructor_error_paths() { NemoRelayStatus::InvalidArg ); let mut openinference = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_create( c"openinference".as_ptr(), ptr::null(), diff --git a/crates/ffi/tests/unit/api/execution_tests.rs b/crates/ffi/tests/unit/api/execution_tests.rs index 0d158544d..83f45c1c1 100644 --- a/crates/ffi/tests/unit/api/execution_tests.rs +++ b/crates/ffi/tests/unit/api/execution_tests.rs @@ -13,7 +13,7 @@ fn test_ffi_tool_execute_parent_data_and_error_paths() { unsafe { let stack = fresh_scope_stack(); let mut parent = ptr::null_mut(); - assert_eq!(nemo_relay_get_handle(&mut parent), NemoRelayStatus::Ok); + assert_status!(nemo_relay_get_handle(&mut parent), NemoRelayStatus::Ok); let name = cstring("ffi_tool_execute_parent"); let args = cstring(r#"{"value":2}"#); @@ -22,7 +22,7 @@ fn test_ffi_tool_execute_parent_data_and_error_paths() { let invalid_json = cstring("{"); let mut out_json = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_tool_call_execute( name.as_ptr(), args.as_ptr(), @@ -40,7 +40,7 @@ fn test_ffi_tool_execute_parent_data_and_error_paths() { let executed = returned_json(out_json); assert_eq!(executed["executed"], json!(true)); - assert_eq!( + assert_status!( nemo_relay_tool_call_execute( name.as_ptr(), args.as_ptr(), @@ -56,7 +56,7 @@ fn test_ffi_tool_execute_parent_data_and_error_paths() { NemoRelayStatus::InvalidJson ); - assert_eq!( + assert_status!( nemo_relay_tool_call_execute( name.as_ptr(), args.as_ptr(), @@ -90,7 +90,7 @@ fn test_ffi_llm_execute_codec_parent_and_error_paths() { unsafe { let stack = fresh_scope_stack(); let mut parent = ptr::null_mut(); - assert_eq!(nemo_relay_get_handle(&mut parent), NemoRelayStatus::Ok); + assert_status!(nemo_relay_get_handle(&mut parent), NemoRelayStatus::Ok); let name = cstring("ffi_llm_execute_codec"); let request = cstring( @@ -103,7 +103,7 @@ fn test_ffi_llm_execute_codec_parent_and_error_paths() { let invalid_utf8 = [0xffu8, 0]; let mut out_json = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_llm_call_execute( name.as_ptr(), request.as_ptr(), @@ -128,7 +128,7 @@ fn test_ffi_llm_execute_codec_parent_and_error_paths() { assert_eq!(executed["model_seen"], json!("codec-model")); assert_eq!(executed["content"], json!("hello from ffi")); - assert_eq!( + assert_status!( nemo_relay_llm_call_execute( name.as_ptr(), request.as_ptr(), @@ -149,7 +149,7 @@ fn test_ffi_llm_execute_codec_parent_and_error_paths() { ), NemoRelayStatus::InvalidJson ); - assert_eq!( + assert_status!( nemo_relay_llm_call_execute( name.as_ptr(), request.as_ptr(), @@ -170,7 +170,7 @@ fn test_ffi_llm_execute_codec_parent_and_error_paths() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_llm_call_execute( name.as_ptr(), request.as_ptr(), @@ -210,7 +210,7 @@ fn test_ffi_llm_stream_execute_response_codec_defaults_and_error_paths() { unsafe { let stack = fresh_scope_stack(); let mut parent = ptr::null_mut(); - assert_eq!(nemo_relay_get_handle(&mut parent), NemoRelayStatus::Ok); + assert_status!(nemo_relay_get_handle(&mut parent), NemoRelayStatus::Ok); let name = cstring("ffi_llm_stream_defaults"); let request = cstring( @@ -224,7 +224,7 @@ fn test_ffi_llm_stream_execute_response_codec_defaults_and_error_paths() { let mut stream = ptr::null_mut(); let mut chunk = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_llm_stream_call_execute( name.as_ptr(), request.as_ptr(), @@ -250,14 +250,14 @@ fn test_ffi_llm_stream_execute_response_codec_defaults_and_error_paths() { assert_eq!(nemo_relay_stream_next(stream, &mut chunk), 1); let stream_chunk = returned_json(chunk); assert_eq!(stream_chunk["id"], json!("chatcmpl-ffi")); - assert_eq!(nemo_relay_stream_close(stream), NemoRelayStatus::Ok); - assert_eq!(nemo_relay_stream_close(stream), NemoRelayStatus::Ok); + assert_status!(nemo_relay_stream_close(stream), NemoRelayStatus::Ok); + assert_status!(nemo_relay_stream_close(stream), NemoRelayStatus::Ok); assert_eq!(nemo_relay_stream_next(stream, &mut chunk), 0); assert!(lock_unpoisoned(collected_chunks()).is_empty()); assert_eq!(*lock_unpoisoned(finalizer_calls()), 0); nemo_relay_stream_free(stream); - assert_eq!( + assert_status!( nemo_relay_llm_stream_call_execute( name.as_ptr(), request.as_ptr(), @@ -280,7 +280,7 @@ fn test_ffi_llm_stream_execute_response_codec_defaults_and_error_paths() { ), NemoRelayStatus::InvalidJson ); - assert_eq!( + assert_status!( nemo_relay_llm_stream_call_execute( name.as_ptr(), request.as_ptr(), @@ -321,11 +321,11 @@ fn test_ffi_registration_and_exporter_error_paths() { reset_globals(); unsafe { - assert_eq!( + assert_status!( nemo_relay_scope_stack_create(ptr::null_mut()), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_scope_stack_set_thread(ptr::null()), NemoRelayStatus::NullPointer ); @@ -333,7 +333,7 @@ fn test_ffi_registration_and_exporter_error_paths() { let stack = fresh_scope_stack(); let scope_name = cstring("ffi_scope_local"); let mut scope = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_push_scope( scope_name.as_ptr(), NemoRelayScopeType::Function, @@ -350,7 +350,7 @@ fn test_ffi_registration_and_exporter_error_paths() { let invalid_uuid = cstring("not-a-uuid"); let global_tool_san_req = cstring(&unique_name("ffi_tool_san_req")); - assert_eq!( + assert_status!( nemo_relay_register_tool_sanitize_request_guardrail( global_tool_san_req.as_ptr(), 1, @@ -360,7 +360,7 @@ fn test_ffi_registration_and_exporter_error_paths() { ), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_register_tool_sanitize_request_guardrail( global_tool_san_req.as_ptr(), 1, @@ -370,17 +370,17 @@ fn test_ffi_registration_and_exporter_error_paths() { ), NemoRelayStatus::AlreadyExists ); - assert_eq!( + assert_status!( nemo_relay_deregister_tool_sanitize_request_guardrail(global_tool_san_req.as_ptr()), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_deregister_tool_sanitize_request_guardrail(global_tool_san_req.as_ptr()), NemoRelayStatus::Ok ); let global_tool_san_resp = cstring(&unique_name("ffi_tool_san_resp")); - assert_eq!( + assert_status!( nemo_relay_register_tool_sanitize_response_guardrail( global_tool_san_resp.as_ptr(), 1, @@ -390,13 +390,13 @@ fn test_ffi_registration_and_exporter_error_paths() { ), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_deregister_tool_sanitize_response_guardrail(global_tool_san_resp.as_ptr()), NemoRelayStatus::Ok ); let global_tool_exec = cstring(&unique_name("ffi_tool_exec")); - assert_eq!( + assert_status!( nemo_relay_register_tool_execution_intercept( global_tool_exec.as_ptr(), 1, @@ -406,13 +406,13 @@ fn test_ffi_registration_and_exporter_error_paths() { ), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_deregister_tool_execution_intercept(global_tool_exec.as_ptr()), NemoRelayStatus::Ok ); let global_llm_san_req = cstring(&unique_name("ffi_llm_san_req")); - assert_eq!( + assert_status!( nemo_relay_register_llm_sanitize_request_guardrail( global_llm_san_req.as_ptr(), 1, @@ -422,13 +422,13 @@ fn test_ffi_registration_and_exporter_error_paths() { ), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_deregister_llm_sanitize_request_guardrail(global_llm_san_req.as_ptr()), NemoRelayStatus::Ok ); let global_llm_exec = cstring(&unique_name("ffi_llm_exec")); - assert_eq!( + assert_status!( nemo_relay_register_llm_execution_intercept( global_llm_exec.as_ptr(), 1, @@ -438,13 +438,13 @@ fn test_ffi_registration_and_exporter_error_paths() { ), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_deregister_llm_execution_intercept(global_llm_exec.as_ptr()), NemoRelayStatus::Ok ); let global_llm_stream_exec = cstring(&unique_name("ffi_llm_stream_exec")); - assert_eq!( + assert_status!( nemo_relay_register_llm_stream_execution_intercept( global_llm_stream_exec.as_ptr(), 1, @@ -454,13 +454,13 @@ fn test_ffi_registration_and_exporter_error_paths() { ), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_deregister_llm_stream_execution_intercept(global_llm_stream_exec.as_ptr()), NemoRelayStatus::Ok ); let scope_tool_san_req = cstring(&unique_name("scope_tool_san_req")); - assert_eq!( + assert_status!( nemo_relay_scope_register_tool_sanitize_request_guardrail( invalid_uuid.as_ptr(), scope_tool_san_req.as_ptr(), @@ -471,7 +471,7 @@ fn test_ffi_registration_and_exporter_error_paths() { ), NemoRelayStatus::InvalidArg ); - assert_eq!( + assert_status!( nemo_relay_scope_register_tool_sanitize_request_guardrail( scope_uuid.as_ptr(), scope_tool_san_req.as_ptr(), @@ -482,7 +482,7 @@ fn test_ffi_registration_and_exporter_error_paths() { ), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_tool_sanitize_request_guardrail( scope_uuid.as_ptr(), scope_tool_san_req.as_ptr(), @@ -491,7 +491,7 @@ fn test_ffi_registration_and_exporter_error_paths() { ); let scope_tool_san_resp = cstring(&unique_name("scope_tool_san_resp")); - assert_eq!( + assert_status!( nemo_relay_scope_register_tool_sanitize_response_guardrail( scope_uuid.as_ptr(), scope_tool_san_resp.as_ptr(), @@ -502,7 +502,7 @@ fn test_ffi_registration_and_exporter_error_paths() { ), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_tool_sanitize_response_guardrail( scope_uuid.as_ptr(), scope_tool_san_resp.as_ptr(), @@ -511,7 +511,7 @@ fn test_ffi_registration_and_exporter_error_paths() { ); let scope_tool_cond = cstring(&unique_name("scope_tool_cond")); - assert_eq!( + assert_status!( nemo_relay_scope_register_tool_conditional_execution_guardrail( scope_uuid.as_ptr(), scope_tool_cond.as_ptr(), @@ -522,7 +522,7 @@ fn test_ffi_registration_and_exporter_error_paths() { ), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_tool_conditional_execution_guardrail( scope_uuid.as_ptr(), scope_tool_cond.as_ptr(), @@ -531,7 +531,7 @@ fn test_ffi_registration_and_exporter_error_paths() { ); let scope_tool_req = cstring(&unique_name("scope_tool_req")); - assert_eq!( + assert_status!( nemo_relay_scope_register_tool_request_intercept( scope_uuid.as_ptr(), scope_tool_req.as_ptr(), @@ -543,7 +543,7 @@ fn test_ffi_registration_and_exporter_error_paths() { ), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_tool_request_intercept( scope_uuid.as_ptr(), scope_tool_req.as_ptr(), @@ -552,7 +552,7 @@ fn test_ffi_registration_and_exporter_error_paths() { ); let scope_tool_exec = cstring(&unique_name("scope_tool_exec")); - assert_eq!( + assert_status!( nemo_relay_scope_register_tool_execution_intercept( scope_uuid.as_ptr(), scope_tool_exec.as_ptr(), @@ -563,7 +563,7 @@ fn test_ffi_registration_and_exporter_error_paths() { ), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_tool_execution_intercept( scope_uuid.as_ptr(), scope_tool_exec.as_ptr(), @@ -572,7 +572,7 @@ fn test_ffi_registration_and_exporter_error_paths() { ); let scope_llm_san_req = cstring(&unique_name("scope_llm_san_req")); - assert_eq!( + assert_status!( nemo_relay_scope_register_llm_sanitize_request_guardrail( scope_uuid.as_ptr(), scope_llm_san_req.as_ptr(), @@ -583,7 +583,7 @@ fn test_ffi_registration_and_exporter_error_paths() { ), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_llm_sanitize_request_guardrail( scope_uuid.as_ptr(), scope_llm_san_req.as_ptr(), @@ -592,7 +592,7 @@ fn test_ffi_registration_and_exporter_error_paths() { ); let scope_llm_san_resp = cstring(&unique_name("scope_llm_san_resp")); - assert_eq!( + assert_status!( nemo_relay_scope_register_llm_sanitize_response_guardrail( scope_uuid.as_ptr(), scope_llm_san_resp.as_ptr(), @@ -603,7 +603,7 @@ fn test_ffi_registration_and_exporter_error_paths() { ), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_llm_sanitize_response_guardrail( scope_uuid.as_ptr(), scope_llm_san_resp.as_ptr(), @@ -612,7 +612,7 @@ fn test_ffi_registration_and_exporter_error_paths() { ); let scope_llm_cond = cstring(&unique_name("scope_llm_cond")); - assert_eq!( + assert_status!( nemo_relay_scope_register_llm_conditional_execution_guardrail( scope_uuid.as_ptr(), scope_llm_cond.as_ptr(), @@ -623,7 +623,7 @@ fn test_ffi_registration_and_exporter_error_paths() { ), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_llm_conditional_execution_guardrail( scope_uuid.as_ptr(), scope_llm_cond.as_ptr(), @@ -632,7 +632,7 @@ fn test_ffi_registration_and_exporter_error_paths() { ); let scope_llm_req = cstring(&unique_name("scope_llm_req")); - assert_eq!( + assert_status!( nemo_relay_scope_register_llm_request_intercept( scope_uuid.as_ptr(), scope_llm_req.as_ptr(), @@ -644,7 +644,7 @@ fn test_ffi_registration_and_exporter_error_paths() { ), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_llm_request_intercept( scope_uuid.as_ptr(), scope_llm_req.as_ptr(), @@ -653,7 +653,7 @@ fn test_ffi_registration_and_exporter_error_paths() { ); let scope_llm_exec = cstring(&unique_name("scope_llm_exec")); - assert_eq!( + assert_status!( nemo_relay_scope_register_llm_execution_intercept( scope_uuid.as_ptr(), scope_llm_exec.as_ptr(), @@ -664,7 +664,7 @@ fn test_ffi_registration_and_exporter_error_paths() { ), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_llm_execution_intercept( scope_uuid.as_ptr(), scope_llm_exec.as_ptr(), @@ -673,7 +673,7 @@ fn test_ffi_registration_and_exporter_error_paths() { ); let scope_llm_stream_exec = cstring(&unique_name("scope_llm_stream_exec")); - assert_eq!( + assert_status!( nemo_relay_scope_register_llm_stream_execution_intercept( scope_uuid.as_ptr(), scope_llm_stream_exec.as_ptr(), @@ -684,7 +684,7 @@ fn test_ffi_registration_and_exporter_error_paths() { ), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_llm_stream_execution_intercept( scope_uuid.as_ptr(), scope_llm_stream_exec.as_ptr(), @@ -693,7 +693,7 @@ fn test_ffi_registration_and_exporter_error_paths() { ); let scope_subscriber = cstring(&unique_name("scope_subscriber")); - assert_eq!( + assert_status!( nemo_relay_scope_register_subscriber( scope_uuid.as_ptr(), scope_subscriber.as_ptr(), @@ -703,11 +703,11 @@ fn test_ffi_registration_and_exporter_error_paths() { ), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_subscriber(scope_uuid.as_ptr(), scope_subscriber.as_ptr(),), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_subscriber(scope_uuid.as_ptr(), scope_subscriber.as_ptr(),), NemoRelayStatus::Ok ); @@ -716,7 +716,7 @@ fn test_ffi_registration_and_exporter_error_paths() { let session = cstring("ffi-session"); let agent = cstring("ffi-agent"); let version = cstring("1.0.0"); - assert_eq!( + assert_status!( nemo_relay_atif_exporter_create( session.as_ptr(), agent.as_ptr(), @@ -726,7 +726,7 @@ fn test_ffi_registration_and_exporter_error_paths() { ), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_atif_exporter_create( session.as_ptr(), agent.as_ptr(), @@ -736,38 +736,38 @@ fn test_ffi_registration_and_exporter_error_paths() { ), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_atif_exporter_register(ptr::null(), scope_subscriber.as_ptr()), NemoRelayStatus::NullPointer ); let mut null_export = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_atif_exporter_export(ptr::null(), &mut null_export), NemoRelayStatus::NullPointer ); let exporter_name = cstring(&unique_name("ffi_exporter_sub")); - assert_eq!( + assert_status!( nemo_relay_atif_exporter_register(exporter, exporter_name.as_ptr()), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_atif_exporter_register(exporter, exporter_name.as_ptr()), NemoRelayStatus::AlreadyExists ); - assert_eq!( + assert_status!( nemo_relay_atif_exporter_export(exporter, ptr::null_mut()), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_atif_exporter_clear(ptr::null()), NemoRelayStatus::NullPointer ); let missing_exporter = cstring("missing_exporter"); - assert_eq!( + assert_status!( nemo_relay_atif_exporter_deregister(missing_exporter.as_ptr()), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_atif_exporter_deregister(exporter_name.as_ptr()), NemoRelayStatus::Ok ); @@ -777,7 +777,7 @@ fn test_ffi_registration_and_exporter_error_paths() { assert_eq!(nemo_relay_stream_next(ptr::null_mut(), &mut chunk), -1); assert_eq!(nemo_relay_stream_next(ptr::null_mut(), ptr::null_mut()), -1); - assert_eq!( + assert_status!( nemo_relay_pop_scope(scope, ptr::null()), NemoRelayStatus::Ok ); diff --git a/crates/ffi/tests/unit/api/plugin_tests.rs b/crates/ffi/tests/unit/api/plugin_tests.rs index f938a5d38..f038cdf28 100644 --- a/crates/ffi/tests/unit/api/plugin_tests.rs +++ b/crates/ffi/tests/unit/api/plugin_tests.rs @@ -17,7 +17,7 @@ fn test_ffi_dynamic_plugin_activation_rejects_empty_specs_without_outputs() { let mut activation = std::ptr::dangling_mut::(); let mut report_json = std::ptr::dangling_mut::(); unsafe { - assert_eq!( + assert_status!( nemo_relay_initialize_with_dynamic_plugins( config.as_ptr(), specs.as_ptr(), @@ -48,7 +48,7 @@ fn test_ffi_dynamic_plugin_activation_rejects_invalid_inputs_without_outputs() { let invalid = cstring("not-json"); unsafe { let mut report = std::ptr::dangling_mut::(); - assert_eq!( + assert_status!( nemo_relay_initialize_with_dynamic_plugins( config.as_ptr(), specs.as_ptr(), @@ -60,7 +60,7 @@ fn test_ffi_dynamic_plugin_activation_rejects_invalid_inputs_without_outputs() { assert!(report.is_null()); let mut activation = std::ptr::dangling_mut::(); - assert_eq!( + assert_status!( nemo_relay_initialize_with_dynamic_plugins( config.as_ptr(), specs.as_ptr(), @@ -77,7 +77,7 @@ fn test_ffi_dynamic_plugin_activation_rejects_invalid_inputs_without_outputs() { ] { let mut activation = std::ptr::dangling_mut::(); let mut report = std::ptr::dangling_mut::(); - assert_eq!( + assert_status!( nemo_relay_initialize_with_dynamic_plugins( config_json, specs_json, @@ -98,7 +98,7 @@ fn test_ffi_dynamic_plugin_activation_rejects_invalid_inputs_without_outputs() { let invalid_shape = cstring(r#"{"plugin_id":"not-an-array"}"#); let mut activation = ptr::null_mut(); let mut report = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_initialize_with_dynamic_plugins( config.as_ptr(), invalid_shape.as_ptr(), @@ -139,7 +139,7 @@ fn test_ffi_dynamic_plugin_activation_surfaces_load_failures_and_releases_owner( unsafe { let mut activation = ptr::null_mut(); let mut report = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_initialize_with_dynamic_plugins( config.as_ptr(), specs.as_ptr(), @@ -158,7 +158,7 @@ fn test_ffi_dynamic_plugin_activation_surfaces_load_failures_and_releases_owner( let mut retry_activation = ptr::null_mut(); let mut retry_report = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_initialize_with_dynamic_plugins( config.as_ptr(), specs.as_ptr(), @@ -194,7 +194,7 @@ fn test_ffi_plugin_registration_validation_and_cleanup() { let user_data = Box::into_raw(Box::new(7usize)) as *mut libc::c_void; unsafe { - assert_eq!( + assert_status!( nemo_relay_register_plugin( plugin_kind_c.as_ptr(), Some(plugin_validate_warn), @@ -206,7 +206,7 @@ fn test_ffi_plugin_registration_validation_and_cleanup() { ); let mut report_json = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_validate_plugin_config(config.as_ptr(), &mut report_json), NemoRelayStatus::Ok ); @@ -220,7 +220,7 @@ fn test_ffi_plugin_registration_validation_and_cleanup() { ); let mut init_json = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_initialize_plugins(config.as_ptr(), &mut init_json), NemoRelayStatus::Ok ); @@ -234,19 +234,19 @@ fn test_ffi_plugin_registration_validation_and_cleanup() { ); let mut active_json = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_active_plugin_report_json(&mut active_json), NemoRelayStatus::Ok ); let active = returned_json(active_json); assert_eq!(active["diagnostics"], initialized["diagnostics"]); - assert_eq!(nemo_relay_clear_plugin_configuration(), NemoRelayStatus::Ok); - assert_eq!( + assert_status!(nemo_relay_clear_plugin_configuration(), NemoRelayStatus::Ok); + assert_status!( nemo_relay_deregister_plugin(plugin_kind_c.as_ptr()), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_deregister_plugin(plugin_kind_c.as_ptr()), NemoRelayStatus::NotFound ); @@ -289,7 +289,7 @@ fn test_ffi_plugin_validation_failure_modes_are_reported() { let user_data = Box::into_raw(Box::new(9usize)) as *mut libc::c_void; unsafe { - assert_eq!( + assert_status!( nemo_relay_register_plugin( plugin_kind_c.as_ptr(), validate_cb, @@ -301,7 +301,7 @@ fn test_ffi_plugin_validation_failure_modes_are_reported() { ); let mut report_json = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_validate_plugin_config(config.as_ptr(), &mut report_json), NemoRelayStatus::Ok ); @@ -317,7 +317,7 @@ fn test_ffi_plugin_validation_failure_modes_are_reported() { "missing expected plugin validation diagnostic: {expected_fragment}" ); - assert_eq!( + assert_status!( nemo_relay_deregister_plugin(plugin_kind_c.as_ptr()), NemoRelayStatus::Ok ); @@ -349,7 +349,7 @@ fn test_ffi_plugin_without_validate_callback_uses_registration_fallback_error() let user_data = Box::into_raw(Box::new(11usize)) as *mut libc::c_void; unsafe { - assert_eq!( + assert_status!( nemo_relay_register_plugin( plugin_kind_c.as_ptr(), None, @@ -361,7 +361,7 @@ fn test_ffi_plugin_without_validate_callback_uses_registration_fallback_error() ); let mut report_json = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_validate_plugin_config(config.as_ptr(), &mut report_json), NemoRelayStatus::Ok ); @@ -369,7 +369,7 @@ fn test_ffi_plugin_without_validate_callback_uses_registration_fallback_error() assert_eq!(report["diagnostics"], json!([])); let mut init_json = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_initialize_plugins(config.as_ptr(), &mut init_json), NemoRelayStatus::Internal ); @@ -377,13 +377,13 @@ fn test_ffi_plugin_without_validate_callback_uses_registration_fallback_error() assert!(err.contains("register callback failed with status Internal")); let mut active_json = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_active_plugin_report_json(&mut active_json), NemoRelayStatus::Ok ); assert_eq!(returned_json(active_json), Json::Null); - assert_eq!( + assert_status!( nemo_relay_deregister_plugin(plugin_kind_c.as_ptr()), NemoRelayStatus::Ok ); @@ -414,7 +414,7 @@ fn test_ffi_plugin_registration_failure_prefers_last_error_message() { let user_data = Box::into_raw(Box::new(13usize)) as *mut libc::c_void; unsafe { - assert_eq!( + assert_status!( nemo_relay_register_plugin( plugin_kind_c.as_ptr(), None, @@ -426,7 +426,7 @@ fn test_ffi_plugin_registration_failure_prefers_last_error_message() { ); let mut init_json = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_initialize_plugins(config.as_ptr(), &mut init_json), NemoRelayStatus::Internal ); @@ -436,7 +436,7 @@ fn test_ffi_plugin_registration_failure_prefers_last_error_message() { .contains("plugin register callback set last error explicitly") ); - assert_eq!( + assert_status!( nemo_relay_deregister_plugin(plugin_kind_c.as_ptr()), NemoRelayStatus::Ok ); @@ -455,7 +455,7 @@ fn test_ffi_plugin_context_helpers_cover_null_and_success_paths() { let tool_name = cstring("tool"); unsafe { - assert_eq!( + assert_status!( nemo_relay_plugin_context_register_subscriber( ptr::null_mut(), name.as_ptr(), @@ -465,7 +465,7 @@ fn test_ffi_plugin_context_helpers_cover_null_and_success_paths() { ), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_plugin_context_register_tool_sanitize_request_guardrail( ptr::null_mut(), tool_name.as_ptr(), @@ -476,7 +476,7 @@ fn test_ffi_plugin_context_helpers_cover_null_and_success_paths() { ), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_plugin_context_register_tool_sanitize_response_guardrail( ptr::null_mut(), tool_name.as_ptr(), @@ -487,7 +487,7 @@ fn test_ffi_plugin_context_helpers_cover_null_and_success_paths() { ), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_plugin_context_register_tool_conditional_execution_guardrail( ptr::null_mut(), tool_name.as_ptr(), @@ -498,7 +498,7 @@ fn test_ffi_plugin_context_helpers_cover_null_and_success_paths() { ), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_plugin_context_register_llm_sanitize_request_guardrail( ptr::null_mut(), llm_name.as_ptr(), @@ -509,7 +509,7 @@ fn test_ffi_plugin_context_helpers_cover_null_and_success_paths() { ), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_plugin_context_register_llm_sanitize_response_guardrail( ptr::null_mut(), llm_name.as_ptr(), @@ -520,7 +520,7 @@ fn test_ffi_plugin_context_helpers_cover_null_and_success_paths() { ), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_plugin_context_register_llm_conditional_execution_guardrail( ptr::null_mut(), llm_name.as_ptr(), @@ -531,7 +531,7 @@ fn test_ffi_plugin_context_helpers_cover_null_and_success_paths() { ), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_plugin_context_register_llm_request_intercept( ptr::null_mut(), llm_name.as_ptr(), @@ -543,7 +543,7 @@ fn test_ffi_plugin_context_helpers_cover_null_and_success_paths() { ), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_plugin_context_register_tool_request_intercept( ptr::null_mut(), tool_name.as_ptr(), @@ -555,7 +555,7 @@ fn test_ffi_plugin_context_helpers_cover_null_and_success_paths() { ), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_plugin_context_register_llm_execution_intercept( ptr::null_mut(), llm_name.as_ptr(), @@ -566,7 +566,7 @@ fn test_ffi_plugin_context_helpers_cover_null_and_success_paths() { ), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_plugin_context_register_llm_stream_execution_intercept( ptr::null_mut(), llm_name.as_ptr(), @@ -577,7 +577,7 @@ fn test_ffi_plugin_context_helpers_cover_null_and_success_paths() { ), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_plugin_context_register_tool_execution_intercept( ptr::null_mut(), tool_name.as_ptr(), @@ -594,7 +594,7 @@ fn test_ffi_plugin_context_helpers_cover_null_and_success_paths() { let mut ctx = FfiPluginContext(&mut inner as *mut _); unsafe { - assert_eq!( + assert_status!( nemo_relay_plugin_context_register_subscriber( &mut ctx, name.as_ptr(), @@ -604,7 +604,7 @@ fn test_ffi_plugin_context_helpers_cover_null_and_success_paths() { ), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_plugin_context_register_tool_sanitize_request_guardrail( &mut ctx, tool_name.as_ptr(), @@ -615,7 +615,7 @@ fn test_ffi_plugin_context_helpers_cover_null_and_success_paths() { ), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_plugin_context_register_tool_sanitize_response_guardrail( &mut ctx, tool_name.as_ptr(), @@ -626,7 +626,7 @@ fn test_ffi_plugin_context_helpers_cover_null_and_success_paths() { ), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_plugin_context_register_tool_conditional_execution_guardrail( &mut ctx, tool_name.as_ptr(), @@ -637,7 +637,7 @@ fn test_ffi_plugin_context_helpers_cover_null_and_success_paths() { ), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_plugin_context_register_llm_sanitize_request_guardrail( &mut ctx, llm_name.as_ptr(), @@ -648,7 +648,7 @@ fn test_ffi_plugin_context_helpers_cover_null_and_success_paths() { ), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_plugin_context_register_llm_sanitize_response_guardrail( &mut ctx, llm_name.as_ptr(), @@ -659,7 +659,7 @@ fn test_ffi_plugin_context_helpers_cover_null_and_success_paths() { ), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_plugin_context_register_llm_conditional_execution_guardrail( &mut ctx, llm_name.as_ptr(), @@ -670,7 +670,7 @@ fn test_ffi_plugin_context_helpers_cover_null_and_success_paths() { ), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_plugin_context_register_llm_request_intercept( &mut ctx, llm_name.as_ptr(), @@ -682,7 +682,7 @@ fn test_ffi_plugin_context_helpers_cover_null_and_success_paths() { ), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_plugin_context_register_tool_request_intercept( &mut ctx, tool_name.as_ptr(), @@ -694,7 +694,7 @@ fn test_ffi_plugin_context_helpers_cover_null_and_success_paths() { ), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_plugin_context_register_llm_execution_intercept( &mut ctx, llm_name.as_ptr(), @@ -705,7 +705,7 @@ fn test_ffi_plugin_context_helpers_cover_null_and_success_paths() { ), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_plugin_context_register_llm_stream_execution_intercept( &mut ctx, llm_name.as_ptr(), @@ -716,7 +716,7 @@ fn test_ffi_plugin_context_helpers_cover_null_and_success_paths() { ), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_plugin_context_register_tool_execution_intercept( &mut ctx, tool_name.as_ptr(), @@ -768,7 +768,7 @@ fn test_ffi_plugin_context_helpers_reject_duplicate_names() { let mut ctx = FfiPluginContext(&mut inner as *mut _); unsafe { - assert_eq!( + assert_status!( nemo_relay_plugin_context_register_subscriber( &mut ctx, subscriber_name.as_ptr(), @@ -778,7 +778,7 @@ fn test_ffi_plugin_context_helpers_reject_duplicate_names() { ), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_plugin_context_register_subscriber( &mut ctx, subscriber_name.as_ptr(), @@ -794,7 +794,7 @@ fn test_ffi_plugin_context_helpers_reject_duplicate_names() { .contains("already exists") ); - assert_eq!( + assert_status!( nemo_relay_plugin_context_register_tool_sanitize_request_guardrail( &mut ctx, tool_name.as_ptr(), @@ -805,7 +805,7 @@ fn test_ffi_plugin_context_helpers_reject_duplicate_names() { ), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_plugin_context_register_tool_sanitize_request_guardrail( &mut ctx, tool_name.as_ptr(), @@ -822,7 +822,7 @@ fn test_ffi_plugin_context_helpers_reject_duplicate_names() { .contains("already exists") ); - assert_eq!( + assert_status!( nemo_relay_plugin_context_register_llm_request_intercept( &mut ctx, llm_name.as_ptr(), @@ -834,7 +834,7 @@ fn test_ffi_plugin_context_helpers_reject_duplicate_names() { ), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_plugin_context_register_llm_request_intercept( &mut ctx, llm_name.as_ptr(), @@ -852,7 +852,7 @@ fn test_ffi_plugin_context_helpers_reject_duplicate_names() { .contains("already exists") ); - assert_eq!( + assert_status!( nemo_relay_plugin_context_register_tool_execution_intercept( &mut ctx, tool_name.as_ptr(), @@ -863,7 +863,7 @@ fn test_ffi_plugin_context_helpers_reject_duplicate_names() { ), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_plugin_context_register_tool_execution_intercept( &mut ctx, tool_name.as_ptr(), @@ -895,7 +895,7 @@ fn test_ffi_plugin_context_helpers_reject_invalid_utf8_names_in_bulk() { macro_rules! assert_invalid_name_status { ($call:expr) => {{ - assert_eq!($call, NemoRelayStatus::InvalidUtf8); + assert_status!($call, NemoRelayStatus::InvalidUtf8); assert!(read_last_error().unwrap_or_default().contains("utf-8")); }}; } @@ -1028,7 +1028,7 @@ fn test_ffi_plugin_context_helpers_reject_duplicate_names_in_bulk() { macro_rules! assert_duplicate { ($call:expr) => {{ - assert_eq!($call, NemoRelayStatus::Internal); + assert_status!($call, NemoRelayStatus::Internal); assert!( read_last_error() .unwrap_or_default() @@ -1039,7 +1039,7 @@ fn test_ffi_plugin_context_helpers_reject_duplicate_names_in_bulk() { unsafe { let subscriber_name = cstring("duplicate-subscriber-bulk"); - assert_eq!( + assert_status!( nemo_relay_plugin_context_register_subscriber( &mut ctx, subscriber_name.as_ptr(), @@ -1058,7 +1058,7 @@ fn test_ffi_plugin_context_helpers_reject_duplicate_names_in_bulk() { )); let tool_sanitize_req = cstring("duplicate-tool-sanitize-req-bulk"); - assert_eq!( + assert_status!( nemo_relay_plugin_context_register_tool_sanitize_request_guardrail( &mut ctx, tool_sanitize_req.as_ptr(), @@ -1081,7 +1081,7 @@ fn test_ffi_plugin_context_helpers_reject_duplicate_names_in_bulk() { ); let tool_sanitize_resp = cstring("duplicate-tool-sanitize-resp-bulk"); - assert_eq!( + assert_status!( nemo_relay_plugin_context_register_tool_sanitize_response_guardrail( &mut ctx, tool_sanitize_resp.as_ptr(), @@ -1104,7 +1104,7 @@ fn test_ffi_plugin_context_helpers_reject_duplicate_names_in_bulk() { ); let tool_conditional = cstring("duplicate-tool-conditional-bulk"); - assert_eq!( + assert_status!( nemo_relay_plugin_context_register_tool_conditional_execution_guardrail( &mut ctx, tool_conditional.as_ptr(), @@ -1127,7 +1127,7 @@ fn test_ffi_plugin_context_helpers_reject_duplicate_names_in_bulk() { ); let llm_sanitize_req = cstring("duplicate-llm-sanitize-req-bulk"); - assert_eq!( + assert_status!( nemo_relay_plugin_context_register_llm_sanitize_request_guardrail( &mut ctx, llm_sanitize_req.as_ptr(), @@ -1150,7 +1150,7 @@ fn test_ffi_plugin_context_helpers_reject_duplicate_names_in_bulk() { ); let llm_sanitize_resp = cstring("duplicate-llm-sanitize-resp-bulk"); - assert_eq!( + assert_status!( nemo_relay_plugin_context_register_llm_sanitize_response_guardrail( &mut ctx, llm_sanitize_resp.as_ptr(), @@ -1173,7 +1173,7 @@ fn test_ffi_plugin_context_helpers_reject_duplicate_names_in_bulk() { ); let llm_conditional = cstring("duplicate-llm-conditional-bulk"); - assert_eq!( + assert_status!( nemo_relay_plugin_context_register_llm_conditional_execution_guardrail( &mut ctx, llm_conditional.as_ptr(), @@ -1196,7 +1196,7 @@ fn test_ffi_plugin_context_helpers_reject_duplicate_names_in_bulk() { ); let llm_request = cstring("duplicate-llm-request-bulk"); - assert_eq!( + assert_status!( nemo_relay_plugin_context_register_llm_request_intercept( &mut ctx, llm_request.as_ptr(), @@ -1219,7 +1219,7 @@ fn test_ffi_plugin_context_helpers_reject_duplicate_names_in_bulk() { )); let tool_request = cstring("duplicate-tool-request-bulk"); - assert_eq!( + assert_status!( nemo_relay_plugin_context_register_tool_request_intercept( &mut ctx, tool_request.as_ptr(), @@ -1242,7 +1242,7 @@ fn test_ffi_plugin_context_helpers_reject_duplicate_names_in_bulk() { )); let llm_exec = cstring("duplicate-llm-exec-bulk"); - assert_eq!( + assert_status!( nemo_relay_plugin_context_register_llm_execution_intercept( &mut ctx, llm_exec.as_ptr(), @@ -1263,7 +1263,7 @@ fn test_ffi_plugin_context_helpers_reject_duplicate_names_in_bulk() { )); let llm_stream_exec = cstring("duplicate-llm-stream-exec-bulk"); - assert_eq!( + assert_status!( nemo_relay_plugin_context_register_llm_stream_execution_intercept( &mut ctx, llm_stream_exec.as_ptr(), @@ -1286,7 +1286,7 @@ fn test_ffi_plugin_context_helpers_reject_duplicate_names_in_bulk() { ); let tool_exec = cstring("duplicate-tool-exec-bulk"); - assert_eq!( + assert_status!( nemo_relay_plugin_context_register_tool_execution_intercept( &mut ctx, tool_exec.as_ptr(), @@ -1323,7 +1323,7 @@ fn test_ffi_specialized_subscriber_and_exporter_default_and_invalid_name_paths() let endpoint = c"http://localhost:4318/v1/traces"; let mut otel_subscriber: *mut FfiOpenTelemetrySubscriber = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_create( c"full".as_ptr(), ptr::null(), @@ -1340,42 +1340,42 @@ fn test_ffi_specialized_subscriber_and_exporter_default_and_invalid_name_paths() NemoRelayStatus::Ok ); let otel_name = cstring(&unique_name("ffi_otel_defaults")); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_register(otel_subscriber, invalid_name), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_register(otel_subscriber, otel_name.as_ptr()), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_register(otel_subscriber, otel_name.as_ptr()), NemoRelayStatus::Internal ); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_deregister(ptr::null()), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_deregister(invalid_name), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_deregister(otel_name.as_ptr()), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_force_flush(otel_subscriber), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_shutdown(otel_subscriber), NemoRelayStatus::Ok ); nemo_relay_otel_subscriber_free(otel_subscriber); let mut oi_subscriber: *mut FfiOpenTelemetrySubscriber = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_create( c"openinference".as_ptr(), ptr::null(), @@ -1392,35 +1392,35 @@ fn test_ffi_specialized_subscriber_and_exporter_default_and_invalid_name_paths() NemoRelayStatus::Ok ); let oi_name = cstring(&unique_name("ffi_oi_defaults")); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_register(oi_subscriber, invalid_name), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_register(oi_subscriber, oi_name.as_ptr()), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_register(oi_subscriber, oi_name.as_ptr()), NemoRelayStatus::Internal ); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_deregister(ptr::null()), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_deregister(invalid_name), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_deregister(oi_name.as_ptr()), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_force_flush(oi_subscriber), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_shutdown(oi_subscriber), NemoRelayStatus::Ok ); @@ -1430,7 +1430,7 @@ fn test_ffi_specialized_subscriber_and_exporter_default_and_invalid_name_paths() let agent = cstring("specialized-agent"); let version = cstring("1.0.0"); let mut exporter = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_atif_exporter_create( session.as_ptr(), agent.as_ptr(), @@ -1441,27 +1441,27 @@ fn test_ffi_specialized_subscriber_and_exporter_default_and_invalid_name_paths() NemoRelayStatus::Ok ); let exporter_name = cstring(&unique_name("ffi_exporter_defaults")); - assert_eq!( + assert_status!( nemo_relay_atif_exporter_register(exporter, invalid_name), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_atif_exporter_register(exporter, exporter_name.as_ptr()), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_atif_exporter_register(exporter, exporter_name.as_ptr()), NemoRelayStatus::AlreadyExists ); - assert_eq!( + assert_status!( nemo_relay_atif_exporter_deregister(ptr::null()), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_atif_exporter_deregister(invalid_name), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_atif_exporter_deregister(exporter_name.as_ptr()), NemoRelayStatus::Ok ); @@ -1479,7 +1479,7 @@ fn test_ffi_otel_projection_options_accept_and_validate_legacy_controls() { let exclusions = c"[\"custom.mark\"]"; let mappings = c"[{\"key\":\"nemo_relay.model_name\",\"alias\":\"model.alias\"}]"; let mut subscriber: *mut FfiOpenTelemetrySubscriber = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_create_with_projection_options( c"full".as_ptr(), ptr::null(), @@ -1502,7 +1502,7 @@ fn test_ffi_otel_projection_options_accept_and_validate_legacy_controls() { let invalid_mappings = c"[{\"key\":\"\",\"alias\":\"model.alias\"}]"; let mut invalid_subscriber: *mut FfiOpenTelemetrySubscriber = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_create_with_projection_options( c"full".as_ptr(), ptr::null(), @@ -1544,7 +1544,7 @@ fn test_ffi_specialized_constructor_invalid_utf8_and_malformed_json_sweep() { let grpc = cstring("grpc"); let mut otel = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_create( c"full".as_ptr(), ptr::null(), @@ -1560,7 +1560,7 @@ fn test_ffi_specialized_constructor_invalid_utf8_and_malformed_json_sweep() { ), NemoRelayStatus::InvalidJson ); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_create( c"full".as_ptr(), ptr::null(), @@ -1626,7 +1626,7 @@ fn test_ffi_specialized_constructor_invalid_utf8_and_malformed_json_sweep() { invalid, ), ] { - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_create( c"full".as_ptr(), transport, @@ -1645,7 +1645,7 @@ fn test_ffi_specialized_constructor_invalid_utf8_and_malformed_json_sweep() { } let mut openinference = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_create( c"openinference".as_ptr(), ptr::null(), @@ -1661,7 +1661,7 @@ fn test_ffi_specialized_constructor_invalid_utf8_and_malformed_json_sweep() { ), NemoRelayStatus::InvalidJson ); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_create( c"openinference".as_ptr(), ptr::null(), @@ -1727,7 +1727,7 @@ fn test_ffi_specialized_constructor_invalid_utf8_and_malformed_json_sweep() { invalid, ), ] { - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_create( c"openinference".as_ptr(), transport, @@ -1765,7 +1765,7 @@ fn test_ffi_specialized_constructor_invalid_utf8_and_malformed_json_sweep() { invalid, ), ] { - assert_eq!( + assert_status!( nemo_relay_atif_exporter_create( session_ptr, agent_ptr, @@ -1778,7 +1778,7 @@ fn test_ffi_specialized_constructor_invalid_utf8_and_malformed_json_sweep() { } let plugin_kind = invalid; - assert_eq!( + assert_status!( nemo_relay_register_plugin( plugin_kind, None, @@ -1788,7 +1788,7 @@ fn test_ffi_specialized_constructor_invalid_utf8_and_malformed_json_sweep() { ), NemoRelayStatus::InvalidUtf8 ); - assert_eq!( + assert_status!( nemo_relay_deregister_plugin(plugin_kind), NemoRelayStatus::InvalidUtf8 ); diff --git a/crates/ffi/tests/unit/api/registry_tests.rs b/crates/ffi/tests/unit/api/registry_tests.rs index ae1bc656b..eefc5f9ef 100644 --- a/crates/ffi/tests/unit/api/registry_tests.rs +++ b/crates/ffi/tests/unit/api/registry_tests.rs @@ -6,7 +6,7 @@ use super::*; use nemo_relay::plugin::rollback_registrations; use std::io::{Read, Write}; -use std::net::TcpListener; +use std::net::{TcpListener, TcpStream}; use std::sync::mpsc::{self, Receiver}; use std::thread::{self, JoinHandle}; use std::time::{Duration, Instant}; @@ -26,61 +26,8 @@ fn start_otlp_http_collector() -> (String, Receiver>, JoinHandle<()>) { stream .set_read_timeout(Some(Duration::from_secs(1))) .unwrap(); - let mut request = Vec::new(); - let mut buffer = [0_u8; 4096]; - let (header_end, content_length) = loop { - let read = match stream.read(&mut buffer) { - Ok(read) => read, - Err(error) - if matches!( - error.kind(), - std::io::ErrorKind::WouldBlock | std::io::ErrorKind::TimedOut - ) && Instant::now() < deadline => - { - continue; - } - Err(error) => panic!("collector header read failed: {error}"), - }; - assert!( - read > 0, - "collector connection closed before request headers" - ); - request.extend_from_slice(&buffer[..read]); - if let Some(header_end) = - request.windows(4).position(|value| value == b"\r\n\r\n") - { - let header_end = header_end + 4; - let headers = String::from_utf8_lossy(&request[..header_end]); - let content_length = headers - .lines() - .find_map(|line| { - line.split_once(':').and_then(|(name, value)| { - name.eq_ignore_ascii_case("content-length") - .then(|| value.trim().parse::().unwrap()) - }) - }) - .expect("OTLP request must include content-length"); - break (header_end, content_length); - } - }; - while request.len() < header_end + content_length { - let read = match stream.read(&mut buffer) { - Ok(read) => read, - Err(error) - if matches!( - error.kind(), - std::io::ErrorKind::WouldBlock | std::io::ErrorKind::TimedOut - ) && Instant::now() < deadline => - { - continue; - } - Err(error) => panic!("collector body read failed: {error}"), - }; - assert!(read > 0, "collector connection closed before request body"); - request.extend_from_slice(&buffer[..read]); - } sender - .send(request[header_end..header_end + content_length].to_vec()) + .send(read_otlp_request(&mut stream, deadline)) .unwrap(); stream .write_all( @@ -103,6 +50,56 @@ fn start_otlp_http_collector() -> (String, Receiver>, JoinHandle<()>) { (format!("http://{address}/v1/traces"), receiver, handle) } +fn read_otlp_request(stream: &mut TcpStream, deadline: Instant) -> Vec { + let mut request = Vec::new(); + let mut buffer = [0_u8; 4096]; + let header_end = loop { + let read = read_collector_chunk(stream, &mut buffer, deadline, "headers"); + request.extend_from_slice(&buffer[..read]); + if let Some(position) = request.windows(4).position(|value| value == b"\r\n\r\n") { + break position + 4; + } + }; + let content_length = otlp_content_length(&request[..header_end]); + while request.len() < header_end + content_length { + let read = read_collector_chunk(stream, &mut buffer, deadline, "body"); + request.extend_from_slice(&buffer[..read]); + } + request[header_end..header_end + content_length].to_vec() +} + +fn read_collector_chunk( + stream: &mut TcpStream, + buffer: &mut [u8], + deadline: Instant, + phase: &str, +) -> usize { + loop { + match stream.read(buffer) { + Ok(0) => panic!("collector connection closed before request {phase}"), + Ok(read) => return read, + Err(error) + if matches!( + error.kind(), + std::io::ErrorKind::WouldBlock | std::io::ErrorKind::TimedOut + ) && Instant::now() < deadline => {} + Err(error) => panic!("collector {phase} read failed: {error}"), + } + } +} + +fn otlp_content_length(headers: &[u8]) -> usize { + String::from_utf8_lossy(headers) + .lines() + .find_map(|line| { + line.split_once(':').and_then(|(name, value)| { + name.eq_ignore_ascii_case("content-length") + .then(|| value.trim().parse::().unwrap()) + }) + }) + .expect("OTLP request must include content-length (chunked bodies are not supported)") +} + unsafe extern "C" fn event_sanitize_cb( _user_data: *mut libc::c_void, event: *const FfiEvent, @@ -137,7 +134,7 @@ fn test_ffi_event_sanitizer_registries_and_error_paths() { unsafe { let stack = fresh_scope_stack(); let subscriber_name = cstring(&unique_name("ffi_event_sanitize_subscriber")); - assert_eq!( + assert_status!( nemo_relay_register_subscriber( subscriber_name.as_ptr(), subscriber_cb, @@ -173,13 +170,13 @@ fn test_ffi_event_sanitizer_registries_and_error_paths() { Some(plugin_free), ), ] { - assert_eq!(status, NemoRelayStatus::Ok); + assert_status!(status, NemoRelayStatus::Ok); } let scope_name = cstring("ffi-global-scope"); let original = cstring(r#"{"secret":true}"#); let mut scope = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_push_scope( scope_name.as_ptr(), NemoRelayScopeType::Custom, @@ -193,7 +190,7 @@ fn test_ffi_event_sanitizer_registries_and_error_paths() { NemoRelayStatus::Ok ); let mark_name = cstring("ffi-global-mark"); - assert_eq!( + assert_status!( nemo_relay_event( mark_name.as_ptr(), scope, @@ -202,47 +199,33 @@ fn test_ffi_event_sanitizer_registries_and_error_paths() { ), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_pop_scope(scope, ptr::null()), NemoRelayStatus::Ok ); nemo_relay_scope_handle_free(scope); - assert_eq!(nemo_relay_flush_subscribers(), NemoRelayStatus::Ok); + assert_status!(nemo_relay_flush_subscribers(), NemoRelayStatus::Ok); let events = lock_unpoisoned(event_log()); - for name in ["ffi-global-scope", "ffi-global-mark"] { - for event in events.iter().filter(|event| event["name"] == name) { - assert_eq!(event["data"], json!({"sanitized_by": name})); - assert_eq!( - event["json"]["category_profile"]["subtype"], - "ffi.sanitized" - ); - assert_eq!(event["metadata"], Json::Null); - } - } - for phase in ["start", "end"] { - assert!(events.iter().any(|event| { - event["name"] == "ffi-global-scope" && event["json"]["scope_category"] == phase - })); - } + assert_sanitized_event_log(&events); drop(events); - assert_eq!( + assert_status!( nemo_relay_deregister_mark_sanitize_guardrail(mark_guard.as_ptr()), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_deregister_scope_sanitize_start_guardrail(start_guard.as_ptr()), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_deregister_scope_sanitize_end_guardrail(end_guard.as_ptr()), NemoRelayStatus::Ok ); - assert_eq!(*lock_unpoisoned(plugin_frees()), 3); + assert_plugin_free_count(3); let invalid_guard = cstring(&unique_name("ffi_invalid_event_sanitize")); - assert_eq!( + assert_status!( nemo_relay_register_mark_sanitize_guardrail( invalid_guard.as_ptr(), 1, @@ -253,7 +236,7 @@ fn test_ffi_event_sanitizer_registries_and_error_paths() { NemoRelayStatus::Ok ); let invalid_mark = cstring("ffi-invalid-callback-mark"); - assert_eq!( + assert_status!( nemo_relay_event( invalid_mark.as_ptr(), ptr::null(), @@ -262,17 +245,17 @@ fn test_ffi_event_sanitizer_registries_and_error_paths() { ), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_deregister_mark_sanitize_guardrail(invalid_guard.as_ptr()), NemoRelayStatus::Ok ); // The queued event retains its sanitizer snapshot after deregistration. - assert_eq!(nemo_relay_flush_subscribers(), NemoRelayStatus::Ok); - assert_eq!(*lock_unpoisoned(plugin_frees()), 4); + assert_status!(nemo_relay_flush_subscribers(), NemoRelayStatus::Ok); + assert_plugin_free_count(4); let mut owner = ptr::null_mut(); let owner_name = cstring("ffi-event-owner"); - assert_eq!( + assert_status!( nemo_relay_push_scope( owner_name.as_ptr(), NemoRelayScopeType::Agent, @@ -315,11 +298,11 @@ fn test_ffi_event_sanitizer_registries_and_error_paths() { Some(plugin_free), ), ] { - assert_eq!(status, NemoRelayStatus::Ok); + assert_status!(status, NemoRelayStatus::Ok); } let child_name = cstring("ffi-local-child"); let mut child = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_push_scope( child_name.as_ptr(), NemoRelayScopeType::Function, @@ -333,7 +316,7 @@ fn test_ffi_event_sanitizer_registries_and_error_paths() { NemoRelayStatus::Ok ); let local_mark_name = cstring("ffi-local-mark"); - assert_eq!( + assert_status!( nemo_relay_event( local_mark_name.as_ptr(), child, @@ -342,26 +325,26 @@ fn test_ffi_event_sanitizer_registries_and_error_paths() { ), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_pop_scope(child, ptr::null()), NemoRelayStatus::Ok ); nemo_relay_scope_handle_free(child); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_mark_sanitize_guardrail( owner_uuid.as_ptr(), local_mark.as_ptr(), ), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_scope_sanitize_start_guardrail( owner_uuid.as_ptr(), local_start.as_ptr(), ), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_scope_sanitize_end_guardrail( owner_uuid.as_ptr(), local_end.as_ptr(), @@ -369,12 +352,12 @@ fn test_ffi_event_sanitizer_registries_and_error_paths() { NemoRelayStatus::Ok ); // Scope removal does not alter the sanitizer snapshots already queued. - assert_eq!(nemo_relay_flush_subscribers(), NemoRelayStatus::Ok); - assert_eq!(*lock_unpoisoned(plugin_frees()), 7); + assert_status!(nemo_relay_flush_subscribers(), NemoRelayStatus::Ok); + assert_plugin_free_count(7); let invalid_uuid = cstring("not-a-uuid"); let invalid_name = cstring("invalid-scope-event-sanitizer"); - assert_eq!( + assert_status!( nemo_relay_scope_register_mark_sanitize_guardrail( invalid_uuid.as_ptr(), invalid_name.as_ptr(), @@ -385,14 +368,14 @@ fn test_ffi_event_sanitizer_registries_and_error_paths() { ), NemoRelayStatus::InvalidArg ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_scope_sanitize_start_guardrail( invalid_uuid.as_ptr(), invalid_name.as_ptr(), ), NemoRelayStatus::InvalidArg ); - assert_eq!( + assert_status!( nemo_relay_register_scope_sanitize_end_guardrail( ptr::null(), 1, @@ -402,11 +385,11 @@ fn test_ffi_event_sanitizer_registries_and_error_paths() { ), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_deregister_mark_sanitize_guardrail(ptr::null()), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_scope_register_scope_sanitize_end_guardrail( owner_uuid.as_ptr(), ptr::null(), @@ -417,7 +400,7 @@ fn test_ffi_event_sanitizer_registries_and_error_paths() { ), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_scope_sanitize_end_guardrail( owner_uuid.as_ptr(), ptr::null(), @@ -425,12 +408,12 @@ fn test_ffi_event_sanitizer_registries_and_error_paths() { NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_pop_scope(owner, ptr::null()), NemoRelayStatus::Ok ); nemo_relay_scope_handle_free(owner); - assert_eq!(nemo_relay_flush_subscribers(), NemoRelayStatus::Ok); + assert_status!(nemo_relay_flush_subscribers(), NemoRelayStatus::Ok); let events = lock_unpoisoned(event_log()); let invalid_callback_event = events .iter() @@ -497,9 +480,9 @@ fn test_ffi_event_sanitizer_registries_and_error_paths() { ), ] { assert!(!name.as_bytes().is_empty()); - assert_eq!(status, NemoRelayStatus::Ok); + assert_status!(status, NemoRelayStatus::Ok); } - assert_eq!( + assert_status!( nemo_relay_plugin_context_register_mark_sanitize_guardrail( ptr::null_mut(), invalid_name.as_ptr(), @@ -510,7 +493,7 @@ fn test_ffi_event_sanitizer_registries_and_error_paths() { ), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_plugin_context_register_mark_sanitize_guardrail( &mut ffi_context, ptr::null(), @@ -525,7 +508,7 @@ fn test_ffi_event_sanitizer_registries_and_error_paths() { rollback_registrations(&mut registrations); assert_eq!(*lock_unpoisoned(plugin_frees()), 10); - assert_eq!( + assert_status!( nemo_relay_deregister_subscriber(subscriber_name.as_ptr()), NemoRelayStatus::Ok ); @@ -533,6 +516,28 @@ fn test_ffi_event_sanitizer_registries_and_error_paths() { } } +fn assert_sanitized_event_log(events: &[Json]) { + for name in ["ffi-global-scope", "ffi-global-mark"] { + for event in events.iter().filter(|event| event["name"] == name) { + assert_eq!(event["data"], json!({"sanitized_by": name})); + assert_eq!( + event["json"]["category_profile"]["subtype"], + "ffi.sanitized" + ); + assert_eq!(event["metadata"], Json::Null); + } + } + for phase in ["start", "end"] { + assert!(events.iter().any(|event| { + event["name"] == "ffi-global-scope" && event["json"]["scope_category"] == phase + })); + } +} + +fn assert_plugin_free_count(expected: usize) { + assert_eq!(*lock_unpoisoned(plugin_frees()), expected); +} + #[test] fn test_ffi_open_telemetry_subscriber_lifecycle_and_errors() { let _lock = TEST_MUTEX.lock().unwrap_or_else(|e| e.into_inner()); @@ -552,7 +557,7 @@ fn test_ffi_open_telemetry_subscriber_lifecycle_and_errors() { let invalid_headers = cstring(r#"{"authorization":1}"#); let invalid_resource_attributes = cstring(r#"["not-an-object"]"#); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_create( c"full".as_ptr(), ptr::null(), @@ -568,7 +573,7 @@ fn test_ffi_open_telemetry_subscriber_lifecycle_and_errors() { ), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_create( c"full".as_ptr(), invalid_transport.as_ptr(), @@ -584,7 +589,7 @@ fn test_ffi_open_telemetry_subscriber_lifecycle_and_errors() { ), NemoRelayStatus::InvalidArg ); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_create( c"full".as_ptr(), ptr::null(), @@ -600,7 +605,7 @@ fn test_ffi_open_telemetry_subscriber_lifecycle_and_errors() { ), NemoRelayStatus::InvalidArg ); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_create( c"full".as_ptr(), ptr::null(), @@ -616,7 +621,7 @@ fn test_ffi_open_telemetry_subscriber_lifecycle_and_errors() { ), NemoRelayStatus::InvalidArg ); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_create( c"full".as_ptr(), grpc_transport.as_ptr(), @@ -635,7 +640,7 @@ fn test_ffi_open_telemetry_subscriber_lifecycle_and_errors() { assert!(!subscriber.is_null()); nemo_relay_otel_subscriber_free(subscriber); subscriber = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_create( c"full".as_ptr(), ptr::null(), @@ -654,36 +659,36 @@ fn test_ffi_open_telemetry_subscriber_lifecycle_and_errors() { assert!(!subscriber.is_null()); let name = cstring(&unique_name("ffi_otel")); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_register(ptr::null(), name.as_ptr()), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_force_flush(ptr::null()), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_shutdown(ptr::null()), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_register(subscriber, name.as_ptr()), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_deregister(name.as_ptr()), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_deregister(name.as_ptr()), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_force_flush(subscriber), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_shutdown(subscriber), NemoRelayStatus::Ok ); @@ -709,7 +714,7 @@ fn test_ffi_open_telemetry_typed_required_fields_and_gen_ai_wire_output() { (c"full".as_ptr(), ptr::null()), (c"full".as_ptr(), blank.as_ptr()), ] { - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_create( otel_type, ptr::null(), @@ -730,7 +735,7 @@ fn test_ffi_open_telemetry_typed_required_fields_and_gen_ai_wire_output() { let (collector_endpoint, request, collector) = start_otlp_http_collector(); let collector_endpoint = cstring(&collector_endpoint); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_create( c"gen_ai".as_ptr(), ptr::null(), @@ -747,7 +752,7 @@ fn test_ffi_open_telemetry_typed_required_fields_and_gen_ai_wire_output() { NemoRelayStatus::Ok ); let subscriber_name = cstring(&unique_name("ffi_gen_ai")); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_register(subscriber, subscriber_name.as_ptr()), NemoRelayStatus::Ok ); @@ -755,7 +760,7 @@ fn test_ffi_open_telemetry_typed_required_fields_and_gen_ai_wire_output() { let stack = fresh_scope_stack(); let scope_name = cstring("research-agent"); let mut scope = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_push_scope( scope_name.as_ptr(), NemoRelayScopeType::Agent, @@ -768,12 +773,12 @@ fn test_ffi_open_telemetry_typed_required_fields_and_gen_ai_wire_output() { ), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_pop_scope(scope, ptr::null()), NemoRelayStatus::Ok ); - assert_eq!(nemo_relay_flush_subscribers(), NemoRelayStatus::Ok); - assert_eq!( + assert_status!(nemo_relay_flush_subscribers(), NemoRelayStatus::Ok); + assert_status!( nemo_relay_otel_subscriber_force_flush(subscriber), NemoRelayStatus::Ok ); @@ -793,11 +798,11 @@ fn test_ffi_open_telemetry_typed_required_fields_and_gen_ai_wire_output() { .any(|value| value == b"nemo_relay.") ); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_deregister(subscriber_name.as_ptr()), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_shutdown(subscriber), NemoRelayStatus::Ok ); @@ -827,7 +832,7 @@ fn test_ffi_open_inference_subscriber_lifecycle_and_errors() { let invalid_headers = cstring(r#"{"authorization":1}"#); let invalid_resource_attributes = cstring(r#"["not-an-object"]"#); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_create( c"openinference".as_ptr(), ptr::null(), @@ -843,7 +848,7 @@ fn test_ffi_open_inference_subscriber_lifecycle_and_errors() { ), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_create( c"openinference".as_ptr(), invalid_transport.as_ptr(), @@ -859,7 +864,7 @@ fn test_ffi_open_inference_subscriber_lifecycle_and_errors() { ), NemoRelayStatus::InvalidArg ); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_create( c"openinference".as_ptr(), ptr::null(), @@ -875,7 +880,7 @@ fn test_ffi_open_inference_subscriber_lifecycle_and_errors() { ), NemoRelayStatus::InvalidArg ); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_create( c"openinference".as_ptr(), ptr::null(), @@ -891,7 +896,7 @@ fn test_ffi_open_inference_subscriber_lifecycle_and_errors() { ), NemoRelayStatus::InvalidArg ); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_create( c"openinference".as_ptr(), grpc_transport.as_ptr(), @@ -910,7 +915,7 @@ fn test_ffi_open_inference_subscriber_lifecycle_and_errors() { assert!(!subscriber.is_null()); nemo_relay_otel_subscriber_free(subscriber); subscriber = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_create( c"openinference".as_ptr(), ptr::null(), @@ -929,36 +934,36 @@ fn test_ffi_open_inference_subscriber_lifecycle_and_errors() { assert!(!subscriber.is_null()); let name = cstring(&unique_name("ffi_openinference")); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_register(ptr::null(), name.as_ptr()), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_force_flush(ptr::null()), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_shutdown(ptr::null()), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_register(subscriber, name.as_ptr()), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_deregister(name.as_ptr()), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_deregister(name.as_ptr()), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_force_flush(subscriber), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_otel_subscriber_shutdown(subscriber), NemoRelayStatus::Ok ); @@ -981,7 +986,7 @@ fn test_ffi_helper_rejection_and_null_name_paths() { let mut tool_out = ptr::null_mut(); let mut llm_error_out = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_tool_request_intercepts(tool_name.as_ptr(), args.as_ptr(), ptr::null_mut()), NemoRelayStatus::NullPointer ); @@ -990,11 +995,11 @@ fn test_ffi_helper_rejection_and_null_name_paths() { .unwrap_or_default() .contains("out pointer is null") ); - assert_eq!( + assert_status!( nemo_relay_tool_request_intercepts(ptr::null(), args.as_ptr(), &mut tool_out), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_tool_request_intercepts( tool_name.as_ptr(), invalid_json.as_ptr(), @@ -1003,17 +1008,17 @@ fn test_ffi_helper_rejection_and_null_name_paths() { NemoRelayStatus::InvalidJson ); assert!(tool_out.is_null()); - assert_eq!( + assert_status!( nemo_relay_tool_conditional_execution(ptr::null(), args.as_ptr()), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_tool_conditional_execution(tool_name.as_ptr(), invalid_json.as_ptr()), NemoRelayStatus::InvalidJson ); let tool_guard = cstring(&unique_name("ffi_tool_reject")); - assert_eq!( + assert_status!( nemo_relay_register_tool_conditional_execution_guardrail( tool_guard.as_ptr(), 1, @@ -1023,24 +1028,24 @@ fn test_ffi_helper_rejection_and_null_name_paths() { ), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_tool_conditional_execution(tool_name.as_ptr(), args.as_ptr()), NemoRelayStatus::GuardrailRejected ); - assert_eq!( + assert_status!( nemo_relay_deregister_tool_conditional_execution_guardrail(tool_guard.as_ptr()), NemoRelayStatus::Ok ); let mut llm_out = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_llm_request_intercepts(ptr::null(), request.as_ptr(), &mut llm_out), NemoRelayStatus::Ok ); let llm_json = returned_json(llm_out); assert_eq!(llm_json["request"]["content"]["model"], json!("ffi-model")); - assert_eq!( + assert_status!( nemo_relay_llm_request_intercepts(llm_name.as_ptr(), request.as_ptr(), ptr::null_mut()), NemoRelayStatus::NullPointer ); @@ -1049,7 +1054,7 @@ fn test_ffi_helper_rejection_and_null_name_paths() { .unwrap_or_default() .contains("out pointer is null") ); - assert_eq!( + assert_status!( nemo_relay_llm_request_intercepts( llm_name.as_ptr(), invalid_json.as_ptr(), @@ -1058,13 +1063,13 @@ fn test_ffi_helper_rejection_and_null_name_paths() { NemoRelayStatus::InvalidJson ); assert!(llm_error_out.is_null()); - assert_eq!( + assert_status!( nemo_relay_llm_conditional_execution(invalid_json.as_ptr()), NemoRelayStatus::InvalidJson ); let llm_guard = cstring(&unique_name("ffi_llm_reject")); - assert_eq!( + assert_status!( nemo_relay_register_llm_conditional_execution_guardrail( llm_guard.as_ptr(), 1, @@ -1074,11 +1079,11 @@ fn test_ffi_helper_rejection_and_null_name_paths() { ), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_llm_conditional_execution(request.as_ptr()), NemoRelayStatus::GuardrailRejected ); - assert_eq!( + assert_status!( nemo_relay_deregister_llm_conditional_execution_guardrail(llm_guard.as_ptr()), NemoRelayStatus::Ok ); @@ -1094,12 +1099,12 @@ fn test_ffi_registration_name_and_uuid_error_sweep() { macro_rules! assert_invalid_arg { ($expr:expr_2021) => { - assert_eq!($expr, NemoRelayStatus::InvalidArg); + assert_status!($expr, NemoRelayStatus::InvalidArg); }; } macro_rules! assert_null_pointer { ($expr:expr_2021) => { - assert_eq!($expr, NemoRelayStatus::NullPointer); + assert_status!($expr, NemoRelayStatus::NullPointer); }; } @@ -1107,7 +1112,7 @@ fn test_ffi_registration_name_and_uuid_error_sweep() { let stack = fresh_scope_stack(); let scope_name = cstring("ffi_error_sweep_scope"); let mut scope = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_push_scope( scope_name.as_ptr(), NemoRelayScopeType::Function, @@ -1392,7 +1397,7 @@ fn test_ffi_registration_name_and_uuid_error_sweep() { ptr::null(), )); - assert_eq!( + assert_status!( nemo_relay_pop_scope(scope, ptr::null()), NemoRelayStatus::Ok ); @@ -1408,7 +1413,7 @@ fn test_ffi_duplicate_registration_sweep_and_helper_callbacks() { macro_rules! assert_already_exists { ($expr:expr_2021) => { - assert_eq!($expr, NemoRelayStatus::AlreadyExists); + assert_status!($expr, NemoRelayStatus::AlreadyExists); }; } @@ -1435,7 +1440,7 @@ fn test_ffi_duplicate_registration_sweep_and_helper_callbacks() { let stack = fresh_scope_stack(); let scope_name = cstring("ffi_duplicate_scope"); let mut scope = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_push_scope( scope_name.as_ptr(), NemoRelayScopeType::Function, @@ -1451,7 +1456,7 @@ fn test_ffi_duplicate_registration_sweep_and_helper_callbacks() { let scope_uuid = cstring(&take_string(nemo_relay_scope_handle_uuid(scope)).unwrap()); let tool_cond = cstring(&unique_name("dup_tool_cond")); - assert_eq!( + assert_status!( nemo_relay_register_tool_conditional_execution_guardrail( tool_cond.as_ptr(), 1, @@ -1468,13 +1473,13 @@ fn test_ffi_duplicate_registration_sweep_and_helper_callbacks() { ptr::null_mut(), None, )); - assert_eq!( + assert_status!( nemo_relay_deregister_tool_conditional_execution_guardrail(tool_cond.as_ptr()), NemoRelayStatus::Ok ); let tool_req = cstring(&unique_name("dup_tool_req")); - assert_eq!( + assert_status!( nemo_relay_register_tool_request_intercept( tool_req.as_ptr(), 1, @@ -1493,13 +1498,13 @@ fn test_ffi_duplicate_registration_sweep_and_helper_callbacks() { ptr::null_mut(), None, )); - assert_eq!( + assert_status!( nemo_relay_deregister_tool_request_intercept(tool_req.as_ptr()), NemoRelayStatus::Ok ); let llm_san_resp = cstring(&unique_name("dup_llm_san_resp")); - assert_eq!( + assert_status!( nemo_relay_register_llm_sanitize_response_guardrail( llm_san_resp.as_ptr(), 1, @@ -1516,13 +1521,13 @@ fn test_ffi_duplicate_registration_sweep_and_helper_callbacks() { ptr::null_mut(), None, )); - assert_eq!( + assert_status!( nemo_relay_deregister_llm_sanitize_response_guardrail(llm_san_resp.as_ptr()), NemoRelayStatus::Ok ); let llm_cond = cstring(&unique_name("dup_llm_cond")); - assert_eq!( + assert_status!( nemo_relay_register_llm_conditional_execution_guardrail( llm_cond.as_ptr(), 1, @@ -1539,13 +1544,13 @@ fn test_ffi_duplicate_registration_sweep_and_helper_callbacks() { ptr::null_mut(), None, )); - assert_eq!( + assert_status!( nemo_relay_deregister_llm_conditional_execution_guardrail(llm_cond.as_ptr()), NemoRelayStatus::Ok ); let llm_req = cstring(&unique_name("dup_llm_req")); - assert_eq!( + assert_status!( nemo_relay_register_llm_request_intercept( llm_req.as_ptr(), 1, @@ -1564,13 +1569,13 @@ fn test_ffi_duplicate_registration_sweep_and_helper_callbacks() { ptr::null_mut(), None, )); - assert_eq!( + assert_status!( nemo_relay_deregister_llm_request_intercept(llm_req.as_ptr()), NemoRelayStatus::Ok ); let subscriber = cstring(&unique_name("dup_subscriber")); - assert_eq!( + assert_status!( nemo_relay_register_subscriber( subscriber.as_ptr(), subscriber_cb, @@ -1585,14 +1590,14 @@ fn test_ffi_duplicate_registration_sweep_and_helper_callbacks() { ptr::null_mut(), None, )); - assert_eq!(nemo_relay_flush_subscribers(), NemoRelayStatus::Ok); - assert_eq!( + assert_status!(nemo_relay_flush_subscribers(), NemoRelayStatus::Ok); + assert_status!( nemo_relay_deregister_subscriber(subscriber.as_ptr()), NemoRelayStatus::Ok ); let scope_tool_cond = cstring(&unique_name("dup_scope_tool_cond")); - assert_eq!( + assert_status!( nemo_relay_scope_register_tool_conditional_execution_guardrail( scope_uuid.as_ptr(), scope_tool_cond.as_ptr(), @@ -1613,7 +1618,7 @@ fn test_ffi_duplicate_registration_sweep_and_helper_callbacks() { None, ) ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_tool_conditional_execution_guardrail( scope_uuid.as_ptr(), scope_tool_cond.as_ptr(), @@ -1622,7 +1627,7 @@ fn test_ffi_duplicate_registration_sweep_and_helper_callbacks() { ); let scope_tool_req = cstring(&unique_name("dup_scope_tool_req")); - assert_eq!( + assert_status!( nemo_relay_scope_register_tool_request_intercept( scope_uuid.as_ptr(), scope_tool_req.as_ptr(), @@ -1643,7 +1648,7 @@ fn test_ffi_duplicate_registration_sweep_and_helper_callbacks() { ptr::null_mut(), None, )); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_tool_request_intercept( scope_uuid.as_ptr(), scope_tool_req.as_ptr(), @@ -1652,7 +1657,7 @@ fn test_ffi_duplicate_registration_sweep_and_helper_callbacks() { ); let scope_llm_cond = cstring(&unique_name("dup_scope_llm_cond")); - assert_eq!( + assert_status!( nemo_relay_scope_register_llm_conditional_execution_guardrail( scope_uuid.as_ptr(), scope_llm_cond.as_ptr(), @@ -1673,7 +1678,7 @@ fn test_ffi_duplicate_registration_sweep_and_helper_callbacks() { None, ) ); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_llm_conditional_execution_guardrail( scope_uuid.as_ptr(), scope_llm_cond.as_ptr(), @@ -1682,7 +1687,7 @@ fn test_ffi_duplicate_registration_sweep_and_helper_callbacks() { ); let scope_llm_req = cstring(&unique_name("dup_scope_llm_req")); - assert_eq!( + assert_status!( nemo_relay_scope_register_llm_request_intercept( scope_uuid.as_ptr(), scope_llm_req.as_ptr(), @@ -1703,7 +1708,7 @@ fn test_ffi_duplicate_registration_sweep_and_helper_callbacks() { ptr::null_mut(), None, )); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_llm_request_intercept( scope_uuid.as_ptr(), scope_llm_req.as_ptr(), @@ -1712,7 +1717,7 @@ fn test_ffi_duplicate_registration_sweep_and_helper_callbacks() { ); let scope_subscriber = cstring(&unique_name("dup_scope_subscriber")); - assert_eq!( + assert_status!( nemo_relay_scope_register_subscriber( scope_uuid.as_ptr(), scope_subscriber.as_ptr(), @@ -1729,7 +1734,7 @@ fn test_ffi_duplicate_registration_sweep_and_helper_callbacks() { ptr::null_mut(), None, )); - assert_eq!( + assert_status!( nemo_relay_scope_deregister_subscriber(scope_uuid.as_ptr(), scope_subscriber.as_ptr(),), NemoRelayStatus::Ok ); @@ -1738,7 +1743,7 @@ fn test_ffi_duplicate_registration_sweep_and_helper_callbacks() { let agent = cstring("dup-agent"); let version = cstring("1.0.0"); let mut exporter = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_atif_exporter_create( ptr::null(), agent.as_ptr(), @@ -1748,7 +1753,7 @@ fn test_ffi_duplicate_registration_sweep_and_helper_callbacks() { ), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_atif_exporter_create( session.as_ptr(), ptr::null(), @@ -1758,7 +1763,7 @@ fn test_ffi_duplicate_registration_sweep_and_helper_callbacks() { ), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_atif_exporter_create( session.as_ptr(), agent.as_ptr(), @@ -1768,7 +1773,7 @@ fn test_ffi_duplicate_registration_sweep_and_helper_callbacks() { ), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_atif_exporter_create( session.as_ptr(), agent.as_ptr(), @@ -1778,12 +1783,12 @@ fn test_ffi_duplicate_registration_sweep_and_helper_callbacks() { ), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_atif_exporter_register(exporter, ptr::null()), NemoRelayStatus::NullPointer ); let exporter_name = cstring(&unique_name("dup_exporter_subscriber")); - assert_eq!( + assert_status!( nemo_relay_atif_exporter_register(exporter, exporter_name.as_ptr()), NemoRelayStatus::Ok ); @@ -1791,11 +1796,11 @@ fn test_ffi_duplicate_registration_sweep_and_helper_callbacks() { exporter, exporter_name.as_ptr(), )); - assert_eq!( + assert_status!( nemo_relay_atif_exporter_deregister(ptr::null()), NemoRelayStatus::NullPointer ); - assert_eq!( + assert_status!( nemo_relay_atif_exporter_deregister(exporter_name.as_ptr()), NemoRelayStatus::Ok ); @@ -1827,7 +1832,7 @@ fn test_ffi_duplicate_registration_sweep_and_helper_callbacks() { json!({"role":"assistant","content":"next","tool_calls":[]}) ); - assert_eq!( + assert_status!( nemo_relay_pop_scope(scope, ptr::null()), NemoRelayStatus::Ok ); @@ -1844,39 +1849,39 @@ fn test_ffi_registration_table_sweep_for_remaining_wrappers() { macro_rules! assert_global_guardrail_sweep { ($prefix:literal, $register:ident, $deregister:ident, $cb:expr) => {{ let name = cstring(&unique_name($prefix)); - assert_eq!( + assert_status!( $register(name.as_ptr(), 1, $cb, ptr::null_mut(), None), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( $register(name.as_ptr(), 1, $cb, ptr::null_mut(), None), NemoRelayStatus::AlreadyExists ); - assert_eq!($deregister(name.as_ptr()), NemoRelayStatus::Ok); - assert_eq!($deregister(name.as_ptr()), NemoRelayStatus::Ok); + assert_status!($deregister(name.as_ptr()), NemoRelayStatus::Ok); + assert_status!($deregister(name.as_ptr()), NemoRelayStatus::Ok); }}; } macro_rules! assert_global_execution_sweep { ($prefix:literal, $register:ident, $deregister:ident, $cb:expr) => {{ let name = cstring(&unique_name($prefix)); - assert_eq!( + assert_status!( $register(name.as_ptr(), 1, $cb, ptr::null_mut(), None), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( $register(name.as_ptr(), 1, $cb, ptr::null_mut(), None), NemoRelayStatus::AlreadyExists ); - assert_eq!($deregister(name.as_ptr()), NemoRelayStatus::Ok); - assert_eq!($deregister(name.as_ptr()), NemoRelayStatus::Ok); + assert_status!($deregister(name.as_ptr()), NemoRelayStatus::Ok); + assert_status!($deregister(name.as_ptr()), NemoRelayStatus::Ok); }}; } macro_rules! assert_scope_guardrail_sweep { ($scope_uuid:expr, $prefix:literal, $register:ident, $deregister:ident, $cb:expr) => {{ let name = cstring(&unique_name($prefix)); - assert_eq!( + assert_status!( $register( $scope_uuid.as_ptr(), name.as_ptr(), @@ -1887,7 +1892,7 @@ fn test_ffi_registration_table_sweep_for_remaining_wrappers() { ), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( $register( $scope_uuid.as_ptr(), name.as_ptr(), @@ -1898,11 +1903,11 @@ fn test_ffi_registration_table_sweep_for_remaining_wrappers() { ), NemoRelayStatus::AlreadyExists ); - assert_eq!( + assert_status!( $deregister($scope_uuid.as_ptr(), name.as_ptr()), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( $deregister($scope_uuid.as_ptr(), name.as_ptr()), NemoRelayStatus::Ok ); @@ -1912,7 +1917,7 @@ fn test_ffi_registration_table_sweep_for_remaining_wrappers() { macro_rules! assert_scope_execution_sweep { ($scope_uuid:expr, $prefix:literal, $register:ident, $deregister:ident, $cb:expr) => {{ let name = cstring(&unique_name($prefix)); - assert_eq!( + assert_status!( $register( $scope_uuid.as_ptr(), name.as_ptr(), @@ -1923,7 +1928,7 @@ fn test_ffi_registration_table_sweep_for_remaining_wrappers() { ), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( $register( $scope_uuid.as_ptr(), name.as_ptr(), @@ -1934,11 +1939,11 @@ fn test_ffi_registration_table_sweep_for_remaining_wrappers() { ), NemoRelayStatus::AlreadyExists ); - assert_eq!( + assert_status!( $deregister($scope_uuid.as_ptr(), name.as_ptr()), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( $deregister($scope_uuid.as_ptr(), name.as_ptr()), NemoRelayStatus::Ok ); @@ -1949,7 +1954,7 @@ fn test_ffi_registration_table_sweep_for_remaining_wrappers() { let stack = fresh_scope_stack(); let scope_name = cstring("ffi_table_sweep_scope"); let mut scope = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_push_scope( scope_name.as_ptr(), NemoRelayScopeType::Function, @@ -2043,7 +2048,7 @@ fn test_ffi_registration_table_sweep_for_remaining_wrappers() { let agent = cstring("table-sweep-agent"); let version = cstring("1.0.0"); let exporter_name = cstring(&unique_name("table_exporter_subscriber")); - assert_eq!( + assert_status!( nemo_relay_atif_exporter_create( session.as_ptr(), agent.as_ptr(), @@ -2053,21 +2058,21 @@ fn test_ffi_registration_table_sweep_for_remaining_wrappers() { ), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_atif_exporter_register(exporter, exporter_name.as_ptr()), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_atif_exporter_deregister(exporter_name.as_ptr()), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_atif_exporter_deregister(exporter_name.as_ptr()), NemoRelayStatus::Ok ); nemo_relay_atif_exporter_free(exporter); - assert_eq!( + assert_status!( nemo_relay_pop_scope(scope, ptr::null()), NemoRelayStatus::Ok ); @@ -2086,7 +2091,7 @@ fn test_ffi_llm_execute_stream_and_atif_exporter() { let subscriber_name = unique_name("ffi_llm_subscriber"); let subscriber_name_c = cstring(&subscriber_name); - assert_eq!( + assert_status!( nemo_relay_register_subscriber( subscriber_name_c.as_ptr(), subscriber_cb, @@ -2097,12 +2102,12 @@ fn test_ffi_llm_execute_stream_and_atif_exporter() { ); let mut root = ptr::null_mut(); - assert_eq!(nemo_relay_get_handle(&mut root), NemoRelayStatus::Ok); + assert_status!(nemo_relay_get_handle(&mut root), NemoRelayStatus::Ok); nemo_relay_scope_handle_free(root); let intercept_name = unique_name("ffi_llm_intercept"); let intercept_name_c = cstring(&intercept_name); - assert_eq!( + assert_status!( nemo_relay_register_llm_request_intercept( intercept_name_c.as_ptr(), 1, @@ -2116,7 +2121,7 @@ fn test_ffi_llm_execute_stream_and_atif_exporter() { let conditional_name = unique_name("ffi_llm_conditional"); let conditional_name_c = cstring(&conditional_name); - assert_eq!( + assert_status!( nemo_relay_register_llm_conditional_execution_guardrail( conditional_name_c.as_ptr(), 1, @@ -2129,7 +2134,7 @@ fn test_ffi_llm_execute_stream_and_atif_exporter() { let sanitize_name = unique_name("ffi_llm_sanitize"); let sanitize_name_c = cstring(&sanitize_name); - assert_eq!( + assert_status!( nemo_relay_register_llm_sanitize_response_guardrail( sanitize_name_c.as_ptr(), 1, @@ -2145,7 +2150,7 @@ fn test_ffi_llm_execute_stream_and_atif_exporter() { let agent = cstring("ffi-agent"); let version = cstring("1.0.0"); let model_name = cstring("ffi-model"); - assert_eq!( + assert_status!( nemo_relay_atif_exporter_create( session.as_ptr(), agent.as_ptr(), @@ -2158,7 +2163,7 @@ fn test_ffi_llm_execute_stream_and_atif_exporter() { let exporter_sub = unique_name("ffi_exporter"); let exporter_sub_c = cstring(&exporter_sub); - assert_eq!( + assert_status!( nemo_relay_atif_exporter_register(exporter, exporter_sub_c.as_ptr()), NemoRelayStatus::Ok ); @@ -2188,7 +2193,7 @@ fn test_ffi_llm_execute_stream_and_atif_exporter() { nemo_relay_llm_request_free(llm_request); let mut helper_out = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_llm_request_intercepts(llm_name.as_ptr(), request.as_ptr(), &mut helper_out), NemoRelayStatus::Ok ); @@ -2198,13 +2203,13 @@ fn test_ffi_llm_execute_stream_and_atif_exporter() { json!(true) ); - assert_eq!( + assert_status!( nemo_relay_llm_conditional_execution(request.as_ptr()), NemoRelayStatus::Ok ); let mut handle: *mut FfiLLMHandle = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_llm_call( llm_name.as_ptr(), request.as_ptr(), @@ -2226,14 +2231,14 @@ fn test_ffi_llm_execute_stream_and_atif_exporter() { assert!(take_string(nemo_relay_llm_handle_parent_uuid(handle)).is_some()); let response = cstring(r#"{"content":"manual end","role":"assistant","tool_calls":[]}"#); - assert_eq!( + assert_status!( nemo_relay_llm_call_end(handle, response.as_ptr(), ptr::null(), ptr::null()), NemoRelayStatus::Ok ); nemo_relay_llm_handle_free(handle); let mut execute_out = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_llm_call_execute( llm_name.as_ptr(), request.as_ptr(), @@ -2257,21 +2262,12 @@ fn test_ffi_llm_execute_stream_and_atif_exporter() { let execute_json = returned_json(execute_out); assert_eq!(execute_json["content"], json!("hello from ffi")); assert_eq!(execute_json["model_seen"], json!("ffi-model")); - assert_eq!(nemo_relay_flush_subscribers(), NemoRelayStatus::Ok); + assert_status!(nemo_relay_flush_subscribers(), NemoRelayStatus::Ok); let events = lock_unpoisoned(event_log()).clone(); - assert!( - events - .iter() - .any(|event| event["output"]["sanitized"] == json!(true)) - ); - assert!( - events - .iter() - .any(|event| event["model_name"] == "ffi-model") - ); + assert_llm_execution_events(&events); let mut stream = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_llm_stream_call_execute( llm_name.as_ptr(), request.as_ptr(), @@ -2305,48 +2301,65 @@ fn test_ffi_llm_execute_stream_and_atif_exporter() { assert_eq!(*lock_unpoisoned(finalizer_calls()), 1); let mut exported = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_atif_exporter_export(exporter, &mut exported), NemoRelayStatus::Ok ); let trajectory = returned_json(exported); - assert_eq!(trajectory["schema_version"], json!("ATIF-v1.7")); - assert!(trajectory["steps"].as_array().unwrap().len() >= 4); + assert_atif_trajectory(&trajectory); - assert_eq!( + assert_status!( nemo_relay_atif_exporter_clear(exporter), NemoRelayStatus::Ok ); let mut cleared = ptr::null_mut(); - assert_eq!( + assert_status!( nemo_relay_atif_exporter_export(exporter, &mut cleared), NemoRelayStatus::Ok ); let cleared_json = returned_json(cleared); assert_eq!(cleared_json["steps"].as_array().unwrap().len(), 0); - assert_eq!( + assert_status!( nemo_relay_atif_exporter_deregister(exporter_sub_c.as_ptr()), NemoRelayStatus::Ok ); nemo_relay_atif_exporter_free(exporter); - assert_eq!( + assert_status!( nemo_relay_deregister_subscriber(subscriber_name_c.as_ptr()), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_deregister_llm_request_intercept(intercept_name_c.as_ptr()), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_deregister_llm_conditional_execution_guardrail(conditional_name_c.as_ptr()), NemoRelayStatus::Ok ); - assert_eq!( + assert_status!( nemo_relay_deregister_llm_sanitize_response_guardrail(sanitize_name_c.as_ptr()), NemoRelayStatus::Ok ); nemo_relay_scope_stack_free(stack); } } + +fn assert_llm_execution_events(events: &[Json]) { + assert!( + events + .iter() + .any(|event| event["output"]["sanitized"] == json!(true)) + ); + assert!( + events + .iter() + .any(|event| event["model_name"] == "ffi-model") + ); +} + +fn assert_atif_trajectory(trajectory: &Json) { + assert_eq!(trajectory["schema_version"], json!("ATIF-v1.7")); + assert!(trajectory["steps"].as_array().unwrap().len() >= 4); +} diff --git a/crates/ffi/tests/unit/api_tests.rs b/crates/ffi/tests/unit/api_tests.rs index 2a7f41246..6f9eac085 100644 --- a/crates/ffi/tests/unit/api_tests.rs +++ b/crates/ffi/tests/unit/api_tests.rs @@ -46,6 +46,17 @@ static COLLECTED_CHUNKS: OnceLock>> = OnceLock::new(); static FINALIZER_CALLS: OnceLock> = OnceLock::new(); static PLUGIN_FREES: OnceLock> = OnceLock::new(); +#[track_caller] +fn assert_native_status(actual: NemoRelayStatus, expected: NemoRelayStatus) { + assert_eq!(actual, expected); +} + +macro_rules! assert_status { + ($actual:expr, $expected:expr $(,)?) => { + assert_native_status($actual, $expected) + }; +} + fn event_log() -> &'static Mutex> { EVENT_LOG.get_or_init(|| Mutex::new(Vec::new())) } diff --git a/crates/ffi/tests/unit/callable_tests.rs b/crates/ffi/tests/unit/callable_tests.rs index 4542bb551..2bc0427db 100644 --- a/crates/ffi/tests/unit/callable_tests.rs +++ b/crates/ffi/tests/unit/callable_tests.rs @@ -577,6 +577,13 @@ fn test_wrap_llm_exec_stream_and_event_callbacks() { .build() .unwrap(); + assert_llm_exec_callbacks(&runtime); + assert_llm_stream_callbacks(&runtime); + assert_collector_and_finalizer_callbacks(); + assert_event_callbacks(); +} + +fn assert_llm_exec_callbacks(runtime: &tokio::runtime::Runtime) { let exec = wrap_llm_exec_fn(llm_exec_cb, std::ptr::null_mut(), None); let result = runtime.block_on(exec(make_request())).unwrap(); assert_eq!(result["ok"], json!(true)); @@ -593,7 +600,9 @@ fn test_wrap_llm_exec_stream_and_event_callbacks() { .block_on(intercept("llm", make_request(), next)) .unwrap(); assert_eq!(intercepted["intercepted"], json!(true)); +} +fn assert_llm_stream_callbacks(runtime: &tokio::runtime::Runtime) { let stream_exec = wrap_llm_stream_exec_fn(llm_exec_cb, std::ptr::null_mut(), None); let mut stream = runtime.block_on(stream_exec(make_request())).unwrap(); let first = runtime.block_on(async { stream.next().await.unwrap().unwrap() }); @@ -635,7 +644,9 @@ fn test_wrap_llm_exec_stream_and_event_callbacks() { let first = runtime.block_on(async { intercepted_stream.next().await.unwrap().unwrap() }); assert_eq!(first["intercepted"], json!(true)); assert_eq!(first["model"], json!("test-model")); +} +fn assert_collector_and_finalizer_callbacks() { COLLECTED_COUNT.store(0, Ordering::SeqCst); let mut collector = wrap_collector_fn(collector_cb); collector(json!({"chunk": 1})).unwrap(); @@ -643,7 +654,9 @@ fn test_wrap_llm_exec_stream_and_event_callbacks() { let finalizer = wrap_finalizer_fn(finalizer_cb); assert_eq!(finalizer(), json!({"done": true})); +} +fn assert_event_callbacks() { let (user_data, seen) = user_data_counter(); let subscriber = wrap_event_subscriber(subscriber_cb, user_data, Some(free_arc_counter)); let event = Event::Scope(nemo_relay::api::event::ScopeEvent::new( diff --git a/crates/ffi/tests/unit/types_tests.rs b/crates/ffi/tests/unit/types_tests.rs index 9c9a008c2..074ddf60d 100644 --- a/crates/ffi/tests/unit/types_tests.rs +++ b/crates/ffi/tests/unit/types_tests.rs @@ -218,6 +218,12 @@ fn test_tool_and_llm_handle_accessors_and_null_guards() { #[test] fn test_llm_request_null_inputs_event_null_guards_and_free_nulls() { + assert_llm_request_null_defaults(); + assert_null_event_accessors(); + free_null_ffi_handles(); +} + +fn assert_llm_request_null_defaults() { let request_ptr = unsafe { nemo_relay_llm_request_new(std::ptr::null(), std::ptr::null()) }; assert!(!request_ptr.is_null()); assert_eq!( @@ -229,7 +235,14 @@ fn test_llm_request_null_inputs_event_null_guards_and_free_nulls() { Some("null".into()) ); unsafe { nemo_relay_llm_request_free(request_ptr) }; +} + +fn assert_null_event_accessors() { + assert_null_event_identity_accessors(); + assert_null_event_payload_accessors(); +} +fn assert_null_event_identity_accessors() { assert!(unsafe { nemo_relay_llm_request_headers(std::ptr::null()) }.is_null()); assert!(unsafe { nemo_relay_llm_request_content(std::ptr::null()) }.is_null()); assert!(unsafe { nemo_relay_event_uuid(std::ptr::null()) }.is_null()); @@ -239,6 +252,9 @@ fn test_llm_request_null_inputs_event_null_guards_and_free_nulls() { assert!(unsafe { nemo_relay_event_data(std::ptr::null()) }.is_null()); assert!(unsafe { nemo_relay_event_metadata(std::ptr::null()) }.is_null()); assert!(unsafe { nemo_relay_event_timestamp(std::ptr::null()) }.is_null()); +} + +fn assert_null_event_payload_accessors() { assert!(unsafe { nemo_relay_event_input(std::ptr::null()) }.is_null()); assert!(unsafe { nemo_relay_event_output(std::ptr::null()) }.is_null()); assert!(unsafe { nemo_relay_event_model_name(std::ptr::null()) }.is_null()); @@ -251,7 +267,9 @@ fn test_llm_request_null_inputs_event_null_guards_and_free_nulls() { assert!(unsafe { nemo_relay_event_attributes_json(std::ptr::null()) }.is_null()); assert!(unsafe { nemo_relay_event_category_profile(std::ptr::null()) }.is_null()); assert!(unsafe { nemo_relay_event_data_schema(std::ptr::null()) }.is_null()); +} +fn free_null_ffi_handles() { unsafe { nemo_relay_scope_handle_free(std::ptr::null_mut()); nemo_relay_tool_handle_free(std::ptr::null_mut()); @@ -387,172 +405,200 @@ fn test_llm_request_and_event_accessors() { }); let ffi_event = FfiEvent(scope_event.clone()); + assert_scope_event_accessors(&ffi_event, &scope_event, parent_uuid); + + let llm_event = make_scope_event(ScopeEventFixture { + scope_category: ScopeCategory::Start, + scope_type: ScopeType::Llm, + name: "ffi-llm", + parent_uuid: Some(parent_uuid), + data: Some(json!({"input": true})), + metadata: None, + attributes: llm_attributes_to_strings(LlmAttributes::empty()), + category_profile: Some(CategoryProfile::builder().model_name("model").build()), + }); + assert_llm_event_accessors(&FfiEvent(llm_event)); + + let tool_event = make_scope_event(ScopeEventFixture { + scope_category: ScopeCategory::End, + scope_type: ScopeType::Tool, + name: "ffi-tool", + parent_uuid: Some(parent_uuid), + data: Some(json!({"output": true})), + metadata: None, + attributes: tool_attributes_to_strings(ToolAttributes::empty()), + category_profile: Some( + CategoryProfile::builder() + .tool_call_id("tool-call-id") + .build(), + ), + }); + assert_tool_event_accessors(&FfiEvent(tool_event)); + + let mark_event = mark_event("ffi-mark", Some(parent_uuid), None, None); + assert_mark_event_accessors(&FfiEvent(mark_event)); +} + +fn assert_scope_event_accessors(ffi_event: &FfiEvent, scope_event: &Event, parent_uuid: Uuid) { + assert_scope_event_shape_accessors(ffi_event); + assert_scope_event_identity_accessors(ffi_event, scope_event, parent_uuid); + assert_scope_event_payload_accessors(ffi_event); +} + +fn assert_scope_event_shape_accessors(ffi_event: &FfiEvent) { assert_eq!( - take_string(unsafe { nemo_relay_event_kind(&ffi_event) }), + take_string(unsafe { nemo_relay_event_kind(ffi_event) }), Some("scope".into()) ); assert_eq!( - take_string(unsafe { nemo_relay_event_scope_category(&ffi_event) }), + take_string(unsafe { nemo_relay_event_scope_category(ffi_event) }), Some("start".into()) ); assert_eq!( - take_string(unsafe { nemo_relay_event_atof_version(&ffi_event) }), + take_string(unsafe { nemo_relay_event_atof_version(ffi_event) }), Some("0.1".into()) ); assert_eq!( - take_string(unsafe { nemo_relay_event_category(&ffi_event) }), + take_string(unsafe { nemo_relay_event_category(ffi_event) }), Some("guardrail".into()) ); assert_eq!( - take_string(unsafe { nemo_relay_event_attributes_json(&ffi_event) }), + take_string(unsafe { nemo_relay_event_attributes_json(ffi_event) }), Some("[]".into()) ); assert_eq!( - take_string(unsafe { nemo_relay_event_category_profile(&ffi_event) }), + take_string(unsafe { nemo_relay_event_category_profile(ffi_event) }), None ); assert_eq!( - take_string(unsafe { nemo_relay_event_data_schema(&ffi_event) }), + take_string(unsafe { nemo_relay_event_data_schema(ffi_event) }), None ); +} + +fn assert_scope_event_identity_accessors( + ffi_event: &FfiEvent, + scope_event: &Event, + parent_uuid: Uuid, +) { assert_eq!( - take_string(unsafe { nemo_relay_event_uuid(&ffi_event) }), + take_string(unsafe { nemo_relay_event_uuid(ffi_event) }), Some(scope_event.uuid().to_string()) ); assert_eq!( - take_string(unsafe { nemo_relay_event_name(&ffi_event) }), + take_string(unsafe { nemo_relay_event_name(ffi_event) }), Some("ffi-event".into()) ); assert_eq!( - take_string(unsafe { nemo_relay_event_data(&ffi_event) }), + take_string(unsafe { nemo_relay_event_data(ffi_event) }), Some(r#"{"data":1}"#.into()) ); assert_eq!( - take_string(unsafe { nemo_relay_event_metadata(&ffi_event) }), + take_string(unsafe { nemo_relay_event_metadata(ffi_event) }), Some(r#"{"meta":2}"#.into()) ); assert_eq!( - take_string(unsafe { nemo_relay_event_scope_type(&ffi_event) }), + take_string(unsafe { nemo_relay_event_scope_type(ffi_event) }), Some("guardrail".into()) ); assert_eq!( - unsafe { nemo_relay_event_attributes(&ffi_event) }, + unsafe { nemo_relay_event_attributes(ffi_event) }, ScopeAttributes::empty().bits() ); assert_eq!( - take_string(unsafe { nemo_relay_event_parent_uuid(&ffi_event) }), + take_string(unsafe { nemo_relay_event_parent_uuid(ffi_event) }), Some(parent_uuid.to_string()) ); assert!( - take_string(unsafe { nemo_relay_event_timestamp(&ffi_event) }) + take_string(unsafe { nemo_relay_event_timestamp(ffi_event) }) .unwrap() .contains('T') ); +} +fn assert_scope_event_payload_accessors(ffi_event: &FfiEvent) { assert_eq!( - take_string(unsafe { nemo_relay_event_input(&ffi_event) }), + take_string(unsafe { nemo_relay_event_input(ffi_event) }), Some(r#"{"data":1}"#.into()) ); assert_eq!( - take_string(unsafe { nemo_relay_event_output(&ffi_event) }), + take_string(unsafe { nemo_relay_event_output(ffi_event) }), None ); assert_eq!( - take_string(unsafe { nemo_relay_event_model_name(&ffi_event) }), + take_string(unsafe { nemo_relay_event_model_name(ffi_event) }), None ); assert_eq!( - take_string(unsafe { nemo_relay_event_tool_call_id(&ffi_event) }), + take_string(unsafe { nemo_relay_event_tool_call_id(ffi_event) }), None ); +} - let llm_event = make_scope_event(ScopeEventFixture { - scope_category: ScopeCategory::Start, - scope_type: ScopeType::Llm, - name: "ffi-llm", - parent_uuid: Some(parent_uuid), - data: Some(json!({"input": true})), - metadata: None, - attributes: llm_attributes_to_strings(LlmAttributes::empty()), - category_profile: Some(CategoryProfile::builder().model_name("model").build()), - }); - let ffi_llm_event = FfiEvent(llm_event); +fn assert_llm_event_accessors(ffi_llm_event: &FfiEvent) { assert_eq!( - take_string(unsafe { nemo_relay_event_input(&ffi_llm_event) }), + take_string(unsafe { nemo_relay_event_input(ffi_llm_event) }), Some(r#"{"input":true}"#.into()) ); assert_eq!( - unsafe { nemo_relay_event_attributes(&ffi_llm_event) }, + unsafe { nemo_relay_event_attributes(ffi_llm_event) }, LlmAttributes::empty().bits() ); assert_eq!( - take_string(unsafe { nemo_relay_event_model_name(&ffi_llm_event) }), + take_string(unsafe { nemo_relay_event_model_name(ffi_llm_event) }), Some("model".into()) ); - let llm_profile = take_string(unsafe { nemo_relay_event_category_profile(&ffi_llm_event) }) + let llm_profile = take_string(unsafe { nemo_relay_event_category_profile(ffi_llm_event) }) .expect("expected category profile"); let llm_profile: serde_json::Value = serde_json::from_str(&llm_profile).unwrap(); assert_eq!(llm_profile["model_name"], json!("model")); assert_eq!( - take_string(unsafe { nemo_relay_event_scope_type(&ffi_llm_event) }), + take_string(unsafe { nemo_relay_event_scope_type(ffi_llm_event) }), Some("llm".into()) ); +} - let tool_event = make_scope_event(ScopeEventFixture { - scope_category: ScopeCategory::End, - scope_type: ScopeType::Tool, - name: "ffi-tool", - parent_uuid: Some(parent_uuid), - data: Some(json!({"output": true})), - metadata: None, - attributes: tool_attributes_to_strings(ToolAttributes::empty()), - category_profile: Some( - CategoryProfile::builder() - .tool_call_id("tool-call-id") - .build(), - ), - }); - let ffi_tool_event = FfiEvent(tool_event); +fn assert_tool_event_accessors(ffi_tool_event: &FfiEvent) { assert_eq!( - take_string(unsafe { nemo_relay_event_output(&ffi_tool_event) }), + take_string(unsafe { nemo_relay_event_output(ffi_tool_event) }), Some(r#"{"output":true}"#.into()) ); assert_eq!( - unsafe { nemo_relay_event_attributes(&ffi_tool_event) }, + unsafe { nemo_relay_event_attributes(ffi_tool_event) }, ToolAttributes::empty().bits() ); assert_eq!( - take_string(unsafe { nemo_relay_event_tool_call_id(&ffi_tool_event) }), + take_string(unsafe { nemo_relay_event_tool_call_id(ffi_tool_event) }), Some("tool-call-id".into()) ); assert_eq!( - take_string(unsafe { nemo_relay_event_scope_type(&ffi_tool_event) }), + take_string(unsafe { nemo_relay_event_scope_type(ffi_tool_event) }), Some("tool".into()) ); +} - let mark_event = mark_event("ffi-mark", Some(parent_uuid), None, None); - let ffi_mark_event = FfiEvent(mark_event); +fn assert_mark_event_accessors(ffi_mark_event: &FfiEvent) { assert_eq!( - take_string(unsafe { nemo_relay_event_scope_type(&ffi_mark_event) }), + take_string(unsafe { nemo_relay_event_scope_type(ffi_mark_event) }), None ); assert_eq!( - take_string(unsafe { nemo_relay_event_atof_version(&ffi_mark_event) }), + take_string(unsafe { nemo_relay_event_atof_version(ffi_mark_event) }), Some("0.1".into()) ); assert_eq!( - take_string(unsafe { nemo_relay_event_scope_category(&ffi_mark_event) }), + take_string(unsafe { nemo_relay_event_scope_category(ffi_mark_event) }), None ); assert_eq!( - take_string(unsafe { nemo_relay_event_category(&ffi_mark_event) }), + take_string(unsafe { nemo_relay_event_category(ffi_mark_event) }), None ); assert_eq!( - take_string(unsafe { nemo_relay_event_attributes_json(&ffi_mark_event) }), + take_string(unsafe { nemo_relay_event_attributes_json(ffi_mark_event) }), None ); - assert_eq!(unsafe { nemo_relay_event_attributes(&ffi_mark_event) }, 0); + assert_eq!(unsafe { nemo_relay_event_attributes(ffi_mark_event) }, 0); } #[test] diff --git a/crates/node/adaptive.js b/crates/node/adaptive.js index 836e60ab4..c96082012 100644 --- a/crates/node/adaptive.js +++ b/crates/node/adaptive.js @@ -157,7 +157,7 @@ function acgConfig(config = {}) { * @param {object} [config={}] - Partial response-cache settings to override. * @returns {object} A normalized response-cache config object. * @remarks The default backend is in-memory; pass a `backend` (e.g. - * `redisBackend(url)`) for a shared cache. `bypassRate` defaults to `0.0`, + * `redisBackend(url)`) for a shared cache. `bypassRate` defaults to `0`, * while caching nondeterministic requests is opt-in. Set a non-empty * `namespace` identifying one trusted cache-sharing domain before validation; * the empty helper default is an unconfigured sentinel. @@ -168,7 +168,7 @@ function responseCacheConfig(config = {}) { ttlSeconds: 3600, namespace: '', priority: 50, - bypassRate: 0.0, + bypassRate: 0, cacheNondeterministic: false, keyStrategy: 'exact_request', headerAllowlist: [], diff --git a/crates/node/tests/llm_tests.mjs b/crates/node/tests/llm_tests.mjs index afeabd0fd..a2254af50 100644 --- a/crates/node/tests/llm_tests.mjs +++ b/crates/node/tests/llm_tests.mjs @@ -1337,7 +1337,7 @@ describe('LLM intercepts', () => { let providerSideEffects = 0; registerLlmExecutionIntercept('node_llm_exec_abort_started_provider', 10, async (native, next) => { downstream = next(native); - void downstream.catch(() => {}); + downstream.catch(() => undefined); await started; return { source: 'intercept' }; }); @@ -1950,7 +1950,9 @@ describe('LLM intercepts', () => { ]; for (const registration of registrations) { - const declaration = declarations.match(new RegExp(`export declare function ${registration}\\([^\\n]+`))?.[0]; + const declaration = declarations.match( + new RegExp(String.raw`export declare function ${registration}\([^\n]+`), + )?.[0]; assert.ok(declaration, `missing declaration for ${registration}`); assert.doesNotMatch(declaration, /\.\.\.args: any\[\]/, `${registration} must not expose an any callback`); assert.match(declaration, /Promise { let providerSideEffects = 0; registerToolExecutionIntercept('node_tool_exec_abort_started_provider', 10, async (args, next) => { downstream = next(args); - void downstream.catch(() => {}); + downstream.catch(() => undefined); await started; return { result: { source: 'intercept' } }; }); diff --git a/crates/node/tests/typed_tests.mjs b/crates/node/tests/typed_tests.mjs index 668e738e1..189cb6a43 100644 --- a/crates/node/tests/typed_tests.mjs +++ b/crates/node/tests/typed_tests.mjs @@ -420,7 +420,7 @@ describe('typedToolExecute', () => { let providerSideEffects = 0; registerToolExecutionIntercept('typed_tool_abort_started_provider', 10, async (args, next) => { downstream = next(args); - void downstream.catch(() => {}); + downstream.catch(() => undefined); await started; return { result: { source: 'intercept' } }; }); @@ -545,7 +545,7 @@ describe('typedLlmExecute', () => { let providerSideEffects = 0; registerLlmExecutionIntercept('typed_llm_abort_started_provider', 10, async (request, next) => { downstream = next(request); - void downstream.catch(() => {}); + downstream.catch(() => undefined); await started; return { source: 'intercept' }; }); diff --git a/crates/plugin/tests/typed_callbacks.rs b/crates/plugin/tests/typed_callbacks.rs index 39a6ece8b..0161400fb 100644 --- a/crates/plugin/tests/typed_callbacks.rs +++ b/crates/plugin/tests/typed_callbacks.rs @@ -1887,6 +1887,15 @@ fn plugin_runtime_scope_mark_and_stack_helpers_call_host() { drop(stack); let calls = RUNTIME_CALLS.lock().unwrap().clone(); + assert_scope_runtime_calls(&calls); + assert_stack_runtime_calls(&calls); + assert_eq!(SCOPE_HANDLE_FREES.load(Ordering::SeqCst), 2); + assert_eq!(SCOPE_STACK_FREES.load(Ordering::SeqCst), 1); + assert_eq!(SCOPE_STACK_BINDING_RESTORES.load(Ordering::SeqCst), 1); + assert_eq!(SCOPE_STACK_BINDING_FREES.load(Ordering::SeqCst), 0); +} + +fn assert_scope_runtime_calls(calls: &[String]) { assert!(calls.iter().any(|call| call == "current_scope")); assert!(calls.iter().any(|call| { call.starts_with("push:work:Tool:0:parent=false") @@ -1904,15 +1913,14 @@ fn plugin_runtime_scope_mark_and_stack_helpers_call_host() { && call.contains(r#""output":true"#) && call.contains(r#""closed":true"#) })); +} + +fn assert_stack_runtime_calls(calls: &[String]) { assert!(calls.iter().any(|call| call == "stack_create")); assert!(calls.iter().any(|call| call == "stack_with_current")); assert!(calls.iter().any(|call| call == "stack_capture")); assert!(calls.iter().any(|call| call == "stack_set_thread")); assert!(calls.iter().any(|call| call == "stack_restore")); - assert_eq!(SCOPE_HANDLE_FREES.load(Ordering::SeqCst), 2); - assert_eq!(SCOPE_STACK_FREES.load(Ordering::SeqCst), 1); - assert_eq!(SCOPE_STACK_BINDING_RESTORES.load(Ordering::SeqCst), 1); - assert_eq!(SCOPE_STACK_BINDING_FREES.load(Ordering::SeqCst), 0); } #[test] @@ -2750,7 +2758,7 @@ fn typed_callback_free_catches_drop_panics() { } #[test] -fn typed_callbacks_reject_null_abi_pointers_before_decoding_inputs() { +fn typed_event_and_tool_callbacks_reject_null_abi_pointers_before_decoding_inputs() { let _guard = begin_test(); let host = test_host(); @@ -2953,6 +2961,14 @@ fn typed_callbacks_reject_null_abi_pointers_before_decoding_inputs() { drop(Box::from_raw(next_state)); registration.free(); } +} + +#[test] +fn typed_llm_callbacks_reject_null_abi_pointers_before_decoding_inputs() { + let _guard = begin_test(); + let host = test_host(); + let mut out = ptr::null_mut(); + let mut reason = ptr::null_mut(); let mut ctx = test_context(&host); ctx.register_llm_sanitize_request_guardrail("llm-request", 0, |request, _context| { @@ -3180,7 +3196,7 @@ fn typed_callbacks_reject_null_abi_pointers_before_decoding_inputs() { } #[test] -fn typed_callbacks_report_invalid_json_for_each_decoder_family() { +fn typed_subscriber_event_and_tool_sanitize_callbacks_report_invalid_json() { let _guard = begin_test(); let host = test_host(); @@ -3277,6 +3293,12 @@ fn typed_callbacks_report_invalid_json_for_each_decoder_family() { (host.string_free)(payload); registration.free(); } +} + +#[test] +fn typed_conditional_execution_and_llm_callbacks_report_invalid_json() { + let _guard = begin_test(); + let host = test_host(); let mut ctx = test_context(&host); ctx.register_tool_conditional_execution_guardrail("tool-conditional", 0, |_name, _value| { @@ -4701,14 +4723,8 @@ fn typed_llm_stream_execution_wraps_next_chunks() { "llm-stream", 31, |_name, request, next: LlmStreamNext<'_>| { - let stream = next.call(request)?; - let stream: LlmJsonStream = Box::new(stream.map(|chunk| { - chunk.map(|mut chunk| { - chunk["wrapped"] = json!(true); - chunk - }) - })); - Ok(stream) + let stream: LlmJsonStream = Box::new(next.call(request)?); + Ok(wrap_stream_chunks(stream)) }, ) .unwrap(); @@ -4748,20 +4764,7 @@ fn typed_llm_stream_execution_wraps_next_chunks() { NemoRelayStatus::NullPointer ); - let (status, chunk) = poll_stream_chunk(&host, &stream); - assert_eq!(status, NemoRelayStatus::Ok); - assert_eq!(chunk.unwrap()["wrapped"], json!(true)); - let (status, chunk) = poll_stream_chunk(&host, &stream); - assert_eq!(status, NemoRelayStatus::Ok); - let chunk = chunk.unwrap(); - assert_eq!(chunk["chunk"], json!(2)); - assert_eq!(chunk["wrapped"], json!(true)); - let (status, chunk) = poll_stream_chunk(&host, &stream); - assert_eq!(status, NemoRelayStatus::StreamEnd); - assert!(chunk.is_none()); - let (status, chunk) = poll_stream_chunk(&host, &stream); - assert_eq!(status, NemoRelayStatus::StreamEnd); - assert!(chunk.is_none()); + assert_wrapped_stream_chunks(&host, &stream); assert_eq!( unsafe { stream.cancel.unwrap()(stream.user_data) }, NemoRelayStatus::Ok @@ -4782,6 +4785,36 @@ fn typed_llm_stream_execution_wraps_next_chunks() { assert_eq!(dropped.load(Ordering::SeqCst), 1); } +fn wrap_stream_chunks(stream: LlmJsonStream) -> LlmJsonStream { + Box::new(stream.map(|chunk| { + chunk.map(|mut chunk| { + chunk["wrapped"] = json!(true); + chunk + }) + })) +} + +fn assert_wrapped_stream_chunks( + host: &NemoRelayNativeHostApiV1, + stream: &NemoRelayNativeLlmStreamV1, +) { + let (status, chunk) = poll_stream_chunk(host, stream); + assert_eq!(status, NemoRelayStatus::Ok); + assert_eq!(chunk.unwrap()["wrapped"], json!(true)); + + let (status, chunk) = poll_stream_chunk(host, stream); + assert_eq!(status, NemoRelayStatus::Ok); + let chunk = chunk.unwrap(); + assert_eq!(chunk["chunk"], json!(2)); + assert_eq!(chunk["wrapped"], json!(true)); + + for _ in 0..2 { + let (status, chunk) = poll_stream_chunk(host, stream); + assert_eq!(status, NemoRelayStatus::StreamEnd); + assert!(chunk.is_none()); + } +} + #[test] fn typed_llm_stream_drop_catches_stream_state_panics() { let _guard = begin_test(); diff --git a/crates/python/src/py_callable.rs b/crates/python/src/py_callable.rs index db7b727cb..50772a3e6 100644 --- a/crates/python/src/py_callable.rs +++ b/crates/python/src/py_callable.rs @@ -1727,6 +1727,115 @@ pub fn wrap_py_event_subscriber(py_fn: Py) -> EventSubscriberFn { }) } +fn prepare_event_sanitizer_invocation<'py>( + py: Python<'py>, + publication_context: Option<&PythonPublicationContext>, + task_locals: Option, + publication_buffer: Option, +) -> FlowResult<(Option>, Option)> { + match publication_context { + Some(context) => { + let (context, task_locals) = copy_publication_invocation_with_buffer( + py, + context, + task_locals, + publication_buffer, + ) + .map_err(|error| FlowError::Internal(error.to_string()))?; + Ok((Some(context), task_locals)) + } + None => copy_middleware_invocation(py, task_locals) + .map_err(|error| FlowError::Internal(error.to_string())), + } +} + +fn py_event_object(py: Python<'_>, event: &Event) -> PyResult> { + 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()), + } +} + +fn call_event_sanitizer( + py: Python<'_>, + invoke: &Bound<'_, PyAny>, + callback: &Py, + invocation_context: Option<&Bound<'_, PyAny>>, + loop_affine: bool, + py_event: Py, + py_fields: Py, +) -> PyResult> { + let result = match (invocation_context, loop_affine) { + (Some(context), false) => { + context.call_method1("run", (invoke, callback.bind(py), py_event, py_fields)) + } + (None, false) => invoke.call1((callback.bind(py), py_event, py_fields)), + (Some(context), true) => { + context.call_method1("run", (callback.bind(py), py_event, py_fields)) + } + (None, true) => callback.bind(py).call1((py_event, py_fields)), + }?; + Ok(result.unbind()) +} + +fn start_py_event_sanitizer( + py: Python<'_>, + py_fn: &Py, + event: &Event, + fields: &EventSanitizeFields, + 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 fields_json = + serde_json::to_value(fields).map_err(|error| FlowError::Internal(error.to_string()))?; + let py_fields = + json_to_py(py, &fields_json).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_sanitizer( + py, + &invoke, + &callback, + invocation_context.as_ref(), + loop_affine, + py_event, + py_fields, + ) + .map_err(|error| FlowError::Internal(error.to_string()))?; + split_py_object_or_future_with_locals( + py, + result, + task_locals.as_ref(), + invocation_context.as_ref(), + ) +} + /// Wrap a Python callable ``(Event, EventSanitizeFields) -> EventSanitizeFields``. pub fn wrap_py_event_sanitize_fn(py_fn: Py) -> EventSanitizeFn { let py_fn = Arc::new(py_fn); @@ -1737,83 +1846,17 @@ pub fn wrap_py_event_sanitize_fn(py_fn: Py) -> EventSanitizeFn { let publication_context = publication_context::(); let publication_buffer = capture_nested_publication_buffer(); Box::pin(async move { - let result = Python::attach( - |py| -> FlowResult, PyValueFuture>> { - let (invocation_context, task_locals) = match publication_context.as_ref() { - Some(context) => { - let (context, publication_task_locals) = - copy_publication_invocation_with_buffer( - py, - context, - task_locals, - publication_buffer.clone(), - ) - .map_err(|error| FlowError::Internal(error.to_string()))?; - (Some(context), publication_task_locals) - } - None => { copy_middleware_invocation(py, task_locals) } - .map_err(|error| FlowError::Internal(error.to_string()))?, - }; - let py_event = match event.as_ref() { - 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) => { - return Err(FlowError::Internal(error.to_string())); - } - }; - let fields_json = match serde_json::to_value(&fields) { - Ok(value) => value, - Err(error) => { - return Err(FlowError::Internal(error.to_string())); - } - }; - let py_fields = match json_to_py(py, &fields_json) { - Ok(value) => value, - Err(error) => { - return Err(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 = match (invocation_context.as_ref(), !loop_affine) { - (Some(context), true) => context - .call_method1("run", (invoke, callback.bind(py), py_event, py_fields)), - (None, true) => invoke.call1((callback.bind(py), py_event, py_fields)), - (Some(context), false) => { - context.call_method1("run", (callback.bind(py), py_event, py_fields)) - } - (None, false) => callback.bind(py).call1((py_event, py_fields)), - } - .map_err(|error| FlowError::Internal(error.to_string()))?; - split_py_object_or_future_with_locals( - py, - result.unbind(), - task_locals.as_ref(), - invocation_context.as_ref(), - ) - }, - ); + let result = Python::attach(|py| { + start_py_event_sanitizer( + py, + py_fn.as_ref(), + event.as_ref(), + &fields, + publication_context.as_deref(), + task_locals, + publication_buffer, + ) + }); let result = resolve_py_object_or_future(result) .await .and_then(|result| { diff --git a/crates/python/tests/coverage/py_adaptive_coverage_tests.rs b/crates/python/tests/coverage/py_adaptive_coverage_tests.rs index aa5f8ef95..5361a6fd9 100644 --- a/crates/python/tests/coverage/py_adaptive_coverage_tests.rs +++ b/crates/python/tests/coverage/py_adaptive_coverage_tests.rs @@ -393,126 +393,96 @@ fn adaptive_runtime_locking_and_helper_errors_are_covered() { }), ) .unwrap(); - assert!( - locked_runtime - .deregister() - .unwrap_err() - .to_string() - .contains("locked by an async operation") - ); - assert!( - locked_runtime - .wait_for_idle() - .unwrap_err() - .to_string() - .contains("locked by an async operation") - ); - assert!( - locked_runtime - .report(py) - .unwrap_err() - .to_string() - .contains("locked by an async operation") + assert_py_error_contains(locked_runtime.deregister(), "locked by an async operation"); + assert_py_error_contains( + locked_runtime.wait_for_idle(), + "locked by an async operation", ); + assert_py_error_contains(locked_runtime.report(py), "locked by an async operation"); let scope_handle = crate::py_types::PyScopeHandle { inner: ScopeHandle::builder() .name("adaptive-runtime-cov-locked") .scope_type(CoreScopeType::Agent) .build(), }; - assert!(match locked_runtime.bind_scope(&scope_handle) { - Ok(_) => panic!("expected locked runtime rewrite to fail"), - Err(err) => err.to_string().contains("locked by an async operation"), - }); - assert!( - locked_runtime - .build_cache_request_facts( - py, - "openai", - "00000000-0000-0000-0000-000000000204", - annotated_request.bind(py), - "adaptive-runtime-cov", - None, - ) - .unwrap_err() - .to_string() - .contains("locked by an async operation") + assert_py_error_contains( + locked_runtime.bind_scope(&scope_handle), + "locked by an async operation", + ); + assert_py_error_contains( + locked_runtime.build_cache_request_facts( + py, + "openai", + "00000000-0000-0000-0000-000000000204", + annotated_request.bind(py), + "adaptive-runtime-cov", + None, + ), + "locked by an async operation", ); drop(guard); let empty_runtime = PyAdaptiveRuntime { inner: std::sync::Arc::new(tokio::sync::Mutex::new(None)), }; - assert!( - empty_runtime - .deregister() - .unwrap_err() - .to_string() - .contains("already shut down") - ); - assert!( - empty_runtime - .wait_for_idle() - .unwrap_err() - .to_string() - .contains("already shut down") - ); + assert_py_error_contains(empty_runtime.deregister(), "already shut down"); + assert_py_error_contains(empty_runtime.wait_for_idle(), "already shut down"); let scope_handle = crate::py_types::PyScopeHandle { inner: ScopeHandle::builder() .name("adaptive-runtime-cov-empty") .scope_type(CoreScopeType::Agent) .build(), }; - assert!(match empty_runtime.bind_scope(&scope_handle) { - Ok(_) => panic!("expected shut down runtime rewrite to fail"), - Err(err) => err.to_string().contains("already shut down"), - }); - - let valid_report = validate_adaptive_config_or_err(&parsed_config).unwrap(); - assert!(!valid_report.has_errors()); - - let invalid_config: AdaptiveConfig = serde_json::from_value(json!({ - "version": 1, - "policy": { - "unknown_component": "warn", - "unknown_field": "warn", - "unsupported_value": "error" - }, - "tool_parallelism": { - "mode": "definitely_not_supported" - } - })) - .unwrap(); - let invalid_err = validate_adaptive_config_or_err(&invalid_config).unwrap_err(); - assert!(invalid_err.to_string().contains("unsupported")); - - assert!(matches!( - parse_cache_telemetry_provider("anthropic").unwrap(), - CacheTelemetryProvider::Anthropic - )); - assert!(matches!( - parse_cache_telemetry_provider("openai").unwrap(), - CacheTelemetryProvider::OpenAI - )); - assert!( - parse_cache_telemetry_provider("bogus") - .unwrap_err() - .to_string() - .contains("unsupported provider") - ); - assert!( - parse_cache_telemetry_request_id("not-a-uuid") - .unwrap_err() - .to_string() - .contains("invalid request_id UUID") - ); - assert!( - parse_cache_telemetry_timestamp(Some("not-a-timestamp")) - .unwrap_err() - .to_string() - .contains("invalid timestamp") - ); - assert!(parse_cache_telemetry_timestamp(None).is_ok()); + assert_py_error_contains(empty_runtime.bind_scope(&scope_handle), "already shut down"); + + fn assert_adaptive_config_and_telemetry_parsing(parsed_config: &AdaptiveConfig) { + let valid_report = validate_adaptive_config_or_err(parsed_config).unwrap(); + assert!(!valid_report.has_errors()); + + let invalid_config: AdaptiveConfig = serde_json::from_value(json!({ + "version": 1, + "policy": { + "unknown_component": "warn", + "unknown_field": "warn", + "unsupported_value": "error" + }, + "tool_parallelism": { + "mode": "definitely_not_supported" + } + })) + .unwrap(); + let invalid_err = validate_adaptive_config_or_err(&invalid_config).unwrap_err(); + assert!(invalid_err.to_string().contains("unsupported")); + + assert!(matches!( + parse_cache_telemetry_provider("anthropic").unwrap(), + CacheTelemetryProvider::Anthropic + )); + assert!(matches!( + parse_cache_telemetry_provider("openai").unwrap(), + CacheTelemetryProvider::OpenAI + )); + assert!( + parse_cache_telemetry_provider("bogus") + .unwrap_err() + .to_string() + .contains("unsupported provider") + ); + assert!( + parse_cache_telemetry_request_id("not-a-uuid") + .unwrap_err() + .to_string() + .contains("invalid request_id UUID") + ); + assert!( + parse_cache_telemetry_timestamp(Some("not-a-timestamp")) + .unwrap_err() + .to_string() + .contains("invalid timestamp") + ); + assert!(parse_cache_telemetry_timestamp(None).is_ok()); + } + assert_adaptive_config_and_telemetry_parsing(&parsed_config); let types_module = PyModule::new(py, "_adaptive_types").unwrap(); crate::py_types::register(&types_module).unwrap(); @@ -607,6 +577,11 @@ fn adaptive_runtime_locking_and_helper_errors_are_covered() { }); } +fn assert_py_error_contains(result: PyResult, expected: &str) { + let error = result.err().expect("operation should fail"); + assert!(error.to_string().contains(expected), "{error}"); +} + #[test] fn adaptive_runtime_shutdown_and_register_error_paths_are_covered() { let _python = crate::test_support::init_python_test(); diff --git a/crates/python/tests/coverage/py_api_coverage_tests.rs b/crates/python/tests/coverage/py_api_coverage_tests.rs index 7f2af7995..5a449b960 100644 --- a/crates/python/tests/coverage/py_api_coverage_tests.rs +++ b/crates/python/tests/coverage/py_api_coverage_tests.rs @@ -494,407 +494,482 @@ async def run_stream(api, request, func, collector, finalizer, handle, attribute ) .unwrap(); - let tool_intercepted = tool_request_intercepts( - py, - "demo-tool".to_string(), - &py_dict(py, json!({"value": 1})), - ) - .unwrap(); - assert_eq!( - crate::convert::py_to_json(&tool_intercepted).unwrap(), - json!({"value": 3}) - ); - tool_conditional_execution( - py, - "demo-tool".to_string(), - &py_dict(py, json!({"value": 1})), - ) - .unwrap(); - assert!( - tool_conditional_execution( + fn assert_python_api_execution_paths( + py: Python<'_>, + helpers: Bound<'_, PyModule>, + runner: Bound<'_, PyModule>, + api_module: Bound<'_, PyModule>, + types_module: Bound<'_, PyModule>, + child: PyScopeHandle, + ) { + let tool_intercepted = tool_request_intercepts( py, "demo-tool".to_string(), - &py_dict(py, json!({"value": -1})) + &py_dict(py, json!({"value": 1})), ) - .unwrap_err() - .to_string() - .contains("blocked") - ); - let async_sync_rejection_name = format!("async-sync-{}", Uuid::now_v7()); - register_tool_conditional_execution_guardrail( - &async_sync_rejection_name, - 20, - helpers.getattr("async_tool_conditional").unwrap().unbind(), - ) - .unwrap(); - assert!( + .unwrap(); + assert_eq!( + crate::convert::py_to_json(&tool_intercepted).unwrap(), + json!({"value": 3}) + ); tool_conditional_execution( py, "demo-tool".to_string(), &py_dict(py, json!({"value": 1})), ) - .unwrap_err() - .to_string() - .contains("requires an async caller") - ); - let llm_request = PyLLMRequest { - inner: nemo_relay::api::llm::LlmRequest { - headers: serde_json::Map::new(), - content: json!({"messages": [{"role": "user", "content": "hello"}], "model": "demo-model"}), - }, - }; - let intercepted_request = - llm_request_intercepts(py, "demo-llm".to_string(), llm_request.clone()).unwrap(); - let intercepted_request: PyRef<'_, crate::py_types::PyLLMRequestInterceptOutcome> = - intercepted_request.extract().unwrap(); - assert_eq!( - intercepted_request - .inner - .request - .headers - .get("x-intercepted"), - Some(&json!("1")) - ); - llm_conditional_execution(py, llm_request.clone()).unwrap(); - assert!( - llm_conditional_execution( - py, - PyLLMRequest { - inner: nemo_relay::api::llm::LlmRequest { - headers: serde_json::Map::new(), - content: json!({"messages": [], "model": "blocked"}), - }, - } + .unwrap(); + assert!( + tool_conditional_execution( + py, + "demo-tool".to_string(), + &py_dict(py, json!({"value": -1})) + ) + .unwrap_err() + .to_string() + .contains("blocked") + ); + let async_sync_rejection_name = format!("async-sync-{}", Uuid::now_v7()); + register_tool_conditional_execution_guardrail( + &async_sync_rejection_name, + 20, + helpers.getattr("async_tool_conditional").unwrap().unbind(), ) - .unwrap_err() - .to_string() - .contains("blocked") - ); - - with_event_loop(py, |event_loop| { - let standalone = event_loop - .call_method1( - "run_until_complete", - (runner - .getattr("run_standalone") - .unwrap() - .call1((api_module.clone(), llm_request.clone())) - .unwrap(),), + .unwrap(); + assert!( + tool_conditional_execution( + py, + "demo-tool".to_string(), + &py_dict(py, json!({"value": 1})), ) - .unwrap(); + .unwrap_err() + .to_string() + .contains("requires an async caller") + ); + let llm_request = PyLLMRequest { + inner: nemo_relay::api::llm::LlmRequest { + headers: serde_json::Map::new(), + content: json!({"messages": [{"role": "user", "content": "hello"}], "model": "demo-model"}), + }, + }; + let intercepted_request = + llm_request_intercepts(py, "demo-llm".to_string(), llm_request.clone()).unwrap(); + let intercepted_request: PyRef<'_, crate::py_types::PyLLMRequestInterceptOutcome> = + intercepted_request.extract().unwrap(); assert_eq!( - crate::convert::py_to_json(&standalone).unwrap(), - json!({"tool_value": 3, "conditional_allowed": true, "llm_header": "1"}) + intercepted_request + .inner + .request + .headers + .get("x-intercepted"), + Some(&json!("1")) ); - - let tool_result = event_loop - .call_method1( - "run_until_complete", - (runner - .getattr("run_tool") - .unwrap() - .call1(( - api_module.clone(), - helpers.getattr("tool_exec").unwrap(), - child.clone(), - PyToolAttributes { - inner: nemo_relay::api::tool::ToolAttributes::REMOTE, - }, - )) - .unwrap(),), + llm_conditional_execution(py, llm_request.clone()).unwrap(); + assert!( + llm_conditional_execution( + py, + PyLLMRequest { + inner: nemo_relay::api::llm::LlmRequest { + headers: serde_json::Map::new(), + content: json!({"messages": [], "model": "blocked"}), + }, + } ) - .unwrap(); - let tool_json = crate::convert::py_to_json(&tool_result).unwrap(); - assert_eq!(tool_json["tool_result"], json!(6)); - assert_eq!(tool_json["tool_intercepted"], json!(true)); - - let codec = helpers.getattr("EchoCodec").unwrap().call0().unwrap(); - let response_codec = types_module - .getattr("OpenAIChatCodec") - .unwrap() - .call0() - .unwrap(); - let llm_result = event_loop - .call_method1( - "run_until_complete", - (runner - .getattr("run_llm") - .unwrap() - .call1(( - api_module.clone(), - llm_request.clone(), - helpers.getattr("llm_exec").unwrap(), - child.clone(), - PyLLMAttributes { - inner: nemo_relay::api::llm::LlmAttributes::STATEFUL, - }, - codec, - response_codec, - )) - .unwrap(),), - ) - .unwrap(); - let llm_json = crate::convert::py_to_json(&llm_result).unwrap(); - assert_eq!(llm_json["id"], json!("chatcmpl-test")); - assert_eq!(llm_json["from_intercept"], json!(true)); - - let stream_codec = helpers.getattr("EchoCodec").unwrap().call0().unwrap(); - let stream_response_codec = types_module - .getattr("OpenAIChatCodec") - .unwrap() - .call0() - .unwrap(); - let stream_items = event_loop - .call_method1( - "run_until_complete", - (runner - .getattr("run_stream") - .unwrap() - .call1(( - api_module.clone(), - llm_request.clone(), - helpers.getattr("llm_stream_exec").unwrap(), - helpers.getattr("collector").unwrap(), - helpers.getattr("finalizer").unwrap(), - child.clone(), - PyLLMAttributes { - inner: nemo_relay::api::llm::LlmAttributes::STREAMING, - }, - stream_codec, - stream_response_codec, - )) - .unwrap(),), - ) - .unwrap(); - assert_eq!( - crate::convert::py_to_json(&stream_items).unwrap(), - json!([{"delta": 11}, {"delta": 12}]) + .unwrap_err() + .to_string() + .contains("blocked") ); - }); - assert!( - deregister_tool_conditional_execution_guardrail(&async_sync_rejection_name).unwrap() - ); - let events = helpers.getattr("events").unwrap(); - let events_json = crate::convert::py_to_json(events.as_any()).unwrap(); - assert!( - events_json - .as_array() - .unwrap() - .iter() - .any(|event| event[0] == "scope" && event[1] == "tool" && event[2] == "start") - ); - assert!( - events_json - .as_array() - .unwrap() - .iter() - .any(|event| event[0] == "scope" && event[1] == "llm" && event[2] == "end") - ); - - let chunks = helpers.getattr("chunks").unwrap(); - assert_eq!( - crate::convert::py_to_json(chunks.as_any()).unwrap(), - json!([11, 12]) + with_event_loop(py, |event_loop| { + let standalone = event_loop + .call_method1( + "run_until_complete", + (runner + .getattr("run_standalone") + .unwrap() + .call1((api_module.clone(), llm_request.clone())) + .unwrap(),), + ) + .unwrap(); + assert_eq!( + crate::convert::py_to_json(&standalone).unwrap(), + json!({"tool_value": 3, "conditional_allowed": true, "llm_header": "1"}) + ); + + let tool_result = event_loop + .call_method1( + "run_until_complete", + (runner + .getattr("run_tool") + .unwrap() + .call1(( + api_module.clone(), + helpers.getattr("tool_exec").unwrap(), + child.clone(), + PyToolAttributes { + inner: nemo_relay::api::tool::ToolAttributes::REMOTE, + }, + )) + .unwrap(),), + ) + .unwrap(); + let tool_json = crate::convert::py_to_json(&tool_result).unwrap(); + assert_eq!(tool_json["tool_result"], json!(6)); + assert_eq!(tool_json["tool_intercepted"], json!(true)); + + let codec = helpers.getattr("EchoCodec").unwrap().call0().unwrap(); + let response_codec = types_module + .getattr("OpenAIChatCodec") + .unwrap() + .call0() + .unwrap(); + let llm_result = event_loop + .call_method1( + "run_until_complete", + (runner + .getattr("run_llm") + .unwrap() + .call1(( + api_module.clone(), + llm_request.clone(), + helpers.getattr("llm_exec").unwrap(), + child.clone(), + PyLLMAttributes { + inner: nemo_relay::api::llm::LlmAttributes::STATEFUL, + }, + codec, + response_codec, + )) + .unwrap(),), + ) + .unwrap(); + let llm_json = crate::convert::py_to_json(&llm_result).unwrap(); + assert_eq!(llm_json["id"], json!("chatcmpl-test")); + assert_eq!(llm_json["from_intercept"], json!(true)); + + let stream_codec = helpers.getattr("EchoCodec").unwrap().call0().unwrap(); + let stream_response_codec = types_module + .getattr("OpenAIChatCodec") + .unwrap() + .call0() + .unwrap(); + let stream_items = event_loop + .call_method1( + "run_until_complete", + (runner + .getattr("run_stream") + .unwrap() + .call1(( + api_module.clone(), + llm_request.clone(), + helpers.getattr("llm_stream_exec").unwrap(), + helpers.getattr("collector").unwrap(), + helpers.getattr("finalizer").unwrap(), + child.clone(), + PyLLMAttributes { + inner: nemo_relay::api::llm::LlmAttributes::STREAMING, + }, + stream_codec, + stream_response_codec, + )) + .unwrap(),), + ) + .unwrap(); + assert_eq!( + crate::convert::py_to_json(&stream_items).unwrap(), + json!([{"delta": 11}, {"delta": 12}]) + ); + }); + assert!( + deregister_tool_conditional_execution_guardrail(&async_sync_rejection_name) + .unwrap() + ); + } + assert_python_api_execution_paths( + py, + helpers.clone(), + runner, + api_module, + types_module.clone(), + child.clone(), ); - let scope_tool_sanitize_request_name = format!("scope-tsrq-{}", Uuid::now_v7()); - let scope_tool_sanitize_response_name = format!("scope-tsrs-{}", Uuid::now_v7()); - let scope_tool_conditional_name = format!("scope-tcond-{}", Uuid::now_v7()); - let scope_tool_request_name = format!("scope-treq-{}", Uuid::now_v7()); - let scope_tool_exec_name = format!("scope-texec-{}", Uuid::now_v7()); - let scope_llm_sanitize_request_name = format!("scope-lsrq-{}", Uuid::now_v7()); - let scope_llm_sanitize_response_name = format!("scope-lsrs-{}", Uuid::now_v7()); - let scope_llm_conditional_name = format!("scope-lcond-{}", Uuid::now_v7()); - let scope_llm_request_name = format!("scope-lreq-{}", Uuid::now_v7()); - let scope_llm_exec_name = format!("scope-lexec-{}", Uuid::now_v7()); - let scope_llm_stream_name = format!("scope-lstream-{}", Uuid::now_v7()); - let scope_subscriber = format!("scope-sub-{}", Uuid::now_v7()); - - scope_register_tool_sanitize_request_guardrail( - &child_uuid, - &scope_tool_sanitize_request_name, - 5, - helpers.getattr("tool_sanitize_request").unwrap().unbind(), - ) - .unwrap(); - scope_register_tool_sanitize_response_guardrail( - &child_uuid, - &scope_tool_sanitize_response_name, - 5, - helpers.getattr("tool_sanitize_response").unwrap().unbind(), - ) - .unwrap(); - scope_register_tool_conditional_execution_guardrail( - &child_uuid, - &scope_tool_conditional_name, - 5, - helpers.getattr("tool_conditional").unwrap().unbind(), - ) - .unwrap(); - scope_register_tool_request_intercept( - &child_uuid, - &scope_tool_request_name, - 5, - false, - helpers.getattr("tool_request_intercept").unwrap().unbind(), - ) - .unwrap(); - scope_register_tool_execution_intercept( - &child_uuid, - &scope_tool_exec_name, - 5, - helpers.getattr("tool_exec_intercept").unwrap().unbind(), - ) - .unwrap(); - scope_register_llm_sanitize_request_guardrail( - &child_uuid, - &scope_llm_sanitize_request_name, - 5, - helpers.getattr("llm_sanitize_request").unwrap().unbind(), - ) - .unwrap(); - scope_register_llm_sanitize_response_guardrail( - &child_uuid, - &scope_llm_sanitize_response_name, - 5, - helpers.getattr("llm_sanitize_response").unwrap().unbind(), - ) - .unwrap(); - scope_register_llm_conditional_execution_guardrail( - &child_uuid, - &scope_llm_conditional_name, - 5, - helpers.getattr("llm_conditional").unwrap().unbind(), - ) - .unwrap(); - scope_register_llm_request_intercept( - &child_uuid, - &scope_llm_request_name, - 5, - false, - helpers.getattr("llm_request_intercept").unwrap().unbind(), - ) - .unwrap(); - scope_register_llm_execution_intercept( - &child_uuid, - &scope_llm_exec_name, - 5, - helpers.getattr("llm_exec_intercept").unwrap().unbind(), - ) - .unwrap(); - scope_register_llm_stream_execution_intercept( - &child_uuid, - &scope_llm_stream_name, - 5, - helpers.getattr("llm_stream_intercept").unwrap().unbind(), - ) - .unwrap(); - scope_register_subscriber( - &child_uuid, - &scope_subscriber, - helpers.getattr("subscriber").unwrap().unbind(), - ) - .unwrap(); + fn assert_python_api_emitted_events(helpers: &Bound<'_, PyModule>) { + let events = helpers.getattr("events").unwrap(); + let events_json = crate::convert::py_to_json(events.as_any()).unwrap(); + assert!( + events_json + .as_array() + .unwrap() + .iter() + .any(|event| event[0] == "scope" && event[1] == "tool" && event[2] == "start") + ); + assert!( + events_json + .as_array() + .unwrap() + .iter() + .any(|event| event[0] == "scope" && event[1] == "llm" && event[2] == "end") + ); - assert!( - scope_register_subscriber( - "not-a-uuid", - "bad", - helpers.getattr("subscriber").unwrap().unbind(), + let chunks = helpers.getattr("chunks").unwrap(); + assert_eq!( + crate::convert::py_to_json(chunks.as_any()).unwrap(), + json!([11, 12]) + ); + } + assert_python_api_emitted_events(&helpers); + + fn assert_scope_registry_paths(child_uuid: &str, helpers: &Bound<'_, PyModule>) { + let scope_tool_sanitize_request_name = format!("scope-tsrq-{}", Uuid::now_v7()); + let scope_tool_sanitize_response_name = format!("scope-tsrs-{}", Uuid::now_v7()); + let scope_tool_conditional_name = format!("scope-tcond-{}", Uuid::now_v7()); + let scope_tool_request_name = format!("scope-treq-{}", Uuid::now_v7()); + let scope_tool_exec_name = format!("scope-texec-{}", Uuid::now_v7()); + let scope_llm_sanitize_request_name = format!("scope-lsrq-{}", Uuid::now_v7()); + let scope_llm_sanitize_response_name = format!("scope-lsrs-{}", Uuid::now_v7()); + let scope_llm_conditional_name = format!("scope-lcond-{}", Uuid::now_v7()); + let scope_llm_request_name = format!("scope-lreq-{}", Uuid::now_v7()); + let scope_llm_exec_name = format!("scope-lexec-{}", Uuid::now_v7()); + let scope_llm_stream_name = format!("scope-lstream-{}", Uuid::now_v7()); + let scope_subscriber = format!("scope-sub-{}", Uuid::now_v7()); + + scope_register_tool_sanitize_request_guardrail( + child_uuid, + &scope_tool_sanitize_request_name, + 5, + helpers.getattr("tool_sanitize_request").unwrap().unbind(), ) - .unwrap_err() - .to_string() - .contains("invalid UUID") - ); - - assert!( - scope_deregister_tool_sanitize_request_guardrail( - &child_uuid, - &scope_tool_sanitize_request_name + .unwrap(); + scope_register_tool_sanitize_response_guardrail( + child_uuid, + &scope_tool_sanitize_response_name, + 5, + helpers.getattr("tool_sanitize_response").unwrap().unbind(), ) - .unwrap() - ); - assert!( - scope_deregister_tool_sanitize_response_guardrail( - &child_uuid, - &scope_tool_sanitize_response_name + .unwrap(); + scope_register_tool_conditional_execution_guardrail( + child_uuid, + &scope_tool_conditional_name, + 5, + helpers.getattr("tool_conditional").unwrap().unbind(), ) - .unwrap() - ); - assert!( - scope_deregister_tool_conditional_execution_guardrail( - &child_uuid, - &scope_tool_conditional_name + .unwrap(); + scope_register_tool_request_intercept( + child_uuid, + &scope_tool_request_name, + 5, + false, + helpers.getattr("tool_request_intercept").unwrap().unbind(), ) - .unwrap() - ); - assert!( - scope_deregister_tool_request_intercept(&child_uuid, &scope_tool_request_name).unwrap() - ); - assert!( - scope_deregister_tool_execution_intercept(&child_uuid, &scope_tool_exec_name).unwrap() - ); - assert!( - scope_deregister_llm_sanitize_request_guardrail( - &child_uuid, - &scope_llm_sanitize_request_name + .unwrap(); + scope_register_tool_execution_intercept( + child_uuid, + &scope_tool_exec_name, + 5, + helpers.getattr("tool_exec_intercept").unwrap().unbind(), ) - .unwrap() - ); - assert!( - scope_deregister_llm_sanitize_response_guardrail( - &child_uuid, - &scope_llm_sanitize_response_name + .unwrap(); + scope_register_llm_sanitize_request_guardrail( + child_uuid, + &scope_llm_sanitize_request_name, + 5, + helpers.getattr("llm_sanitize_request").unwrap().unbind(), ) - .unwrap() - ); - assert!( - scope_deregister_llm_conditional_execution_guardrail( - &child_uuid, - &scope_llm_conditional_name + .unwrap(); + scope_register_llm_sanitize_response_guardrail( + child_uuid, + &scope_llm_sanitize_response_name, + 5, + helpers.getattr("llm_sanitize_response").unwrap().unbind(), ) - .unwrap() - ); - assert!( - scope_deregister_llm_request_intercept(&child_uuid, &scope_llm_request_name).unwrap() - ); - assert!( - scope_deregister_llm_execution_intercept(&child_uuid, &scope_llm_exec_name).unwrap() - ); - assert!( - scope_deregister_llm_stream_execution_intercept(&child_uuid, &scope_llm_stream_name) + .unwrap(); + scope_register_llm_conditional_execution_guardrail( + child_uuid, + &scope_llm_conditional_name, + 5, + helpers.getattr("llm_conditional").unwrap().unbind(), + ) + .unwrap(); + scope_register_llm_request_intercept( + child_uuid, + &scope_llm_request_name, + 5, + false, + helpers.getattr("llm_request_intercept").unwrap().unbind(), + ) + .unwrap(); + scope_register_llm_execution_intercept( + child_uuid, + &scope_llm_exec_name, + 5, + helpers.getattr("llm_exec_intercept").unwrap().unbind(), + ) + .unwrap(); + scope_register_llm_stream_execution_intercept( + child_uuid, + &scope_llm_stream_name, + 5, + helpers.getattr("llm_stream_intercept").unwrap().unbind(), + ) + .unwrap(); + scope_register_subscriber( + child_uuid, + &scope_subscriber, + helpers.getattr("subscriber").unwrap().unbind(), + ) + .unwrap(); + + assert!( + scope_register_subscriber( + "not-a-uuid", + "bad", + helpers.getattr("subscriber").unwrap().unbind(), + ) + .unwrap_err() + .to_string() + .contains("invalid UUID") + ); + + assert!( + scope_deregister_tool_sanitize_request_guardrail( + child_uuid, + &scope_tool_sanitize_request_name + ) + .unwrap() + ); + assert!( + scope_deregister_tool_sanitize_response_guardrail( + child_uuid, + &scope_tool_sanitize_response_name + ) + .unwrap() + ); + assert!( + scope_deregister_tool_conditional_execution_guardrail( + child_uuid, + &scope_tool_conditional_name + ) .unwrap() + ); + assert!( + scope_deregister_tool_request_intercept(child_uuid, &scope_tool_request_name) + .unwrap() + ); + assert!( + scope_deregister_tool_execution_intercept(child_uuid, &scope_tool_exec_name) + .unwrap() + ); + assert!( + scope_deregister_llm_sanitize_request_guardrail( + child_uuid, + &scope_llm_sanitize_request_name + ) + .unwrap() + ); + assert!( + scope_deregister_llm_sanitize_response_guardrail( + child_uuid, + &scope_llm_sanitize_response_name + ) + .unwrap() + ); + assert!( + scope_deregister_llm_conditional_execution_guardrail( + child_uuid, + &scope_llm_conditional_name + ) + .unwrap() + ); + assert!( + scope_deregister_llm_request_intercept(child_uuid, &scope_llm_request_name) + .unwrap() + ); + assert!( + scope_deregister_llm_execution_intercept(child_uuid, &scope_llm_exec_name).unwrap() + ); + assert!( + scope_deregister_llm_stream_execution_intercept(child_uuid, &scope_llm_stream_name) + .unwrap() + ); + assert!(scope_deregister_subscriber(child_uuid, &scope_subscriber).unwrap()); + } + assert_scope_registry_paths(&child_uuid, &helpers); + + fn assert_global_tool_deregistration( + tool_sanitize_request_name: &str, + tool_sanitize_response_name: &str, + tool_conditional_name: &str, + tool_request_name: &str, + tool_exec_name: &str, + ) { + assert!( + deregister_tool_sanitize_request_guardrail(tool_sanitize_request_name).unwrap() + ); + assert!( + !deregister_tool_sanitize_request_guardrail(tool_sanitize_request_name).unwrap() + ); + assert!( + deregister_tool_sanitize_response_guardrail(tool_sanitize_response_name).unwrap() + ); + assert!( + !deregister_tool_sanitize_response_guardrail(tool_sanitize_response_name).unwrap() + ); + assert!( + deregister_tool_conditional_execution_guardrail(tool_conditional_name).unwrap() + ); + assert!( + !deregister_tool_conditional_execution_guardrail(tool_conditional_name).unwrap() + ); + assert!(deregister_tool_request_intercept(tool_request_name).unwrap()); + assert!(!deregister_tool_request_intercept(tool_request_name).unwrap()); + assert!(deregister_tool_execution_intercept(tool_exec_name).unwrap()); + assert!(!deregister_tool_execution_intercept(tool_exec_name).unwrap()); + } + assert_global_tool_deregistration( + &tool_sanitize_request_name, + &tool_sanitize_response_name, + &tool_conditional_name, + &tool_request_name, + &tool_exec_name, ); - assert!(scope_deregister_subscriber(&child_uuid, &scope_subscriber).unwrap()); - assert!(deregister_tool_sanitize_request_guardrail(&tool_sanitize_request_name).unwrap()); - assert!(!deregister_tool_sanitize_request_guardrail(&tool_sanitize_request_name).unwrap()); - assert!(deregister_tool_sanitize_response_guardrail(&tool_sanitize_response_name).unwrap()); - assert!( - !deregister_tool_sanitize_response_guardrail(&tool_sanitize_response_name).unwrap() + fn assert_global_llm_and_subscriber_deregistration( + llm_sanitize_request_name: &str, + llm_sanitize_response_name: &str, + llm_conditional_name: &str, + llm_request_name: &str, + llm_exec_name: &str, + llm_stream_name: &str, + global_subscriber: &str, + ) { + assert!(deregister_llm_sanitize_request_guardrail(llm_sanitize_request_name).unwrap()); + assert!(!deregister_llm_sanitize_request_guardrail(llm_sanitize_request_name).unwrap()); + assert!( + deregister_llm_sanitize_response_guardrail(llm_sanitize_response_name).unwrap() + ); + assert!( + !deregister_llm_sanitize_response_guardrail(llm_sanitize_response_name).unwrap() + ); + assert!(deregister_llm_conditional_execution_guardrail(llm_conditional_name).unwrap()); + assert!(!deregister_llm_conditional_execution_guardrail(llm_conditional_name).unwrap()); + assert!(deregister_llm_request_intercept(llm_request_name).unwrap()); + assert!(!deregister_llm_request_intercept(llm_request_name).unwrap()); + assert!(deregister_llm_execution_intercept(llm_exec_name).unwrap()); + assert!(!deregister_llm_execution_intercept(llm_exec_name).unwrap()); + assert!(deregister_llm_stream_execution_intercept(llm_stream_name).unwrap()); + assert!(!deregister_llm_stream_execution_intercept(llm_stream_name).unwrap()); + assert!(deregister_subscriber(global_subscriber).unwrap()); + assert!(!deregister_subscriber(global_subscriber).unwrap()); + } + assert_global_llm_and_subscriber_deregistration( + &llm_sanitize_request_name, + &llm_sanitize_response_name, + &llm_conditional_name, + &llm_request_name, + &llm_exec_name, + &llm_stream_name, + &global_subscriber, ); - assert!(deregister_tool_conditional_execution_guardrail(&tool_conditional_name).unwrap()); - assert!(!deregister_tool_conditional_execution_guardrail(&tool_conditional_name).unwrap()); - assert!(deregister_tool_request_intercept(&tool_request_name).unwrap()); - assert!(!deregister_tool_request_intercept(&tool_request_name).unwrap()); - assert!(deregister_tool_execution_intercept(&tool_exec_name).unwrap()); - assert!(!deregister_tool_execution_intercept(&tool_exec_name).unwrap()); - - assert!(deregister_llm_sanitize_request_guardrail(&llm_sanitize_request_name).unwrap()); - assert!(!deregister_llm_sanitize_request_guardrail(&llm_sanitize_request_name).unwrap()); - assert!(deregister_llm_sanitize_response_guardrail(&llm_sanitize_response_name).unwrap()); - assert!(!deregister_llm_sanitize_response_guardrail(&llm_sanitize_response_name).unwrap()); - assert!(deregister_llm_conditional_execution_guardrail(&llm_conditional_name).unwrap()); - assert!(!deregister_llm_conditional_execution_guardrail(&llm_conditional_name).unwrap()); - assert!(deregister_llm_request_intercept(&llm_request_name).unwrap()); - assert!(!deregister_llm_request_intercept(&llm_request_name).unwrap()); - assert!(deregister_llm_execution_intercept(&llm_exec_name).unwrap()); - assert!(!deregister_llm_execution_intercept(&llm_exec_name).unwrap()); - assert!(deregister_llm_stream_execution_intercept(&llm_stream_name).unwrap()); - assert!(!deregister_llm_stream_execution_intercept(&llm_stream_name).unwrap()); - assert!(deregister_subscriber(&global_subscriber).unwrap()); - assert!(!deregister_subscriber(&global_subscriber).unwrap()); pop_scope(py, &child, None, None, None).unwrap(); }); diff --git a/crates/python/tests/coverage/py_plugin_coverage_tests.rs b/crates/python/tests/coverage/py_plugin_coverage_tests.rs index 9353d0bcf..88127cdea 100644 --- a/crates/python/tests/coverage/py_plugin_coverage_tests.rs +++ b/crates/python/tests/coverage/py_plugin_coverage_tests.rs @@ -1006,175 +1006,193 @@ async def tool_execution_intercept(name, value, next): namespace_prefix: "poison.".to_string(), }; - assert!( - context - .drain_registrations() - .unwrap_err() - .to_string() - .contains("lock poisoned") - ); + fn assert_poisoned_tool_registrations( + context: &PyPluginContext, + helpers: &Bound<'_, PyModule>, + ) { + assert!( + context + .drain_registrations() + .unwrap_err() + .to_string() + .contains("lock poisoned") + ); - assert!( - context - .register_subscriber( - "subscriber", - helpers.getattr("subscriber").unwrap().unbind() - ) - .unwrap_err() - .to_string() - .contains("lock poisoned") - ); - assert!(deregister_subscriber("poison.subscriber").unwrap()); + assert!( + context + .register_subscriber( + "subscriber", + helpers.getattr("subscriber").unwrap().unbind() + ) + .unwrap_err() + .to_string() + .contains("lock poisoned") + ); + assert!(deregister_subscriber("poison.subscriber").unwrap()); - assert!( - context - .register_tool_sanitize_request_guardrail( - "tool_req", - 1, - helpers.getattr("tool_fn").unwrap().unbind(), - ) - .unwrap_err() - .to_string() - .contains("lock poisoned") - ); - assert!(deregister_tool_sanitize_request_guardrail("poison.tool_req").unwrap()); + assert!( + context + .register_tool_sanitize_request_guardrail( + "tool_req", + 1, + helpers.getattr("tool_fn").unwrap().unbind(), + ) + .unwrap_err() + .to_string() + .contains("lock poisoned") + ); + assert!(deregister_tool_sanitize_request_guardrail("poison.tool_req").unwrap()); - assert!( - context - .register_tool_sanitize_response_guardrail( - "tool_resp", - 1, - helpers.getattr("tool_fn").unwrap().unbind(), - ) - .unwrap_err() - .to_string() - .contains("lock poisoned") - ); - assert!(deregister_tool_sanitize_response_guardrail("poison.tool_resp").unwrap()); + assert!( + context + .register_tool_sanitize_response_guardrail( + "tool_resp", + 1, + helpers.getattr("tool_fn").unwrap().unbind(), + ) + .unwrap_err() + .to_string() + .contains("lock poisoned") + ); + assert!(deregister_tool_sanitize_response_guardrail("poison.tool_resp").unwrap()); - assert!( - context - .register_tool_conditional_execution_guardrail( - "tool_cond", - 1, - helpers.getattr("tool_conditional").unwrap().unbind(), - ) - .unwrap_err() - .to_string() - .contains("lock poisoned") - ); - assert!(deregister_tool_conditional_execution_guardrail("poison.tool_cond").unwrap()); + assert!( + context + .register_tool_conditional_execution_guardrail( + "tool_cond", + 1, + helpers.getattr("tool_conditional").unwrap().unbind(), + ) + .unwrap_err() + .to_string() + .contains("lock poisoned") + ); + assert!(deregister_tool_conditional_execution_guardrail("poison.tool_cond").unwrap()); + } + assert_poisoned_tool_registrations(&context, &helpers); - assert!( - context - .register_llm_sanitize_request_guardrail( - "llm_req", - 1, - helpers.getattr("llm_sanitize_request").unwrap().unbind(), - ) - .unwrap_err() - .to_string() - .contains("lock poisoned") - ); - assert!(deregister_llm_sanitize_request_guardrail("poison.llm_req").unwrap()); + fn assert_poisoned_llm_guardrail_registrations( + context: &PyPluginContext, + helpers: &Bound<'_, PyModule>, + ) { + assert!( + context + .register_llm_sanitize_request_guardrail( + "llm_req", + 1, + helpers.getattr("llm_sanitize_request").unwrap().unbind(), + ) + .unwrap_err() + .to_string() + .contains("lock poisoned") + ); + assert!(deregister_llm_sanitize_request_guardrail("poison.llm_req").unwrap()); - assert!( - context - .register_llm_sanitize_response_guardrail( - "llm_resp", - 1, - helpers.getattr("llm_sanitize_response").unwrap().unbind(), - ) - .unwrap_err() - .to_string() - .contains("lock poisoned") - ); - assert!(deregister_llm_sanitize_response_guardrail("poison.llm_resp").unwrap()); + assert!( + context + .register_llm_sanitize_response_guardrail( + "llm_resp", + 1, + helpers.getattr("llm_sanitize_response").unwrap().unbind(), + ) + .unwrap_err() + .to_string() + .contains("lock poisoned") + ); + assert!(deregister_llm_sanitize_response_guardrail("poison.llm_resp").unwrap()); - assert!( - context - .register_llm_conditional_execution_guardrail( - "llm_cond", - 1, - helpers.getattr("llm_conditional").unwrap().unbind(), - ) - .unwrap_err() - .to_string() - .contains("lock poisoned") - ); - assert!(deregister_llm_conditional_execution_guardrail("poison.llm_cond").unwrap()); + assert!( + context + .register_llm_conditional_execution_guardrail( + "llm_cond", + 1, + helpers.getattr("llm_conditional").unwrap().unbind(), + ) + .unwrap_err() + .to_string() + .contains("lock poisoned") + ); + assert!(deregister_llm_conditional_execution_guardrail("poison.llm_cond").unwrap()); + } + assert_poisoned_llm_guardrail_registrations(&context, &helpers); - assert!( - context - .register_llm_request_intercept( - "llm_request", - 1, - false, - helpers.getattr("llm_request_intercept").unwrap().unbind(), - ) - .unwrap_err() - .to_string() - .contains("lock poisoned") - ); - assert!(deregister_llm_request_intercept("poison.llm_request").unwrap()); + fn assert_poisoned_intercept_registrations( + context: &PyPluginContext, + helpers: &Bound<'_, PyModule>, + ) { + assert!( + context + .register_llm_request_intercept( + "llm_request", + 1, + false, + helpers.getattr("llm_request_intercept").unwrap().unbind(), + ) + .unwrap_err() + .to_string() + .contains("lock poisoned") + ); + assert!(deregister_llm_request_intercept("poison.llm_request").unwrap()); - assert!( - context - .register_llm_execution_intercept( - "llm_exec", - 1, - helpers.getattr("llm_execution_intercept").unwrap().unbind(), - ) - .unwrap_err() - .to_string() - .contains("lock poisoned") - ); - assert!(deregister_llm_execution_intercept("poison.llm_exec").unwrap()); + assert!( + context + .register_llm_execution_intercept( + "llm_exec", + 1, + helpers.getattr("llm_execution_intercept").unwrap().unbind(), + ) + .unwrap_err() + .to_string() + .contains("lock poisoned") + ); + assert!(deregister_llm_execution_intercept("poison.llm_exec").unwrap()); - assert!( - context - .register_llm_stream_execution_intercept( - "llm_stream", - 1, - helpers - .getattr("llm_stream_execution_intercept") - .unwrap() - .unbind(), - ) - .unwrap_err() - .to_string() - .contains("lock poisoned") - ); - assert!(deregister_llm_stream_execution_intercept("poison.llm_stream").unwrap()); + assert!( + context + .register_llm_stream_execution_intercept( + "llm_stream", + 1, + helpers + .getattr("llm_stream_execution_intercept") + .unwrap() + .unbind(), + ) + .unwrap_err() + .to_string() + .contains("lock poisoned") + ); + assert!(deregister_llm_stream_execution_intercept("poison.llm_stream").unwrap()); - assert!( - context - .register_tool_request_intercept( - "tool_request", - 1, - false, - helpers.getattr("tool_request_intercept").unwrap().unbind(), - ) - .unwrap_err() - .to_string() - .contains("lock poisoned") - ); - assert!(deregister_tool_request_intercept("poison.tool_request").unwrap()); + assert!( + context + .register_tool_request_intercept( + "tool_request", + 1, + false, + helpers.getattr("tool_request_intercept").unwrap().unbind(), + ) + .unwrap_err() + .to_string() + .contains("lock poisoned") + ); + assert!(deregister_tool_request_intercept("poison.tool_request").unwrap()); - assert!( - context - .register_tool_execution_intercept( - "tool_exec", - 1, - helpers - .getattr("tool_execution_intercept") - .unwrap() - .unbind(), - ) - .unwrap_err() - .to_string() - .contains("lock poisoned") - ); - assert!(deregister_tool_execution_intercept("poison.tool_exec").unwrap()); + assert!( + context + .register_tool_execution_intercept( + "tool_exec", + 1, + helpers + .getattr("tool_execution_intercept") + .unwrap() + .unbind(), + ) + .unwrap_err() + .to_string() + .contains("lock poisoned") + ); + assert!(deregister_tool_execution_intercept("poison.tool_exec").unwrap()); + } + assert_poisoned_intercept_registrations(&context, &helpers); }); } diff --git a/crates/python/tests/coverage/py_types_coverage_tests.rs b/crates/python/tests/coverage/py_types_coverage_tests.rs index abebfbc3f..c539d8703 100644 --- a/crates/python/tests/coverage/py_types_coverage_tests.rs +++ b/crates/python/tests/coverage/py_types_coverage_tests.rs @@ -54,6 +54,20 @@ fn base_event( .build() } +fn py_scope_event(event: Event) -> PyScopeEvent { + match event { + Event::Scope(inner) => PyScopeEvent { inner }, + _ => panic!("expected scope event"), + } +} + +fn py_mark_event(event: Event) -> PyMarkEvent { + match event { + Event::Mark(inner) => PyMarkEvent { inner }, + _ => panic!("expected mark event"), + } +} + #[test] fn test_register_exposes_all_type_bindings() { let _python = crate::test_support::init_python_test(); @@ -61,26 +75,31 @@ fn test_register_exposes_all_type_bindings() { let module = PyModule::new(py, "_types_test").unwrap(); register(&module).unwrap(); - assert!(module.getattr("ScopeStack").is_ok()); - assert!(module.getattr("LlmStream").is_ok()); - assert!(module.getattr("ScopeAttributes").is_ok()); - assert!(module.getattr("ToolAttributes").is_ok()); - assert!(module.getattr("LLMAttributes").is_ok()); - assert!(module.getattr("ScopeType").is_ok()); - assert!(module.getattr("ScopeHandle").is_ok()); - assert!(module.getattr("ToolHandle").is_ok()); - assert!(module.getattr("LLMHandle").is_ok()); - assert!(module.getattr("LLMRequest").is_ok()); - assert!(module.getattr("ScopeEvent").is_ok()); - assert!(module.getattr("MarkEvent").is_ok()); - assert!(module.getattr("AtifExporter").is_ok()); - assert!(module.getattr("OpenInferenceConfig").is_err()); - assert!(module.getattr("OpenInferenceSubscriber").is_err()); - assert!(module.getattr("OpenTelemetryConfig").is_ok()); - assert!(module.getattr("OpenTelemetrySubscriber").is_ok()); - assert!(module.getattr("OpenAIChatCodec").is_ok()); - assert!(module.getattr("OpenAIResponsesCodec").is_ok()); - assert!(module.getattr("AnthropicMessagesCodec").is_ok()); + for name in [ + "ScopeStack", + "LlmStream", + "ScopeAttributes", + "ToolAttributes", + "LLMAttributes", + "ScopeType", + "ScopeHandle", + "ToolHandle", + "LLMHandle", + "LLMRequest", + "ScopeEvent", + "MarkEvent", + "AtifExporter", + "OpenTelemetryConfig", + "OpenTelemetrySubscriber", + "OpenAIChatCodec", + "OpenAIResponsesCodec", + "AnthropicMessagesCodec", + ] { + assert!(module.getattr(name).is_ok(), "{name} should be registered"); + } + for name in ["OpenInferenceConfig", "OpenInferenceSubscriber"] { + assert!(module.getattr(name).is_err(), "{name} should be absent"); + } }); } @@ -133,169 +152,188 @@ fn test_bitflags_handles_and_event_wrappers_expose_expected_fields() { Python::attach(|py| { let parent_uuid = Uuid::now_v7(); - let scope = PyScopeHandle::from( - ScopeHandle::builder() - .name("scope") - .scope_type(ScopeType::Tool) - .attributes(ScopeAttributes::PARALLEL) - .parent_uuid(parent_uuid) - .data(json!({"scope": true})) - .metadata(json!({"meta": "scope"})) - .build(), - ); - assert!(scope.scope_type() == PyScopeType::Tool); - assert_eq!(scope.parent_uuid(), Some(parent_uuid.to_string())); - assert_eq!( - py_to_json(scope.data(py).unwrap().bind(py)).unwrap(), - json!({"scope": true}) - ); - assert_eq!( - py_to_json(scope.metadata(py).unwrap().bind(py)).unwrap(), - json!({"meta": "scope"}) - ); - assert!(scope.__repr__().contains("ScopeHandle")); - - let tool = PyToolHandle::from( - ToolHandle::builder() - .name("tool") - .attributes(ToolAttributes::REMOTE) - .parent_uuid(parent_uuid) - .data(json!({"tool": true})) - .metadata(json!({"meta": "tool"})) - .build(), - ); - assert_eq!(tool.parent_uuid(), Some(parent_uuid.to_string())); - assert_eq!(tool.attributes().value(), PyToolAttributes::REMOTE); - assert_eq!( - py_to_json(tool.data(py).unwrap().bind(py)).unwrap(), - json!({"tool": true}) - ); - assert_eq!( - py_to_json(tool.metadata(py).unwrap().bind(py)).unwrap(), - json!({"meta": "tool"}) - ); - assert!(tool.__repr__().contains("ToolHandle")); - - let llm = PyLLMHandle::from( - LlmHandle::builder() - .name("llm") - .attributes(LlmAttributes::STATEFUL | LlmAttributes::STREAMING) - .parent_uuid(parent_uuid) - .data(json!({"llm": true})) - .metadata(json!({"meta": "llm"})) - .build(), - ); - assert_eq!(llm.parent_uuid(), Some(parent_uuid.to_string())); - assert_eq!( - llm.attributes().value(), - PyLLMAttributes::STATEFUL | PyLLMAttributes::STREAMING - ); - assert_eq!( - py_to_json(llm.data(py).unwrap().bind(py)).unwrap(), - json!({"llm": true}) - ); - assert_eq!( - py_to_json(llm.metadata(py).unwrap().bind(py)).unwrap(), - json!({"meta": "llm"}) - ); - assert!(llm.__repr__().contains("LLMHandle")); - let request = PyLLMRequest { - inner: LlmRequest { - headers: serde_json::Map::from_iter([("x-trace".into(), json!("1"))]), - content: json!({"prompt": "hello"}), - }, - }; - assert_eq!( - py_to_json(request.headers(py).unwrap().bind(py)).unwrap(), - json!({"x-trace": "1"}) - ); - assert_eq!( - py_to_json(request.content(py).unwrap().bind(py)).unwrap(), - json!({"prompt": "hello"}) - ); - assert_eq!(request.__repr__(), "LLMRequest(...)"); - - let event = match Event::Mark(MarkEvent::new( - base_event( - parent_uuid, - "event", - json!({"event": true}), - json!({"meta": "event"}), - ), - None, - None, - )) { - Event::Mark(inner) => PyMarkEvent { inner }, - _ => unreachable!(), - }; - assert_eq!(event.kind(), "mark"); - assert_eq!(event.parent_uuid(), Some(parent_uuid.to_string())); - assert_eq!( - py_to_json(event.data(py).unwrap().bind(py)).unwrap(), - json!({"event": true}) - ); - assert_eq!( - py_to_json(event.metadata(py).unwrap().bind(py)).unwrap(), - json!({"meta": "event"}) - ); - assert!(event.timestamp().contains('T')); - - let tool_event = match Event::Scope(ScopeEvent::new( - base_event( - parent_uuid, - "tool-event", - json!({"input": true}), - json!({"meta": "event"}), - ), - ScopeCategory::Start, - tool_attributes_to_strings(ToolAttributes::REMOTE), - EventCategory::tool(), - Some(CategoryProfile::builder().tool_call_id("tool-1").build()), - )) { - Event::Scope(inner) => PyScopeEvent { inner }, - _ => unreachable!(), - }; - assert_eq!(tool_event.kind(), "scope"); - assert_eq!(tool_event.scope_category(), "start"); - assert_eq!(tool_event.category(), "tool"); - assert_eq!( - py_to_json(tool_event.data(py).unwrap().bind(py)).unwrap(), - json!({"input": true}) - ); - assert_eq!( - py_to_json(tool_event.category_profile(py).unwrap().bind(py)).unwrap(), - json!({"tool_call_id": "tool-1"}) - ); - assert_eq!(tool_event.attributes(), vec!["remote"]); - - let llm_event = match Event::Scope(ScopeEvent::new( - base_event( - parent_uuid, - "llm-event", - json!({"output": true}), - json!({"meta": "event"}), - ), - ScopeCategory::End, - llm_attributes_to_strings(LlmAttributes::STATEFUL), - EventCategory::llm(), - Some(CategoryProfile::builder().model_name("model").build()), - )) { - Event::Scope(inner) => PyScopeEvent { inner }, - _ => unreachable!(), - }; - assert_eq!(llm_event.kind(), "scope"); - assert_eq!(llm_event.scope_category(), "end"); - assert_eq!(llm_event.category(), "llm"); - assert_eq!( - py_to_json(llm_event.data(py).unwrap().bind(py)).unwrap(), - json!({"output": true}) - ); - assert_eq!( - py_to_json(llm_event.category_profile(py).unwrap().bind(py)).unwrap(), - json!({"model_name": "model"}) - ); - assert_eq!(llm_event.attributes(), vec!["stateful"]); + fn assert_scope_handle_fields(py: Python<'_>, parent_uuid: Uuid) { + let scope = PyScopeHandle::from( + ScopeHandle::builder() + .name("scope") + .scope_type(ScopeType::Tool) + .attributes(ScopeAttributes::PARALLEL) + .parent_uuid(parent_uuid) + .data(json!({"scope": true})) + .metadata(json!({"meta": "scope"})) + .build(), + ); + assert!(scope.scope_type() == PyScopeType::Tool); + assert_eq!(scope.parent_uuid(), Some(parent_uuid.to_string())); + assert_eq!( + py_to_json(scope.data(py).unwrap().bind(py)).unwrap(), + json!({"scope": true}) + ); + assert_eq!( + py_to_json(scope.metadata(py).unwrap().bind(py)).unwrap(), + json!({"meta": "scope"}) + ); + assert!(scope.__repr__().contains("ScopeHandle")); + } + assert_scope_handle_fields(py, parent_uuid); + + fn assert_tool_handle_fields(py: Python<'_>, parent_uuid: Uuid) { + let tool = PyToolHandle::from( + ToolHandle::builder() + .name("tool") + .attributes(ToolAttributes::REMOTE) + .parent_uuid(parent_uuid) + .data(json!({"tool": true})) + .metadata(json!({"meta": "tool"})) + .build(), + ); + assert_eq!(tool.parent_uuid(), Some(parent_uuid.to_string())); + assert_eq!(tool.attributes().value(), PyToolAttributes::REMOTE); + assert_eq!( + py_to_json(tool.data(py).unwrap().bind(py)).unwrap(), + json!({"tool": true}) + ); + assert_eq!( + py_to_json(tool.metadata(py).unwrap().bind(py)).unwrap(), + json!({"meta": "tool"}) + ); + assert!(tool.__repr__().contains("ToolHandle")); + } + assert_tool_handle_fields(py, parent_uuid); + + fn assert_llm_handle_fields(py: Python<'_>, parent_uuid: Uuid) { + let llm = PyLLMHandle::from( + LlmHandle::builder() + .name("llm") + .attributes(LlmAttributes::STATEFUL | LlmAttributes::STREAMING) + .parent_uuid(parent_uuid) + .data(json!({"llm": true})) + .metadata(json!({"meta": "llm"})) + .build(), + ); + assert_eq!(llm.parent_uuid(), Some(parent_uuid.to_string())); + assert_eq!( + llm.attributes().value(), + PyLLMAttributes::STATEFUL | PyLLMAttributes::STREAMING + ); + assert_eq!( + py_to_json(llm.data(py).unwrap().bind(py)).unwrap(), + json!({"llm": true}) + ); + assert_eq!( + py_to_json(llm.metadata(py).unwrap().bind(py)).unwrap(), + json!({"meta": "llm"}) + ); + assert!(llm.__repr__().contains("LLMHandle")); + } + assert_llm_handle_fields(py, parent_uuid); + + fn assert_llm_request_fields(py: Python<'_>) { + let request = PyLLMRequest { + inner: LlmRequest { + headers: serde_json::Map::from_iter([("x-trace".into(), json!("1"))]), + content: json!({"prompt": "hello"}), + }, + }; + assert_eq!( + py_to_json(request.headers(py).unwrap().bind(py)).unwrap(), + json!({"x-trace": "1"}) + ); + assert_eq!( + py_to_json(request.content(py).unwrap().bind(py)).unwrap(), + json!({"prompt": "hello"}) + ); + assert_eq!(request.__repr__(), "LLMRequest(...)"); + } + assert_llm_request_fields(py); + + fn assert_mark_and_tool_event_fields(py: Python<'_>, parent_uuid: Uuid) { + let event = match Event::Mark(MarkEvent::new( + base_event( + parent_uuid, + "event", + json!({"event": true}), + json!({"meta": "event"}), + ), + None, + None, + )) { + Event::Mark(inner) => PyMarkEvent { inner }, + _ => unreachable!(), + }; + assert_eq!(event.kind(), "mark"); + assert_eq!(event.parent_uuid(), Some(parent_uuid.to_string())); + assert_eq!( + py_to_json(event.data(py).unwrap().bind(py)).unwrap(), + json!({"event": true}) + ); + assert_eq!( + py_to_json(event.metadata(py).unwrap().bind(py)).unwrap(), + json!({"meta": "event"}) + ); + assert!(event.timestamp().contains('T')); + + let tool_event = match Event::Scope(ScopeEvent::new( + base_event( + parent_uuid, + "tool-event", + json!({"input": true}), + json!({"meta": "event"}), + ), + ScopeCategory::Start, + tool_attributes_to_strings(ToolAttributes::REMOTE), + EventCategory::tool(), + Some(CategoryProfile::builder().tool_call_id("tool-1").build()), + )) { + Event::Scope(inner) => PyScopeEvent { inner }, + _ => unreachable!(), + }; + assert_eq!(tool_event.kind(), "scope"); + assert_eq!(tool_event.scope_category(), "start"); + assert_eq!(tool_event.category(), "tool"); + assert_eq!( + py_to_json(tool_event.data(py).unwrap().bind(py)).unwrap(), + json!({"input": true}) + ); + assert_eq!( + py_to_json(tool_event.category_profile(py).unwrap().bind(py)).unwrap(), + json!({"tool_call_id": "tool-1"}) + ); + assert_eq!(tool_event.attributes(), vec!["remote"]); + } + assert_mark_and_tool_event_fields(py, parent_uuid); + + fn assert_llm_event_fields(py: Python<'_>, parent_uuid: Uuid) { + let llm_event = match Event::Scope(ScopeEvent::new( + base_event( + parent_uuid, + "llm-event", + json!({"output": true}), + json!({"meta": "event"}), + ), + ScopeCategory::End, + llm_attributes_to_strings(LlmAttributes::STATEFUL), + EventCategory::llm(), + Some(CategoryProfile::builder().model_name("model").build()), + )) { + Event::Scope(inner) => PyScopeEvent { inner }, + _ => unreachable!(), + }; + assert_eq!(llm_event.kind(), "scope"); + assert_eq!(llm_event.scope_category(), "end"); + assert_eq!(llm_event.category(), "llm"); + assert_eq!( + py_to_json(llm_event.data(py).unwrap().bind(py)).unwrap(), + json!({"output": true}) + ); + assert_eq!( + py_to_json(llm_event.category_profile(py).unwrap().bind(py)).unwrap(), + json!({"model_name": "model"}) + ); + assert_eq!(llm_event.attributes(), vec!["stateful"]); + } + assert_llm_event_fields(py, parent_uuid); }); } @@ -501,7 +539,7 @@ fn test_openinference_typed_otel_config_rejects_invalid_inputs() { } #[test] -fn test_stream_request_event_and_handle_wrappers_cover_remaining_methods() { +fn test_attribute_wrappers_cover_remaining_bitwise_methods() { let _python = crate::test_support::init_python_test(); let scope_or = PyScopeAttributes::new(PyScopeAttributes::PARALLEL) @@ -526,372 +564,383 @@ fn test_stream_request_event_and_handle_wrappers_cover_remaining_methods() { ); let llm_and = llm_or.__and__(&PyLLMAttributes::new(PyLLMAttributes::STREAMING)); assert_eq!(llm_and.value(), PyLLMAttributes::STREAMING); +} +#[test] +fn test_request_and_handle_wrappers_cover_remaining_methods() { + let _python = crate::test_support::init_python_test(); Python::attach(|py| { - let stack = PyScopeStack { - inner: nemo_relay::api::runtime::create_scope_stack(), - publication_buffer: None, - }; - assert_eq!(stack.__repr__(), ""); - - let parent_uuid = Uuid::now_v7(); - let scope = PyScopeHandle::from( - ScopeHandle::builder() - .name("scope") - .scope_type(ScopeType::Agent) - .attributes(ScopeAttributes::PARALLEL | ScopeAttributes::RELOCATABLE) - .parent_uuid(parent_uuid) - .data(json!({"scope": true})) - .metadata(json!({"scope_meta": true})) - .build(), - ); - assert!(!scope.uuid().is_empty()); - assert_eq!(scope.name(), "scope"); - assert_eq!( - scope.attributes().value(), - PyScopeAttributes::PARALLEL | PyScopeAttributes::RELOCATABLE - ); + fn assert_remaining_handle_methods(py: Python<'_>) { + let stack = PyScopeStack { + inner: nemo_relay::api::runtime::create_scope_stack(), + publication_buffer: None, + }; + assert_eq!(stack.__repr__(), ""); + + let parent_uuid = Uuid::now_v7(); + let scope = PyScopeHandle::from( + ScopeHandle::builder() + .name("scope") + .scope_type(ScopeType::Agent) + .attributes(ScopeAttributes::PARALLEL | ScopeAttributes::RELOCATABLE) + .parent_uuid(parent_uuid) + .data(json!({"scope": true})) + .metadata(json!({"scope_meta": true})) + .build(), + ); + assert!(!scope.uuid().is_empty()); + assert_eq!(scope.name(), "scope"); + assert_eq!( + scope.attributes().value(), + PyScopeAttributes::PARALLEL | PyScopeAttributes::RELOCATABLE + ); - let tool = PyToolHandle::from( - ToolHandle::builder() - .name("tool") - .attributes(ToolAttributes::REMOTE) - .parent_uuid(parent_uuid) - .data(json!({"tool": true})) - .metadata(json!({"tool_meta": true})) - .build(), - ); - assert!(!tool.uuid().is_empty()); - assert_eq!(tool.name(), "tool"); - - let llm = PyLLMHandle::from( - LlmHandle::builder() - .name("llm") - .attributes(LlmAttributes::STATEFUL | LlmAttributes::STREAMING) - .parent_uuid(parent_uuid) - .data(json!({"llm": true})) - .metadata(json!({"llm_meta": true})) - .build(), - ); - assert!(!llm.uuid().is_empty()); - assert_eq!(llm.name(), "llm"); + let tool = PyToolHandle::from( + ToolHandle::builder() + .name("tool") + .attributes(ToolAttributes::REMOTE) + .parent_uuid(parent_uuid) + .data(json!({"tool": true})) + .metadata(json!({"tool_meta": true})) + .build(), + ); + assert!(!tool.uuid().is_empty()); + assert_eq!(tool.name(), "tool"); + + let llm = PyLLMHandle::from( + LlmHandle::builder() + .name("llm") + .attributes(LlmAttributes::STATEFUL | LlmAttributes::STREAMING) + .parent_uuid(parent_uuid) + .data(json!({"llm": true})) + .metadata(json!({"llm_meta": true})) + .build(), + ); + assert!(!llm.uuid().is_empty()); + assert_eq!(llm.name(), "llm"); - let headers = PyDict::new(py); - headers.set_item("x-trace", "1").unwrap(); - let content = json_to_py(py, &json!({"model": "demo", "messages": []})).unwrap(); - let request = PyLLMRequest::new(headers.as_any(), content.bind(py)).unwrap(); - assert_eq!( - py_to_json(request.headers(py).unwrap().bind(py)).unwrap(), - json!({"x-trace": "1"}) - ); - assert_eq!( - py_to_json(request.content(py).unwrap().bind(py)).unwrap(), - json!({"model": "demo", "messages": []}) - ); - assert_eq!(request.__repr__(), "LLMRequest(...)"); - - let annotated_request = AnnotatedLLMRequest { - instructions: None, - api_specific: None, - messages: vec![ - Message::System { - content: MessageContent::Text("system".into()), - name: None, - }, - Message::User { - content: MessageContent::Text("user".into()), - name: None, - }, - ], - model: Some("codec-model".into()), - params: None, - tools: None, - tool_choice: None, - store: None, - previous_response_id: None, - truncation: None, - reasoning: None, - include: None, - user: None, - metadata: None, - service_tier: None, - parallel_tool_calls: None, - max_output_tokens: None, - max_tool_calls: None, - top_logprobs: None, - stream: None, - extra: serde_json::Map::new(), - }; - let annotated_response = AnnotatedLLMResponse { - id: Some("resp-1".into()), - model: Some("codec-model".into()), - message: Some(MessageContent::Text("done".into())), - tool_calls: Some(vec![ResponseToolCall { - id: "call-1".into(), - name: "lookup".into(), - arguments: json!({"city": "NYC"}), - }]), - finish_reason: Some(FinishReason::Complete), - usage: Some(Usage { - prompt_tokens: Some(1), - completion_tokens: Some(2), - total_tokens: Some(3), - cache_read_tokens: None, - cache_write_tokens: None, - cost: None, - }), - api_specific: Some(ApiSpecificResponse::Custom { - api_name: "custom".into(), - data: json!({"ok": true}), - }), - optimization_summary: None, - extra: serde_json::Map::from_iter([("extra".into(), json!(true))]), - }; + let headers = PyDict::new(py); + headers.set_item("x-trace", "1").unwrap(); + let content = json_to_py(py, &json!({"model": "demo", "messages": []})).unwrap(); + let request = PyLLMRequest::new(headers.as_any(), content.bind(py)).unwrap(); + assert_eq!( + py_to_json(request.headers(py).unwrap().bind(py)).unwrap(), + json!({"x-trace": "1"}) + ); + assert_eq!( + py_to_json(request.content(py).unwrap().bind(py)).unwrap(), + json!({"model": "demo", "messages": []}) + ); + assert_eq!(request.__repr__(), "LLMRequest(...)"); + } + assert_remaining_handle_methods(py); + }); +} - let scope_start = match Event::Scope(ScopeEvent::new( - base_event( - parent_uuid, - "scope-start", - json!({"phase": "start"}), - json!({"meta": true}), - ), - ScopeCategory::Start, - scope_attributes_to_strings(ScopeAttributes::PARALLEL), - EventCategory::agent(), - None, - )) { - Event::Scope(inner) => PyScopeEvent { inner }, - _ => unreachable!(), - }; - assert_eq!(scope_start.kind(), "scope"); - assert_eq!(scope_start.scope_category(), "start"); - assert_eq!(scope_start.name(), "scope-start"); - assert_eq!(scope_start.category(), "agent"); - assert_eq!(scope_start.attributes(), vec!["parallel".to_string()]); - - let scope_end = match Event::Scope(ScopeEvent::new( - base_event( - parent_uuid, - "scope-end", - json!({"phase": "end"}), - json!({"meta": true}), - ), - ScopeCategory::End, - scope_attributes_to_strings(ScopeAttributes::RELOCATABLE), - EventCategory::tool(), - None, - )) { - Event::Scope(inner) => PyScopeEvent { inner }, - _ => unreachable!(), - }; - assert_eq!(scope_end.kind(), "scope"); - assert_eq!(scope_end.scope_category(), "end"); - assert_eq!(scope_end.category(), "tool"); - - let tool_end = match Event::Scope(ScopeEvent::new( - base_event( - parent_uuid, - "tool-end", - json!({"output": 1}), - json!({"meta": true}), - ), - ScopeCategory::End, - tool_attributes_to_strings(ToolAttributes::REMOTE), - EventCategory::tool(), - Some(CategoryProfile::builder().tool_call_id("call-1").build()), - )) { - Event::Scope(inner) => PyScopeEvent { inner }, - _ => unreachable!(), - }; - assert_eq!(tool_end.kind(), "scope"); - assert_eq!(tool_end.scope_category(), "end"); - assert_eq!( - py_to_json(tool_end.data(py).unwrap().bind(py)).unwrap(), - json!({"output": 1}) - ); - assert_eq!( - py_to_json(tool_end.category_profile(py).unwrap().bind(py)).unwrap(), - json!({"tool_call_id": "call-1"}) - ); +#[test] +fn test_event_wrappers_cover_remaining_methods() { + let _python = crate::test_support::init_python_test(); + Python::attach(|py| { + fn assert_remaining_event_methods(py: Python<'_>) { + let parent_uuid = Uuid::now_v7(); + let annotated_request = AnnotatedLLMRequest { + instructions: None, + api_specific: None, + messages: vec![ + Message::System { + content: MessageContent::Text("system".into()), + name: None, + }, + Message::User { + content: MessageContent::Text("user".into()), + name: None, + }, + ], + model: Some("codec-model".into()), + params: None, + tools: None, + tool_choice: None, + store: None, + previous_response_id: None, + truncation: None, + reasoning: None, + include: None, + user: None, + metadata: None, + service_tier: None, + parallel_tool_calls: None, + max_output_tokens: None, + max_tool_calls: None, + top_logprobs: None, + stream: None, + extra: serde_json::Map::new(), + }; + let annotated_response = AnnotatedLLMResponse { + id: Some("resp-1".into()), + model: Some("codec-model".into()), + message: Some(MessageContent::Text("done".into())), + tool_calls: Some(vec![ResponseToolCall { + id: "call-1".into(), + name: "lookup".into(), + arguments: json!({"city": "NYC"}), + }]), + finish_reason: Some(FinishReason::Complete), + usage: Some(Usage { + prompt_tokens: Some(1), + completion_tokens: Some(2), + total_tokens: Some(3), + cache_read_tokens: None, + cache_write_tokens: None, + cost: None, + }), + api_specific: Some(ApiSpecificResponse::Custom { + api_name: "custom".into(), + data: json!({"ok": true}), + }), + optimization_summary: None, + extra: serde_json::Map::from_iter([("extra".into(), json!(true))]), + }; + + fn assert_scope_and_tool_event_methods(py: Python<'_>, parent_uuid: Uuid) { + let scope_start = py_scope_event(Event::Scope(ScopeEvent::new( + base_event( + parent_uuid, + "scope-start", + json!({"phase": "start"}), + json!({"meta": true}), + ), + ScopeCategory::Start, + scope_attributes_to_strings(ScopeAttributes::PARALLEL), + EventCategory::agent(), + None, + ))); + assert_eq!(scope_start.kind(), "scope"); + assert_eq!(scope_start.scope_category(), "start"); + assert_eq!(scope_start.name(), "scope-start"); + assert_eq!(scope_start.category(), "agent"); + assert_eq!(scope_start.attributes(), vec!["parallel".to_string()]); + + let scope_end = py_scope_event(Event::Scope(ScopeEvent::new( + base_event( + parent_uuid, + "scope-end", + json!({"phase": "end"}), + json!({"meta": true}), + ), + ScopeCategory::End, + scope_attributes_to_strings(ScopeAttributes::RELOCATABLE), + EventCategory::tool(), + None, + ))); + assert_eq!(scope_end.kind(), "scope"); + assert_eq!(scope_end.scope_category(), "end"); + assert_eq!(scope_end.category(), "tool"); + + let tool_end = py_scope_event(Event::Scope(ScopeEvent::new( + base_event( + parent_uuid, + "tool-end", + json!({"output": 1}), + json!({"meta": true}), + ), + ScopeCategory::End, + tool_attributes_to_strings(ToolAttributes::REMOTE), + EventCategory::tool(), + Some(CategoryProfile::builder().tool_call_id("call-1").build()), + ))); + assert_eq!(tool_end.kind(), "scope"); + assert_eq!(tool_end.scope_category(), "end"); + assert_eq!( + py_to_json(tool_end.data(py).unwrap().bind(py)).unwrap(), + json!({"output": 1}) + ); + assert_eq!( + py_to_json(tool_end.category_profile(py).unwrap().bind(py)).unwrap(), + json!({"tool_call_id": "call-1"}) + ); + } + assert_scope_and_tool_event_methods(py, parent_uuid); + + let llm_start = py_scope_event(Event::Scope(ScopeEvent::new( + base_event( + parent_uuid, + "llm-start", + json!({"input": true}), + json!({"meta": true}), + ), + ScopeCategory::Start, + llm_attributes_to_strings(LlmAttributes::STATEFUL), + EventCategory::llm(), + Some( + CategoryProfile::builder() + .model_name("demo-model") + .annotated_request(std::sync::Arc::new(annotated_request.clone())) + .build(), + ), + ))); + let mut expected_start_profile = json!({"model_name": "demo-model"}); + expected_start_profile.as_object_mut().unwrap().insert( + "annotated_request".into(), + serde_json::to_value(&annotated_request).unwrap(), + ); + assert_eq!( + py_to_json(llm_start.category_profile(py).unwrap().bind(py)).unwrap(), + expected_start_profile + ); - let llm_start = match Event::Scope(ScopeEvent::new( - base_event( - parent_uuid, - "llm-start", - json!({"input": true}), - json!({"meta": true}), - ), - ScopeCategory::Start, - llm_attributes_to_strings(LlmAttributes::STATEFUL), - EventCategory::llm(), - Some( - CategoryProfile::builder() - .model_name("demo-model") - .annotated_request(std::sync::Arc::new(annotated_request.clone())) - .build(), - ), - )) { - Event::Scope(inner) => PyScopeEvent { inner }, - _ => unreachable!(), - }; - let mut expected_start_profile = json!({"model_name": "demo-model"}); - expected_start_profile.as_object_mut().unwrap().insert( - "annotated_request".into(), - serde_json::to_value(&annotated_request).unwrap(), - ); - assert_eq!( - py_to_json(llm_start.category_profile(py).unwrap().bind(py)).unwrap(), - expected_start_profile - ); + let llm_end = py_scope_event(Event::Scope(ScopeEvent::new( + base_event( + parent_uuid, + "llm-end", + json!({"output": true}), + json!({"meta": true}), + ), + ScopeCategory::End, + llm_attributes_to_strings(LlmAttributes::STREAMING), + EventCategory::llm(), + Some( + CategoryProfile::builder() + .model_name("demo-model") + .annotated_response(std::sync::Arc::new(annotated_response.clone())) + .build(), + ), + ))); + let mut expected_end_profile = json!({"model_name": "demo-model"}); + expected_end_profile.as_object_mut().unwrap().insert( + "annotated_response".into(), + serde_json::to_value(&annotated_response).unwrap(), + ); + assert_eq!( + py_to_json(llm_end.category_profile(py).unwrap().bind(py)).unwrap(), + expected_end_profile + ); - let llm_end = match Event::Scope(ScopeEvent::new( - base_event( - parent_uuid, - "llm-end", - json!({"output": true}), - json!({"meta": true}), - ), - ScopeCategory::End, - llm_attributes_to_strings(LlmAttributes::STREAMING), - EventCategory::llm(), - Some( - CategoryProfile::builder() - .model_name("demo-model") - .annotated_response(std::sync::Arc::new(annotated_response.clone())) - .build(), - ), - )) { - Event::Scope(inner) => PyScopeEvent { inner }, - _ => unreachable!(), - }; - let mut expected_end_profile = json!({"model_name": "demo-model"}); - expected_end_profile.as_object_mut().unwrap().insert( - "annotated_response".into(), - serde_json::to_value(&annotated_response).unwrap(), - ); - assert_eq!( - py_to_json(llm_end.category_profile(py).unwrap().bind(py)).unwrap(), - expected_end_profile - ); + let mark = py_mark_event(Event::Mark(MarkEvent::new( + base_event( + parent_uuid, + "mark", + json!({"mark": true}), + json!({"meta": true}), + ), + None, + None, + ))); + assert_eq!(mark.kind(), "mark"); + assert_eq!(mark.name(), "mark"); + } + assert_remaining_event_methods(py); + }); +} - let mark = match Event::Mark(MarkEvent::new( - base_event( - parent_uuid, - "mark", - json!({"mark": true}), - json!({"meta": true}), - ), - None, - None, - )) { - Event::Mark(inner) => PyMarkEvent { inner }, - _ => unreachable!(), - }; - assert_eq!(mark.kind(), "mark"); - assert_eq!(mark.name(), "mark"); - - with_event_loop(py, |event_loop| { - let runner = PyModule::from_code( - py, - &CString::new( - r#" +#[test] +fn test_llm_stream_wrapper_covers_remaining_methods() { + let _python = crate::test_support::init_python_test(); + Python::attach(|py| { + fn assert_llm_stream_methods(py: Python<'_>) { + with_event_loop(py, |event_loop| { + let runner = PyModule::from_code( + py, + &CString::new( + r#" async def next_item(stream): return await stream.__anext__() "#, - ) - .unwrap(), - &CString::new("py_types_stream_runner.py").unwrap(), - &CString::new("py_types_stream_runner").unwrap(), - ) - .unwrap(); - let (tx_ok, rx_ok) = tokio::sync::mpsc::channel(2); - tx_ok.blocking_send(Ok(json!({"chunk": 1}))).unwrap(); - drop(tx_ok); - let (cancel_ok, _cancel_ok_rx) = tokio::sync::watch::channel(false); - let (_closed_ok, closed_ok_rx) = tokio::sync::watch::channel(Some(Ok(()))); - let stream_ok = pyo3::Py::new( - py, - PyLlmStream { - receiver: Arc::new(tokio::sync::Mutex::new(rx_ok)), - cancel: cancel_ok, - closed: closed_ok_rx, - }, - ) - .unwrap(); - { - let ok_ref = stream_ok.bind(py).borrow(); - let _ = PyLlmStream::__aiter__(ok_ref); - } - let ok_chunk = event_loop - .call_method1( - "run_until_complete", - (runner - .getattr("next_item") - .unwrap() - .call1((stream_ok.clone_ref(py),)) - .unwrap(),), + ) + .unwrap(), + &CString::new("py_types_stream_runner.py").unwrap(), + &CString::new("py_types_stream_runner").unwrap(), ) .unwrap(); - assert_eq!( - crate::convert::py_to_json(&ok_chunk).unwrap(), - json!({"chunk": 1}) - ); - - let (tx_err, rx_err) = tokio::sync::mpsc::channel(1); - tx_err - .blocking_send(Err(nemo_relay::error::FlowError::Internal( - "stream boom".into(), - ))) + let (tx_ok, rx_ok) = tokio::sync::mpsc::channel(2); + tx_ok.blocking_send(Ok(json!({"chunk": 1}))).unwrap(); + drop(tx_ok); + let (cancel_ok, _cancel_ok_rx) = tokio::sync::watch::channel(false); + let (_closed_ok, closed_ok_rx) = tokio::sync::watch::channel(Some(Ok(()))); + let stream_ok = pyo3::Py::new( + py, + PyLlmStream { + receiver: Arc::new(tokio::sync::Mutex::new(rx_ok)), + cancel: cancel_ok, + closed: closed_ok_rx, + }, + ) .unwrap(); - drop(tx_err); - let (cancel_err, _cancel_err_rx) = tokio::sync::watch::channel(false); - let (_closed_err, closed_err_rx) = tokio::sync::watch::channel(Some(Ok(()))); - let stream_err = pyo3::Py::new( - py, - PyLlmStream { - receiver: Arc::new(tokio::sync::Mutex::new(rx_err)), - cancel: cancel_err, - closed: closed_err_rx, - }, - ) - .unwrap(); - let err = event_loop - .call_method1( - "run_until_complete", - (runner - .getattr("next_item") - .unwrap() - .call1((stream_err.clone_ref(py),)) - .unwrap(),), + { + let ok_ref = stream_ok.bind(py).borrow(); + let _ = PyLlmStream::__aiter__(ok_ref); + } + let ok_chunk = event_loop + .call_method1( + "run_until_complete", + (runner + .getattr("next_item") + .unwrap() + .call1((stream_ok.clone_ref(py),)) + .unwrap(),), + ) + .unwrap(); + assert_eq!( + crate::convert::py_to_json(&ok_chunk).unwrap(), + json!({"chunk": 1}) + ); + + let (tx_err, rx_err) = tokio::sync::mpsc::channel(1); + tx_err + .blocking_send(Err(nemo_relay::error::FlowError::Internal( + "stream boom".into(), + ))) + .unwrap(); + drop(tx_err); + let (cancel_err, _cancel_err_rx) = tokio::sync::watch::channel(false); + let (_closed_err, closed_err_rx) = tokio::sync::watch::channel(Some(Ok(()))); + let stream_err = pyo3::Py::new( + py, + PyLlmStream { + receiver: Arc::new(tokio::sync::Mutex::new(rx_err)), + cancel: cancel_err, + closed: closed_err_rx, + }, ) - .unwrap_err(); - assert!(err.to_string().contains("stream boom")); - - let (tx_done, rx_done) = tokio::sync::mpsc::channel(1); - drop(tx_done); - let (cancel_done, _cancel_done_rx) = tokio::sync::watch::channel(false); - let (_closed_done, closed_done_rx) = tokio::sync::watch::channel(Some(Ok(()))); - let stream_done = pyo3::Py::new( - py, - PyLlmStream { - receiver: Arc::new(tokio::sync::Mutex::new(rx_done)), - cancel: cancel_done, - closed: closed_done_rx, - }, - ) - .unwrap(); - let stop = event_loop - .call_method1( - "run_until_complete", - (runner - .getattr("next_item") - .unwrap() - .call1((stream_done.clone_ref(py),)) - .unwrap(),), + .unwrap(); + let err = event_loop + .call_method1( + "run_until_complete", + (runner + .getattr("next_item") + .unwrap() + .call1((stream_err.clone_ref(py),)) + .unwrap(),), + ) + .unwrap_err(); + assert!(err.to_string().contains("stream boom")); + + let (tx_done, rx_done) = tokio::sync::mpsc::channel(1); + drop(tx_done); + let (cancel_done, _cancel_done_rx) = tokio::sync::watch::channel(false); + let (_closed_done, closed_done_rx) = tokio::sync::watch::channel(Some(Ok(()))); + let stream_done = pyo3::Py::new( + py, + PyLlmStream { + receiver: Arc::new(tokio::sync::Mutex::new(rx_done)), + cancel: cancel_done, + closed: closed_done_rx, + }, ) - .unwrap_err(); - assert!(stop.to_string().contains("StopAsyncIteration")); - }); + .unwrap(); + let stop = event_loop + .call_method1( + "run_until_complete", + (runner + .getattr("next_item") + .unwrap() + .call1((stream_done.clone_ref(py),)) + .unwrap(),), + ) + .unwrap_err(); + assert!(stop.to_string().contains("StopAsyncIteration")); + }); + } + assert_llm_stream_methods(py); }); } @@ -1106,55 +1155,63 @@ fn test_annotated_llm_types_and_builtin_codecs_cover_mutators_and_codecs() { Some(extra.bind(py)), ) .unwrap(); - assert_eq!(annotated.model(), Some("demo-model".into())); - assert_eq!( - py_to_json(annotated.instructions(py).unwrap().bind(py)).unwrap(), - json!("Initial policy") - ); - assert_eq!( - py_to_json(annotated.api_specific(py).unwrap().bind(py)).unwrap(), - json!({"api": "openai_chat", "seed": 7}) - ); - assert_eq!(annotated.system_prompt(), Some("Initial policy".into())); - assert_eq!( - annotated.last_user_message(), - Some("Where is the weather?".into()) - ); - assert!(annotated.has_tool_calls()); - assert!(annotated.__repr__().contains("AnnotatedLLMRequest")); - assert_eq!( - py_to_json(annotated.messages(py).unwrap().bind(py)).unwrap()[0]["role"], - json!("system") - ); - assert_eq!( - py_to_json(annotated.params(py).unwrap().bind(py)).unwrap()["max_tokens"], - json!(64) - ); - assert_eq!( - py_to_json(annotated.tools(py).unwrap().bind(py)).unwrap()[0]["function"]["name"], - json!("lookup") - ); - assert_eq!( - py_to_json(annotated.tool_choice(py).unwrap().bind(py)).unwrap()["function"]["name"], - json!("lookup") - ); - assert_eq!( - py_to_json(annotated.extra(py).unwrap().bind(py)).unwrap()["provider"], - json!("test") - ); - assert_eq!(annotated.store(), None); - assert_eq!(annotated.previous_response_id(), None); - assert!(annotated.truncation(py).unwrap().bind(py).is_none()); - assert!(annotated.reasoning(py).unwrap().bind(py).is_none()); - assert!(annotated.include(py).unwrap().bind(py).is_none()); - assert_eq!(annotated.user(), None); - assert!(annotated.metadata(py).unwrap().bind(py).is_none()); - assert_eq!(annotated.service_tier(), None); - assert_eq!(annotated.parallel_tool_calls(), None); - assert_eq!(annotated.max_output_tokens(), None); - assert_eq!(annotated.max_tool_calls(), None); - assert_eq!(annotated.top_logprobs(), None); - assert_eq!(annotated.stream(), None); + + fn assert_initial_annotated_content(py: Python<'_>, annotated: &PyAnnotatedLLMRequest) { + assert_eq!(annotated.model(), Some("demo-model".into())); + assert_eq!( + py_to_json(annotated.instructions(py).unwrap().bind(py)).unwrap(), + json!("Initial policy") + ); + assert_eq!( + py_to_json(annotated.api_specific(py).unwrap().bind(py)).unwrap(), + json!({"api": "openai_chat", "seed": 7}) + ); + assert_eq!(annotated.system_prompt(), Some("Initial policy".into())); + assert_eq!( + annotated.last_user_message(), + Some("Where is the weather?".into()) + ); + assert!(annotated.has_tool_calls()); + assert!(annotated.__repr__().contains("AnnotatedLLMRequest")); + assert_eq!( + py_to_json(annotated.messages(py).unwrap().bind(py)).unwrap()[0]["role"], + json!("system") + ); + assert_eq!( + py_to_json(annotated.params(py).unwrap().bind(py)).unwrap()["max_tokens"], + json!(64) + ); + assert_eq!( + py_to_json(annotated.tools(py).unwrap().bind(py)).unwrap()[0]["function"]["name"], + json!("lookup") + ); + assert_eq!( + py_to_json(annotated.tool_choice(py).unwrap().bind(py)).unwrap()["function"]["name"], + json!("lookup") + ); + assert_eq!( + py_to_json(annotated.extra(py).unwrap().bind(py)).unwrap()["provider"], + json!("test") + ); + } + assert_initial_annotated_content(py, &annotated); + + fn assert_initial_annotated_options(py: Python<'_>, annotated: &PyAnnotatedLLMRequest) { + assert_eq!(annotated.store(), None); + assert_eq!(annotated.previous_response_id(), None); + assert!(annotated.truncation(py).unwrap().bind(py).is_none()); + assert!(annotated.reasoning(py).unwrap().bind(py).is_none()); + assert!(annotated.include(py).unwrap().bind(py).is_none()); + assert_eq!(annotated.user(), None); + assert!(annotated.metadata(py).unwrap().bind(py).is_none()); + assert_eq!(annotated.service_tier(), None); + assert_eq!(annotated.parallel_tool_calls(), None); + assert_eq!(annotated.max_output_tokens(), None); + assert_eq!(annotated.max_tool_calls(), None); + assert_eq!(annotated.top_logprobs(), None); + assert_eq!(annotated.stream(), None); + } + assert_initial_annotated_options(py, &annotated); let updated_messages = json_to_py(py, &json!([{"role": "user", "content": "updated"}])).unwrap(); @@ -1204,45 +1261,53 @@ fn test_annotated_llm_types_and_builtin_codecs_cover_mutators_and_codecs() { annotated.set_stream(Some(true)); let updated_extra = json_to_py(py, &json!({"updated": true})).unwrap(); annotated.set_extra(updated_extra.bind(py)).unwrap(); - assert_eq!(annotated.model(), Some("updated-model".into())); - assert_eq!(annotated.last_user_message(), Some("updated".into())); - assert_eq!( - py_to_json(annotated.instructions(py).unwrap().bind(py)).unwrap(), - json!([{"type": "text", "text": "Updated policy"}]) - ); - assert_eq!( - py_to_json(annotated.api_specific(py).unwrap().bind(py)).unwrap(), - json!({"api": "openai_responses", "background": true}) - ); - assert_eq!(annotated.store(), Some(true)); - assert_eq!(annotated.previous_response_id(), Some("resp_1".into())); - assert_eq!( - py_to_json(annotated.truncation(py).unwrap().bind(py)).unwrap(), - json!("disabled") - ); - assert_eq!( - py_to_json(annotated.reasoning(py).unwrap().bind(py)).unwrap(), - json!({"effort": "low"}) - ); - assert_eq!( - py_to_json(annotated.include(py).unwrap().bind(py)).unwrap(), - json!(["reasoning.encrypted_content"]) - ); - assert_eq!(annotated.user(), Some("user-1".into())); - assert_eq!( - py_to_json(annotated.metadata(py).unwrap().bind(py)).unwrap(), - json!({"tenant": "qa"}) - ); - assert_eq!(annotated.service_tier(), Some("default".into())); - assert_eq!(annotated.parallel_tool_calls(), Some(false)); - assert_eq!(annotated.max_output_tokens(), Some(128)); - assert_eq!(annotated.max_tool_calls(), Some(3)); - assert_eq!(annotated.top_logprobs(), Some(2)); - assert_eq!(annotated.stream(), Some(true)); - assert_eq!( - py_to_json(annotated.extra(py).unwrap().bind(py)).unwrap(), - json!({"updated": true}) - ); + + fn assert_updated_annotated_content(py: Python<'_>, annotated: &PyAnnotatedLLMRequest) { + assert_eq!(annotated.model(), Some("updated-model".into())); + assert_eq!(annotated.last_user_message(), Some("updated".into())); + assert_eq!( + py_to_json(annotated.instructions(py).unwrap().bind(py)).unwrap(), + json!([{"type": "text", "text": "Updated policy"}]) + ); + assert_eq!( + py_to_json(annotated.api_specific(py).unwrap().bind(py)).unwrap(), + json!({"api": "openai_responses", "background": true}) + ); + assert_eq!(annotated.store(), Some(true)); + assert_eq!(annotated.previous_response_id(), Some("resp_1".into())); + assert_eq!( + py_to_json(annotated.truncation(py).unwrap().bind(py)).unwrap(), + json!("disabled") + ); + assert_eq!( + py_to_json(annotated.reasoning(py).unwrap().bind(py)).unwrap(), + json!({"effort": "low"}) + ); + assert_eq!( + py_to_json(annotated.include(py).unwrap().bind(py)).unwrap(), + json!(["reasoning.encrypted_content"]) + ); + } + assert_updated_annotated_content(py, &annotated); + + fn assert_updated_annotated_options(py: Python<'_>, annotated: &PyAnnotatedLLMRequest) { + assert_eq!(annotated.user(), Some("user-1".into())); + assert_eq!( + py_to_json(annotated.metadata(py).unwrap().bind(py)).unwrap(), + json!({"tenant": "qa"}) + ); + assert_eq!(annotated.service_tier(), Some("default".into())); + assert_eq!(annotated.parallel_tool_calls(), Some(false)); + assert_eq!(annotated.max_output_tokens(), Some(128)); + assert_eq!(annotated.max_tool_calls(), Some(3)); + assert_eq!(annotated.top_logprobs(), Some(2)); + assert_eq!(annotated.stream(), Some(true)); + assert_eq!( + py_to_json(annotated.extra(py).unwrap().bind(py)).unwrap(), + json!({"updated": true}) + ); + } + assert_updated_annotated_options(py, &annotated); annotated.set_params(py.None().bind(py)).unwrap(); annotated.set_instructions(py.None().bind(py)).unwrap(); @@ -1253,15 +1318,19 @@ fn test_annotated_llm_types_and_builtin_codecs_cover_mutators_and_codecs() { annotated.set_reasoning(py.None().bind(py)).unwrap(); annotated.set_include(py.None().bind(py)).unwrap(); annotated.set_metadata(py.None().bind(py)).unwrap(); - assert!(annotated.params(py).unwrap().bind(py).is_none()); - assert!(annotated.instructions(py).unwrap().bind(py).is_none()); - assert!(annotated.api_specific(py).unwrap().bind(py).is_none()); - assert!(annotated.tools(py).unwrap().bind(py).is_none()); - assert!(annotated.tool_choice(py).unwrap().bind(py).is_none()); - assert!(annotated.truncation(py).unwrap().bind(py).is_none()); - assert!(annotated.reasoning(py).unwrap().bind(py).is_none()); - assert!(annotated.include(py).unwrap().bind(py).is_none()); - assert!(annotated.metadata(py).unwrap().bind(py).is_none()); + + fn assert_cleared_annotated_options(py: Python<'_>, annotated: &PyAnnotatedLLMRequest) { + assert!(annotated.params(py).unwrap().bind(py).is_none()); + assert!(annotated.instructions(py).unwrap().bind(py).is_none()); + assert!(annotated.api_specific(py).unwrap().bind(py).is_none()); + assert!(annotated.tools(py).unwrap().bind(py).is_none()); + assert!(annotated.tool_choice(py).unwrap().bind(py).is_none()); + assert!(annotated.truncation(py).unwrap().bind(py).is_none()); + assert!(annotated.reasoning(py).unwrap().bind(py).is_none()); + assert!(annotated.include(py).unwrap().bind(py).is_none()); + assert!(annotated.metadata(py).unwrap().bind(py).is_none()); + } + assert_cleared_annotated_options(py, &annotated); let bad_messages = json_to_py(py, &json!([{"content": "missing role"}])).unwrap(); let err = PyAnnotatedLLMRequest::new( @@ -1298,237 +1367,249 @@ fn test_annotated_llm_types_and_builtin_codecs_cover_mutators_and_codecs() { let bad_extra = PyList::empty(py); assert!(annotated.set_extra(&bad_extra.into_any()).is_err()); - let response = PyAnnotatedLLMResponse { - inner: AnnotatedLLMResponse { - id: Some("resp-42".into()), - model: Some("demo-model".into()), - message: Some(MessageContent::Text("hello".into())), - tool_calls: Some(vec![ResponseToolCall { - id: "call-1".into(), - name: "lookup".into(), - arguments: json!({"city": "NYC"}), - }]), - finish_reason: Some(FinishReason::Complete), - usage: Some(Usage { - prompt_tokens: Some(2), - completion_tokens: Some(3), - total_tokens: Some(5), - cache_read_tokens: Some(1), - cache_write_tokens: None, - cost: Some(CostEstimate { - total: Some(0.000_001), - currency: "USD".into(), - input: Some(0.000_000_2), - output: Some(0.000_000_8), - cache_read: None, - cache_write: None, - source: CostSource::ProviderReported, - pricing_provider: Some("test-provider".into()), - pricing_model: Some("demo-model".into()), - pricing_as_of: Some("2026-06-04".into()), - pricing_source: Some("https://example.test/pricing".into()), + fn assert_annotated_response_fields(py: Python<'_>) { + let response = PyAnnotatedLLMResponse { + inner: AnnotatedLLMResponse { + id: Some("resp-42".into()), + model: Some("demo-model".into()), + message: Some(MessageContent::Text("hello".into())), + tool_calls: Some(vec![ResponseToolCall { + id: "call-1".into(), + name: "lookup".into(), + arguments: json!({"city": "NYC"}), + }]), + finish_reason: Some(FinishReason::Complete), + usage: Some(Usage { + prompt_tokens: Some(2), + completion_tokens: Some(3), + total_tokens: Some(5), + cache_read_tokens: Some(1), + cache_write_tokens: None, + cost: Some(CostEstimate { + total: Some(0.000_001), + currency: "USD".into(), + input: Some(0.000_000_2), + output: Some(0.000_000_8), + cache_read: None, + cache_write: None, + source: CostSource::ProviderReported, + pricing_provider: Some("test-provider".into()), + pricing_model: Some("demo-model".into()), + pricing_as_of: Some("2026-06-04".into()), + pricing_source: Some("https://example.test/pricing".into()), + }), }), - }), - api_specific: Some(ApiSpecificResponse::Custom { - api_name: "custom".into(), - data: json!({"debug": true}), - }), - optimization_summary: Some( - serde_json::from_value(json!({ - "schema_version": "1", - "calculation_version": "1", - "status": "partial", - "limitations": ["missing_pricing"], - "tokens_saved": {"prompt_tokens": 2, "total_tokens": 2}, - "contributions": [] - })) - .unwrap(), - ), - extra: serde_json::Map::from_iter([("trace".into(), json!("abc"))]), - }, - }; - assert_eq!(response.id(), Some("resp-42".into())); - assert_eq!(response.model(), Some("demo-model".into())); - assert_eq!( - py_to_json(response.message(py).unwrap().bind(py)).unwrap(), - json!("hello") - ); - assert_eq!( - py_to_json(response.tool_calls(py).unwrap().bind(py)).unwrap()[0]["name"], - json!("lookup") - ); - assert_eq!(response.finish_reason(), Some("complete".into())); - assert_eq!( - py_to_json(response.usage(py).unwrap().bind(py)).unwrap()["total_tokens"], - json!(5) - ); - assert_eq!( - py_to_json(response.usage(py).unwrap().bind(py)).unwrap()["cost"]["pricing_provider"], - json!("test-provider") - ); - assert_eq!( - py_to_json(response.optimization_summary(py).unwrap().bind(py)).unwrap()["tokens_saved"] - ["prompt_tokens"], - json!(2) - ); - assert_eq!( - py_to_json(response.api_specific(py).unwrap().bind(py)).unwrap()["api_name"], - json!("custom") - ); - assert_eq!( - py_to_json(response.extra(py).unwrap().bind(py)).unwrap()["trace"], - json!("abc") - ); - assert_eq!(response.response_text(), Some("hello".into())); - assert!(response.has_tool_calls()); - assert!(response.__repr__().contains("AnnotatedLLMResponse")); - - let response_without_api_specific = PyAnnotatedLLMResponse { - inner: AnnotatedLLMResponse { - id: None, - model: None, - message: None, - tool_calls: None, - finish_reason: None, - usage: None, - api_specific: None, - optimization_summary: None, - extra: serde_json::Map::new(), - }, - }; - assert!( - response_without_api_specific - .api_specific(py) - .unwrap() - .bind(py) - .is_none() - ); - assert!( - response_without_api_specific - .optimization_summary(py) - .unwrap() - .bind(py) - .is_none() - ); + api_specific: Some(ApiSpecificResponse::Custom { + api_name: "custom".into(), + data: json!({"debug": true}), + }), + optimization_summary: Some( + serde_json::from_value(json!({ + "schema_version": "1", + "calculation_version": "1", + "status": "partial", + "limitations": ["missing_pricing"], + "tokens_saved": {"prompt_tokens": 2, "total_tokens": 2}, + "contributions": [] + })) + .unwrap(), + ), + extra: serde_json::Map::from_iter([("trace".into(), json!("abc"))]), + }, + }; + assert_eq!(response.id(), Some("resp-42".into())); + assert_eq!(response.model(), Some("demo-model".into())); + assert_eq!( + py_to_json(response.message(py).unwrap().bind(py)).unwrap(), + json!("hello") + ); + assert_eq!( + py_to_json(response.tool_calls(py).unwrap().bind(py)).unwrap()[0]["name"], + json!("lookup") + ); + assert_eq!(response.finish_reason(), Some("complete".into())); + assert_eq!( + py_to_json(response.usage(py).unwrap().bind(py)).unwrap()["total_tokens"], + json!(5) + ); + assert_eq!( + py_to_json(response.usage(py).unwrap().bind(py)).unwrap()["cost"]["pricing_provider"], + json!("test-provider") + ); + assert_eq!( + py_to_json(response.optimization_summary(py).unwrap().bind(py)).unwrap()["tokens_saved"] + ["prompt_tokens"], + json!(2) + ); + assert_eq!( + py_to_json(response.api_specific(py).unwrap().bind(py)).unwrap()["api_name"], + json!("custom") + ); + assert_eq!( + py_to_json(response.extra(py).unwrap().bind(py)).unwrap()["trace"], + json!("abc") + ); + assert_eq!(response.response_text(), Some("hello".into())); + assert!(response.has_tool_calls()); + assert!(response.__repr__().contains("AnnotatedLLMResponse")); + + let response_without_api_specific = PyAnnotatedLLMResponse { + inner: AnnotatedLLMResponse { + id: None, + model: None, + message: None, + tool_calls: None, + finish_reason: None, + usage: None, + api_specific: None, + optimization_summary: None, + extra: serde_json::Map::new(), + }, + }; + assert!( + response_without_api_specific + .api_specific(py) + .unwrap() + .bind(py) + .is_none() + ); + assert!( + response_without_api_specific + .optimization_summary(py) + .unwrap() + .bind(py) + .is_none() + ); + } + assert_annotated_response_fields(py); - let chat_request = PyLLMRequest { - inner: nemo_relay::api::llm::LlmRequest { - headers: serde_json::Map::new(), - content: json!({ - "model": "gpt-4o-mini", - "messages": [{"role": "user", "content": "hi"}], - "max_tokens": 16 - }), - }, - }; - let chat_codec = PyOpenAIChatCodec::new(); - let chat_decoded = chat_codec.decode(&chat_request).unwrap(); - assert_eq!(chat_decoded.model(), Some("gpt-4o-mini".into())); - let chat_encoded = chat_codec.encode(&chat_decoded, &chat_request).unwrap(); - assert_eq!(chat_encoded.inner.content["model"], json!("gpt-4o-mini")); - let chat_response = chat_codec - .decode_response( - json_to_py( - py, - &json!({ - "id": "chatcmpl-1", + fn assert_openai_chat_codec(py: Python<'_>) { + let chat_request = PyLLMRequest { + inner: nemo_relay::api::llm::LlmRequest { + headers: serde_json::Map::new(), + content: json!({ "model": "gpt-4o-mini", - "choices": [{ - "message": {"role": "assistant", "content": "hello"}, - "finish_reason": "stop" - }] + "messages": [{"role": "user", "content": "hi"}], + "max_tokens": 16 }), + }, + }; + let chat_codec = PyOpenAIChatCodec::new(); + let chat_decoded = chat_codec.decode(&chat_request).unwrap(); + assert_eq!(chat_decoded.model(), Some("gpt-4o-mini".into())); + let chat_encoded = chat_codec.encode(&chat_decoded, &chat_request).unwrap(); + assert_eq!(chat_encoded.inner.content["model"], json!("gpt-4o-mini")); + let chat_response = chat_codec + .decode_response( + json_to_py( + py, + &json!({ + "id": "chatcmpl-1", + "model": "gpt-4o-mini", + "choices": [{ + "message": {"role": "assistant", "content": "hello"}, + "finish_reason": "stop" + }] + }), + ) + .unwrap() + .bind(py), ) - .unwrap() - .bind(py), - ) - .unwrap(); - assert_eq!(chat_response.response_text(), Some("hello".into())); - assert_eq!(chat_codec.__repr__(), ""); - - let responses_request = PyLLMRequest { - inner: nemo_relay::api::llm::LlmRequest { - headers: serde_json::Map::new(), - content: json!({ - "model": "gpt-4o-mini", - "instructions": "Be helpful", - "input": [{"role": "user", "content": "hi"}], - "max_output_tokens": 32 - }), - }, - }; - let responses_codec = PyOpenAIResponsesCodec::new(); - let responses_decoded = responses_codec.decode(&responses_request).unwrap(); - assert_eq!(responses_decoded.system_prompt(), Some("Be helpful".into())); - let responses_encoded = responses_codec - .encode(&responses_decoded, &responses_request) - .unwrap(); - assert_eq!( - responses_encoded.inner.content["instructions"], - json!("Be helpful") - ); - let responses_response = responses_codec - .decode_response( - json_to_py( - py, - &json!({ - "id": "resp-1", + .unwrap(); + assert_eq!(chat_response.response_text(), Some("hello".into())); + assert_eq!(chat_codec.__repr__(), ""); + } + assert_openai_chat_codec(py); + + fn assert_openai_responses_codec(py: Python<'_>) { + let responses_request = PyLLMRequest { + inner: nemo_relay::api::llm::LlmRequest { + headers: serde_json::Map::new(), + content: json!({ "model": "gpt-4o-mini", - "status": "completed", - "output": [{ - "type": "message", - "role": "assistant", - "status": "completed", - "content": [{"type": "output_text", "text": "done"}] - }] + "instructions": "Be helpful", + "input": [{"role": "user", "content": "hi"}], + "max_output_tokens": 32 }), + }, + }; + let responses_codec = PyOpenAIResponsesCodec::new(); + let responses_decoded = responses_codec.decode(&responses_request).unwrap(); + assert_eq!(responses_decoded.system_prompt(), Some("Be helpful".into())); + let responses_encoded = responses_codec + .encode(&responses_decoded, &responses_request) + .unwrap(); + assert_eq!( + responses_encoded.inner.content["instructions"], + json!("Be helpful") + ); + let responses_response = responses_codec + .decode_response( + json_to_py( + py, + &json!({ + "id": "resp-1", + "model": "gpt-4o-mini", + "status": "completed", + "output": [{ + "type": "message", + "role": "assistant", + "status": "completed", + "content": [{"type": "output_text", "text": "done"}] + }] + }), + ) + .unwrap() + .bind(py), ) - .unwrap() - .bind(py), - ) - .unwrap(); - assert_eq!(responses_response.response_text(), Some("done".into())); - assert_eq!(responses_codec.__repr__(), ""); - - let anthropic_request = PyLLMRequest { - inner: nemo_relay::api::llm::LlmRequest { - headers: serde_json::Map::new(), - content: json!({ - "model": "claude-sonnet-4-20250514", - "system": "Be careful", - "messages": [{"role": "user", "content": "hi"}], - "max_tokens": 64 - }), - }, - }; - let anthropic_codec = PyAnthropicMessagesCodec::new(); - let anthropic_decoded = anthropic_codec.decode(&anthropic_request).unwrap(); - assert_eq!(anthropic_decoded.system_prompt(), Some("Be careful".into())); - let anthropic_encoded = anthropic_codec - .encode(&anthropic_decoded, &anthropic_request) - .unwrap(); - assert_eq!( - anthropic_encoded.inner.content["system"], - json!("Be careful") - ); - let anthropic_response = anthropic_codec - .decode_response( - json_to_py( - py, - &json!({ - "id": "msg-1", + .unwrap(); + assert_eq!(responses_response.response_text(), Some("done".into())); + assert_eq!(responses_codec.__repr__(), ""); + } + assert_openai_responses_codec(py); + + fn assert_anthropic_messages_codec(py: Python<'_>) { + let anthropic_request = PyLLMRequest { + inner: nemo_relay::api::llm::LlmRequest { + headers: serde_json::Map::new(), + content: json!({ "model": "claude-sonnet-4-20250514", - "content": [{"type": "text", "text": "done"}], - "stop_reason": "end_turn", - "usage": {"input_tokens": 1, "output_tokens": 2} + "system": "Be careful", + "messages": [{"role": "user", "content": "hi"}], + "max_tokens": 64 }), + }, + }; + let anthropic_codec = PyAnthropicMessagesCodec::new(); + let anthropic_decoded = anthropic_codec.decode(&anthropic_request).unwrap(); + assert_eq!(anthropic_decoded.system_prompt(), Some("Be careful".into())); + let anthropic_encoded = anthropic_codec + .encode(&anthropic_decoded, &anthropic_request) + .unwrap(); + assert_eq!( + anthropic_encoded.inner.content["system"], + json!("Be careful") + ); + let anthropic_response = anthropic_codec + .decode_response( + json_to_py( + py, + &json!({ + "id": "msg-1", + "model": "claude-sonnet-4-20250514", + "content": [{"type": "text", "text": "done"}], + "stop_reason": "end_turn", + "usage": {"input_tokens": 1, "output_tokens": 2} + }), + ) + .unwrap() + .bind(py), ) - .unwrap() - .bind(py), - ) - .unwrap(); - assert_eq!(anthropic_response.response_text(), Some("done".into())); - assert_eq!(anthropic_codec.__repr__(), ""); + .unwrap(); + assert_eq!(anthropic_response.response_text(), Some("done".into())); + assert_eq!(anthropic_codec.__repr__(), ""); + } + assert_anthropic_messages_codec(py); }); } @@ -1893,60 +1974,44 @@ def run(types): .unwrap(); let result_json = py_to_json(&result).unwrap(); - assert!(result_json["request_error"].as_str().is_some()); + for key in [ + "request_error", + "invalid_params_error", + "invalid_tools_error", + "invalid_tool_choice_error", + "invalid_extra_error", + "invalid_messages_setter_error", + "invalid_params_setter_error", + "invalid_tools_setter_error", + "invalid_tool_choice_setter_error", + "invalid_extra_setter_error", + ] { + assert!(result_json[key].as_str().is_some(), "{key}"); + } assert!(result_json["otel_grpc_error"].is_null()); assert!(result_json["oi_grpc_error"].is_null()); - assert_eq!(result_json["params_is_none"], json!(true)); - assert_eq!(result_json["tools_is_none"], json!(true)); - assert_eq!(result_json["tool_choice_is_none"], json!(true)); - assert!(result_json["invalid_params_error"].as_str().is_some()); - assert!(result_json["invalid_tools_error"].as_str().is_some()); - assert!(result_json["invalid_tool_choice_error"].as_str().is_some()); - assert!(result_json["invalid_extra_error"].as_str().is_some()); - assert!( - result_json["invalid_messages_setter_error"] - .as_str() - .is_some() - ); - assert!( - result_json["invalid_params_setter_error"] - .as_str() - .is_some() - ); - assert!(result_json["invalid_tools_setter_error"].as_str().is_some()); - assert!( - result_json["invalid_tool_choice_setter_error"] - .as_str() - .is_some() - ); - assert!(result_json["invalid_extra_setter_error"].as_str().is_some()); - assert!(result_json["chat_tool_calls_is_none"].as_bool().is_some()); - assert!(result_json["chat_usage_is_none"].as_bool().is_some()); - assert!(result_json["chat_api_specific_is_none"].as_bool().is_some()); - assert!(result_json["responses_message_is_none"].as_bool().is_some()); - assert!( - result_json["responses_tool_calls_is_none"] - .as_bool() - .is_some() - ); - assert!(result_json["responses_usage_is_none"].as_bool().is_some()); - assert!( - result_json["responses_api_specific_is_none"] - .as_bool() - .is_some() - ); - assert!(result_json["anthropic_message_is_none"].as_bool().is_some()); - assert!( - result_json["anthropic_tool_calls_is_none"] - .as_bool() - .is_some() - ); - assert!( - result_json["anthropic_api_specific_is_none"] - .as_bool() - .is_some() - ); - assert!(result_json["anthropic_usage_is_none"].as_bool().is_some()); + for key in ["params_is_none", "tools_is_none", "tool_choice_is_none"] { + assert_eq!(result_json[key].as_bool(), Some(true), "{key}"); + } + for key in [ + "chat_tool_calls_is_none", + "chat_usage_is_none", + "responses_message_is_none", + "responses_tool_calls_is_none", + "responses_usage_is_none", + "anthropic_message_is_none", + "anthropic_tool_calls_is_none", + "anthropic_usage_is_none", + ] { + assert_eq!(result_json[key].as_bool(), Some(true), "{key}"); + } + for key in [ + "chat_api_specific_is_none", + "responses_api_specific_is_none", + "anthropic_api_specific_is_none", + ] { + assert_eq!(result_json[key].as_bool(), Some(false), "{key}"); + } }); } diff --git a/go/nemo_relay/adaptive_plugin_test.go b/go/nemo_relay/adaptive_plugin_test.go index 9931d0187..0af7ce9e9 100644 --- a/go/nemo_relay/adaptive_plugin_test.go +++ b/go/nemo_relay/adaptive_plugin_test.go @@ -302,7 +302,9 @@ func assertClosedContextRegistrationFails(t *testing.T, name string, err error) } func TestTopLevelPluginValidationAndLifecycle(t *testing.T) { - runTestWithScopeStack(t, testTopLevelPluginValidationAndLifecycle) + runTestInIsolatedWorkingDirectory(t, func(t *testing.T) { + runTestWithScopeStack(t, testTopLevelPluginValidationAndLifecycle) + }) } func testTopLevelPluginValidationAndLifecycle(t *testing.T) { diff --git a/go/nemo_relay/adaptive_runtime_test.go b/go/nemo_relay/adaptive_runtime_test.go index 9b0e9220a..705779f60 100644 --- a/go/nemo_relay/adaptive_runtime_test.go +++ b/go/nemo_relay/adaptive_runtime_test.go @@ -11,8 +11,12 @@ import ( ) const ( - testAgentID = "go-agent" - newAdaptiveRuntimeFailedMsg = "NewAdaptiveRuntime failed: %v" + adaptiveRuntimeClosedMessage = "adaptive runtime is nil or shut down" + forcedAdaptiveMarshalFailure = "forced adaptive JSON marshal failure" + newAdaptiveRuntimeFailedMsg = "NewAdaptiveRuntime failed: %v" + responseCacheTestNamespace = "go-harness" + testAgentID = "go-agent" + validateAdaptiveConfigFailedMsg = "ValidateAdaptiveConfig failed: %v" ) func testAdaptiveRuntimeConfig(provider string) AdaptiveConfig { @@ -34,7 +38,7 @@ func uint64Ptr(value uint64) *uint64 { func TestValidateAdaptiveConfigAndOwnedRuntime(t *testing.T) { report, err := ValidateAdaptiveConfig(NewAdaptiveConfig()) if err != nil { - t.Fatalf("ValidateAdaptiveConfig failed: %v", err) + t.Fatalf(validateAdaptiveConfigFailedMsg, err) } if len(report.Diagnostics) != 0 { t.Fatalf("expected clean report, got %#v", report.Diagnostics) @@ -175,22 +179,28 @@ func TestSetLatencySensitivityRejectsInvalidValue(t *testing.T) { func TestResponseCacheConfigReachesTypedSurface(t *testing.T) { backend := NewInMemoryResponseCacheBackend() rc := NewResponseCacheConfig() - if rc.TTLSeconds == nil || *rc.TTLSeconds != 3600 { - t.Fatalf("constructor TTL default mismatch: %#v", rc.TTLSeconds) - } - if rc.Priority == nil || *rc.Priority != 50 { - t.Fatalf("constructor priority default mismatch: %#v", rc.Priority) - } - rc.Namespace = "go-harness" + assertResponseCacheConstructorDefaults(t, rc) + rc.Namespace = responseCacheTestNamespace rc.CacheNondeterministic = true rc.Backend = &backend + assertResponseCacheJSONSurface(t, rc) + assertResponseCacheValidation(t, rc) +} - config := NewAdaptiveConfig() - config.ResponseCache = &rc +func assertResponseCacheConstructorDefaults(t *testing.T, config ResponseCacheConfig) { + t.Helper() + if config.TTLSeconds == nil || *config.TTLSeconds != 3600 { + t.Fatalf("constructor TTL default mismatch: %#v", config.TTLSeconds) + } + if config.Priority == nil || *config.Priority != 50 { + t.Fatalf("constructor priority default mismatch: %#v", config.Priority) + } +} - // 1. The typed AdaptiveConfig must carry response_cache through json.Marshal. - // The bug this guards is the struct silently DROPPING the section because it - // enumerates fields by name with no response_cache field and no catch-all. +func assertResponseCacheJSONSurface(t *testing.T, responseCache ResponseCacheConfig) { + t.Helper() + config := NewAdaptiveConfig() + config.ResponseCache = &responseCache payload, err := json.Marshal(config) if err != nil { t.Fatalf("marshal failed: %v", err) @@ -203,7 +213,7 @@ func TestResponseCacheConfigReachesTypedSurface(t *testing.T) { if !ok { t.Fatalf("response_cache missing from marshaled config: %s", payload) } - if section["namespace"] != "go-harness" { + if section["namespace"] != responseCacheTestNamespace { t.Fatalf("response_cache fields not preserved: %#v", section) } if _, ok := section["skip_keys"]; ok { @@ -215,21 +225,22 @@ func TestResponseCacheConfigReachesTypedSurface(t *testing.T) { if b, ok := section["backend"].(map[string]any); !ok || b["kind"] != "in_memory" { t.Fatalf("backend not preserved: %#v", section["backend"]) } +} - // 2. A valid section validates clean through the FFI -> Rust adaptive validator. +func assertResponseCacheValidation(t *testing.T, responseCache ResponseCacheConfig) { + t.Helper() + config := NewAdaptiveConfig() + config.ResponseCache = &responseCache report, err := ValidateAdaptiveConfig(config) if err != nil { - t.Fatalf("ValidateAdaptiveConfig failed: %v", err) + t.Fatalf(validateAdaptiveConfigFailedMsg, err) } if len(report.Diagnostics) != 0 { t.Fatalf("expected clean report, got %#v", report.Diagnostics) } - // 3. An invalid section produces a response_cache diagnostic. This proves the - // section is actually validated end-to-end (a dropped section would yield no - // diagnostic at all), not merely carried in the struct. bad := NewResponseCacheConfig() - bad.Namespace = "go-harness" + bad.Namespace = responseCacheTestNamespace bad.BypassRate = 2.0 badConfig := NewAdaptiveConfig() badConfig.ResponseCache = &bad @@ -237,100 +248,97 @@ func TestResponseCacheConfigReachesTypedSurface(t *testing.T) { if err != nil { t.Fatalf("ValidateAdaptiveConfig (invalid bypass_rate) returned error: %v", err) } - found := false - for _, d := range badReport.Diagnostics { - if d.Code == "response_cache.invalid_bypass_rate" { - found = true - } - } - if !found { + if !hasAdaptiveDiagnostic(badReport, "response_cache.invalid_bypass_rate") { t.Fatalf("expected response_cache.invalid_bypass_rate diagnostic, got %#v", badReport.Diagnostics) } } func TestResponseCacheConfigPreservesOmissionAndExplicitZero(t *testing.T) { - marshal := func(t *testing.T, responseCache ResponseCacheConfig) map[string]any { - t.Helper() - payload, err := json.Marshal(responseCache) - if err != nil { - t.Fatalf("marshal failed: %v", err) - } - var decoded map[string]any - if err := json.Unmarshal(payload, &decoded); err != nil { - t.Fatalf("unmarshal failed: %v", err) - } - return decoded - } - validate := func(t *testing.T, responseCache ResponseCacheConfig) ConfigReport { - t.Helper() - config := NewAdaptiveConfig() - config.ResponseCache = &responseCache - report, err := ValidateAdaptiveConfig(config) - if err != nil { - t.Fatalf("ValidateAdaptiveConfig failed: %v", err) - } - return report + t.Run("partial config delegates to Rust defaults", testPartialResponseCacheConfig) + t.Run("missing namespace remains invalid", testMissingResponseCacheNamespace) + t.Run("explicit TTL zero remains invalid", testExplicitZeroResponseCacheTTL) + t.Run("explicit priority zero remains valid", testExplicitZeroResponseCachePriority) +} + +func marshalResponseCacheConfig(t *testing.T, responseCache ResponseCacheConfig) map[string]any { + t.Helper() + payload, err := json.Marshal(responseCache) + if err != nil { + t.Fatalf("marshal failed: %v", err) + } + var decoded map[string]any + if err := json.Unmarshal(payload, &decoded); err != nil { + t.Fatalf("unmarshal failed: %v", err) } + return decoded +} - t.Run("partial config delegates to Rust defaults", func(t *testing.T) { - responseCache := ResponseCacheConfig{Namespace: "dev"} - decoded := marshal(t, responseCache) - if _, ok := decoded["ttl_seconds"]; ok { - t.Fatalf("partial config must omit ttl_seconds: %#v", decoded) - } - if _, ok := decoded["priority"]; ok { - t.Fatalf("partial config must omit priority: %#v", decoded) - } - if report := validate(t, responseCache); len(report.Diagnostics) != 0 { - t.Fatalf("expected Rust defaults to validate cleanly, got %#v", report.Diagnostics) - } - }) +func validateResponseCacheConfig(t *testing.T, responseCache ResponseCacheConfig) ConfigReport { + t.Helper() + config := NewAdaptiveConfig() + config.ResponseCache = &responseCache + report, err := ValidateAdaptiveConfig(config) + if err != nil { + t.Fatalf(validateAdaptiveConfigFailedMsg, err) + } + return report +} - t.Run("missing namespace remains invalid", func(t *testing.T) { - report := validate(t, ResponseCacheConfig{}) - found := false - for _, diagnostic := range report.Diagnostics { - if diagnostic.Code == "response_cache.missing_namespace" { - found = true - } +func hasAdaptiveDiagnostic(report ConfigReport, code string) bool { + for _, diagnostic := range report.Diagnostics { + if diagnostic.Code == code { + return true } - if !found { - t.Fatalf("expected response_cache.missing_namespace, got %#v", report.Diagnostics) - } - }) + } + return false +} - t.Run("explicit TTL zero remains invalid", func(t *testing.T) { - zero := uint64(0) - responseCache := ResponseCacheConfig{TTLSeconds: &zero, Namespace: "dev"} - if got := marshal(t, responseCache)["ttl_seconds"]; got != float64(0) { - t.Fatalf("explicit ttl_seconds=0 was not preserved: %#v", got) - } - report := validate(t, responseCache) - found := false - for _, diagnostic := range report.Diagnostics { - if diagnostic.Code == "response_cache.invalid_ttl" { - found = true - } - } - if !found { - t.Fatalf("expected response_cache.invalid_ttl, got %#v", report.Diagnostics) - } - }) +func testPartialResponseCacheConfig(t *testing.T) { + responseCache := ResponseCacheConfig{Namespace: "dev"} + decoded := marshalResponseCacheConfig(t, responseCache) + if _, ok := decoded["ttl_seconds"]; ok { + t.Fatalf("partial config must omit ttl_seconds: %#v", decoded) + } + if _, ok := decoded["priority"]; ok { + t.Fatalf("partial config must omit priority: %#v", decoded) + } + if report := validateResponseCacheConfig(t, responseCache); len(report.Diagnostics) != 0 { + t.Fatalf("expected Rust defaults to validate cleanly, got %#v", report.Diagnostics) + } +} - t.Run("explicit priority zero remains valid", func(t *testing.T) { - zero := int32(0) - responseCache := ResponseCacheConfig{Priority: &zero, Namespace: "dev"} - decoded := marshal(t, responseCache) - if got := decoded["priority"]; got != float64(0) { - t.Fatalf("explicit priority=0 was not preserved: %#v", got) - } - if _, ok := decoded["ttl_seconds"]; ok { - t.Fatalf("unconfigured ttl_seconds must remain omitted: %#v", decoded) - } - if report := validate(t, responseCache); len(report.Diagnostics) != 0 { - t.Fatalf("expected priority=0 to validate cleanly, got %#v", report.Diagnostics) - } - }) +func testMissingResponseCacheNamespace(t *testing.T) { + report := validateResponseCacheConfig(t, ResponseCacheConfig{}) + if !hasAdaptiveDiagnostic(report, "response_cache.missing_namespace") { + t.Fatalf("expected response_cache.missing_namespace, got %#v", report.Diagnostics) + } +} + +func testExplicitZeroResponseCacheTTL(t *testing.T) { + zero := uint64(0) + responseCache := ResponseCacheConfig{TTLSeconds: &zero, Namespace: "dev"} + if got := marshalResponseCacheConfig(t, responseCache)["ttl_seconds"]; got != float64(0) { + t.Fatalf("explicit ttl_seconds=0 was not preserved: %#v", got) + } + report := validateResponseCacheConfig(t, responseCache) + if !hasAdaptiveDiagnostic(report, "response_cache.invalid_ttl") { + t.Fatalf("expected response_cache.invalid_ttl, got %#v", report.Diagnostics) + } +} + +func testExplicitZeroResponseCachePriority(t *testing.T) { + zero := int32(0) + responseCache := ResponseCacheConfig{Priority: &zero, Namespace: "dev"} + decoded := marshalResponseCacheConfig(t, responseCache) + if got := decoded["priority"]; got != float64(0) { + t.Fatalf("explicit priority=0 was not preserved: %#v", got) + } + if _, ok := decoded["ttl_seconds"]; ok { + t.Fatalf("unconfigured ttl_seconds must remain omitted: %#v", decoded) + } + if report := validateResponseCacheConfig(t, responseCache); len(report.Diagnostics) != 0 { + t.Fatalf("expected priority=0 to validate cleanly, got %#v", report.Diagnostics) + } } func TestAdaptiveRuntimeLifecycleRejectsUseAfterShutdown(t *testing.T) { @@ -358,18 +366,18 @@ func assertAdaptiveRuntimeClosed(t *testing.T, runtime *AdaptiveRuntime) { {name: "WaitForIdle", err: runtime.WaitForIdle()}, {name: "BindScope", err: runtime.BindScope(nil)}, } { - if test.err == nil || !strings.Contains(test.err.Error(), "adaptive runtime is nil or shut down") { + if test.err == nil || !strings.Contains(test.err.Error(), adaptiveRuntimeClosedMessage) { t.Fatalf("expected %s to reject a shut down runtime, got %v", test.name, test.err) } } - if _, err := runtime.Report(); err == nil || !strings.Contains(err.Error(), "adaptive runtime is nil or shut down") { + if _, err := runtime.Report(); err == nil || !strings.Contains(err.Error(), adaptiveRuntimeClosedMessage) { t.Fatalf("expected Report to reject a shut down runtime, got %v", err) } - if _, err := runtime.BuildCacheRequestFacts(CacheRequestFactsInput{}); err == nil || !strings.Contains(err.Error(), "adaptive runtime is nil or shut down") { + if _, err := runtime.BuildCacheRequestFacts(CacheRequestFactsInput{}); err == nil || !strings.Contains(err.Error(), adaptiveRuntimeClosedMessage) { t.Fatalf("expected BuildCacheRequestFacts to reject a shut down runtime, got %v", err) } - if err := runtime.Shutdown(); err == nil || !strings.Contains(err.Error(), "adaptive runtime is nil or shut down") { + if err := runtime.Shutdown(); err == nil || !strings.Contains(err.Error(), adaptiveRuntimeClosedMessage) { t.Fatalf("expected repeated Shutdown to reject a shut down runtime, got %v", err) } } @@ -379,16 +387,16 @@ func TestAdaptiveRuntimePublicHelpersPropagateJSONMarshalFailures(t *testing.T) t.Cleanup(func() { jsonMarshal = oldMarshal }) jsonMarshal = func(any) ([]byte, error) { - return nil, errors.New("forced adaptive JSON marshal failure") + return nil, errors.New(forcedAdaptiveMarshalFailure) } - if _, err := ValidateAdaptiveConfig(NewAdaptiveConfig()); err == nil || !strings.Contains(err.Error(), "forced adaptive JSON marshal failure") { + if _, err := ValidateAdaptiveConfig(NewAdaptiveConfig()); err == nil || !strings.Contains(err.Error(), forcedAdaptiveMarshalFailure) { t.Fatalf("expected ValidateAdaptiveConfig to return marshal failure, got %v", err) } - if _, err := NewAdaptiveRuntime(NewAdaptiveConfig()); err == nil || !strings.Contains(err.Error(), "forced adaptive JSON marshal failure") { + if _, err := NewAdaptiveRuntime(NewAdaptiveConfig()); err == nil || !strings.Contains(err.Error(), forcedAdaptiveMarshalFailure) { t.Fatalf("expected NewAdaptiveRuntime to return marshal failure, got %v", err) } - if _, err := BuildCacheTelemetryEvent(CacheTelemetryEventInput{}); err == nil || !strings.Contains(err.Error(), "forced adaptive JSON marshal failure") { + if _, err := BuildCacheTelemetryEvent(CacheTelemetryEventInput{}); err == nil || !strings.Contains(err.Error(), forcedAdaptiveMarshalFailure) { t.Fatalf("expected BuildCacheTelemetryEvent to return marshal failure, got %v", err) } } diff --git a/go/nemo_relay/adaptive_test.go b/go/nemo_relay/adaptive_test.go index fba86b6ff..f20a9cd50 100644 --- a/go/nemo_relay/adaptive_test.go +++ b/go/nemo_relay/adaptive_test.go @@ -169,6 +169,10 @@ func TestValidatePluginConfigWarnsMissingStateForTelemetry(t *testing.T) { } func TestConfigureAdaptiveComponentLifecycle(t *testing.T) { + runTestInIsolatedWorkingDirectory(t, testConfigureAdaptiveComponentLifecycle) +} + +func testConfigureAdaptiveComponentLifecycle(t *testing.T) { config := NewAdaptiveConfig() config.State = &AdaptiveStateConfig{ Backend: NewInMemoryAdaptiveBackend(), diff --git a/go/nemo_relay/callbacks_test.go b/go/nemo_relay/callbacks_test.go index 4be9b5637..5394b5618 100644 --- a/go/nemo_relay/callbacks_test.go +++ b/go/nemo_relay/callbacks_test.go @@ -68,18 +68,37 @@ func TestRegisterAndUnregisterClosure(t *testing.T) { } } +type codecIdentityTestCase struct { + name string + kind uint32 + id *string + want LLMCodecKind +} + +func assertCodecIdentity(t *testing.T, test codecIdentityTestCase) { + t.Helper() + codec := llmCodecIdentity(test.kind, test.id) + if codec.CodecKind != test.want { + t.Fatalf("codec kind = %q, want %q", codec.CodecKind, test.want) + } + if codec.CodecID == nil && test.id != nil { + t.Fatal("codec ID was lost") + } + if codec.CodecID != nil && test.id == nil { + t.Fatalf("unexpected codec ID %q", *codec.CodecID) + } + if codec.CodecID != nil && test.id != nil && *codec.CodecID != *test.id { + t.Fatalf("codec ID = %q, want %q", *codec.CodecID, *test.id) + } +} + func TestLlmSanitizeDirectionalContextsPreserveEveryCodecIdentity(t *testing.T) { openAIChat := "openai_chat" openAIResponses := "openai_responses" anthropicMessages := "anthropic_messages" runtimeCodec := "com.example.chat.v1" - cases := []struct { - name string - kind uint32 - id *string - want LLMCodecKind - }{ + cases := []codecIdentityTestCase{ {"none", 0, nil, LLMCodecNone}, {"openai chat", 1, &openAIChat, LLMCodecBuiltin}, {"openai responses", 1, &openAIResponses, LLMCodecBuiltin}, @@ -91,19 +110,7 @@ func TestLlmSanitizeDirectionalContextsPreserveEveryCodecIdentity(t *testing.T) for _, test := range cases { t.Run(test.name, func(t *testing.T) { - codec := llmCodecIdentity(test.kind, test.id) - if codec.CodecKind != test.want { - t.Fatalf("codec kind = %q, want %q", codec.CodecKind, test.want) - } - if codec.CodecID == nil && test.id != nil { - t.Fatal("codec ID was lost") - } - if codec.CodecID != nil && test.id == nil { - t.Fatalf("unexpected codec ID %q", *codec.CodecID) - } - if codec.CodecID != nil && test.id != nil && *codec.CodecID != *test.id { - t.Fatalf("codec ID = %q, want %q", *codec.CodecID, *test.id) - } + assertCodecIdentity(t, test) }) } } diff --git a/go/nemo_relay/context_test.go b/go/nemo_relay/context_test.go index e4c39fd92..f58910486 100644 --- a/go/nemo_relay/context_test.go +++ b/go/nemo_relay/context_test.go @@ -11,7 +11,14 @@ import ( "testing" ) -const newScopeStackFailed = "NewScopeStack failed: %v" +const ( + getHandleFailed = "GetHandle failed: %v" + invalidUUID = "not-a-uuid" + newScopeStackFailed = "NewScopeStack failed: %v" + newStackFromPropagationFailed = "NewScopeStackFromPropagation failed: %v" + propagationParentUUID = "018f13f0-7c1a-7a80-8000-000000000002" + propagationRootUUID = "018f13f0-7c1a-7a80-8000-000000000001" +) func scopeNameInStack(stack *ScopeStack, scopeName string, scopeType ScopeType) (string, error) { var currentName string @@ -260,22 +267,22 @@ func TestCreateScopeStackCreatesFreshStack(t *testing.T) { } func TestNewScopeStackFromPropagationUsesParentAsCurrentHandle(t *testing.T) { - rootUUID := "018f13f0-7c1a-7a80-8000-000000000001" - parentUUID := "018f13f0-7c1a-7a80-8000-000000000002" + rootUUID := propagationRootUUID + parentUUID := propagationParentUUID stack, err := NewScopeStackFromPropagation(PropagationContext{ Version: 1, RootUUID: &rootUUID, ParentUUID: parentUUID, }) if err != nil { - t.Fatalf("NewScopeStackFromPropagation failed: %v", err) + t.Fatalf(newStackFromPropagationFailed, err) } defer stack.Close() stack.Run(func() { handle, err := GetHandle() if err != nil { - t.Fatalf("GetHandle failed: %v", err) + t.Fatalf(getHandleFailed, err) } if handle.UUID() != parentUUID { t.Fatalf("expected parent UUID %s, got %s", parentUUID, handle.UUID()) @@ -289,42 +296,50 @@ func TestPropagationContextCaptureAndValidation(t *testing.T) { t.Fatalf(newScopeStackFailed, err) } defer stack.Close() - stack.Run(func() { - context, err := CapturePropagationContext() - if err != nil { - t.Fatalf("CapturePropagationContext failed: %v", err) - } - if context.Version != 1 || context.RootUUID != nil || context.ParentUUID == "" { - t.Fatalf("unexpected rootless context: %+v", context) - } + assertCapturedPropagationContexts(t) + }) + assertInvalidPropagationContexts(t) +} - rootUUID := "018f13f0-7c1a-7a80-8000-000000000001" - withRoot, err := CapturePropagationContextWithRoot(&rootUUID) - if err != nil { - t.Fatalf("CapturePropagationContextWithRoot failed: %v", err) - } - if withRoot.RootUUID == nil || *withRoot.RootUUID != rootUUID { - t.Fatalf("expected root UUID %s, got %+v", rootUUID, withRoot.RootUUID) - } +func assertCapturedPropagationContexts(t *testing.T) { + t.Helper() + context, err := CapturePropagationContext() + if err != nil { + t.Fatalf("CapturePropagationContext failed: %v", err) + } + if context.Version != 1 || context.RootUUID != nil || context.ParentUUID == "" { + t.Fatalf("unexpected rootless context: %+v", context) + } - withNilRoot, err := CapturePropagationContextWithRoot(nil) - if err != nil { - t.Fatalf("CapturePropagationContextWithRoot(nil) failed: %v", err) - } - if withNilRoot.RootUUID != nil || withNilRoot.ParentUUID != context.ParentUUID { - t.Fatalf("unexpected nil-root context: %+v", withNilRoot) - } - }) + rootUUID := propagationRootUUID + withRoot, err := CapturePropagationContextWithRoot(&rootUUID) + if err != nil { + t.Fatalf("CapturePropagationContextWithRoot failed: %v", err) + } + if withRoot.RootUUID == nil || *withRoot.RootUUID != rootUUID { + t.Fatalf("expected root UUID %s, got %+v", rootUUID, withRoot.RootUUID) + } + + withNilRoot, err := CapturePropagationContextWithRoot(nil) + if err != nil { + t.Fatalf("CapturePropagationContextWithRoot(nil) failed: %v", err) + } + if withNilRoot.RootUUID != nil || withNilRoot.ParentUUID != context.ParentUUID { + t.Fatalf("unexpected nil-root context: %+v", withNilRoot) + } +} - invalidRoot := "not-a-uuid" +func assertInvalidPropagationContexts(t *testing.T) { + t.Helper() + invalidRoot := invalidUUID if _, err := CapturePropagationContextWithRoot(&invalidRoot); err == nil { t.Fatal("expected invalid root UUID to be rejected") } for _, context := range []PropagationContext{ - {Version: 2, ParentUUID: "018f13f0-7c1a-7a80-8000-000000000002"}, - {Version: 1, ParentUUID: "not-a-uuid"}, + {Version: 2, ParentUUID: propagationParentUUID}, + {Version: 1, ParentUUID: invalidUUID}, } { if _, err := NewScopeStackFromPropagation(context); err == nil { t.Fatalf("expected invalid context to be rejected: %+v", context) @@ -333,11 +348,11 @@ func TestPropagationContextCaptureAndValidation(t *testing.T) { } func TestPropagationContextJSONRoundTripAndValidation(t *testing.T) { - rootUUID := "018f13f0-7c1a-7a80-8000-000000000001" + rootUUID := propagationRootUUID context := PropagationContext{ Version: 1, RootUUID: &rootUUID, - ParentUUID: "018f13f0-7c1a-7a80-8000-000000000002", + ParentUUID: propagationParentUUID, } payload, err := context.ToJSON() @@ -370,7 +385,7 @@ func TestPropagationContextJSONRoundTripAndValidation(t *testing.T) { } } - if _, err := (PropagationContext{Version: 1, ParentUUID: "not-a-uuid"}).ToJSON(); err == nil { + if _, err := (PropagationContext{Version: 1, ParentUUID: invalidUUID}).ToJSON(); err == nil { t.Fatal("expected ToJSON to reject an invalid propagation context") } } @@ -383,13 +398,13 @@ func TestNewScopeStackFromRootlessAndRootParentPropagation(t *testing.T) { } { stack, err := NewScopeStackFromPropagation(context) if err != nil { - t.Fatalf("NewScopeStackFromPropagation failed: %v", err) + t.Fatalf(newStackFromPropagationFailed, err) } stack.Run(func() { handle, err := GetHandle() if err != nil { - t.Fatalf("GetHandle failed: %v", err) + t.Fatalf(getHandleFailed, err) } if handle.UUID() != parentUUID { t.Fatalf("expected propagated parent %s, got %s", parentUUID, handle.UUID()) @@ -409,19 +424,19 @@ func TestPropagatedScopeStackRunRestoresOuterBinding(t *testing.T) { parentUUID := "018f13f0-7c1a-7a80-8000-000000000005" propagated, err := NewScopeStackFromPropagation(PropagationContext{Version: 1, ParentUUID: parentUUID}) if err != nil { - t.Fatalf("NewScopeStackFromPropagation failed: %v", err) + t.Fatalf(newStackFromPropagationFailed, err) } defer propagated.Close() outer.Run(func() { outerHandle, err := GetHandle() if err != nil { - t.Fatalf("GetHandle failed: %v", err) + t.Fatalf(getHandleFailed, err) } propagated.Run(func() { handle, err := GetHandle() if err != nil { - t.Fatalf("GetHandle failed: %v", err) + t.Fatalf(getHandleFailed, err) } if handle.UUID() != parentUUID { t.Fatalf("expected propagated parent %s, got %s", parentUUID, handle.UUID()) @@ -438,8 +453,8 @@ func TestPropagatedScopeStackRunRestoresOuterBinding(t *testing.T) { } func TestPropagatedScopeStacksRemainIsolated(t *testing.T) { - rootUUID := "018f13f0-7c1a-7a80-8000-000000000001" - first, err := NewScopeStackFromPropagation(PropagationContext{Version: 1, RootUUID: &rootUUID, ParentUUID: "018f13f0-7c1a-7a80-8000-000000000002"}) + rootUUID := propagationRootUUID + first, err := NewScopeStackFromPropagation(PropagationContext{Version: 1, RootUUID: &rootUUID, ParentUUID: propagationParentUUID}) if err != nil { t.Fatal(err) } @@ -452,7 +467,7 @@ func TestPropagatedScopeStacksRemainIsolated(t *testing.T) { first.Run(func() { handle, _ := GetHandle() - if handle.UUID() != "018f13f0-7c1a-7a80-8000-000000000002" { + if handle.UUID() != propagationParentUUID { t.Fatalf("unexpected first propagated parent: %s", handle.UUID()) } }) diff --git a/go/nemo_relay/coverage_gap_test.go b/go/nemo_relay/coverage_gap_test.go index 7b9f480aa..f53a9d82e 100644 --- a/go/nemo_relay/coverage_gap_test.go +++ b/go/nemo_relay/coverage_gap_test.go @@ -14,7 +14,9 @@ import ( type failingAtofSink struct{} -func (failingAtofSink) atofExporterSink() {} +func (failingAtofSink) atofExporterSink() { + // This marker method intentionally has no runtime behavior. +} func (failingAtofSink) MarshalJSON() ([]byte, error) { return nil, errors.New("forced ATOF sink marshal failure") @@ -98,7 +100,7 @@ func assertInvalidScopePayloads(t *testing.T) { handle, err := PushScope("invalid_scope_end_metadata", ScopeTypeAgent) if err != nil { - t.Fatalf("PushScope failed: %v", err) + t.Fatalf(pushScopeFailed, err) } if PopScope(handle, WithScopeEndMetadata(json.RawMessage("{"))) == nil { t.Fatal("expected PopScope to fail on invalid end metadata JSON") @@ -200,7 +202,7 @@ func assertAdaptiveJSONUnmarshalFailures(t *testing.T) { } scope, err := PushScope("adaptive_bind_scope", ScopeTypeAgent) if err != nil { - t.Fatalf("PushScope failed: %v", err) + t.Fatalf(pushScopeFailed, err) } if err := runtime.BindScope(scope); err != nil { t.Fatalf("BindScope failed: %v", err) @@ -395,7 +397,9 @@ func assertClosedHandleErrorPaths(t *testing.T) { t.Fatal("expected closed scope stack Run to panic") } }() - stack.Run(func() {}) + stack.Run(func() { + // The empty callback isolates the closed-stack panic path. + }) }() exporter, err := NewAtofExporter(NewAtofExporterConfig()) @@ -424,7 +428,7 @@ func assertOptimizationContributionValidation(t *testing.T) { `{"producer":"test"}`, } { var contribution LLMOptimizationContribution - if err := json.Unmarshal([]byte(raw), &contribution); err == nil { + if json.Unmarshal([]byte(raw), &contribution) == nil { t.Fatalf("expected invalid optimization contribution %s to fail", raw) } } @@ -433,7 +437,7 @@ func assertOptimizationContributionValidation(t *testing.T) { func testWrapperAndCodecFinalizersRun(t *testing.T) { scopeHandle, err := PushScope("finalizer_scope", ScopeTypeAgent) if err != nil { - t.Fatalf("PushScope failed: %v", err) + t.Fatalf(pushScopeFailed, err) } if err := PopScope(scopeHandle); err != nil { t.Fatalf("PopScope failed: %v", err) diff --git a/go/nemo_relay/llm_test.go b/go/nemo_relay/llm_test.go index 5aa1eb837..f9ab2de5d 100644 --- a/go/nemo_relay/llm_test.go +++ b/go/nemo_relay/llm_test.go @@ -282,6 +282,126 @@ func TestLlmCallExecuteWithRequestAndResponseCodecs(t *testing.T) { } } +type resolvedCodecCallbackState struct { + sync.Mutex + requestResolved bool + responseResolved bool + retainedRequestCodec *LLMRequestSanitizeCodec + retainedResponseCodec *LLMResponseSanitizeCodec + retainedRequest LLMRequestDTO + errors []string +} + +type resolvedCodecSnapshot struct { + requestResolved bool + responseResolved bool + retainedRequestCodec *LLMRequestSanitizeCodec + retainedResponseCodec *LLMResponseSanitizeCodec + retainedRequest LLMRequestDTO + errors []string +} + +func (state *resolvedCodecCallbackState) recordError(format string, args ...any) { + state.Lock() + defer state.Unlock() + state.errors = append(state.errors, fmt.Sprintf(format, args...)) +} + +func (state *resolvedCodecCallbackState) sanitizeRequest( + request LLMRequestDTO, + context LLMSanitizeRequestContext, +) (LLMRequestDTO, bool) { + if context.Codec.CodecKind != LLMCodecOpaque || context.Codec.CodecID != nil { + state.recordError("unexpected request codec identity: %#v", context.Codec) + } + codec := context.ResolveCodec() + if codec == nil { + state.recordError("active request codec did not resolve") + return request, false + } + state.Lock() + state.retainedRequestCodec = codec + state.retainedRequest = request + state.Unlock() + annotated, err := codec.Decode(request) + if err != nil { + state.recordError("request codec decode failed: %v", err) + return request, false + } + encoded, err := codec.Encode(annotated, request) + if err != nil { + state.recordError("request codec encode failed: %v", err) + return request, false + } + state.Lock() + state.requestResolved = true + state.Unlock() + return encoded, false +} + +func (state *resolvedCodecCallbackState) sanitizeResponse( + response json.RawMessage, + context LLMSanitizeResponseContext, +) (json.RawMessage, bool) { + if context.Codec.CodecKind != LLMCodecBuiltin || + context.Codec.CodecID == nil || + *context.Codec.CodecID != "openai_chat" { + state.recordError("unexpected response codec identity: %#v", context.Codec) + } + codec := context.ResolveCodec() + if codec == nil { + state.recordError("active response codec did not resolve") + return response, false + } + state.Lock() + state.retainedResponseCodec = codec + state.Unlock() + if _, err := codec.Decode(response); err != nil { + state.recordError("response codec decode failed: %v", err) + return response, false + } + state.Lock() + state.responseResolved = true + state.Unlock() + return response, false +} + +func (state *resolvedCodecCallbackState) snapshot() resolvedCodecSnapshot { + state.Lock() + defer state.Unlock() + return resolvedCodecSnapshot{ + requestResolved: state.requestResolved, + responseResolved: state.responseResolved, + retainedRequestCodec: state.retainedRequestCodec, + retainedResponseCodec: state.retainedResponseCodec, + retainedRequest: state.retainedRequest, + errors: append([]string(nil), state.errors...), + } +} + +func assertResolvedCodecsExpire(t *testing.T, snapshot resolvedCodecSnapshot, response json.RawMessage) { + t.Helper() + if len(snapshot.errors) != 0 { + t.Fatalf("sanitizer callbacks failed: %v", snapshot.errors) + } + if !snapshot.requestResolved || !snapshot.responseResolved { + t.Fatalf( + "expected both codec capabilities to resolve, request=%t response=%t", + snapshot.requestResolved, + snapshot.responseResolved, + ) + } + if _, err := snapshot.retainedRequestCodec.Decode(snapshot.retainedRequest); !errors.Is(err, ErrLLMSanitizeCodecExpired) { + t.Fatalf("retained request codec must expire after callback, got %v", err) + } + if _, err := snapshot.retainedRequestCodec.Encode(json.RawMessage(`{}`), snapshot.retainedRequest); !errors.Is(err, ErrLLMSanitizeCodecExpired) { + t.Fatalf("retained request codec encode must expire after callback, got %v", err) + } + if _, err := snapshot.retainedResponseCodec.Decode(response); !errors.Is(err, ErrLLMSanitizeCodecExpired) { + t.Fatalf("retained response codec must expire after callback, got %v", err) + } +} + func TestLlmSanitizersResolveDirectionalCodecs(t *testing.T) { const requestGuard = "go_llm_resolved_request_codec" const responseGuard = "go_llm_resolved_response_codec" @@ -290,81 +410,18 @@ func TestLlmSanitizersResolveDirectionalCodecs(t *testing.T) { defer DeregisterLlmSanitizeRequestGuardrail(requestGuard) defer DeregisterLlmSanitizeResponseGuardrail(responseGuard) - var callbackState struct { - sync.Mutex - requestResolved bool - responseResolved bool - retainedRequestCodec *LLMRequestSanitizeCodec - retainedResponseCodec *LLMResponseSanitizeCodec - retainedRequest LLMRequestDTO - errors []string - } - recordCallbackError := func(format string, args ...any) { - callbackState.Lock() - defer callbackState.Unlock() - callbackState.errors = append(callbackState.errors, fmt.Sprintf(format, args...)) - } + callbackState := &resolvedCodecCallbackState{} if err := RegisterLlmSanitizeRequestGuardrail( requestGuard, 0, - func(request LLMRequestDTO, context LLMSanitizeRequestContext) (LLMRequestDTO, bool) { - if context.Codec.CodecKind != LLMCodecOpaque || - context.Codec.CodecID != nil { - recordCallbackError("unexpected request codec identity: %#v", context.Codec) - } - codec := context.ResolveCodec() - if codec == nil { - recordCallbackError("active request codec did not resolve") - return request, false - } - callbackState.Lock() - callbackState.retainedRequestCodec = codec - callbackState.retainedRequest = request - callbackState.Unlock() - annotated, err := codec.Decode(request) - if err != nil { - recordCallbackError("request codec decode failed: %v", err) - return request, false - } - encoded, err := codec.Encode(annotated, request) - if err != nil { - recordCallbackError("request codec encode failed: %v", err) - return request, false - } - callbackState.Lock() - callbackState.requestResolved = true - callbackState.Unlock() - return encoded, false - }, + callbackState.sanitizeRequest, ); err != nil { t.Fatalf("request sanitizer registration failed: %v", err) } if err := RegisterLlmSanitizeResponseGuardrail( responseGuard, 0, - func(response json.RawMessage, context LLMSanitizeResponseContext) (json.RawMessage, bool) { - if context.Codec.CodecKind != LLMCodecBuiltin || - context.Codec.CodecID == nil || - *context.Codec.CodecID != "openai_chat" { - recordCallbackError("unexpected response codec identity: %#v", context.Codec) - } - codec := context.ResolveCodec() - if codec == nil { - recordCallbackError("active response codec did not resolve") - return response, false - } - callbackState.Lock() - callbackState.retainedResponseCodec = codec - callbackState.Unlock() - if _, err := codec.Decode(response); err != nil { - recordCallbackError("response codec decode failed: %v", err) - return response, false - } - callbackState.Lock() - callbackState.responseResolved = true - callbackState.Unlock() - return response, false - }, + callbackState.sanitizeResponse, ); err != nil { t.Fatalf("response sanitizer registration failed: %v", err) } @@ -388,33 +445,7 @@ func TestLlmSanitizersResolveDirectionalCodecs(t *testing.T) { if err != nil { t.Fatalf(llmCallExecuteFailed, err) } - callbackState.Lock() - sanitizerErrors := append([]string(nil), callbackState.errors...) - requestResolved := callbackState.requestResolved - responseResolved := callbackState.responseResolved - retainedRequestCodec := callbackState.retainedRequestCodec - retainedResponseCodec := callbackState.retainedResponseCodec - retainedRequest := callbackState.retainedRequest - callbackState.Unlock() - if len(sanitizerErrors) != 0 { - t.Fatalf("sanitizer callbacks failed: %v", sanitizerErrors) - } - if !requestResolved || !responseResolved { - t.Fatalf( - "expected both codec capabilities to resolve, request=%t response=%t", - requestResolved, - responseResolved, - ) - } - if _, err := retainedRequestCodec.Decode(retainedRequest); !errors.Is(err, ErrLLMSanitizeCodecExpired) { - t.Fatalf("retained request codec must expire after callback, got %v", err) - } - if _, err := retainedRequestCodec.Encode(json.RawMessage(`{}`), retainedRequest); !errors.Is(err, ErrLLMSanitizeCodecExpired) { - t.Fatalf("retained request codec encode must expire after callback, got %v", err) - } - if _, err := retainedResponseCodec.Decode(response); !errors.Is(err, ErrLLMSanitizeCodecExpired) { - t.Fatalf("retained response codec must expire after callback, got %v", err) - } + assertResolvedCodecsExpire(t, callbackState.snapshot(), response) } func llmRequestResponseCodec() CodecFunc { @@ -522,53 +553,90 @@ func TestLlmSanitizeResponseGuardrail(t *testing.T) { DeregisterLlmSanitizeResponseGuardrail("go_llm_san_resp") } +type contextualLlmEventCapture struct { + sync.Mutex + input json.RawMessage + output json.RawMessage +} + +func (capture *contextualLlmEventCapture) record(event Event) { + if event.Kind() != "scope" || event.Category() != "llm" { + return + } + capture.Lock() + defer capture.Unlock() + switch event.ScopeCategory() { + case "start": + capture.input = append(json.RawMessage(nil), event.Input()...) + case "end": + capture.output = append(json.RawMessage(nil), event.Output()...) + } +} + +func (capture *contextualLlmEventCapture) snapshot() (json.RawMessage, json.RawMessage) { + capture.Lock() + defer capture.Unlock() + return append(json.RawMessage(nil), capture.input...), append(json.RawMessage(nil), capture.output...) +} + +type contextualLlmCallbackErrors struct { + sync.Mutex + errors []string +} + +func (callbackErrors *contextualLlmCallbackErrors) record(message string) { + callbackErrors.Lock() + defer callbackErrors.Unlock() + callbackErrors.errors = append(callbackErrors.errors, message) +} + +func (callbackErrors *contextualLlmCallbackErrors) sanitizeRequest( + request LLMRequestDTO, + context LLMSanitizeRequestContext, +) (LLMRequestDTO, bool) { + if context.Codec.CodecKind != LLMCodecNone { + callbackErrors.record("manual registration received an active codec identity") + } + return request, true +} + +func (callbackErrors *contextualLlmCallbackErrors) sanitizeResponse( + response json.RawMessage, + context LLMSanitizeResponseContext, +) (json.RawMessage, bool) { + if context.Codec.CodecID != nil { + callbackErrors.record("manual registration received a codec ID") + } + return response, true +} + +func (callbackErrors *contextualLlmCallbackErrors) snapshot() []string { + callbackErrors.Lock() + defer callbackErrors.Unlock() + return append([]string(nil), callbackErrors.errors...) +} + func TestLlmSanitizeGuardrailsReceiveContext(t *testing.T) { - var capturedInput, capturedOutput json.RawMessage - var mu sync.Mutex - var callbackErrorsMu sync.Mutex - var callbackErrors []string - recordCallbackError := func(message string) { - callbackErrorsMu.Lock() - defer callbackErrorsMu.Unlock() - callbackErrors = append(callbackErrors, message) - } - if err := RegisterSubscriber("go_contextual_llm_sanitize_events", func(event Event) { - mu.Lock() - defer mu.Unlock() - if event.Kind() == "scope" && event.Category() == "llm" && event.ScopeCategory() == "start" { - capturedInput = append(json.RawMessage(nil), event.Input()...) - } - if event.Kind() == "scope" && event.Category() == "llm" && event.ScopeCategory() == "end" { - capturedOutput = append(json.RawMessage(nil), event.Output()...) - } - }); err != nil { + const subscriberName = "go_contextual_llm_sanitize_events" + const requestGuard = "go_contextual_llm_request" + const responseGuard = "go_contextual_llm_response" + + capture := &contextualLlmEventCapture{} + callbackErrors := &contextualLlmCallbackErrors{} + if err := RegisterSubscriber(subscriberName, capture.record); err != nil { t.Fatalf("RegisterSubscriber failed: %v", err) } - defer DeregisterSubscriber("go_contextual_llm_sanitize_events") + defer DeregisterSubscriber(subscriberName) - if err := RegisterLlmSanitizeRequestGuardrail("go_contextual_llm_request", 1, - func(request LLMRequestDTO, context LLMSanitizeRequestContext) (LLMRequestDTO, bool) { - if context.Codec.CodecKind != LLMCodecNone { - recordCallbackError("manual registration received an active codec identity") - } - return request, true - }, - ); err != nil { + if err := RegisterLlmSanitizeRequestGuardrail(requestGuard, 1, callbackErrors.sanitizeRequest); err != nil { t.Fatalf(llmRegisterFailed, err) } - defer DeregisterLlmSanitizeRequestGuardrail("go_contextual_llm_request") + defer DeregisterLlmSanitizeRequestGuardrail(requestGuard) - if err := RegisterLlmSanitizeResponseGuardrail("go_contextual_llm_response", 1, - func(response json.RawMessage, context LLMSanitizeResponseContext) (json.RawMessage, bool) { - if context.Codec.CodecID != nil { - recordCallbackError("manual registration received a codec ID") - } - return response, true - }, - ); err != nil { + if err := RegisterLlmSanitizeResponseGuardrail(responseGuard, 1, callbackErrors.sanitizeResponse); err != nil { t.Fatalf(llmRegisterFailed, err) } - defer DeregisterLlmSanitizeResponseGuardrail("go_contextual_llm_response") + defer DeregisterLlmSanitizeResponseGuardrail(responseGuard) result, err := LlmCallExecute("go_contextual_llm_sanitize", makeRequest(), func(nativeJSON json.RawMessage) (json.RawMessage, error) { @@ -578,9 +646,7 @@ func TestLlmSanitizeGuardrailsReceiveContext(t *testing.T) { if err != nil { t.Fatalf(llmCallExecuteFailed, err) } - callbackErrorsMu.Lock() - sanitizerErrors := append([]string(nil), callbackErrors...) - callbackErrorsMu.Unlock() + sanitizerErrors := callbackErrors.snapshot() if len(sanitizerErrors) != 0 { t.Fatalf("sanitizer callbacks failed: %v", sanitizerErrors) } @@ -590,8 +656,7 @@ func TestLlmSanitizeGuardrailsReceiveContext(t *testing.T) { if err := FlushSubscribers(); err != nil { t.Fatalf(llmFlushSubscribersFailed, err) } - mu.Lock() - defer mu.Unlock() + capturedInput, capturedOutput := capture.snapshot() if capturedInput != nil || capturedOutput != nil { t.Fatalf("contextual omission must remove observability payloads, got input=%s output=%s", capturedInput, capturedOutput) } diff --git a/go/nemo_relay/nemo_relay.go b/go/nemo_relay/nemo_relay.go index a3e80fd18..3046120f5 100644 --- a/go/nemo_relay/nemo_relay.go +++ b/go/nemo_relay/nemo_relay.go @@ -2010,16 +2010,15 @@ type OpenTelemetrySubscriber struct { ptr unsafe.Pointer } -// NewOpenTelemetrySubscriber creates a new OpenTelemetry subscriber from config. -func NewOpenTelemetrySubscriber(config OpenTelemetryConfig) (*OpenTelemetrySubscriber, error) { +func normalizeOpenTelemetryConfig(config OpenTelemetryConfig) (OpenTelemetryConfig, error) { if config.Transport == "" { config.Transport = OpenTelemetryTransportHTTPBinary } if config.Type == "" { - return nil, fmt.Errorf("type is required") + return config, fmt.Errorf("type is required") } if config.Endpoint == "" { - return nil, fmt.Errorf("endpoint is required") + return config, fmt.Errorf("endpoint is required") } if config.ServiceName == "" { config.ServiceName = "unknown_service" @@ -2045,17 +2044,30 @@ func NewOpenTelemetrySubscriber(config OpenTelemetryConfig) (*OpenTelemetrySubsc if config.AttributeMappings == nil { config.AttributeMappings = []OtlpAttributeMapping{} } + return config, nil +} + +func optionalCString(value string) *C.char { + if value == "" { + return nil + } + return C.CString(value) +} + +// NewOpenTelemetrySubscriber creates a new OpenTelemetry subscriber from config. +func NewOpenTelemetrySubscriber(config OpenTelemetryConfig) (*OpenTelemetrySubscriber, error) { + config, err := normalizeOpenTelemetryConfig(config) + if err != nil { + return nil, err + } cTransport := C.CString(string(config.Transport)) defer C.free(unsafe.Pointer(cTransport)) cType := C.CString(string(config.Type)) defer C.free(unsafe.Pointer(cType)) - var cEndpoint *C.char - if config.Endpoint != "" { - cEndpoint = C.CString(config.Endpoint) - defer C.free(unsafe.Pointer(cEndpoint)) - } + cEndpoint := C.CString(config.Endpoint) + defer C.free(unsafe.Pointer(cEndpoint)) headersJSON, err := jsonMarshal(config.Headers) if err != nil { @@ -2074,17 +2086,11 @@ func NewOpenTelemetrySubscriber(config OpenTelemetryConfig) (*OpenTelemetrySubsc cServiceName := C.CString(config.ServiceName) defer C.free(unsafe.Pointer(cServiceName)) - var cServiceNamespace *C.char - if config.ServiceNamespace != "" { - cServiceNamespace = C.CString(config.ServiceNamespace) - defer C.free(unsafe.Pointer(cServiceNamespace)) - } + cServiceNamespace := optionalCString(config.ServiceNamespace) + defer C.free(unsafe.Pointer(cServiceNamespace)) - var cServiceVersion *C.char - if config.ServiceVersion != "" { - cServiceVersion = C.CString(config.ServiceVersion) - defer C.free(unsafe.Pointer(cServiceVersion)) - } + cServiceVersion := optionalCString(config.ServiceVersion) + defer C.free(unsafe.Pointer(cServiceVersion)) cInstrumentationScope := C.CString(config.InstrumentationScope) defer C.free(unsafe.Pointer(cInstrumentationScope)) diff --git a/go/nemo_relay/otel_test.go b/go/nemo_relay/otel_test.go index 4df954fd2..518a888cd 100644 --- a/go/nemo_relay/otel_test.go +++ b/go/nemo_relay/otel_test.go @@ -14,6 +14,14 @@ import ( "time" ) +const ( + newOpenTelemetrySubscriberFailed = "NewOpenTelemetrySubscriber failed: %v" + otelRegisterFailed = "Register failed: %v" + otelTestEndpoint = "http://localhost:4318/v1/traces" + otelTestPath = "/v1/traces" + otelTimeFormat = "150405.000000" +) + func assertOtlpStringAttribute(t *testing.T, body []byte, key string, value string) { t.Helper() encoded := append([]byte{0x0a}, binary.AppendUvarint(nil, uint64(len(key)))...) @@ -29,7 +37,7 @@ func assertOtlpStringAttribute(t *testing.T, body []byte, key string, value stri } func TestNewOpenTelemetryConfigDefaults(t *testing.T) { - config := NewOpenTelemetryConfig(OpenTelemetryTypeFull, "http://localhost:4318/v1/traces") + config := NewOpenTelemetryConfig(OpenTelemetryTypeFull, otelTestEndpoint) if config.Transport != OpenTelemetryTransportHTTPBinary { t.Fatalf("expected default transport http_binary, got %q", config.Transport) @@ -61,7 +69,7 @@ func TestNewOpenTelemetryConfigDefaults(t *testing.T) { } func TestOpenTelemetrySubscriberAcceptsProjectionControls(t *testing.T) { - config := NewOpenTelemetryConfig(OpenTelemetryTypeFull, "http://localhost:4318/v1/traces") + config := NewOpenTelemetryConfig(OpenTelemetryTypeFull, otelTestEndpoint) config.MarkProjection = MarkProjectionTool config.MarkExcludeNames = []string{"custom.mark"} config.AttributeMappings = []OtlpAttributeMapping{{ @@ -77,7 +85,7 @@ func TestOpenTelemetrySubscriberAcceptsProjectionControls(t *testing.T) { } func TestOpenTelemetrySubscriberRejectsInvalidAttributeMappings(t *testing.T) { - config := NewOpenTelemetryConfig(OpenTelemetryTypeFull, "http://localhost:4318/v1/traces") + config := NewOpenTelemetryConfig(OpenTelemetryTypeFull, otelTestEndpoint) config.AttributeMappings = []OtlpAttributeMapping{{Key: "", Alias: "model.alias"}} if _, err := NewOpenTelemetrySubscriber(config); err == nil { @@ -86,7 +94,7 @@ func TestOpenTelemetrySubscriberRejectsInvalidAttributeMappings(t *testing.T) { } func TestOpenTelemetrySubscriberLifecycle(t *testing.T) { - config := NewOpenTelemetryConfig(OpenTelemetryTypeFull, "http://localhost:4318/v1/traces") + config := NewOpenTelemetryConfig(OpenTelemetryTypeFull, otelTestEndpoint) config.ServiceName = "go-agent" config.ServiceNamespace = "agents" config.ServiceVersion = "1.0.0" @@ -96,13 +104,13 @@ func TestOpenTelemetrySubscriberLifecycle(t *testing.T) { config.ResourceAttributes["deployment.environment"] = "test" subscriber, err := NewOpenTelemetrySubscriber(config) if err != nil { - t.Fatalf("NewOpenTelemetrySubscriber failed: %v", err) + t.Fatalf(newOpenTelemetrySubscriberFailed, err) } defer subscriber.Close() - name := "go_otel_subscriber_" + time.Now().Format("150405.000000") + name := "go_otel_subscriber_" + time.Now().Format(otelTimeFormat) if err := subscriber.Register(name); err != nil { - t.Fatalf("Register failed: %v", err) + t.Fatalf(otelRegisterFailed, err) } if err := subscriber.Deregister(name); err != nil { t.Fatalf("Deregister failed: %v", err) @@ -119,7 +127,7 @@ func TestOpenTelemetrySubscriberLifecycle(t *testing.T) { } func TestOpenTelemetrySubscriberRejectsInvalidTransport(t *testing.T) { - config := NewOpenTelemetryConfig(OpenTelemetryTypeFull, "http://localhost:4318/v1/traces") + config := NewOpenTelemetryConfig(OpenTelemetryTypeFull, otelTestEndpoint) config.Transport = OpenTelemetryTransport("invalid") _, err := NewOpenTelemetrySubscriber(config) @@ -135,11 +143,11 @@ func TestOpenTelemetrySubscriberRejectsInvalidRequiredFields(t *testing.T) { }{ { name: "missing type", - config: NewOpenTelemetryConfig("", "http://localhost:4318/v1/traces"), + config: NewOpenTelemetryConfig("", otelTestEndpoint), }, { name: "unknown type", - config: NewOpenTelemetryConfig(OpenTelemetryType("invalid"), "http://localhost:4318/v1/traces"), + config: NewOpenTelemetryConfig(OpenTelemetryType("invalid"), otelTestEndpoint), }, { name: "missing endpoint", @@ -184,17 +192,17 @@ func TestOpenTelemetrySubscriberExportsScopeLifecycleAndMarks(t *testing.T) { })) defer server.Close() - config := NewOpenTelemetryConfig(OpenTelemetryTypeFull, server.URL+"/v1/traces") + config := NewOpenTelemetryConfig(OpenTelemetryTypeFull, server.URL+otelTestPath) config.ServiceName = "go-agent" subscriber, err := NewOpenTelemetrySubscriber(config) if err != nil { - t.Fatalf("NewOpenTelemetrySubscriber failed: %v", err) + t.Fatalf(newOpenTelemetrySubscriberFailed, err) } defer subscriber.Close() - name := "go_otel_e2e_" + time.Now().Format("150405.000000") + name := "go_otel_e2e_" + time.Now().Format(otelTimeFormat) if err := subscriber.Register(name); err != nil { - t.Fatalf("Register failed: %v", err) + t.Fatalf(otelRegisterFailed, err) } defer func() { _ = subscriber.Deregister(name) }() @@ -221,7 +229,7 @@ func TestOpenTelemetrySubscriberExportsScopeLifecycleAndMarks(t *testing.T) { select { case request := <-requests: - if request.Path != "/v1/traces" { + if request.Path != otelTestPath { t.Fatalf("expected /v1/traces path, got %q", request.Path) } if request.ContentType != "application/x-protobuf" { @@ -241,16 +249,16 @@ func TestOpenTelemetrySubscriberExportsGenAIAgentProjection(t *testing.T) { server := NewOtelTestServer(t, requests) defer server.Close() - config := NewOpenTelemetryConfig(OpenTelemetryTypeGenAI, server.URL+"/v1/traces") + config := NewOpenTelemetryConfig(OpenTelemetryTypeGenAI, server.URL+otelTestPath) subscriber, err := NewOpenTelemetrySubscriber(config) if err != nil { - t.Fatalf("NewOpenTelemetrySubscriber failed: %v", err) + t.Fatalf(newOpenTelemetrySubscriberFailed, err) } defer subscriber.Close() - name := "go_gen_ai_e2e_" + time.Now().Format("150405.000000") + name := "go_gen_ai_e2e_" + time.Now().Format(otelTimeFormat) if err := subscriber.Register(name); err != nil { - t.Fatalf("Register failed: %v", err) + t.Fatalf(otelRegisterFailed, err) } defer func() { _ = subscriber.Deregister(name) }() diff --git a/go/nemo_relay/plugin_activation_test.go b/go/nemo_relay/plugin_activation_test.go index ba1cfd25b..3d31d347c 100644 --- a/go/nemo_relay/plugin_activation_test.go +++ b/go/nemo_relay/plugin_activation_test.go @@ -763,12 +763,18 @@ func assertMissingNativePluginFails(t *testing.T) { } func TestInitializeWithDynamicPluginsLoadsWorkerPluginThroughCgo(t *testing.T) { + executable := goWorkerPluginFixture(t) + manifest := writeGoWorkerPluginManifest(t, executable) + runTestInIsolatedWorkingDirectory(t, func(t *testing.T) { + testInitializeWithDynamicWorkerPlugin(t, manifest) + }) +} + +func testInitializeWithDynamicWorkerPlugin(t *testing.T, manifest string) { if err := ClearPluginConfiguration(); err != nil { t.Fatalf(clearConfigurationErrorFmt, err) } - executable := goWorkerPluginFixture(t) - manifest := writeGoWorkerPluginManifest(t, executable) activation, report, err := InitializeWithDynamicPlugins(NewPluginConfig(), []DynamicPluginActivationSpec{{ PluginID: "fixture_worker", Kind: DynamicPluginKindWorker, @@ -812,9 +818,6 @@ func TestInitializeWithDynamicPluginsLoadsWorkerPluginThroughCgo(t *testing.T) { } func TestPluginActivationFinalizerReleasesHostOwnership(t *testing.T) { - if err := ClearPluginConfiguration(); err != nil { - t.Fatalf(clearConfigurationErrorFmt, err) - } library := goNativePluginFixture(t) manifest := writeGoNativePluginManifest(t, library) specs := []DynamicPluginActivationSpec{{ @@ -823,6 +826,15 @@ func TestPluginActivationFinalizerReleasesHostOwnership(t *testing.T) { ManifestRef: manifest, Config: map[string]any{}, }} + runTestInIsolatedWorkingDirectory(t, func(t *testing.T) { + testPluginActivationFinalizerReleasesHostOwnership(t, specs) + }) +} + +func testPluginActivationFinalizerReleasesHostOwnership(t *testing.T, specs []DynamicPluginActivationSpec) { + if err := ClearPluginConfiguration(); err != nil { + t.Fatalf(clearConfigurationErrorFmt, err) + } createUnclosedPluginActivation(t, specs) deadline := time.Now().Add(10 * time.Second) diff --git a/go/nemo_relay/test_helpers_test.go b/go/nemo_relay/test_helpers_test.go index 4cc5dd5af..05aaa47ee 100644 --- a/go/nemo_relay/test_helpers_test.go +++ b/go/nemo_relay/test_helpers_test.go @@ -3,7 +3,32 @@ package nemo_relay -import "testing" +import ( + "os" + "testing" +) + +// runTestInIsolatedWorkingDirectory changes the process-wide working directory, +// so callers must not use it with t.Parallel(). +func runTestInIsolatedWorkingDirectory(t *testing.T, fn func(*testing.T)) { + t.Helper() + + t.Setenv("XDG_CONFIG_HOME", t.TempDir()) + original, err := os.Getwd() + if err != nil { + t.Fatalf("Getwd failed: %v", err) + } + if err := os.Chdir(t.TempDir()); err != nil { + t.Fatalf("Chdir to temporary directory failed: %v", err) + } + defer func() { + if err := os.Chdir(original); err != nil { + t.Errorf("restore working directory failed: %v", err) + } + }() + + fn(t) +} func runWithTestScopeStack(t *testing.T, fn func()) { t.Helper() diff --git a/python/plugin/build_backend.py b/python/plugin/build_backend.py index 2200d104b..8a7481ccf 100644 --- a/python/plugin/build_backend.py +++ b/python/plugin/build_backend.py @@ -160,7 +160,8 @@ def _sdist_manifest() -> Iterator[None]: if previous is None: _MANIFEST.unlink(missing_ok=True) else: - _MANIFEST.write_bytes(previous) + # `_MANIFEST` is resolved and constrained to a direct child of this backend's directory. + _MANIFEST.write_bytes(previous) # NOSONAR def _setuptools_backend() -> Any: diff --git a/scripts/package_node_musllinux.mjs b/scripts/package_node_musllinux.mjs index 8428370e7..1b2c3146b 100755 --- a/scripts/package_node_musllinux.mjs +++ b/scripts/package_node_musllinux.mjs @@ -20,13 +20,14 @@ function argumentsFrom(args) { let version; let output; let platform; - for (let index = 0; index < args.length; index += 1) { + for (let index = 0; index < args.length; index += 2) { + const value = args[index + 1]; if (args[index] === "--version") { - version = args[++index]; + version = value; } else if (args[index] === "--out") { - output = args[++index]; + output = value; } else if (args[index] === "--platform") { - platform = args[++index]; + platform = value; } else { throw new Error(`Unexpected argument: ${args[index]}`); }