diff --git a/Cargo.toml b/Cargo.toml index dc6d24f33..c5da34540 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -109,7 +109,7 @@ n00n-markdown = { path = "n00n-markdown" } n00n-daemon = { path = "n00n-daemon" } n00n-token-profile = { path = "n00n-token-profile" } serde = { version = "1", features = ["derive", "rc"] } -serde_json = "1" +serde_json = { version = "1", features = ["arbitrary_precision"] } sonic-rs = "0.5.8" toon-format = { version = "0.5", default-features = false } tiktoken-rs = "0.12" diff --git a/changelog.d/320.added.md b/changelog.d/320.added.md new file mode 100644 index 000000000..20a015ec4 --- /dev/null +++ b/changelog.d/320.added.md @@ -0,0 +1 @@ +Added host-owned, root- and session-scoped Lua plugin state with bounded snapshots, lifecycle restoration, and cleanup. diff --git a/n00n-agent/src/agent/run.rs b/n00n-agent/src/agent/run.rs index 70bc03171..93a63f857 100644 --- a/n00n-agent/src/agent/run.rs +++ b/n00n-agent/src/agent/run.rs @@ -19,8 +19,8 @@ use crate::cancel::{CancelMap, CancelToken, PreDispatchGate}; use crate::mcp::McpSession; use crate::permissions::{PermissionAnswer, PermissionManager}; use crate::tools::{ - ActiveTools, Deadline, FileReadTracker, LocalTools, ToolAudience, ToolContext, ToolFilter, - ToolRegistry, + ActiveTools, Deadline, FileReadTracker, LocalTools, SessionIdentity, ToolAudience, ToolContext, + ToolFilter, ToolRegistry, }; use crate::{ AgentConfig, AgentError, AgentEvent, AgentInput, AgentMode, EventSender, ExtractedCommand, @@ -28,6 +28,7 @@ use crate::{ InterruptSource, ToolDoneEvent, TurnCompleteEvent, }; use n00n_config::{ToolKey, ToolOutputLines}; +#[cfg(test)] use n00n_storage::id::SessionRef; use crate::tokenize::{ @@ -148,7 +149,7 @@ pub struct AgentParams { pub config: Arc, pub tool_output_lines: ToolOutputLines, pub permissions: Arc, - pub session_id: Option, + pub identity: Option, pub timeouts: n00n_providers::Timeouts, pub openai_options: OpenAiOptions, pub file_tracker: Arc, @@ -195,7 +196,7 @@ pub struct Agent<'h> { thinking_empty_retried: bool, permissions: Arc, opts: RequestOptions, - session_id: Option, + identity: Option, timeouts: n00n_providers::Timeouts, openai_options: OpenAiOptions, file_tracker: Arc, @@ -254,7 +255,7 @@ impl<'h> Agent<'h> { post_tool_empty_retried: false, thinking_empty_retried: false, opts: RequestOptions::default(), - session_id: params.session_id, + identity: params.identity, file_tracker: params.file_tracker, prompt_slots: params.prompt_slots, subagent_cancels: params.subagent_cancels, @@ -484,7 +485,7 @@ impl<'h> Agent<'h> { event_tx: &self.event_tx, cancel: &self.cancel, opts, - session_id: self.session_id.as_ref(), + session_id: self.identity.as_ref().map(SessionIdentity::session_id), }) .await } @@ -897,12 +898,22 @@ impl<'h> Agent<'h> { } fn effective_tool_filter(&self) -> ToolFilter { - if !self.allow_dynamic_mcp_tools { - return self.tool_filter.clone(); - } let Some(mcp) = self.mcp.as_ref() else { return self.tool_filter.clone(); }; + let mut filter = self.tool_filter.clone(); + let tool_search = crate::mcp::TOOL_SEARCH_TOOL_NAME; + if crate::tools::is_tool_enabled(&self.config.disabled_tools, tool_search) { + if !filter.matches(tool_search) { + filter = filter.including([tool_search.to_owned()]); + } + } else { + filter = filter.excluding(&[tool_search]); + } + if !self.allow_dynamic_mcp_tools { + return filter; + } + let capability_exclusions = crate::tools::capability_exclusions(&self.model); let mut definitions = Value::Array(Vec::new()); mcp.extend_tools(&mut definitions); let names = definitions @@ -910,8 +921,12 @@ impl<'h> Agent<'h> { .into_iter() .flatten() .filter_map(|definition| definition.get("name").and_then(Value::as_str)) + .filter(|name| { + crate::tools::is_tool_enabled(&self.config.disabled_tools, name) + && !capability_exclusions.contains(name) + }) .map(str::to_owned); - self.tool_filter.clone().including(names) + filter.including(names) } fn tool_context(&self) -> ToolContext { @@ -935,6 +950,7 @@ impl<'h> Agent<'h> { prompt_slots: Arc::clone(&self.prompt_slots), opts: self.opts.clone(), subagent_cancels: Arc::clone(&self.subagent_cancels), + identity: self.identity.clone(), registry: Arc::clone(&self.registry), workflow: self.workflow, audience: self.audience, @@ -1097,7 +1113,7 @@ impl<'h> Agent<'h> { &self.event_tx, &self.cancel, CompactionTrigger::Auto, - self.session_id.as_ref(), + self.identity.as_ref().map(SessionIdentity::session_id), &cwd, None, ) @@ -1346,6 +1362,74 @@ mod tests { assert_eq!(names, ["codegraph", "server__search"]); } + #[test] + fn dynamic_mcp_filter_includes_tool_search_with_mcp() { + let mut history = History::new(Vec::new()); + let (mut agent, _) = make_agent(MockProvider::new(Vec::new()), &mut history); + let mcp = crate::mcp::stub_session(&[("srv.fetch_issue", "Fetch a GitHub issue")]); + agent.tool_filter = ToolFilter::Only(vec!["read".into()]); + agent = agent.with_mcp(Some(mcp)).with_dynamic_mcp_tools(false); + + let effective_filter = agent.effective_tool_filter(); + assert!(effective_filter.matches("tool_search")); + assert!(effective_filter.matches("read")); + assert!(!effective_filter.matches("write")); + assert!(!effective_filter.matches("srv__fetch_issue")); + } + + #[test] + fn dynamic_mcp_filter_keeps_disabled_tool_search_blocked() { + let mut history = History::new(Vec::new()); + let (mut agent, _) = make_agent(MockProvider::new(Vec::new()), &mut history); + let mcp = crate::mcp::stub_session(&[("srv.fetch_issue", "Fetch a GitHub issue")]); + let mut config = (*agent.config).clone(); + config + .disabled_tools + .push(crate::mcp::TOOL_SEARCH_TOOL_NAME.into()); + agent.config = Arc::new(config); + agent = agent.with_mcp(Some(mcp)).with_dynamic_mcp_tools(true); + + assert!(!agent.effective_tool_filter().matches("tool_search")); + } + + #[test] + fn dynamic_mcp_filter_keeps_disabled_tools_blocked() { + const DISABLED_MCP_TOOL: &str = "srv__fetch_issue"; + + let mut history = History::new(Vec::new()); + let (mut agent, _) = make_agent(MockProvider::new(Vec::new()), &mut history); + let mcp = crate::mcp::stub_session(&[("srv.fetch_issue", "Fetch a GitHub issue")]); + let mut config = (*agent.config).clone(); + config.disabled_tools.push(DISABLED_MCP_TOOL.into()); + agent.config = Arc::new(config); + agent.tool_filter = ToolFilter::Only(vec![crate::mcp::TOOL_SEARCH_TOOL_NAME.into()]); + agent = agent + .with_mcp(Some(mcp.clone())) + .with_dynamic_mcp_tools(true); + + mcp.search_tools("issue").unwrap(); + + assert!(!agent.effective_tool_filter().matches(DISABLED_MCP_TOOL)); + } + + #[test] + fn dynamic_mcp_filter_includes_loaded_tools_with_flag() { + let mut history = History::new(Vec::new()); + let (mut agent, _) = make_agent(MockProvider::new(Vec::new()), &mut history); + let mcp = crate::mcp::stub_session(&[("srv.fetch_issue", "Fetch a GitHub issue")]); + agent.tool_filter = ToolFilter::Only(vec!["read".into()]); + agent = agent + .with_mcp(Some(mcp.clone())) + .with_dynamic_mcp_tools(true); + + mcp.search_tools("issue").unwrap(); + let effective_filter = agent.effective_tool_filter(); + assert!(effective_filter.matches("tool_search")); + assert!(effective_filter.matches("srv__fetch_issue")); + assert!(effective_filter.matches("read")); + assert!(!effective_filter.matches("write")); + } + #[test] fn estimate_message_tokens_empty_is_zero() { assert_eq!(estimate_message_tokens(&[], ""), 0); @@ -1934,7 +2018,7 @@ mod tests { }, std::path::PathBuf::from("/tmp"), )), - session_id: None, + identity: None, timeouts: n00n_providers::Timeouts::default(), openai_options: OpenAiOptions::default(), file_tracker: FileReadTracker::fresh(), @@ -2700,7 +2784,7 @@ mod tests { }, std::path::PathBuf::from("/tmp"), )), - session_id: None, + identity: None, timeouts: n00n_providers::Timeouts::default(), openai_options: OpenAiOptions::default(), file_tracker: FileReadTracker::fresh(), diff --git a/n00n-agent/src/headless.rs b/n00n-agent/src/headless.rs index b8a5edbd7..862a8e591 100644 --- a/n00n-agent/src/headless.rs +++ b/n00n-agent/src/headless.rs @@ -21,7 +21,9 @@ use crate::cancel::{CancelMap, CancelToken}; use crate::permissions::PermissionManager; use crate::prompt::ResolvedSlots; use crate::template; -use crate::tools::{DescriptionContext, FileReadTracker, ToolAudience, ToolFilter, ToolRegistry}; +use crate::tools::{ + DescriptionContext, FileReadTracker, SessionIdentity, ToolAudience, ToolFilter, ToolRegistry, +}; use crate::{ Agent, AgentConfig, AgentEvent, AgentInput, AgentMode, AgentParams, AgentRunParams, Envelope, EventSender, ImageSource, McpHandle, McpSession, PermissionsConfig, ToolOutput, @@ -263,7 +265,7 @@ pub fn spawn(params: HeadlessParams) -> HeadlessHandle { params.permissions_config, working_dir_path, )), - session_id: Some(session_ref_clone.clone()), + identity: Some(SessionIdentity::root(session_ref_clone.clone())), timeouts: params.timeouts, openai_options: params.openai_options, file_tracker: FileReadTracker::fresh(), @@ -518,7 +520,7 @@ pub fn spawn_interactive(params: InteractiveParams) -> InteractiveHandle { config: Arc::clone(¶ms.config), tool_output_lines: ToolOutputLines::default(), permissions: Arc::clone(&permissions), - session_id: Some(session_ref_clone.clone()), + identity: Some(SessionIdentity::root(session_ref_clone.clone())), timeouts: params.timeouts, openai_options: params.openai_options, file_tracker: Arc::clone(&file_tracker), diff --git a/n00n-agent/src/tools/mod.rs b/n00n-agent/src/tools/mod.rs index 1f17a3909..af86c9fb6 100644 --- a/n00n-agent/src/tools/mod.rs +++ b/n00n-agent/src/tools/mod.rs @@ -300,6 +300,40 @@ pub fn timeout_annotation(secs: u64) -> String { pub type LocalToolFn = Arc Result + Send + Sync>; pub type LocalTools = Arc>; +#[derive(Clone, Debug, PartialEq, Eq)] +pub struct SessionIdentity { + session_id: SessionRef, + root_session_id: SessionRef, +} + +impl SessionIdentity { + #[must_use] + pub fn root(session_id: SessionRef) -> Self { + Self { + root_session_id: session_id.clone(), + session_id, + } + } + + #[must_use] + pub fn child(session_id: SessionRef, root_session_id: SessionRef) -> Self { + Self { + session_id, + root_session_id, + } + } + + #[must_use] + pub fn session_id(&self) -> &SessionRef { + &self.session_id + } + + #[must_use] + pub fn root_session_id(&self) -> &SessionRef { + &self.root_session_id + } +} + #[derive(Clone)] pub struct ToolContext { pub provider: Arc, @@ -321,6 +355,8 @@ pub struct ToolContext { pub prompt_slots: Arc, pub opts: RequestOptions, pub subagent_cancels: Arc>, + /// Immutable session and root identity of the agent executing this tool. + pub identity: Option, pub registry: Arc, pub tool_filter: ToolFilter, pub workflow: bool, @@ -588,6 +624,7 @@ pub fn interpreter_ctx( prompt_slots: Arc::new(crate::prompt::ResolvedSlots::default()), opts: RequestOptions::default(), subagent_cancels: Arc::new(CancelMap::new()), + identity: None, registry, tool_filter: ToolFilter::All, workflow: false, @@ -629,7 +666,8 @@ pub mod test_support { use super::{ AgentMode, Arc, CancelToken, DescriptionContext, FileReadTracker, LazyLock, - PermissionManager, ToolContext, ToolRegistry, Value, interpreter_ctx, registry, + PermissionManager, SessionIdentity, ToolContext, ToolRegistry, Value, interpreter_ctx, + registry, }; pub const GUARDED_TOOL_NAME: &str = "guarded_mock"; @@ -705,6 +743,9 @@ pub mod test_support { Arc::new(ToolRegistry::new()), ); ctx.tool_use_id = tool_use_id.map(String::from); + ctx.identity = Some(SessionIdentity::root( + n00n_storage::id::SessionRef::generate(), + )); ctx } @@ -745,6 +786,33 @@ mod tests { const LINE_LIMIT: usize = 500; + #[test] + fn root_session_identity_uses_one_id() { + let session_id = n00n_storage::id::SessionRef::generate(); + let identity = SessionIdentity::root(session_id.clone()); + + assert_eq!(identity.session_id(), &session_id); + assert_eq!(identity.root_session_id(), &session_id); + } + + #[test] + fn descendant_session_identities_inherit_root_and_remain_distinct() { + let root = SessionIdentity::root(n00n_storage::id::SessionRef::generate()); + let child = SessionIdentity::child( + n00n_storage::id::SessionRef::generate(), + root.root_session_id().clone(), + ); + let grandchild = SessionIdentity::child( + n00n_storage::id::SessionRef::generate(), + child.root_session_id().clone(), + ); + + assert_ne!(child.session_id(), root.session_id()); + assert_ne!(grandchild.session_id(), child.session_id()); + assert_eq!(child.root_session_id(), root.root_session_id()); + assert_eq!(grandchild.root_session_id(), root.root_session_id()); + } + #[test_case(true ; "vision_model_keeps_view_image")] #[test_case(false ; "text_only_model_loses_view_image")] fn from_config_gates_view_image_on_vision(vision: bool) { diff --git a/n00n-lua/src/api/agent.rs b/n00n-lua/src/api/agent.rs index e64fb0fe7..78b436cae 100644 --- a/n00n-lua/src/api/agent.rs +++ b/n00n-lua/src/api/agent.rs @@ -17,8 +17,8 @@ use n00n_agent::cancel::CancelMap; use n00n_agent::tools::interpreter_bridge; use n00n_agent::tools::registry::ToolRegistry; use n00n_agent::tools::{ - Deadline, DescriptionContext, FileReadTracker, LocalToolFn, LocalTools, ToolAudience, - ToolContext, ToolFilter, ToolLive, + Deadline, DescriptionContext, FileReadTracker, LocalToolFn, LocalTools, SessionIdentity, + ToolAudience, ToolContext, ToolFilter, ToolLive, }; use n00n_agent::{ Agent, AgentEvent, AgentInput, AgentMode, AgentParams, AgentRunParams, Envelope, EventSender, @@ -36,6 +36,7 @@ use tracing::info; use crate::api::ui::buf::BufHandle; use crate::api::util::convert::{JSON_ARRAY_META_FIELD, json_to_lua, lua_to_json, lua_tool_result}; use crate::api::util::ctx::{AgentContext, LuaCtx}; +use crate::state::PluginStateStore; const SESSION_CLOSED_ERR: &str = "session closed"; const DEFAULT_SESSION_AUDIENCE: ToolAudience = ToolAudience::GENERAL_SUB; @@ -93,8 +94,7 @@ fn model_to_lua_table(lua: &Lua, model: &Model) -> LuaResult { } fn dispatch_ctx<'a>(ctx: &'a LuaCtx, method: &str) -> Result<&'a AgentContext, String> { - ctx.agent() - .ok_or_else(|| ctx.cap_err(&format!("n00n.agent.{method}"))) + ctx.dispatch_agent(method) } fn parse_session_mode( @@ -642,6 +642,10 @@ async fn session( opts: Table, ) -> LuaResult> { let agent_ctx = try_pair!(dispatch_ctx(&ctx, "session")).clone(); + let Some(parent_identity) = agent_ctx.identity.clone() else { + return Ok(err_pair("session identity is unavailable")); + }; + let plugin_state_store = try_pair!(ctx.plugin_state_store()); drop(ctx); let model_spec: Option = opts.get("model_spec")?; let system: Option = opts.get("system")?; @@ -716,7 +720,7 @@ async fn session( audience, workflow: false, }; - let tools = n00n_agent::tools::ToolRegistry::global().definitions_active( + let tools = agent_ctx.registry.definitions_active( &vars, &ctx, model.supports_tool_examples(), @@ -848,13 +852,16 @@ async fn session( config: session_config(&agent_ctx.config, excluded_tools.clone()), tool_output_lines: n00n_config::ToolOutputLines::default(), permissions: Arc::clone(&agent_ctx.permissions), - session_id: Some(session_id.into()), + identity: Some(SessionIdentity::child( + session_id.into(), + parent_identity.root_session_id().clone(), + )), timeouts: agent_ctx.timeouts, openai_options: agent_ctx.openai_options, file_tracker: FileReadTracker::fresh(), prompt_slots: Arc::clone(&agent_ctx.prompt_slots), subagent_cancels: Arc::new(CancelMap::new()), - registry: Arc::clone(n00n_agent::tools::ToolRegistry::global_arc()), + registry: Arc::clone(&agent_ctx.registry), audience, }, system: system.unwrap_or_else(String::new), @@ -880,6 +887,8 @@ async fn session( prompt_rx, prompt_tx: Some(prompt_tx), parent_cancels: Arc::clone(&agent_ctx.subagent_cancels), + plugin_state_store, + child_state_owner: session_id, child_id, parent_tool_use_id, parent_event_tx: parent_tx, @@ -1381,6 +1390,8 @@ struct SessionState { prompt_rx: flume::Receiver, prompt_tx: Option>, parent_cancels: Arc>, + plugin_state_store: Arc, + child_state_owner: n00nId, child_id: String, parent_tool_use_id: String, parent_event_tx: EventSender, @@ -1404,6 +1415,7 @@ impl SessionState { self.closed = true; self.progress.set_current_done(); self.parent_cancels.remove(&self.child_id); + self.plugin_state_store.drop_owner(self.child_state_owner); let messages = std::mem::replace(&mut self.history, History::new(Vec::new())).into_vec(); let _ = self.parent_event_tx.send(AgentEvent::SubagentHistory { tool_use_id: self.child_id.clone(), diff --git a/n00n-lua/src/api/util/ctx.rs b/n00n-lua/src/api/util/ctx.rs index 7f79feb29..6ee21a63b 100644 --- a/n00n-lua/src/api/util/ctx.rs +++ b/n00n-lua/src/api/util/ctx.rs @@ -1,6 +1,9 @@ use std::ops::Deref; use std::path::{Path, PathBuf}; -use std::sync::Arc; +use std::sync::{ + Arc, + atomic::{AtomicBool, Ordering}, +}; use std::time::{Duration, Instant}; use mlua::{LuaSerdeExt, MultiValue, UserData, UserDataMethods, Value as LuaValue}; @@ -14,9 +17,19 @@ use n00n_config::{AgentConfig, ToolOutputLines}; use crate::api::tool::ToolCallReply; use crate::api::ui::buf::BufHandle; use crate::api::util::convert::json_to_lua; +use crate::api::util::state_convert::{ + json_to_lua as state_json_to_lua, lua_to_json as state_lua_to_json, +}; use crate::runtime::{active_task, lock_cell}; +use crate::state::{PluginStateIdentity, PluginStateScope, PluginStateStore}; +const CONTEXT_INACTIVE_MSG: &str = "state context is no longer active"; const DEADLINE_ALREADY_SET_MSG: &str = "ctx:set_deadline() already called"; +const INVALID_STATE_SCOPE_MSG: &str = "state scope must be 'session' or 'root'"; + +fn parse_state_scope(scope: &str) -> Result { + PluginStateScope::parse(scope).ok_or(INVALID_STATE_SCOPE_MSG) +} fn send_live_buf(lua: &mlua::Lua, buf: &mlua::AnyUserData) -> mlua::Result<()> { let shared = buf.borrow::().map(|h| Arc::clone(&h.buf))?; @@ -82,6 +95,14 @@ pub(crate) struct LuaCtx { pub(crate) cancel: CancelToken, tool_output_lines: ToolOutputLines, pub(crate) finish_tx: Option>, + active: Arc, + plugin_state: Option, +} + +struct PluginStateAccess { + plugin: Arc, + identity: PluginStateIdentity, + store: Arc, } enum Caps { @@ -110,6 +131,8 @@ impl LuaCtx { cancel: ctx.cancel.clone(), tool_output_lines: ctx.tool_output_lines, finish_tx: None, + active: Arc::new(AtomicBool::new(true)), + plugin_state: None, } } @@ -143,6 +166,8 @@ impl LuaCtx { cancel: CancelToken::none(), tool_output_lines, finish_tx: None, + active: Arc::new(AtomicBool::new(true)), + plugin_state: None, } } @@ -207,6 +232,62 @@ impl LuaCtx { } } + pub(crate) fn attach_plugin_state(&mut self, plugin: Arc, store: Arc) { + self.plugin_state = self + .agent() + .and_then(|agent| agent.identity.as_ref()) + .map(PluginStateIdentity::from) + .map(|identity| PluginStateAccess { + plugin, + identity, + store, + }); + } + + pub(crate) fn context_liveness(&self) -> Arc { + Arc::clone(&self.active) + } + + fn ensure_active(&self) -> Result<(), String> { + if self.active.load(Ordering::Acquire) { + Ok(()) + } else { + Err(CONTEXT_INACTIVE_MSG.to_owned()) + } + } + + fn inactive_pair(&self) -> Option<(LuaValue, Option)> { + self.ensure_active() + .err() + .map(|error| (LuaValue::Nil, Some(error))) + } + + pub(crate) fn plugin_state_store(&self) -> Result, String> { + self.plugin_state("n00n.agent.session") + .map(|access| Arc::clone(&access.store)) + } + + pub(crate) fn dispatch_agent(&self, method: &str) -> Result<&AgentContext, String> { + let agent = self + .agent() + .ok_or_else(|| self.cap_err(&format!("n00n.agent.{method}")))?; + self.ensure_active()?; + Ok(agent) + } + + fn plugin_state(&self, method: &str) -> Result<&PluginStateAccess, String> { + if !matches!(self.caps, Caps::Handler { .. }) { + return Err(self.cap_err(method)); + } + + let access = self + .plugin_state + .as_ref() + .ok_or_else(|| "session identity is unavailable".to_owned())?; + self.ensure_active()?; + Ok(access) + } + pub(crate) fn cap_err(&self, method: &str) -> String { format!("{method} not available in {} ctx", self.kind()) } @@ -219,9 +300,15 @@ impl LuaCtx { #[allow(clippy::too_many_lines)] impl UserData for LuaCtx { fn add_methods>(methods: &mut M) { - methods.add_method("cancelled", |_, this, ()| Ok(this.cancel.is_cancelled())); + methods.add_method("cancelled", |_, this, ()| { + this.ensure_active().map_err(mlua::Error::runtime)?; + Ok(this.cancel.is_cancelled()) + }); methods.add_method("workflow", |_, this, ()| { + if let Some(error) = this.inactive_pair() { + return Ok(error); + } let Some(workflow) = this.workflow() else { return Ok(this.cap_err_pair("workflow")); }; @@ -230,6 +317,9 @@ impl UserData for LuaCtx { methods.add_method("audience", |lua, this, ()| { const DEFAULT_AUDIENCE: &str = "main"; + if let Some(error) = this.inactive_pair() { + return Ok(error); + } let Some(audience) = this.audience() else { return Ok(this.cap_err_pair("audience")); }; @@ -238,6 +328,9 @@ impl UserData for LuaCtx { }); methods.add_method("live_buf", |lua, this, buf: mlua::AnyUserData| { + if let Some(error) = this.inactive_pair() { + return Ok(error); + } if matches!(this.caps, Caps::Restore { .. }) { return Ok(this.cap_err_pair("live_buf")); } @@ -246,6 +339,9 @@ impl UserData for LuaCtx { }); methods.add_method("config", |lua, this, args: MultiValue| { + if let Some(error) = this.inactive_pair() { + return Ok(error); + } let Some(config) = this.config() else { return Ok(this.cap_err_pair("config")); }; @@ -270,15 +366,96 @@ impl UserData for LuaCtx { }); methods.add_method("tool_output_lines", |lua, this, ()| { + this.ensure_active().map_err(mlua::Error::runtime)?; lua.to_value(&this.tool_output_lines) }); - methods.add_method("state", |lua, this, ()| match this.state() { - Some(v) => json_to_lua(lua, v), - None => Ok(LuaValue::Nil), + methods.add_method("state", |lua, this, ()| { + this.ensure_active().map_err(mlua::Error::runtime)?; + match this.state() { + Some(v) => json_to_lua(lua, v), + None => Ok(LuaValue::Nil), + } + }); + methods.add_method("state_get", |lua, this, scope: String| { + let scope = match parse_state_scope(&scope) { + Ok(scope) => scope, + Err(error) => return Ok((LuaValue::Nil, Some(error.to_owned()))), + }; + let access = match this.plugin_state("state_get") { + Ok(access) => access, + Err(error) => return Ok((LuaValue::Nil, Some(error))), + }; + let Some(value) = access.store.get(&access.plugin, scope, &access.identity) else { + return Ok((LuaValue::Nil, None)); + }; + match state_json_to_lua(lua, &value) { + Ok(value) => Ok((value, None)), + Err(error) => Ok((LuaValue::Nil, Some(error.to_string()))), + } + }); + + methods.add_method( + "state_replace", + |lua, this, (scope, value): (String, LuaValue)| { + let scope = match parse_state_scope(&scope) { + Ok(scope) => scope, + Err(error) => return Ok((LuaValue::Nil, Some(error.to_owned()))), + }; + let access = match this.plugin_state("state_replace") { + Ok(access) => access, + Err(error) => return Ok((LuaValue::Nil, Some(error))), + }; + let value = match state_lua_to_json(lua, &value) { + Ok(value) => value, + Err(error) => return Ok((LuaValue::Nil, Some(error.to_string()))), + }; + let previous = + match access + .store + .replace(&access.plugin, scope, &access.identity, value) + { + Ok(previous) => previous, + Err(error) => return Ok((LuaValue::Nil, Some(error.to_string()))), + }; + let previous = match previous { + Some(previous) => match state_json_to_lua(lua, &previous) { + Ok(previous) => previous, + Err(error) => return Ok((LuaValue::Nil, Some(error.to_string()))), + }, + None => LuaValue::Nil, + }; + Ok((previous, None)) + }, + ); + + methods.add_method("state_remove", |lua, this, scope: String| { + let scope = match parse_state_scope(&scope) { + Ok(scope) => scope, + Err(error) => return Ok((LuaValue::Nil, Some(error.to_owned()))), + }; + let access = match this.plugin_state("state_remove") { + Ok(access) => access, + Err(error) => return Ok((LuaValue::Nil, Some(error))), + }; + let previous = match access.store.remove(&access.plugin, scope, &access.identity) { + Ok(previous) => previous, + Err(error) => return Ok((LuaValue::Nil, Some(error.to_string()))), + }; + let previous = match previous { + Some(previous) => match state_json_to_lua(lua, &previous) { + Ok(previous) => previous, + Err(error) => return Ok((LuaValue::Nil, Some(error.to_string()))), + }, + None => LuaValue::Nil, + }; + Ok((previous, None)) }); methods.add_method("set_deadline", |lua, this, secs: u64| { + if let Some(error) = this.inactive_pair() { + return Ok(error); + } if !matches!(this.caps, Caps::Handler { .. }) { return Ok(this.cap_err_pair("set_deadline")); } @@ -298,6 +475,9 @@ impl UserData for LuaCtx { }); methods.add_method("record_read", |_, this, path: String| { + if let Some(error) = this.inactive_pair() { + return Ok(error); + } let Some(tracker) = this.file_tracker() else { return Ok(this.cap_err_pair("record_read")); }; @@ -306,6 +486,9 @@ impl UserData for LuaCtx { }); methods.add_method("check_before_edit", |_, this, path: String| { + if let Some(error) = this.inactive_pair() { + return Ok(error); + } let Some(tracker) = this.file_tracker() else { return Ok(this.cap_err_pair("check_before_edit")); }; @@ -318,6 +501,9 @@ impl UserData for LuaCtx { methods.add_async_method( "find_instructions", |lua, this, dir_path: String| async move { + if let Some(error) = this.inactive_pair() { + return Ok(error); + } let Some(loaded) = this.loaded_instructions().cloned() else { return Ok(this.cap_err_pair("find_instructions")); }; @@ -347,6 +533,9 @@ impl UserData for LuaCtx { }); methods.add_method_mut("finish", |lua, this, val: LuaValue| { + if let Some(error) = this.inactive_pair() { + return Ok(error); + } if !matches!(this.caps, Caps::Handler { .. }) { return Ok(this.cap_err_pair("finish")); } diff --git a/n00n-lua/src/api/util/mod.rs b/n00n-lua/src/api/util/mod.rs index 7e911445b..d527a7395 100644 --- a/n00n-lua/src/api/util/mod.rs +++ b/n00n-lua/src/api/util/mod.rs @@ -3,3 +3,4 @@ pub(crate) mod convert; pub(crate) mod ctx; pub(crate) mod dispatch; pub(crate) mod setup; +pub(crate) mod state_convert; diff --git a/n00n-lua/src/api/util/state_convert.rs b/n00n-lua/src/api/util/state_convert.rs new file mode 100644 index 000000000..1c9ab0c37 --- /dev/null +++ b/n00n-lua/src/api/util/state_convert.rs @@ -0,0 +1,542 @@ +use std::{collections::HashSet, ffi::c_void}; + +use mlua::{Lua, LuaSerdeExt, Table, Value}; +use n00n_storage::sessions::MAX_PLUGIN_STATE_BYTES; +use serde_json::Value as JsonValue; + +pub(crate) const MAX_STATE_DEPTH: usize = 64; + +#[derive(Debug, thiserror::Error, PartialEq, Eq)] +pub(crate) enum StateConvertError { + #[error("state contains unsupported Lua value type '{0}'")] + UnsupportedValue(&'static str), + #[error("state contains a non-finite number")] + NonFiniteNumber, + #[error("state contains a string that is not valid UTF-8")] + InvalidUtf8String, + #[error("state object keys must be UTF-8 strings")] + NonStringObjectKey, + #[error("state array keys must be positive integers")] + InvalidArrayKey, + #[error("state arrays must have contiguous keys starting at 1")] + SparseArray, + #[error("state contains a cycle")] + Cycle, + #[error("state exceeds the maximum size of {maximum} bytes")] + MaximumBytesExceeded { maximum: usize }, + #[error("state exceeds the maximum nesting depth of {maximum}")] + MaximumDepthExceeded { maximum: usize }, + #[error("JSON number cannot be represented as a Lua number")] + UnrepresentableNumber, + #[error("Lua operation failed during state conversion: {0}")] + Lua(String), +} + +struct StateBudget { + used: usize, +} + +impl StateBudget { + fn consume(&mut self, bytes: usize) -> Result<(), StateConvertError> { + self.used = self.used.saturating_add(bytes); + if self.used > MAX_PLUGIN_STATE_BYTES { + Err(StateConvertError::MaximumBytesExceeded { + maximum: MAX_PLUGIN_STATE_BYTES, + }) + } else { + Ok(()) + } + } +} + +impl From for StateConvertError { + fn from(error: mlua::Error) -> Self { + Self::Lua(error.to_string()) + } +} + +pub(crate) fn json_to_lua(lua: &Lua, value: &JsonValue) -> Result { + json_to_lua_at_depth(lua, value, 0) +} + +fn json_to_lua_at_depth( + lua: &Lua, + value: &JsonValue, + depth: usize, +) -> Result { + check_depth(depth)?; + match value { + JsonValue::Null => Ok(lua.null()), + JsonValue::Bool(value) => Ok(Value::Boolean(*value)), + JsonValue::Number(value) => { + if let Some(integer) = value.as_i64() { + return Ok(Value::Integer(integer)); + } + if value.as_u64().is_some() { + return Err(StateConvertError::UnrepresentableNumber); + } + let number = value + .as_f64() + .filter(|number| number.is_finite()) + .ok_or(StateConvertError::UnrepresentableNumber)?; + let round_trip = serde_json::Number::from_f64(number) + .ok_or(StateConvertError::UnrepresentableNumber)?; + if &round_trip != value { + return Err(StateConvertError::UnrepresentableNumber); + } + Ok(Value::Number(number)) + } + JsonValue::String(value) => Ok(Value::String(lua.create_string(value)?)), + JsonValue::Array(values) => { + let table = lua.create_table_with_capacity(values.len(), 0)?; + table.set_metatable(Some(lua.array_metatable()))?; + for (offset, value) in values.iter().enumerate() { + table.raw_set(offset + 1, json_to_lua_at_depth(lua, value, depth + 1)?)?; + } + Ok(Value::Table(table)) + } + JsonValue::Object(values) => { + let table = lua.create_table_with_capacity(0, values.len())?; + for (key, value) in values { + table.raw_set(key.as_str(), json_to_lua_at_depth(lua, value, depth + 1)?)?; + } + Ok(Value::Table(table)) + } + } +} + +pub(crate) fn lua_to_json(lua: &Lua, value: &Value) -> Result { + lua_to_json_at_depth( + lua, + value, + 0, + &mut HashSet::new(), + &mut StateBudget { used: 0 }, + ) +} + +fn lua_to_json_at_depth( + lua: &Lua, + value: &Value, + depth: usize, + active_tables: &mut HashSet<*const c_void>, + budget: &mut StateBudget, +) -> Result { + check_depth(depth)?; + match value { + value if value.is_null() => { + budget.consume(4)?; + Ok(JsonValue::Null) + } + Value::Boolean(value) => { + budget.consume(if *value { 4 } else { 5 })?; + Ok(JsonValue::Bool(*value)) + } + Value::Integer(value) => { + budget.consume(value.to_string().len())?; + Ok(JsonValue::Number((*value).into())) + } + Value::Number(value) => { + let number = + serde_json::Number::from_f64(*value).ok_or(StateConvertError::NonFiniteNumber)?; + budget.consume(number.to_string().len())?; + Ok(JsonValue::Number(number)) + } + Value::String(value) => { + let value = value + .to_str() + .map_err(|_| StateConvertError::InvalidUtf8String)?; + budget.consume(serialized_string_len(&value))?; + Ok(JsonValue::String(value.to_owned())) + } + Value::Table(table) => table_to_json(lua, table, depth, active_tables, budget), + value => Err(StateConvertError::UnsupportedValue(value.type_name())), + } +} + +fn table_to_json( + lua: &Lua, + table: &Table, + depth: usize, + active_tables: &mut HashSet<*const c_void>, + budget: &mut StateBudget, +) -> Result { + budget.consume(2)?; + let pointer = table.to_pointer(); + if !active_tables.insert(pointer) { + return Err(StateConvertError::Cycle); + } + + let result = if is_array_table(lua, table)? { + array_to_json(lua, table, depth, active_tables, budget) + } else { + object_to_json(lua, table, depth, active_tables, budget) + }; + active_tables.remove(&pointer); + result +} + +fn is_array_table(lua: &Lua, table: &Table) -> Result { + if table + .metatable() + .is_some_and(|metatable| metatable == lua.array_metatable()) + { + return Ok(true); + } + + let length = table.raw_len(); + if length == 0 { + return Ok(false); + } + let mut entries = 0; + for pair in table.clone().pairs::() { + let (key, _) = pair?; + let Value::Integer(key) = key else { + return Ok(false); + }; + let Ok(index) = usize::try_from(key) else { + return Ok(false); + }; + if !(1..=length).contains(&index) { + return Ok(false); + } + entries += 1; + } + Ok(entries == length) +} + +fn array_to_json( + lua: &Lua, + table: &Table, + depth: usize, + active_tables: &mut HashSet<*const c_void>, + budget: &mut StateBudget, +) -> Result { + let mut entries = Vec::new(); + for pair in table.pairs::() { + let (key, value) = pair?; + let Value::Integer(key) = key else { + return Err(StateConvertError::InvalidArrayKey); + }; + let index = usize::try_from(key).map_err(|_| StateConvertError::InvalidArrayKey)?; + if index == 0 { + return Err(StateConvertError::InvalidArrayKey); + } + if !entries.is_empty() { + budget.consume(1)?; + } + entries.push((index, value)); + } + entries.sort_unstable_by_key(|(index, _)| *index); + + let mut values = Vec::with_capacity(entries.len()); + for (offset, (index, value)) in entries.into_iter().enumerate() { + if index != offset + 1 { + return Err(StateConvertError::SparseArray); + } + values.push(lua_to_json_at_depth( + lua, + &value, + depth + 1, + active_tables, + budget, + )?); + } + Ok(JsonValue::Array(values)) +} + +fn object_to_json( + lua: &Lua, + table: &Table, + depth: usize, + active_tables: &mut HashSet<*const c_void>, + budget: &mut StateBudget, +) -> Result { + let mut values = serde_json::Map::new(); + for pair in table.pairs::() { + let (key, value) = pair?; + let Value::String(key) = key else { + return Err(StateConvertError::NonStringObjectKey); + }; + let key = key + .to_str() + .map_err(|_| StateConvertError::InvalidUtf8String)?; + if !values.is_empty() { + budget.consume(1)?; + } + budget.consume(serialized_string_len(&key).saturating_add(1))?; + values.insert( + key.to_owned(), + lua_to_json_at_depth(lua, &value, depth + 1, active_tables, budget)?, + ); + } + Ok(JsonValue::Object(values)) +} +fn serialized_string_len(value: &str) -> usize { + value.chars().fold(2, |bytes, character| { + bytes.saturating_add(match character { + '"' | '\\' | '\u{0008}' | '\u{0009}' | '\u{000a}' | '\u{000c}' | '\u{000d}' => 2, + '\u{0000}'..='\u{001f}' => 6, + character => character.len_utf8(), + }) + }) +} + +const fn check_depth(depth: usize) -> Result<(), StateConvertError> { + if depth > MAX_STATE_DEPTH { + Err(StateConvertError::MaximumDepthExceeded { + maximum: MAX_STATE_DEPTH, + }) + } else { + Ok(()) + } +} + +#[cfg(test)] +mod tests { + use mlua::{Lua, LuaSerdeExt, Value}; + use n00n_storage::sessions::MAX_PLUGIN_STATE_BYTES; + use serde_json::json; + use test_case::test_case; + + use super::{ + MAX_STATE_DEPTH, StateConvertError, json_to_lua, lua_to_json, serialized_string_len, + }; + + #[test] + fn round_trip_preserves_nested_null_and_array_shape() { + let lua = Lua::new(); + let input = json!({"items": [null, {"enabled": true}], "name": "state"}); + + let lua_value = json_to_lua(&lua, &input).unwrap(); + let items: Value = lua_value.as_table().unwrap().raw_get("items").unwrap(); + let first: Value = items.as_table().unwrap().raw_get(1).unwrap(); + + assert!(first.is_null()); + assert_eq!(lua_to_json(&lua, &lua_value).unwrap(), input); + } + + #[test] + fn ordinary_lua_sequence_converts_to_json_array() { + let lua = Lua::new(); + let value = lua.load("return { items = { 'a', 'b' } }").eval().unwrap(); + + assert_eq!( + lua_to_json(&lua, &value).unwrap(), + json!({"items": ["a", "b"]}) + ); + } + + #[test] + fn round_trip_preserves_empty_array_shape() { + let lua = Lua::new(); + let input = json!({"items": []}); + + let value = json_to_lua(&lua, &input).unwrap(); + + assert_eq!(lua_to_json(&lua, &value).unwrap(), input); + } + + #[test_case(f64::NAN ; "nan")] + #[test_case(f64::INFINITY ; "positive_infinity")] + #[test_case(f64::NEG_INFINITY ; "negative_infinity")] + fn rejects_non_finite_numbers(number: f64) { + let lua = Lua::new(); + + assert_eq!( + lua_to_json(&lua, &Value::Number(number)).unwrap_err(), + StateConvertError::NonFiniteNumber + ); + } + + #[test_case("18446744073709551617" ; "above_u64")] + #[test_case("-9223372036854775809" ; "below_i64")] + #[test_case("0.12345678901234567890123456789" ; "precision_loss")] + fn rejects_json_numbers_that_cannot_round_trip_exactly(raw: &str) { + let lua = Lua::new(); + let value: serde_json::Value = serde_json::from_str(raw).unwrap(); + + assert_eq!( + json_to_lua(&lua, &value).unwrap_err(), + StateConvertError::UnrepresentableNumber + ); + } + + #[test] + fn rejects_non_utf8_strings() { + let lua = Lua::new(); + let value = Value::String(lua.create_string([0xff]).unwrap()); + + assert_eq!( + lua_to_json(&lua, &value).unwrap_err(), + StateConvertError::InvalidUtf8String + ); + } + + #[test] + fn rejects_unsupported_values() { + let lua = Lua::new(); + let value = Value::Function(lua.create_function(|_, ()| Ok(())).unwrap()); + + assert_eq!( + lua_to_json(&lua, &value).unwrap_err(), + StateConvertError::UnsupportedValue("function") + ); + } + + #[test] + fn rejects_nil_but_accepts_mlua_null() { + let lua = Lua::new(); + + assert_eq!( + lua_to_json(&lua, &Value::Nil).unwrap_err(), + StateConvertError::UnsupportedValue("nil") + ); + assert_eq!(lua_to_json(&lua, &lua.null()).unwrap(), json!(null)); + } + + #[test] + fn unmarked_tables_detect_objects_and_sequences() { + let lua = Lua::new(); + let object = lua.create_table().unwrap(); + object.raw_set("name", "plugin").unwrap(); + assert_eq!( + lua_to_json(&lua, &Value::Table(object)).unwrap(), + json!({"name": "plugin"}) + ); + + let sequence = lua.create_table().unwrap(); + sequence.raw_set(1, "item").unwrap(); + assert_eq!( + lua_to_json(&lua, &Value::Table(sequence)).unwrap(), + json!(["item"]) + ); + } + + #[test] + fn marked_arrays_require_only_contiguous_integer_keys() { + let lua = Lua::new(); + let array = lua.create_table().unwrap(); + array.set_metatable(Some(lua.array_metatable())).unwrap(); + array.raw_set(1, "a").unwrap(); + array.raw_set(2, "b").unwrap(); + assert_eq!( + lua_to_json(&lua, &Value::Table(array)).unwrap(), + json!(["a", "b"]) + ); + + let mixed = lua.create_table().unwrap(); + mixed.set_metatable(Some(lua.array_metatable())).unwrap(); + mixed.raw_set(1, "a").unwrap(); + mixed.raw_set("name", "state").unwrap(); + assert_eq!( + lua_to_json(&lua, &Value::Table(mixed)).unwrap_err(), + StateConvertError::InvalidArrayKey + ); + + let sparse = lua.create_table().unwrap(); + sparse.set_metatable(Some(lua.array_metatable())).unwrap(); + sparse.raw_set(1, "a").unwrap(); + sparse.raw_set(3, "c").unwrap(); + assert_eq!( + lua_to_json(&lua, &Value::Table(sparse)).unwrap_err(), + StateConvertError::SparseArray + ); + } + + #[test] + fn rejects_direct_and_indirect_cycles() { + let lua = Lua::new(); + let direct = lua.create_table().unwrap(); + direct.raw_set("self", direct.clone()).unwrap(); + assert_eq!( + lua_to_json(&lua, &Value::Table(direct)).unwrap_err(), + StateConvertError::Cycle + ); + + let first = lua.create_table().unwrap(); + let second = lua.create_table().unwrap(); + first.raw_set("second", second.clone()).unwrap(); + second.raw_set("first", first.clone()).unwrap(); + assert_eq!( + lua_to_json(&lua, &Value::Table(first)).unwrap_err(), + StateConvertError::Cycle + ); + } + + #[test] + fn allows_repeated_acyclic_table_references() { + let lua = Lua::new(); + let shared = lua.create_table().unwrap(); + shared.raw_set("value", 7).unwrap(); + let root = lua.create_table().unwrap(); + root.raw_set("left", shared.clone()).unwrap(); + root.raw_set("right", shared).unwrap(); + + assert_eq!( + lua_to_json(&lua, &Value::Table(root)).unwrap(), + json!({"left": {"value": 7}, "right": {"value": 7}}) + ); + } + + #[test] + fn enforces_serialized_size_during_lua_conversion() { + let lua = Lua::new(); + let exact = Value::String( + lua.create_string("x".repeat(MAX_PLUGIN_STATE_BYTES - 2)) + .unwrap(), + ); + assert!(lua_to_json(&lua, &exact).is_ok()); + + let oversized = Value::String( + lua.create_string("x".repeat(MAX_PLUGIN_STATE_BYTES - 1)) + .unwrap(), + ); + assert_eq!( + lua_to_json(&lua, &oversized).unwrap_err(), + StateConvertError::MaximumBytesExceeded { + maximum: MAX_PLUGIN_STATE_BYTES + } + ); + } + + #[test_case("plain" ; "plain")] + #[test_case("quote\"slash\\" ; "escapes")] + #[test_case("line\ncontrol\u{0001}" ; "controls")] + #[test_case("héllo" ; "utf8")] + fn serialized_string_size_matches_serde_json(value: &str) { + assert_eq!( + serialized_string_len(value), + serde_json::to_vec(value).unwrap().len() + ); + } + + #[test] + fn enforces_maximum_depth_in_both_directions() { + let lua = Lua::new(); + let mut json_value = json!(true); + for _ in 0..=MAX_STATE_DEPTH { + json_value = json!({"child": json_value}); + } + assert_eq!( + json_to_lua(&lua, &json_value).unwrap_err(), + StateConvertError::MaximumDepthExceeded { + maximum: MAX_STATE_DEPTH + } + ); + + let root = lua.create_table().unwrap(); + let mut current = root.clone(); + for _ in 0..=MAX_STATE_DEPTH { + let child = lua.create_table().unwrap(); + current.raw_set("child", child.clone()).unwrap(); + current = child; + } + assert_eq!( + lua_to_json(&lua, &Value::Table(root)).unwrap_err(), + StateConvertError::MaximumDepthExceeded { + maximum: MAX_STATE_DEPTH + } + ); + } +} diff --git a/n00n-lua/src/error.rs b/n00n-lua/src/error.rs index f8a6f4abd..d04d94852 100644 --- a/n00n-lua/src/error.rs +++ b/n00n-lua/src/error.rs @@ -25,4 +25,6 @@ pub enum PluginError { UnknownPlugin { plugin: String }, #[error("plugin host is not running")] HostDead, + #[error("plugin state operation failed: {message}")] + State { message: String }, } diff --git a/n00n-lua/src/lib.rs b/n00n-lua/src/lib.rs index 23ef18d8a..b9e424fd5 100644 --- a/n00n-lua/src/lib.rs +++ b/n00n-lua/src/lib.rs @@ -10,6 +10,7 @@ pub mod language; mod loader; pub(crate) mod plugin_permissions; mod runtime; +mod state; pub use api::keymap::{KeymapEntry, KeymapReader, KeymapSnapshot}; pub use api::options::{OptionSpec, OptionType, PluginOptionSpecs}; diff --git a/n00n-lua/src/loader.rs b/n00n-lua/src/loader.rs index ecabdc2ed..211f96ca7 100644 --- a/n00n-lua/src/loader.rs +++ b/n00n-lua/src/loader.rs @@ -6,8 +6,10 @@ use std::sync::{Arc, LazyLock}; use std::time::Duration; use include_dir::{Dir, include_dir}; -use n00n_agent::tools::ToolRegistry; +use n00n_agent::tools::{SessionIdentity, ToolRegistry}; use n00n_config::{PluginsConfig, RawConfig}; +use n00n_storage::id::n00nId; +use n00n_storage::sessions::StoredSessionStateSnapshot; use crate::api::keymap::KeymapReader; use crate::api::options::{PluginOptionSpecs, PluginOpts}; @@ -15,6 +17,7 @@ use crate::api::util::command::{HintReader, LuaCommandReader, UiAction}; use crate::error::PluginError; use crate::plugin_permissions::{PluginPermissions, load_plugin_permissions}; use crate::runtime::{self, ClickFallback, LuaThread, Request, RestoreItem}; +use crate::state::PluginStateIdentity; use n00n_agent::prompt::ResolvedSlots; const SHUTDOWN_TIMEOUT: Duration = Duration::from_secs(2); @@ -605,6 +608,76 @@ impl EventHandle { .await .unwrap_or_else(|_| ResolvedSlots::default()) } + /// Hydrates host-owned plugin state after all in-flight Lua work drains. + /// + /// # Errors + /// Returns an error when the host is unavailable or the snapshot cannot be applied. + pub fn hydrate_state( + &self, + identity: &SessionIdentity, + snapshot: Option, + ) -> Result<(), PluginError> { + let (reply, recv) = flume::bounded(1); + self.tx + .send(Request::HydrateState { + identity: PluginStateIdentity::from(identity), + snapshot, + reply, + }) + .map_err(|_| PluginError::HostDead)?; + recv.recv() + .map_err(|_| PluginError::HostDead)? + .map_err(|message| PluginError::State { message }) + } + + /// Captures host-owned plugin state after all in-flight Lua work drains. + /// + /// # Errors + /// Returns an error when the host is unavailable or capture validation fails. + pub fn capture_state( + &self, + identity: &SessionIdentity, + revision: u64, + ) -> Result { + let (reply, recv) = flume::bounded(1); + self.tx + .send(Request::CaptureState { + identity: PluginStateIdentity::from(identity), + revision, + reply, + }) + .map_err(|_| PluginError::HostDead)?; + recv.recv() + .map_err(|_| PluginError::HostDead)? + .map_err(|message| PluginError::State { message }) + } + + /// Clears both scopes for an identity and records removals for the next capture. + /// + /// # Errors + /// Returns an error when the host is unavailable. + pub fn reset_state(&self, identity: &SessionIdentity) -> Result<(), PluginError> { + let (reply, recv) = flume::bounded(1); + self.tx + .send(Request::ResetState { + identity: PluginStateIdentity::from(identity), + reply, + }) + .map_err(|_| PluginError::HostDead)?; + recv.recv().map_err(|_| PluginError::HostDead) + } + + /// Drops in-memory state for one canonical owner without persisting removals. + /// + /// # Errors + /// Returns an error when the host is unavailable. + pub fn drop_state_owner(&self, owner: n00nId) -> Result<(), PluginError> { + let (reply, recv) = flume::bounded(1); + self.tx + .send(Request::DropStateOwner { owner, reply }) + .map_err(|_| PluginError::HostDead)?; + recv.recv().map_err(|_| PluginError::HostDead) + } pub fn request_restore(&self, item: RestoreItem, event_tx: n00n_agent::EventSender) { let _ = self.tx.send(Request::RestoreToolAsync { item, event_tx }); diff --git a/n00n-lua/src/runtime.rs b/n00n-lua/src/runtime.rs index 7be7cafa7..24b453162 100644 --- a/n00n-lua/src/runtime.rs +++ b/n00n-lua/src/runtime.rs @@ -38,10 +38,15 @@ use crate::api::util::command::{LuaCommandReader, LuaCommandWriter, UiAction}; use crate::api::util::convert::json_to_lua; use crate::api::util::ctx::LuaCtx; use crate::api::util::setup::ConfigStore; +use crate::api::util::state_convert::json_to_lua as state_json_to_lua; use crate::docs_render; use crate::error::PluginError; use crate::plugin_permissions::{PluginPermissions, load_plugin_permissions}; +use n00n_storage::id::n00nId; +use n00n_storage::sessions::{StoredSessionStateSnapshot, StoredStateScope}; + +use crate::state::{PLUGIN_STATE_SCHEMA_VERSION, PluginStateIdentity, PluginStateStore}; fn register_builtin_tools(registry: &Arc) -> Result<(), PluginError> { let tools: [(Arc, ToolSource); 2] = [ ( @@ -175,6 +180,24 @@ pub enum Request { CollectPluginOptions { reply: flume::Sender, }, + HydrateState { + identity: PluginStateIdentity, + snapshot: Option, + reply: flume::Sender>, + }, + CaptureState { + identity: PluginStateIdentity, + revision: u64, + reply: flume::Sender>, + }, + ResetState { + identity: PluginStateIdentity, + reply: flume::Sender<()>, + }, + DropStateOwner { + owner: n00nId, + reply: flume::Sender<()>, + }, Shutdown, RestoreToolAsync { item: RestoreItem, @@ -281,6 +304,13 @@ pub struct LiveCtx { pub event_tx: n00n_agent::EventSender, pub tool_use_id: String, } +struct ContextLivenessGuard(Arc); + +impl Drop for ContextLivenessGuard { + fn drop(&mut self) { + self.0.store(false, Ordering::Release); + } +} /// Lua is single-threaded so this Mutex never contends, but /// `Lua::app_data` requires `Send + Sync` with the `send` feature. @@ -657,6 +687,40 @@ pub(crate) fn enqueue_async_task(lua: &Lua, work_fn: RegistryKey) -> Result<(), /// Caps concurrent coroutines to avoid blowing the Lua stack. /// Also serves as a drain barrier for load/clear ops. +#[derive(Default)] +struct LifecycleGate { + count: Cell, + event: Event, +} + +impl LifecycleGate { + fn start(self: &Rc) -> LifecycleGuard { + self.count.set(self.count.get() + 1); + LifecycleGuard(Rc::clone(self)) + } + + fn is_idle(&self) -> bool { + self.count.get() == 0 + } + + async fn changed(&self) { + let listener = self.event.listen(); + if self.is_idle() { + return; + } + listener.await; + } +} + +struct LifecycleGuard(Rc); + +impl Drop for LifecycleGuard { + fn drop(&mut self) { + self.0.count.set(self.0.count.get().saturating_sub(1)); + self.0.event.notify(usize::MAX); + } +} + struct InflightGate { lua: Lua, count: Cell, @@ -941,10 +1005,271 @@ async fn drain_barrier( } } +fn spawn_runtime_request( + rt: &LuaRuntime, + ex: &Rc>, + gate: &Rc, + lifecycle: &Rc, + request: Request, +) -> Option { + match request { + Request::CallTool { + plugin, + tool, + input, + mut ctx, + deadline, + reply, + live, + } => { + ctx.attach_plugin_state(Arc::clone(&plugin), Arc::clone(&rt.state)); + let lua = rt.lua.clone(); + let plugins = Rc::clone(&rt.plugins); + let live_tasks = Rc::clone(&rt.live_tasks); + let warm_tools = Rc::clone(&rt.warm_tools); + let shutdown = Arc::clone(&rt.shutdown); + let gate = Rc::clone(gate); + let lifecycle = lifecycle.start(); + ex.spawn(async move { + let result = run_tool_call( + lua, plugin, tool, input, ctx, deadline, live, live_tasks, warm_tools, plugins, + shutdown, gate, lifecycle, + ) + .await; + let _ = reply.send(result); + }) + .detach(); + None + } + Request::ComputeHeader { + plugin, + tool, + input, + reply, + } => { + let lua = rt.lua.clone(); + let plugins = Rc::clone(&rt.plugins); + let lifecycle = lifecycle.start(); + ex.spawn(async move { + let result = compute_header(&lua, &plugins, &plugin, &tool, input).await; + let _ = reply.send(result); + drop(lifecycle); + }) + .detach(); + None + } + Request::ComputePermissionScopes { + plugin, + tool, + input, + reply, + } => { + let lua = rt.lua.clone(); + let plugins = Rc::clone(&rt.plugins); + let lifecycle = lifecycle.start(); + ex.spawn(async move { + let result = + LuaRuntime::compute_permission_scopes(&lua, &plugins, &plugin, &tool, input) + .await; + let _ = reply.send(result); + drop(lifecycle); + }) + .detach(); + None + } + Request::StartTool { + plugin, + tool, + input, + live, + ctx, + reply, + } => { + let func = { + let plugins = rt.plugins.borrow(); + plugins + .get(&*plugin) + .and_then(|tools| tools.get(&*tool)) + .and_then(|keys| keys.start.as_ref()) + .and_then(|key| rt.lua.registry_value::(key).ok()) + }; + let Some(func) = func else { + let _ = reply.send(()); + return None; + }; + let lua = rt.lua.clone(); + let gate = Rc::clone(gate); + ex.spawn(async move { + let _gate_guard = gate.acquire().await; + run_tool_start(&lua, func, &tool, input, live, ctx).await; + let _ = reply.send(()); + }) + .detach(); + None + } + request => Some(request), + } +} + +enum RuntimeWake { + Lifecycle, + Spawn(PendingAsyncTask), + Request(Box), + Priority(Box), + SpawnClosed, + RequestClosed, + PriorityClosed, +} + +// Reentrant agents may dispatch Lua tools from another executor thread. Service every tool +// request while a lifecycle is active; thread-local origin markers cannot classify them safely. +async fn drain_runtime( + rt: &LuaRuntime, + ex: &Rc>, + gate: &Rc, + lifecycle: &Rc, + spawn_rx: &flume::Receiver, + request_rx: &flume::Receiver, + priority_rx: &flume::Receiver, + deferred: &mut VecDeque, +) -> bool { + let mut spawn_closed = false; + let mut request_closed = false; + let mut priority_closed = false; + while !lifecycle.is_idle() { + if rt.shutdown.load(Ordering::Acquire) { + return true; + } + while !priority_closed { + match priority_rx.try_recv() { + Ok(Request::Shutdown) => return true, + Ok(request) => deferred.push_back(request), + Err(flume::TryRecvError::Empty) => break, + Err(flume::TryRecvError::Disconnected) => { + priority_closed = true; + break; + } + } + } + while !spawn_closed { + match spawn_rx.try_recv() { + Ok(task) => spawn_async_task(&rt.lua, ex, gate, task), + Err(flume::TryRecvError::Empty) => break, + Err(flume::TryRecvError::Disconnected) => { + spawn_closed = true; + break; + } + } + } + while !request_closed { + match request_rx.try_recv() { + Ok(request) => { + if let Some(request) = spawn_runtime_request(rt, ex, gate, lifecycle, request) { + deferred.push_back(request); + } + } + Err(flume::TryRecvError::Empty) => break, + Err(flume::TryRecvError::Disconnected) => { + request_closed = true; + break; + } + } + } + if lifecycle.is_idle() { + break; + } + let wake = smol::future::or( + async { + lifecycle.changed().await; + RuntimeWake::Lifecycle + }, + smol::future::or( + async { + if priority_closed { + smol::future::pending::().await + } else { + priority_rx.recv_async().await.map_or_else( + |_| RuntimeWake::PriorityClosed, + |request| RuntimeWake::Priority(Box::new(request)), + ) + } + }, + smol::future::or( + async { + if spawn_closed { + smol::future::pending::().await + } else { + spawn_rx + .recv_async() + .await + .map_or(RuntimeWake::SpawnClosed, RuntimeWake::Spawn) + } + }, + async { + if request_closed { + smol::future::pending::().await + } else { + request_rx.recv_async().await.map_or_else( + |_| RuntimeWake::RequestClosed, + |request| RuntimeWake::Request(Box::new(request)), + ) + } + }, + ), + ), + ) + .await; + match wake { + RuntimeWake::Spawn(task) => spawn_async_task(&rt.lua, ex, gate, task), + RuntimeWake::Request(request) => { + if let Some(request) = spawn_runtime_request(rt, ex, gate, lifecycle, *request) { + deferred.push_back(request); + } + } + RuntimeWake::Priority(request) => { + if matches!(*request, Request::Shutdown) { + return true; + } + deferred.push_back(*request); + } + RuntimeWake::SpawnClosed => spawn_closed = true, + RuntimeWake::RequestClosed => request_closed = true, + RuntimeWake::PriorityClosed => priority_closed = true, + RuntimeWake::Lifecycle => {} + } + } + if rt.shutdown.load(Ordering::Acquire) { + return true; + } + drain_barrier(&rt.lua, ex, gate, spawn_rx).await; + false +} + +fn validate_snapshot_lua_values( + lua: &Lua, + identity: &PluginStateIdentity, + snapshot: Option<&StoredSessionStateSnapshot>, +) -> Result<(), String> { + let Some(snapshot) = snapshot else { + return Ok(()); + }; + let entries = snapshot + .plugin_entries_for_apply(PLUGIN_STATE_SCHEMA_VERSION) + .map_err(|error| error.to_string())?; + for entry in entries { + if !identity.is_root() && matches!(entry.scope, StoredStateScope::Root) { + continue; + } + state_json_to_lua(lua, entry.payload).map_err(|error| error.to_string())?; + } + Ok(()) +} + struct ToolKeys { handler: RegistryKey, header: Option, restore: Option, + start: Option, permission_scopes: Option, describe: Option, @@ -962,6 +1287,7 @@ struct LuaRuntime { live_tasks: LiveTasks, warm_tools: WarmTools, registry: Arc, + state: Arc, tx: flume::Sender, shutdown: Arc, bundled_dirs: &'static [&'static Dir<'static>], @@ -1054,6 +1380,7 @@ impl LuaRuntime { live_tasks: Rc::new(RefCell::new(HashMap::new())), warm_tools: Rc::new(RefCell::new(VecDeque::new())), registry, + state: Arc::new(PluginStateStore::default()), tx, shutdown, bundled_dirs, @@ -1557,21 +1884,22 @@ impl LuaRuntime { } async fn compute_permission_scopes( - &self, + lua: &Lua, + plugins: &PluginMap, plugin: &str, tool: &str, input: Value, ) -> Option { let (func, lua_input) = plugin_fn( - &self.lua, - &self.plugins, + lua, + plugins, plugin, tool, "permission_scopes", |tk| tk.permission_scopes.as_ref(), &input, )?; - let result: LuaValue = match run_detached(&self.lua, func.call_async(lua_input)).await { + let result: LuaValue = match run_detached(lua, func.call_async(lua_input)).await { Ok(v) => v, Err(e) => { tracing::warn!(plugin, tool, error = %e, "permission_scopes callback failed"); @@ -1732,8 +2060,10 @@ async fn restore_item( }), ); + let ctx = LuaCtx::restore(item.tool_output_lines, item.state); + let _context_liveness = ContextLivenessGuard(ctx.context_liveness()); let ctx = lua - .create_userdata(LuaCtx::restore(item.tool_output_lines, item.state)) + .create_userdata(ctx) .map_err(|e| format!("restore context creation failed: {e}"))?; let inner = thread .into_async::((input_lua, &*item.output, item.is_error, ctx)) @@ -2014,6 +2344,7 @@ async fn run_tool_start( live: LiveCtx, ctx: Box, ) { + let _context_liveness = ContextLivenessGuard(ctx.context_liveness()); let scope = TaskScope::new(lua, TaskCell::new(ctx.cancel.clone(), None, Some(live))); let run = async { let input_lua = json_to_lua(lua, &input)?; @@ -2043,7 +2374,9 @@ async fn run_tool_call( plugins: PluginMap, shutdown: Arc, gate: Rc, + _lifecycle: LifecycleGuard, ) -> ToolCallReply { + let _context_liveness = ContextLivenessGuard(ctx.context_liveness()); let handler: Function = { let plugins_ref = plugins.borrow(); let Some(keys) = plugins_ref.get(&*plugin) else { @@ -2239,6 +2572,7 @@ pub fn spawn( let ex = Rc::new(smol::LocalExecutor::new()); let gate = Rc::new(InflightGate::new(rt.lua.clone())); + let lifecycle = Rc::new(LifecycleGate::default()); let restores = Rc::new(RestoreTracker::default()); let spawn_rx = rt .lua @@ -2248,6 +2582,7 @@ pub fn spawn( .clone(); smol::block_on(ex.run(async { + let mut deferred = VecDeque::new(); loop { while let Ok(task) = spawn_rx.try_recv() { spawn_async_task(&rt.lua, &ex, &gate, task); @@ -2256,25 +2591,29 @@ pub fn spawn( // ahead of bulk work like session restores so the UI stays // snappy, and queued `n00n.async.run` tasks jump ahead of // plain requests. - let next = smol::future::or( - async { prio_rx.recv_async().await.map(Some) }, - smol::future::or( - async { - let task = spawn_rx.recv_async().await?; - spawn_async_task(&rt.lua, &ex, &gate, task); - Ok(None) - }, - async { rx.recv_async().await.map(Some) }, - ), - ) - .await; - let msg = match next { - Ok(Some(m)) => m, - Ok(None) => { - smol::future::yield_now().await; - continue; + let msg = if let Some(request) = deferred.pop_front() { + request + } else { + let next = smol::future::or( + async { prio_rx.recv_async().await.map(Some) }, + smol::future::or( + async { + let task = spawn_rx.recv_async().await?; + spawn_async_task(&rt.lua, &ex, &gate, task); + Ok(None) + }, + async { rx.recv_async().await.map(Some) }, + ), + ) + .await; + match next { + Ok(Some(request)) => request, + Ok(None) => { + smol::future::yield_now().await; + continue; + } + Err(_) => break, } - Err(_) => break, }; match msg { Request::Shutdown => break, @@ -2286,47 +2625,43 @@ pub fn spawn( opts, reply, } => { - drain_barrier(&rt.lua, &ex, &gate, &spawn_rx).await; + if drain_runtime( + &rt, + &ex, + &gate, + &lifecycle, + &spawn_rx, + &rx, + &prio_rx, + &mut deferred, + ) + .await + { + break; + } let res = rt.load_source(Arc::clone(&name), &source, plugin_dir, &permissions, opts, None).await; let _ = reply.send(res); } - Request::CallTool { - plugin, - tool, - input, - ctx, - deadline, - reply, - live, - } => { - let lua = rt.lua.clone(); - let plugins = Rc::clone(&rt.plugins); - let live_tasks = Rc::clone(&rt.live_tasks); - let warm_tools = Rc::clone(&rt.warm_tools); - let shutdown_ref = Arc::clone(&rt.shutdown); - let g = Rc::clone(&gate); - ex.spawn(async move { - let res = run_tool_call( - lua.clone(), - plugin, - tool, - input, - ctx, - deadline, - live, - live_tasks, - warm_tools, - plugins, - shutdown_ref, - g, - ) - .await; - let _ = reply.send(res); - }) - .detach(); + request @ (Request::CallTool { .. } | Request::StartTool { .. }) => { + let deferred_request = + spawn_runtime_request(&rt, &ex, &gate, &lifecycle, request); + debug_assert!(deferred_request.is_none()); } Request::ClearPlugin { plugin, reply } => { - drain_barrier(&rt.lua, &ex, &gate, &spawn_rx).await; + if drain_runtime( + &rt, + &ex, + &gate, + &lifecycle, + &spawn_rx, + &rx, + &prio_rx, + &mut deferred, + ) + .await + { + break; + } rt.clear_plugin(&plugin); let _ = reply.send(()); } @@ -2370,7 +2705,14 @@ pub fn spawn( input, reply, } => { - let res = rt.compute_permission_scopes(&plugin, &tool, input).await; + let res = LuaRuntime::compute_permission_scopes( + &rt.lua, + &rt.plugins, + &plugin, + &tool, + input, + ) + .await; let _ = reply.send(res); } Request::RunInitLua { @@ -2379,7 +2721,20 @@ pub fn spawn( plugin_dir, reply, } => { - drain_barrier(&rt.lua, &ex, &gate, &spawn_rx).await; + if drain_runtime( + &rt, + &ex, + &gate, + &lifecycle, + &spawn_rx, + &rx, + &prio_rx, + &mut deferred, + ) + .await + { + break; + } let res = rt.run_init_lua(&source, &source_name, plugin_dir).await; let _ = reply.send(res); } @@ -2390,6 +2745,98 @@ pub fn spawn( Request::CollectPluginOptions { reply } => { let _ = reply.send(collect_plugin_options(&rt.lua)); } + Request::HydrateState { + identity, + snapshot, + reply, + } => { + if drain_runtime( + &rt, + &ex, + &gate, + &lifecycle, + &spawn_rx, + &rx, + &prio_rx, + &mut deferred, + ) + .await + { + break; + } + let result = validate_snapshot_lua_values( + &rt.lua, + &identity, + snapshot.as_ref(), + ) + .and_then(|()| { + rt.state + .hydrate(identity, snapshot) + .map_err(|error| error.to_string()) + }); + let _ = reply.send(result); + } + Request::CaptureState { + identity, + revision, + reply, + } => { + if drain_runtime( + &rt, + &ex, + &gate, + &lifecycle, + &spawn_rx, + &rx, + &prio_rx, + &mut deferred, + ) + .await + { + break; + } + let result = rt + .state + .capture(&identity, revision) + .map_err(|error| error.to_string()); + let _ = reply.send(result); + } + Request::ResetState { identity, reply } => { + if drain_runtime( + &rt, + &ex, + &gate, + &lifecycle, + &spawn_rx, + &rx, + &prio_rx, + &mut deferred, + ) + .await + { + break; + } + rt.state.reset(&identity); + let _ = reply.send(()); + } + Request::DropStateOwner { owner, reply } => { + if drain_runtime( + &rt, + &ex, + &gate, + &lifecycle, + &spawn_rx, + &rx, + &prio_rx, + &mut deferred, + ) + .await + { + break; + } + rt.state.drop_owner(owner); + let _ = reply.send(()); + } Request::RestoreToolAsync { item, event_tx } => { spawn_restore(&ex, &gate, &restores, &rt, item, event_tx); } @@ -2484,35 +2931,6 @@ pub fn spawn( let _ = reply .send(run_describe(&rt.lua, &rt.plugins, &plugin, &tool, &dctx)); } - Request::StartTool { - plugin, - tool, - input, - live, - ctx, - reply, - } => { - let func = { - let plugins = rt.plugins.borrow(); - plugins - .get(&*plugin) - .and_then(|p| p.get(&*tool)) - .and_then(|tk| tk.start.as_ref()) - .and_then(|key| rt.lua.registry_value::(key).ok()) - }; - let Some(func) = func else { - let _ = reply.send(()); - continue; - }; - let lua = rt.lua.clone(); - let g = Rc::clone(&gate); - ex.spawn(async move { - let _gate_guard = g.acquire().await; - run_tool_start(&lua, func, &tool, input, live, ctx).await; - let _ = reply.send(()); - }) - .detach(); - } Request::RunKeybindCallback { id } => { let func = rt.lua.app_data_ref::().and_then(|store| { let key = store.callback_for_id(id)?; diff --git a/n00n-lua/src/state.rs b/n00n-lua/src/state.rs new file mode 100644 index 000000000..172635805 --- /dev/null +++ b/n00n-lua/src/state.rs @@ -0,0 +1,863 @@ +use n00n_agent::tools::SessionIdentity; +use n00n_storage::{ + id::n00nId, + sessions::{ + MAX_PLUGIN_STATE_BYTES, SessionStateError, StoredSessionStateSnapshot, StoredStateScope, + }, +}; +use serde_json::Value; +use std::{ + collections::{HashMap, HashSet}, + sync::{Mutex, MutexGuard}, +}; + +pub(crate) const PLUGIN_STATE_SCHEMA_VERSION: u32 = 1; + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] +pub(crate) enum PluginStateScope { + Session, + Root, +} + +impl PluginStateScope { + pub(crate) fn parse(value: &str) -> Option { + match value { + "session" => Some(Self::Session), + "root" => Some(Self::Root), + _ => None, + } + } + + const fn stored(self) -> StoredStateScope { + match self { + Self::Session => StoredStateScope::Session, + Self::Root => StoredStateScope::Root, + } + } +} + +impl From for PluginStateScope { + fn from(scope: StoredStateScope) -> Self { + match scope { + StoredStateScope::Session => Self::Session, + StoredStateScope::Root => Self::Root, + } + } +} + +#[derive(Debug, Clone, PartialEq, Eq, Hash)] +pub(crate) struct PluginStateIdentity { + session_id: n00nId, + root_session_id: n00nId, +} + +impl PluginStateIdentity { + fn owner(&self, scope: PluginStateScope) -> n00nId { + match scope { + PluginStateScope::Session => self.session_id, + PluginStateScope::Root => self.root_session_id, + } + } + + fn owns_scope(&self, scope: PluginStateScope) -> bool { + matches!(scope, PluginStateScope::Session) || self.is_root() + } + + pub(crate) fn is_root(&self) -> bool { + self.session_id == self.root_session_id + } + + fn scope_identity(&self, scope: PluginStateScope) -> Self { + match scope { + PluginStateScope::Session => self.clone(), + PluginStateScope::Root => Self { + session_id: self.root_session_id, + root_session_id: self.root_session_id, + }, + } + } +} + +impl From<&SessionIdentity> for PluginStateIdentity { + fn from(identity: &SessionIdentity) -> Self { + Self { + session_id: identity.session_id().id(), + root_session_id: identity.root_session_id().id(), + } + } +} + +#[derive(Debug, Clone, PartialEq, Eq, Hash)] +struct StateKey { + plugin: String, + scope: PluginStateScope, + owner: n00nId, +} + +impl StateKey { + fn new(plugin: &str, scope: PluginStateScope, identity: &PluginStateIdentity) -> Self { + Self { + plugin: plugin.to_owned(), + scope, + owner: identity.owner(scope), + } + } +} + +#[derive(Debug, thiserror::Error)] +pub(crate) enum PluginStateError { + #[error("plugin state is {bytes} bytes (maximum {maximum})")] + ValueTooLarge { bytes: usize, maximum: usize }, + #[error(transparent)] + Serialize(#[from] serde_json::Error), + #[error(transparent)] + Snapshot(#[from] SessionStateError), +} + +#[derive(Default)] +struct StateInner { + values: HashMap, + managed: HashSet, + bases: HashMap, +} + +#[derive(Default)] +pub(crate) struct PluginStateStore { + inner: Mutex, +} + +impl PluginStateStore { + fn lock(&self) -> MutexGuard<'_, StateInner> { + self.inner + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + } + + #[must_use] + pub(crate) fn get( + &self, + plugin: &str, + scope: PluginStateScope, + identity: &PluginStateIdentity, + ) -> Option { + self.lock() + .values + .get(&StateKey::new(plugin, scope, identity)) + .cloned() + } + + pub(crate) fn replace( + &self, + plugin: &str, + scope: PluginStateScope, + identity: &PluginStateIdentity, + value: Value, + ) -> Result, PluginStateError> { + validate_value_size(&value)?; + let key = StateKey::new(plugin, scope, identity); + let mut inner = self.lock(); + validate_replacement(&inner, &identity.scope_identity(scope), &key, &value)?; + inner.managed.insert(key.clone()); + Ok(inner.values.insert(key, value)) + } + + pub(crate) fn remove( + &self, + plugin: &str, + scope: PluginStateScope, + identity: &PluginStateIdentity, + ) -> Result, PluginStateError> { + let key = StateKey::new(plugin, scope, identity); + let mut inner = self.lock(); + validate_removal(&inner, &identity.scope_identity(scope), &key)?; + inner.managed.insert(key.clone()); + Ok(inner.values.remove(&key)) + } + + pub(crate) fn hydrate( + &self, + identity: PluginStateIdentity, + mut snapshot: Option, + ) -> Result<(), PluginStateError> { + let entries = match snapshot.as_ref() { + Some(snapshot) => snapshot_entries(snapshot)?, + None => Vec::new(), + }; + if !identity.owns_scope(PluginStateScope::Root) + && let Some(snapshot) = snapshot.as_mut() + { + let root_plugins = snapshot + .plugin_names_with_scope(StoredStateScope::Root)? + .into_iter() + .map(str::to_owned) + .collect::>(); + for plugin in root_plugins { + snapshot.remove_plugin_state(&plugin, StoredStateScope::Root)?; + } + } + + let mut inner = self.lock(); + clear_identity_runtime(&mut inner, &identity, false); + for (plugin, scope, payload) in entries { + if !identity.owns_scope(scope) { + continue; + } + let key = StateKey::new(&plugin, scope, &identity); + inner.managed.insert(key.clone()); + inner.values.insert(key, payload); + } + if let Some(snapshot) = snapshot { + inner.bases.insert(identity, snapshot); + } else { + inner.bases.remove(&identity); + } + Ok(()) + } + + pub(crate) fn capture( + &self, + identity: &PluginStateIdentity, + revision: u64, + ) -> Result { + let mut inner = self.lock(); + let mut candidate = candidate_for(&inner, identity, None)?; + candidate.set_state_revision(revision)?; + inner.bases.insert(identity.clone(), candidate.clone()); + Ok(candidate) + } + + pub(crate) fn reset(&self, identity: &PluginStateIdentity) { + clear_identity_runtime(&mut self.lock(), identity, true); + } + + pub(crate) fn drop_owner(&self, owner: n00nId) { + let mut inner = self.lock(); + inner.values.retain(|key, _| key.owner != owner); + inner.managed.retain(|key| key.owner != owner); + inner.bases.retain(|identity, _| { + identity.session_id != owner && identity.root_session_id != owner + }); + } +} + +fn candidate_for( + inner: &StateInner, + identity: &PluginStateIdentity, + skipped: Option<&StateKey>, +) -> Result { + let mut candidate = inner + .bases + .get(identity) + .cloned() + .unwrap_or_else(|| StoredSessionStateSnapshot::new(0)); + for key in inner + .managed + .iter() + .filter(|key| identity.owns_scope(key.scope) && key.owner == identity.owner(key.scope)) + { + if skipped.map_or(false, |s| key == s) { + continue; + } + if let Some(value) = inner.values.get(key) { + candidate.set_plugin_state( + &key.plugin, + PLUGIN_STATE_SCHEMA_VERSION, + key.scope.stored(), + value.clone(), + )?; + } else { + candidate.remove_plugin_state(&key.plugin, key.scope.stored())?; + } + } + Ok(candidate) +} + +fn validate_replacement( + inner: &StateInner, + identity: &PluginStateIdentity, + replacement_key: &StateKey, + replacement_value: &Value, +) -> Result<(), PluginStateError> { + let mut candidate = candidate_for(inner, identity, Some(replacement_key))?; + candidate.set_plugin_state( + &replacement_key.plugin, + PLUGIN_STATE_SCHEMA_VERSION, + replacement_key.scope.stored(), + replacement_value.clone(), + )?; + Ok(()) +} + +fn validate_value_size(value: &Value) -> Result<(), PluginStateError> { + let bytes = serde_json::to_vec(value)?.len(); + if bytes > MAX_PLUGIN_STATE_BYTES { + return Err(PluginStateError::ValueTooLarge { + bytes, + maximum: MAX_PLUGIN_STATE_BYTES, + }); + } + + Ok(()) +} + +fn validate_removal( + inner: &StateInner, + identity: &PluginStateIdentity, + removal_key: &StateKey, +) -> Result<(), PluginStateError> { + let mut candidate = candidate_for(inner, identity, Some(removal_key))?; + candidate.remove_plugin_state(&removal_key.plugin, removal_key.scope.stored())?; + Ok(()) +} + +fn snapshot_entries( + snapshot: &StoredSessionStateSnapshot, +) -> Result, PluginStateError> { + snapshot + .plugin_entries_for_apply(PLUGIN_STATE_SCHEMA_VERSION)? + .into_iter() + .map(|entry| { + validate_value_size(entry.payload)?; + Ok(( + entry.plugin.to_owned(), + entry.scope.into(), + entry.payload.clone(), + )) + }) + .collect() +} + +fn clear_identity_runtime( + inner: &mut StateInner, + identity: &PluginStateIdentity, + mark_managed: bool, +) { + let keys = inner + .values + .keys() + .filter(|key| identity.owns_scope(key.scope) && key.owner == identity.owner(key.scope)) + .cloned() + .collect::>(); + inner + .values + .retain(|key, _| !identity.owns_scope(key.scope) || key.owner != identity.owner(key.scope)); + if mark_managed { + inner.managed.extend(keys); + } else { + inner.managed.retain(|key| { + !identity.owns_scope(key.scope) || key.owner != identity.owner(key.scope) + }); + } +} + +#[cfg(test)] +mod tests { + use super::{ + PluginStateError, PluginStateIdentity, PluginStateScope, PluginStateStore, StateKey, + }; + use n00n_agent::tools::SessionIdentity; + use n00n_storage::{ + id::{SessionRef, n00nId}, + sessions::{ + MAX_PLUGIN_STATE_BYTES, SESSION_STATE_SCHEMA_VERSION, SessionMeta, + StoredSessionStateSnapshot, StoredStateScope, + }, + }; + use serde_json::json; + use std::{ + sync::{Arc, Barrier}, + thread, + }; + + fn identity() -> PluginStateIdentity { + PluginStateIdentity::from(&SessionIdentity::child( + SessionRef::generate(), + SessionRef::generate(), + )) + } + + #[test] + fn identity_uses_canonical_ids() { + let raw = "550e8400-e29b-41d4-a716-446655440000"; + let id = raw.parse::().unwrap(); + let legacy = SessionIdentity::root(raw.parse::().unwrap()); + let canonical = SessionIdentity::root(SessionRef::from_id(id)); + + assert_eq!( + PluginStateIdentity::from(&legacy), + PluginStateIdentity::from(&canonical) + ); + } + + #[test] + fn plugin_scope_and_owner_are_all_part_of_the_key() { + let store = PluginStateStore::default(); + let root = SessionRef::generate(); + let a = PluginStateIdentity::from(&SessionIdentity::child( + SessionRef::generate(), + root.clone(), + )); + let b = PluginStateIdentity::from(&SessionIdentity::child(SessionRef::generate(), root)); + store + .replace("one", PluginStateScope::Session, &a, json!(1)) + .unwrap(); + store + .replace("two", PluginStateScope::Session, &a, json!(2)) + .unwrap(); + store + .replace("one", PluginStateScope::Session, &b, json!(3)) + .unwrap(); + store + .replace("one", PluginStateScope::Root, &a, json!(4)) + .unwrap(); + + assert_eq!( + store.get("one", PluginStateScope::Session, &a), + Some(json!(1)) + ); + assert_eq!( + store.get("two", PluginStateScope::Session, &a), + Some(json!(2)) + ); + assert_eq!( + store.get("one", PluginStateScope::Session, &b), + Some(json!(3)) + ); + assert_eq!(store.get("one", PluginStateScope::Root, &b), Some(json!(4))); + } + + #[test] + fn oversized_replace_is_atomic() { + let store = PluginStateStore::default(); + let identity = identity(); + store + .replace( + "plugin", + PluginStateScope::Session, + &identity, + json!("kept"), + ) + .unwrap(); + + let error = store + .replace( + "plugin", + PluginStateScope::Session, + &identity, + json!("x".repeat(MAX_PLUGIN_STATE_BYTES)), + ) + .unwrap_err(); + + assert!(matches!(error, PluginStateError::ValueTooLarge { .. })); + assert_eq!( + store.get("plugin", PluginStateScope::Session, &identity), + Some(json!("kept")) + ); + } + + #[test] + fn capture_preserves_opaque_data_and_advances_revision() { + let store = PluginStateStore::default(); + let identity = identity(); + let snapshot = serde_json::from_value(json!({ + "schema_version": SESSION_STATE_SCHEMA_VERSION, + "state_revision": 3, + "future": {"kept": true}, + "plugins": {"unknown": {"future": {"raw": "kept"}}} + })) + .unwrap(); + store.hydrate(identity.clone(), Some(snapshot)).unwrap(); + store + .replace( + "known", + PluginStateScope::Session, + &identity, + json!({"v": 1}), + ) + .unwrap(); + + let captured = store.capture(&identity, 4).unwrap(); + let raw = serde_json::to_value(captured).unwrap(); + assert_eq!(raw["state_revision"], json!(4)); + assert_eq!(raw["future"], json!({"kept": true})); + assert_eq!(raw["plugins"]["unknown"]["future"], json!({"raw": "kept"})); + } + + #[test] + fn future_and_malformed_hydrates_leave_current_state_unchanged() { + let store = PluginStateStore::default(); + let identity = identity(); + store + .replace( + "plugin", + PluginStateScope::Session, + &identity, + json!("kept"), + ) + .unwrap(); + let future = serde_json::from_value(json!({ + "schema_version": SESSION_STATE_SCHEMA_VERSION + 1, + "state_revision": 7, + "opaque": true + })) + .unwrap(); + assert!(store.hydrate(identity.clone(), Some(future)).is_err()); + assert_eq!( + store.get("plugin", PluginStateScope::Session, &identity), + Some(json!("kept")) + ); + + let meta: SessionMeta = serde_json::from_value(json!({ + "state_snapshot": { + "schema_version": SESSION_STATE_SCHEMA_VERSION, + "plugins": {} + } + })) + .unwrap(); + assert!( + store + .hydrate(identity.clone(), meta.state_snapshot) + .is_err() + ); + assert_eq!( + store.get("plugin", PluginStateScope::Session, &identity), + Some(json!("kept")) + ); + } + + #[test] + fn hydrate_replaces_pending_runtime_removals() { + let store = PluginStateStore::default(); + let identity = identity(); + store + .replace("plugin", PluginStateScope::Session, &identity, json!("old")) + .unwrap(); + store + .remove("plugin", PluginStateScope::Session, &identity) + .unwrap(); + + let mut snapshot = StoredSessionStateSnapshot::new(2); + snapshot + .set_plugin_state("plugin", 1, StoredStateScope::Session, json!("hydrated")) + .unwrap(); + store.hydrate(identity.clone(), Some(snapshot)).unwrap(); + + let captured = store.capture(&identity, 3).unwrap(); + assert_eq!( + captured + .plugin_payload_for_apply("plugin", 1, StoredStateScope::Session) + .unwrap(), + Some(&json!("hydrated")) + ); + } + + #[test] + fn failed_revision_regression_does_not_replace_base() { + let store = PluginStateStore::default(); + let identity = identity(); + let snapshot = StoredSessionStateSnapshot::new(5); + store.hydrate(identity.clone(), Some(snapshot)).unwrap(); + + assert!(store.capture(&identity, 4).is_err()); + assert_eq!( + store.capture(&identity, 6).unwrap().state_revision(), + Some(6) + ); + } + + #[test] + fn reset_removes_managed_state_but_preserves_opaque_base() { + let store = PluginStateStore::default(); + let identity = identity(); + let mut snapshot = StoredSessionStateSnapshot::new(1); + snapshot + .set_plugin_state("known", 1, StoredStateScope::Session, json!("old")) + .unwrap(); + snapshot + .set_plugin_state("future", 2, StoredStateScope::Session, json!("opaque")) + .unwrap(); + store.hydrate(identity.clone(), Some(snapshot)).unwrap(); + + store.reset(&identity); + let captured = store.capture(&identity, 2).unwrap(); + assert_eq!( + captured + .plugin_payload_for_apply("known", 1, StoredStateScope::Session) + .unwrap(), + None + ); + assert_eq!( + captured + .plugin_payload_for_apply("future", 2, StoredStateScope::Session) + .unwrap(), + Some(&json!("opaque")) + ); + } + + #[test] + fn replacement_rejects_uncapturable_names_and_entry_counts_atomically() { + let store = PluginStateStore::default(); + let identity = identity(); + assert!( + store + .replace( + "bad/name", + PluginStateScope::Session, + &identity, + json!(true), + ) + .is_err() + ); + assert_eq!( + store.get("bad/name", PluginStateScope::Session, &identity), + None + ); + + for index in 0..64 { + store + .replace( + &format!("plugin_{index}"), + PluginStateScope::Session, + &identity, + json!(index), + ) + .unwrap(); + } + assert!( + store + .replace("plugin_64", PluginStateScope::Session, &identity, json!(64),) + .is_err() + ); + assert_eq!( + store.get("plugin_64", PluginStateScope::Session, &identity), + None + ); + } + + #[test] + fn child_hydration_cannot_clobber_live_root_state() { + let store = PluginStateStore::default(); + let root_ref = SessionRef::generate(); + let root = PluginStateIdentity::from(&SessionIdentity::root(root_ref.clone())); + let child = + PluginStateIdentity::from(&SessionIdentity::child(SessionRef::generate(), root_ref)); + store + .replace("plugin", PluginStateScope::Root, &root, json!("live")) + .unwrap(); + let mut child_snapshot = StoredSessionStateSnapshot::new(1); + child_snapshot + .set_plugin_state("plugin", 1, StoredStateScope::Root, json!("stale")) + .unwrap(); + + store.hydrate(child.clone(), Some(child_snapshot)).unwrap(); + + assert_eq!( + store.get("plugin", PluginStateScope::Root, &child), + Some(json!("live")) + ); + } + + #[test] + fn child_root_replacement_uses_root_snapshot_limits() { + let store = PluginStateStore::default(); + let root_ref = SessionRef::generate(); + let root = PluginStateIdentity::from(&SessionIdentity::root(root_ref.clone())); + let child = + PluginStateIdentity::from(&SessionIdentity::child(SessionRef::generate(), root_ref)); + for index in 0..64 { + store + .replace( + &format!("plugin_{index}"), + PluginStateScope::Root, + &root, + json!(index), + ) + .unwrap(); + } + + assert!( + store + .replace("plugin_64", PluginStateScope::Root, &child, json!(64),) + .is_err() + ); + assert_eq!(store.get("plugin_64", PluginStateScope::Root, &root), None); + } + + #[test] + fn removal_from_malformed_container_is_atomic() { + let store = PluginStateStore::default(); + let root_ref = SessionRef::generate(); + let identity = PluginStateIdentity::from(&SessionIdentity::root(root_ref)); + let snapshot: StoredSessionStateSnapshot = serde_json::from_value(json!({ + "schema_version": SESSION_STATE_SCHEMA_VERSION, + "state_revision": 1, + "plugins": {"plugin": null} + })) + .unwrap(); + store.hydrate(identity.clone(), Some(snapshot)).unwrap(); + + assert!( + store + .remove("plugin", PluginStateScope::Root, &identity) + .is_err() + ); + let captured = store.capture(&identity, 2).unwrap(); + assert_eq!( + serde_json::to_value(captured).unwrap()["plugins"]["plugin"], + json!(null) + ); + } + + #[test] + fn capture_rejects_candidate_mutation_failure_without_replacing_base() { + let store = PluginStateStore::default(); + let identity = identity(); + let snapshot: StoredSessionStateSnapshot = serde_json::from_value(json!({ + "schema_version": SESSION_STATE_SCHEMA_VERSION, + "state_revision": 1, + "plugins": {"plugin": null} + })) + .unwrap(); + store.hydrate(identity.clone(), Some(snapshot)).unwrap(); + + let key = StateKey::new("plugin", PluginStateScope::Session, &identity); + { + let mut inner = store.lock(); + inner.managed.insert(key.clone()); + inner.values.insert(key, json!("new")); + } + + assert!(store.capture(&identity, 2).is_err()); + let inner = store.lock(); + let base = &inner.bases[&identity]; + assert_eq!(base.state_revision(), Some(1)); + assert_eq!( + serde_json::to_value(base).unwrap()["plugins"]["plugin"], + json!(null) + ); + } + + #[test] + fn concurrent_child_mutations_remain_isolated() { + let store = Arc::new(PluginStateStore::default()); + let root_ref = SessionRef::generate(); + let root = PluginStateIdentity::from(&SessionIdentity::root(root_ref.clone())); + let first = PluginStateIdentity::from(&SessionIdentity::child( + SessionRef::generate(), + root_ref.clone(), + )); + let second = + PluginStateIdentity::from(&SessionIdentity::child(SessionRef::generate(), root_ref)); + store + .replace("plugin", PluginStateScope::Root, &root, json!("root")) + .unwrap(); + let barrier = Arc::new(Barrier::new(2)); + + let workers = [first.clone(), second.clone()].map(|identity| { + let store = Arc::clone(&store); + let barrier = Arc::clone(&barrier); + thread::spawn(move || { + barrier.wait(); + for value in 0..100 { + store + .replace("plugin", PluginStateScope::Session, &identity, json!(value)) + .unwrap(); + } + }) + }); + for worker in workers { + worker.join().unwrap(); + } + + assert_eq!( + store.get("plugin", PluginStateScope::Session, &first), + Some(json!(99)) + ); + assert_eq!( + store.get("plugin", PluginStateScope::Session, &second), + Some(json!(99)) + ); + assert_eq!( + store.get("plugin", PluginStateScope::Root, &first), + Some(json!("root")) + ); + } + + #[test] + fn dropping_child_owner_preserves_root_and_sibling_state() { + let store = PluginStateStore::default(); + let root_ref = SessionRef::generate(); + let root = PluginStateIdentity::from(&SessionIdentity::root(root_ref.clone())); + let child = PluginStateIdentity::from(&SessionIdentity::child( + SessionRef::generate(), + root_ref.clone(), + )); + let sibling = + PluginStateIdentity::from(&SessionIdentity::child(SessionRef::generate(), root_ref)); + store + .replace("plugin", PluginStateScope::Root, &root, json!("root")) + .unwrap(); + store + .replace("plugin", PluginStateScope::Session, &child, json!("child")) + .unwrap(); + store + .replace( + "plugin", + PluginStateScope::Session, + &sibling, + json!("sibling"), + ) + .unwrap(); + + store.drop_owner(child.session_id); + + assert_eq!(store.get("plugin", PluginStateScope::Session, &child), None); + assert_eq!( + store.get("plugin", PluginStateScope::Session, &sibling), + Some(json!("sibling")) + ); + assert_eq!( + store.get("plugin", PluginStateScope::Root, &child), + Some(json!("root")) + ); + } + + #[test] + fn child_capture_does_not_emit_supported_root_state() { + let store = PluginStateStore::default(); + let root_ref = SessionRef::generate(); + let child = + PluginStateIdentity::from(&SessionIdentity::child(SessionRef::generate(), root_ref)); + let mut snapshot = StoredSessionStateSnapshot::new(1); + snapshot + .set_plugin_state("plugin", 1, StoredStateScope::Root, json!("stale")) + .unwrap(); + snapshot + .set_plugin_state("plugin", 1, StoredStateScope::Session, json!("session")) + .unwrap(); + snapshot + .set_plugin_state("future", 2, StoredStateScope::Root, json!("opaque")) + .unwrap(); + store.hydrate(child.clone(), Some(snapshot)).unwrap(); + + let captured = store.capture(&child, 2).unwrap(); + assert_eq!( + captured + .plugin_payload_for_apply("plugin", 1, StoredStateScope::Root) + .unwrap(), + None + ); + assert_eq!( + captured + .plugin_payload_for_apply("future", 2, StoredStateScope::Root) + .unwrap(), + None + ); + assert_eq!( + captured + .plugin_payload_for_apply("plugin", 1, StoredStateScope::Session) + .unwrap(), + Some(&json!("session")) + ); + } +} diff --git a/n00n-lua/tests/plugin_host.rs b/n00n-lua/tests/plugin_host.rs index 3e85c99f8..197a0f302 100644 --- a/n00n-lua/tests/plugin_host.rs +++ b/n00n-lua/tests/plugin_host.rs @@ -12,8 +12,8 @@ use std::time::Duration; use n00n_agent::template::env_vars; use n00n_agent::tools::{ - ActiveTools, DescriptionContext, ToolAudience, ToolFilter, ToolRegistry, ToolSource, - timeout_annotation, + ActiveTools, DescriptionContext, SessionIdentity, ToolAudience, ToolFilter, ToolRegistry, + ToolSource, timeout_annotation, }; use n00n_config::{AlwaysThinking, PluginsConfig, ToolOutputLines}; use n00n_lua::{PluginError, PluginHost, WARM_TOOL_CAP}; @@ -235,6 +235,7 @@ const UNKNOWN_FIELD_ERR: &str = "unknown field"; const PERMISSION_DENIED_MSG: &str = "permission denied"; const VALIDATION_PROMPT_NO_PROVIDER_ERR: &str = "validation prompt error: no provider configured — run /login or `n00n auth login`"; +const STALE_CTX_ERR: &str = "state context is no longer active"; const TOOLS_MUST_BE_ARRAY_ERR: &str = "tools must be an array"; #[test] @@ -4061,6 +4062,530 @@ fn session_rejects_nonempty_lua_tools_object() { assert!(error.contains(TOOLS_MUST_BE_ARRAY_ERR), "got: {error}"); } +#[test] +fn lua_session_rejects_missing_identity() { + smol::block_on(async { + let reg = fresh_registry(); + let host = PluginHost::new(Arc::clone(®)).unwrap(); + let src = format!( + r#"n00n.api.register_tool({{ + name = "missing_identity_probe", + description = "test", + schema = {MINIMAL_SCHEMA}, + audiences = {{ "main" }}, + handler = function(input, ctx) + local sess, err = n00n.agent.session(ctx, {{}}) + if sess ~= nil then return "unexpected session" end + return err or "no error" + end + }})"# + ); + host.load_source("missing_identity_plugin", &src).unwrap(); + let entry = reg.get("missing_identity_probe").unwrap(); + let invocation = entry.tool.parse(&serde_json::json!({})).unwrap(); + let mut ctx = n00n_agent::tools::test_support::stub_ctx(&n00n_agent::AgentMode::Build); + ctx.identity = None; + + let result = invocation.execute(&ctx).await.output.unwrap(); + let n00n_agent::ToolOutput::Plain(output) = result else { + panic!("expected plain output"); + }; + assert_eq!(output.text, "session identity is unavailable"); + }); +} + +#[test] +fn plugin_state_capture_waits_for_inflight_handler_callbacks() { + let reg = fresh_registry(); + let host = PluginHost::new(Arc::clone(®)).unwrap(); + let source = format!( + r#" + n00n.api.register_tool({{ + name = "delayed_state", description = "test", schema = {MINIMAL_SCHEMA}, + handler = function(input, ctx) + local buf = n00n.ui.buf() + buf:set_lines({{ "waiting" }}) + ctx:live_buf(buf) + n00n.async.run(function() + local id = n00n.fn.jobstart("sleep 0.5") + n00n.fn.jobwait(id) + return "finished" + end, function(err) + if err then + ctx:finish(err) + return + end + local _, state_err = ctx:state_replace("session", {{ value = "finished" }}) + ctx:finish(state_err or "done") + end) + return nil + end, + }}) + "# + ); + host.load_source("delayed_state", &source).unwrap(); + let identity = SessionIdentity::root(SessionRef::generate()); + let entry = reg.get("delayed_state").unwrap(); + let invocation = entry.tool.parse(&serde_json::json!({})).unwrap(); + let (event_tx, event_rx) = flume::unbounded(); + let sender = n00n_agent::EventSender::new(event_tx, 0); + let mut ctx = n00n_agent::tools::test_support::stub_ctx_with( + &n00n_agent::AgentMode::Build, + Some(&sender), + Some("delayed-state"), + ); + ctx.identity = Some(identity.clone()); + let worker = std::thread::spawn(move || smol::block_on(invocation.execute(&ctx))); + + loop { + let event = event_rx.recv_timeout(Duration::from_secs(5)).unwrap(); + if matches!( + event.event, + n00n_agent::AgentEvent::LiveToolBuf { ref id, .. } if id == "delayed-state" + ) { + break; + } + } + + let snapshot = host + .event_handle() + .unwrap() + .capture_state(&identity, 1) + .unwrap(); + assert_eq!( + snapshot + .plugin_payload_for_apply( + "delayed_state", + 1, + n00n_storage::sessions::StoredStateScope::Session, + ) + .unwrap(), + Some(&serde_json::json!({"value": "finished"})) + ); + assert_eq!(worker.join().unwrap().output.unwrap().as_text(), "done"); +} + +#[test] +fn plugin_state_lifecycle_methods_reject_dead_host() { + let reg = fresh_registry(); + let host = PluginHost::new(reg).unwrap(); + let handle = host.event_handle().unwrap(); + let identity = SessionIdentity::root(SessionRef::generate()); + drop(host); + + assert!(matches!( + handle.capture_state(&identity, 1), + Err(PluginError::HostDead) + )); + assert!(matches!( + handle.hydrate_state(&identity, None), + Err(PluginError::HostDead) + )); + assert!(matches!( + handle.reset_state(&identity), + Err(PluginError::HostDead) + )); + assert!(matches!( + handle.drop_state_owner(identity.session_id().id()), + Err(PluginError::HostDead) + )); +} + +#[test] +fn plugin_state_capture_services_nested_lua_tool_calls() { + let reg = fresh_registry(); + let host = PluginHost::new(Arc::clone(®)).unwrap(); + let source = format!( + r#" + n00n.api.register_tool({{ + name = "nested_state_writer", description = "test", schema = {MINIMAL_SCHEMA}, + header = function() return "nested writer" end, + handler = function(input, ctx) + local _, err = ctx:state_replace("session", {{ value = "nested" }}) + return err or "written" + end, + }}) + n00n.api.register_tool({{ + name = "nested_state_parent", description = "test", schema = {MINIMAL_SCHEMA}, + handler = function(input, ctx) + local buf = n00n.ui.buf() + buf:set_lines({{ "waiting" }}) + ctx:live_buf(buf) + local id = n00n.fn.jobstart("sleep 0.5") + n00n.fn.jobwait(id) + local result, err = n00n.agent.call_tool(ctx, "nested_state_writer", {{}}) + return err or result + end, + }}) + "# + ); + host.load_source("nested_state", &source).unwrap(); + let identity = SessionIdentity::root(SessionRef::generate()); + let entry = reg.get("nested_state_parent").unwrap(); + let invocation = entry.tool.parse(&serde_json::json!({})).unwrap(); + let (event_tx, event_rx) = flume::unbounded(); + let sender = n00n_agent::EventSender::new(event_tx, 0); + let mut ctx = n00n_agent::tools::test_support::stub_ctx_with( + &n00n_agent::AgentMode::Build, + Some(&sender), + Some("nested-state"), + ); + ctx.identity = Some(identity.clone()); + ctx.registry = Arc::clone(®); + let worker = std::thread::spawn(move || smol::block_on(invocation.execute(&ctx))); + + loop { + let event = event_rx.recv_timeout(Duration::from_secs(5)).unwrap(); + if matches!( + event.event, + n00n_agent::AgentEvent::LiveToolBuf { ref id, .. } if id == "nested-state" + ) { + break; + } + } + + let snapshot = host + .event_handle() + .unwrap() + .capture_state(&identity, 1) + .unwrap(); + assert_eq!( + snapshot + .plugin_payload_for_apply( + "nested_state", + 1, + n00n_storage::sessions::StoredStateScope::Session, + ) + .unwrap(), + Some(&serde_json::json!({"value": "nested"})) + ); + assert_eq!(worker.join().unwrap().output.unwrap().as_text(), "written"); +} + +#[test] +fn plugin_state_capture_services_lua_session_tool_calls() { + let reg = fresh_registry(); + let host = PluginHost::new(Arc::clone(®)).unwrap(); + let source = format!( + r#" + n00n.api.register_tool({{ + name = "session_state_writer", description = "test", schema = {MINIMAL_SCHEMA}, + audiences = {{ "main", "general_sub" }}, + handler = function(input, ctx) + local _, err = ctx:state_replace("root", {{ value = "session-nested" }}) + return err or "written" + end, + }}) + n00n.api.register_tool({{ + name = "session_state_parent", description = "test", schema = {MINIMAL_SCHEMA}, + audiences = {{ "main" }}, + handler = function(input, ctx) + local buf = n00n.ui.buf() + buf:set_lines({{ "waiting" }}) + ctx:live_buf(buf) + local id = n00n.fn.jobstart("sleep 0.5") + n00n.fn.jobwait(id) + local session, session_err = n00n.agent.session(ctx, {{}}) + if session_err then return session_err end + local result, prompt_err = session:prompt("write state") + session:close() + if prompt_err then return prompt_err end + local state, state_err = ctx:state_get("root") + if state_err then return state_err end + return state and state.value or result.text + end, + }}) + "# + ); + host.load_source("session_nested_state", &source).unwrap(); + let provider = ScriptedSessionProvider::new([ + Ok(session_response( + vec![ContentBlock::ToolUse { + id: "session-state-call".to_owned(), + name: "session_state_writer".to_owned(), + input: serde_json::json!({}), + }], + TokenUsage::default(), + StopReason::ToolUse, + )), + Ok(session_response( + vec![ContentBlock::Text { + text: "finished".to_owned(), + }], + TokenUsage::default(), + StopReason::EndTurn, + )), + ]); + let identity = SessionIdentity::root(SessionRef::generate()); + let entry = reg.get("session_state_parent").unwrap(); + let invocation = entry.tool.parse(&serde_json::json!({})).unwrap(); + let (event_tx, event_rx) = flume::unbounded(); + let sender = n00n_agent::EventSender::new(event_tx, 0); + let mut ctx = n00n_agent::tools::test_support::stub_ctx_with( + &n00n_agent::AgentMode::Build, + Some(&sender), + Some("session-nested-state"), + ); + ctx.identity = Some(identity.clone()); + ctx.registry = Arc::clone(®); + ctx.provider = Arc::new(provider); + ctx.model = Arc::new(Model::from_spec("anthropic/claude-opus-4-8").unwrap()); + let worker = std::thread::spawn(move || smol::block_on(invocation.execute(&ctx))); + + loop { + let event = event_rx.recv_timeout(Duration::from_secs(5)).unwrap(); + if matches!( + event.event, + n00n_agent::AgentEvent::LiveToolBuf { ref id, .. } if id == "session-nested-state" + ) { + break; + } + } + + let snapshot = host + .event_handle() + .unwrap() + .capture_state(&identity, 1) + .unwrap(); + assert_eq!( + worker.join().unwrap().output.unwrap().as_text(), + "session-nested" + ); + assert_eq!( + snapshot + .plugin_payload_for_apply( + "session_nested_state", + 1, + n00n_storage::sessions::StoredStateScope::Root, + ) + .unwrap(), + Some(&serde_json::json!({"value": "session-nested"})) + ); +} + +#[test] +fn plugin_state_isolates_namespaces_and_session_scope_while_sharing_root_scope() { + let reg = fresh_registry(); + let host = PluginHost::new(Arc::clone(®)).unwrap(); + let plugin_a = format!( + r#" + local function read(ctx, scope) + local value, err = ctx:state_get(scope) + if err then return err end + return value and value.name or "none" + end + n00n.api.register_tool({{ + name = "a_write_root", description = "test", schema = {MINIMAL_SCHEMA}, + handler = function(input, ctx) + local _, err = ctx:state_replace("root", {{ name = "root-a" }}) + return err or "ok" + end, + }}) + n00n.api.register_tool({{ + name = "a_write_session", description = "test", schema = {MINIMAL_SCHEMA}, + handler = function(input, ctx) + local _, err = ctx:state_replace("session", {{ name = "session-a" }}) + return err or "ok" + end, + }}) + n00n.api.register_tool({{ + name = "a_read_root", description = "test", schema = {MINIMAL_SCHEMA}, + handler = function(input, ctx) return read(ctx, "root") end, + }}) + n00n.api.register_tool({{ + name = "a_read_session", description = "test", schema = {MINIMAL_SCHEMA}, + handler = function(input, ctx) return read(ctx, "session") end, + }}) + "# + ); + let plugin_b = format!( + r#"n00n.api.register_tool({{ + name = "b_read_root", description = "test", schema = {MINIMAL_SCHEMA}, + handler = function(input, ctx) + local value, err = ctx:state_get("root") + if err then return err end + return value and value.name or "none" + end, + }})"# + ); + host.load_source("plugin_a", &plugin_a).unwrap(); + host.load_source("plugin_b", &plugin_b).unwrap(); + + let root = SessionIdentity::root(SessionRef::generate()); + let child = SessionIdentity::child(SessionRef::generate(), root.root_session_id().clone()); + let execute = |name: &str, identity: &SessionIdentity| { + let entry = reg.get(name).unwrap(); + let invocation = entry.tool.parse(&serde_json::json!({})).unwrap(); + let mut ctx = n00n_agent::tools::test_support::stub_ctx(&n00n_agent::AgentMode::Build); + ctx.identity = Some(identity.clone()); + let output = smol::block_on(async { invocation.execute(&ctx).await }) + .output + .unwrap(); + let n00n_agent::ToolOutput::Plain(output) = output else { + panic!("expected plain output"); + }; + output.text + }; + + assert_eq!(execute("a_write_root", &root), "ok"); + assert_eq!(execute("a_write_session", &root), "ok"); + assert_eq!(execute("a_read_root", &child), "root-a"); + assert_eq!(execute("a_read_session", &child), "none"); + assert_eq!(execute("b_read_root", &root), "none"); + + let handle = host.event_handle().unwrap(); + let captured = handle.capture_state(&root, 7).unwrap(); + assert_eq!(captured.state_revision(), Some(7)); + assert_eq!( + captured + .plugin_payload_for_apply( + "plugin_a", + 1, + n00n_storage::sessions::StoredStateScope::Root + ) + .unwrap(), + Some(&serde_json::json!({"name": "root-a"})) + ); + handle.reset_state(&root).unwrap(); + let reset = handle.capture_state(&root, 8).unwrap(); + assert!( + reset + .plugin_payload_for_apply( + "plugin_a", + 1, + n00n_storage::sessions::StoredStateScope::Root + ) + .unwrap() + .is_none() + ); + handle.hydrate_state(&root, Some(captured)).unwrap(); + assert_eq!(execute("a_read_root", &root), "root-a"); + assert_eq!(execute("a_read_session", &root), "session-a"); + + host.unload("plugin_a").unwrap(); + let unloaded = handle.capture_state(&root, 9).unwrap(); + assert_eq!( + unloaded + .plugin_payload_for_apply( + "plugin_a", + 1, + n00n_storage::sessions::StoredStateScope::Root, + ) + .unwrap(), + Some(&serde_json::json!({"name": "root-a"})) + ); + assert_eq!( + unloaded + .plugin_payload_for_apply( + "plugin_a", + 1, + n00n_storage::sessions::StoredStateScope::Session, + ) + .unwrap(), + Some(&serde_json::json!({"name": "session-a"})) + ); +} + +#[test] +fn plugin_state_rejects_context_reuse_after_handler_finishes() { + let reg = fresh_registry(); + let host = PluginHost::new(Arc::clone(®)).unwrap(); + let source = format!( + r#" + local saved + n00n.api.register_tool({{ + name = "save_state_ctx", description = "test", schema = {MINIMAL_SCHEMA}, + handler = function(input, ctx) + saved = ctx + local _, err = ctx:state_replace("session", {{ value = "original" }}) + return err or "saved" + end, + }}) + n00n.api.register_tool({{ + name = "save_dispatch_ctx", description = "test", schema = {MINIMAL_SCHEMA}, + handler = function(input, ctx) + saved = ctx + return "saved" + end, + }}) + n00n.api.register_tool({{ + name = "reuse_state_ctx", description = "test", schema = {MINIMAL_SCHEMA}, + handler = function() + local _, err = saved:state_replace("session", {{ value = "stale" }}) + return err or "unexpected success" + end, + }}) + n00n.api.register_tool({{ + name = "reuse_state_ctx_deadline", description = "test", schema = {MINIMAL_SCHEMA}, + handler = function() + local _, err = saved:set_deadline(1) + return err or "unexpected success" + end, + }}) + n00n.api.register_tool({{ + name = "reuse_state_ctx_dispatch", description = "test", schema = {MINIMAL_SCHEMA}, + handler = function() + local _, err = n00n.agent.call_tool(saved, "missing", {{}}) + return err or "unexpected success" + end, + }}) + "# + ); + host.load_source("stale_ctx", &source).unwrap(); + let identity = SessionIdentity::root(SessionRef::generate()); + let execute = |name: &str| { + let entry = reg.get(name).unwrap(); + let invocation = entry.tool.parse(&serde_json::json!({})).unwrap(); + let mut ctx = n00n_agent::tools::test_support::stub_ctx(&n00n_agent::AgentMode::Build); + ctx.identity = Some(identity.clone()); + let output = smol::block_on(async { invocation.execute(&ctx).await }) + .output + .unwrap(); + let n00n_agent::ToolOutput::Plain(output) = output else { + panic!("expected plain output"); + }; + output.text + }; + + assert_eq!(execute("save_state_ctx"), "saved"); + assert_eq!(execute("reuse_state_ctx"), STALE_CTX_ERR); + assert_eq!(execute("reuse_state_ctx_dispatch"), STALE_CTX_ERR); + assert_eq!(execute("reuse_state_ctx_deadline"), STALE_CTX_ERR); + + let execute_without_identity = |name: &str| { + let entry = reg.get(name).unwrap(); + let invocation = entry.tool.parse(&serde_json::json!({})).unwrap(); + let ctx = n00n_agent::tools::test_support::stub_ctx(&n00n_agent::AgentMode::Build); + let output = smol::block_on(async { invocation.execute(&ctx).await }) + .output + .unwrap(); + let n00n_agent::ToolOutput::Plain(output) = output else { + panic!("expected plain output"); + }; + output.text + }; + assert_eq!(execute_without_identity("save_dispatch_ctx"), "saved"); + assert_eq!( + execute_without_identity("reuse_state_ctx_dispatch"), + STALE_CTX_ERR + ); + + let snapshot = host + .event_handle() + .unwrap() + .capture_state(&identity, 1) + .unwrap(); + assert_eq!( + snapshot + .plugin_payload_for_apply( + "stale_ctx", + 1, + n00n_storage::sessions::StoredStateScope::Session, + ) + .unwrap(), + Some(&serde_json::json!({"value": "original"})) + ); +} #[test] fn lua_sessions_under_one_parent_use_unique_identity_everywhere() { diff --git a/n00n-storage/src/sessions.rs b/n00n-storage/src/sessions.rs index 61cbda1dc..77e9e72f2 100644 --- a/n00n-storage/src/sessions.rs +++ b/n00n-storage/src/sessions.rs @@ -51,7 +51,7 @@ const MAX_FIRST_MESSAGE_BYTES: usize = 256 * 1024; pub const SESSION_STATE_SCHEMA_VERSION: u32 = 1; const MAX_PLUGIN_STATE_ENTRIES: usize = 64; const MAX_PLUGIN_STATE_NAME_BYTES: usize = 128; -const MAX_PLUGIN_STATE_BYTES: usize = 256 * 1024; +pub const MAX_PLUGIN_STATE_BYTES: usize = 256 * 1024; const MAX_SESSION_STATE_BYTES: usize = 1024 * 1024; const META_RECORD_PREFIX: &str = r#"{"t":"meta""#; const MSG_RECORD_PREFIX: &str = r#"{"t":"msg""#; @@ -277,11 +277,38 @@ impl Default for StoredPluginScopes { } impl StoredPluginScopes { - fn insert(&mut self, scope: StoredStateScope, state: StoredPluginState) { - self.scopes - .get_or_insert_with(BTreeMap::new) - .insert(scope.as_str().to_owned(), state); - self.malformed = None; + fn set( + &mut self, + plugin: &str, + scope: StoredStateScope, + schema_version: u32, + payload: serde_json::Value, + ) -> Result<(), SessionStateError> { + let Some(scopes) = self.scopes.as_mut() else { + return Err(SessionStateError::InvalidPluginContainer { + plugin: plugin.to_owned(), + }); + }; + if let Some(state) = scopes.get_mut(scope.as_str()) { + let serde_json::Value::Object(fields) = &mut state.raw else { + return Err(SessionStateError::InvalidPluginState { + plugin: plugin.to_owned(), + scope, + }); + }; + fields.insert( + "schema_version".to_owned(), + serde_json::Value::from(schema_version), + ); + fields.insert("payload".to_owned(), payload); + state.schema_version = Some(u64::from(schema_version)); + } else { + scopes.insert( + scope.as_str().to_owned(), + StoredPluginState::new(schema_version, payload), + ); + } + Ok(()) } } @@ -338,6 +365,9 @@ enum StoredSessionStateSnapshotInner { schema_version: u64, raw: serde_json::Value, }, + Malformed { + raw: serde_json::Value, + }, } #[derive(Debug, Clone, PartialEq)] @@ -345,6 +375,13 @@ pub struct StoredSessionStateSnapshot { inner: StoredSessionStateSnapshotInner, } +#[derive(Debug, Clone, PartialEq)] +pub struct StoredPluginStateEntry<'a> { + pub plugin: &'a str, + pub scope: StoredStateScope, + pub payload: &'a serde_json::Value, +} + impl Default for StoredSessionStateSnapshot { fn default() -> Self { Self::new(0) @@ -358,7 +395,8 @@ impl Serialize for StoredSessionStateSnapshot { { match &self.inner { StoredSessionStateSnapshotInner::Supported(snapshot) => snapshot.serialize(serializer), - StoredSessionStateSnapshotInner::Unsupported { raw, .. } => raw.serialize(serializer), + StoredSessionStateSnapshotInner::Unsupported { raw, .. } + | StoredSessionStateSnapshotInner::Malformed { raw } => raw.serialize(serializer), } } } @@ -405,6 +443,8 @@ impl<'de> Deserialize<'de> for StoredSessionStateSnapshot { pub enum SessionStateError { #[error("unsupported session-state schema version {found} (expected {expected})")] UnsupportedSchemaVersion { found: u64, expected: u32 }, + #[error("session-state snapshot envelope is malformed")] + InvalidEnvelope, #[error("session state has {found} plugin entries (maximum {maximum})")] TooManyPlugins { found: usize, maximum: usize }, #[error("invalid session-state plugin name {plugin:?}")] @@ -434,12 +474,8 @@ pub enum SessionStateError { found: u64, expected: u32, }, - #[error("state scope mismatch for plugin {plugin:?}: found {found:?}, expected {expected:?}")] - ScopeMismatch { - plugin: String, - found: StoredStateScope, - expected: StoredStateScope, - }, + #[error("session-state revision cannot regress from {current} to {requested}")] + StateRevisionRegression { current: u64, requested: u64 }, #[error("failed to measure serialized session state: {0}")] Serialize(#[from] serde_json::Error), } @@ -461,15 +497,43 @@ impl StoredSessionStateSnapshot { pub fn state_revision(&self) -> Option { match &self.inner { StoredSessionStateSnapshotInner::Supported(snapshot) => Some(snapshot.state_revision), - StoredSessionStateSnapshotInner::Unsupported { .. } => None, + StoredSessionStateSnapshotInner::Unsupported { .. } + | StoredSessionStateSnapshotInner::Malformed { .. } => None, } } - /// Adds or replaces one scope of a plugin's state after enforcing snapshot bounds. + /// Advances the state revision without allowing regression. /// /// # Errors - /// Returns a typed error for invalid names or exceeded bounds. - pub fn insert_plugin_state( + /// Returns a typed error for unsupported envelopes, revision regression, or exceeded bounds. + pub fn set_state_revision(&mut self, state_revision: u64) -> Result<(), SessionStateError> { + let StoredSessionStateSnapshotInner::Supported(snapshot) = &self.inner else { + return Err(self.unsupported_schema_error()); + }; + if state_revision < snapshot.state_revision { + return Err(SessionStateError::StateRevisionRegression { + current: snapshot.state_revision, + requested: state_revision, + }); + } + let current_size = serde_json::to_vec(snapshot).map(|v| v.len()).ok(); + let mut candidate = snapshot.clone(); + candidate.state_revision = state_revision; + let new_size = serde_json::to_vec(&candidate).map(|v| v.len()).ok(); + if new_size.is_some_and(|new| current_size.is_some_and(|cur| new <= cur)) { + self.inner = StoredSessionStateSnapshotInner::Supported(candidate); + return Ok(()); + } + validate_supported_snapshot(&candidate)?; + self.inner = StoredSessionStateSnapshotInner::Supported(candidate); + Ok(()) + } + + /// Adds or replaces one exact plugin scope after enforcing snapshot bounds. + /// + /// # Errors + /// Returns a typed error for unsupported envelopes, invalid names, or exceeded bounds. + pub fn set_plugin_state( &mut self, plugin: &str, schema_version: u32, @@ -480,16 +544,118 @@ impl StoredSessionStateSnapshot { return Err(self.unsupported_schema_error()); }; let mut candidate = snapshot.clone(); + validate_plugin_name(plugin)?; candidate .plugins .entry(plugin.to_owned()) .or_default() - .insert(scope, StoredPluginState::new(schema_version, payload)); + .set(plugin, scope, schema_version, payload)?; + validate_supported_snapshot(&candidate)?; + self.inner = StoredSessionStateSnapshotInner::Supported(candidate); + Ok(()) + } + + /// Removes one exact plugin scope while preserving every sibling entry. + /// + /// # Errors + /// Returns a typed error for unsupported envelopes, invalid names, or exceeded bounds. + pub fn remove_plugin_state( + &mut self, + plugin: &str, + scope: StoredStateScope, + ) -> Result<(), SessionStateError> { + let StoredSessionStateSnapshotInner::Supported(snapshot) = &self.inner else { + return Err(self.unsupported_schema_error()); + }; + validate_plugin_name(plugin)?; + let mut candidate = snapshot.clone(); + if candidate + .plugins + .get(plugin) + .is_some_and(|stored_scopes| stored_scopes.scopes.is_none()) + { + return Err(SessionStateError::InvalidPluginContainer { + plugin: plugin.to_owned(), + }); + } + let remove_plugin = candidate + .plugins + .get_mut(plugin) + .and_then(|stored_scopes| stored_scopes.scopes.as_mut()) + .is_some_and(|scopes| { + scopes.remove(scope.as_str()); + scopes.is_empty() + }); + if remove_plugin { + candidate.plugins.remove(plugin); + } validate_supported_snapshot(&candidate)?; self.inner = StoredSessionStateSnapshotInner::Supported(candidate); Ok(()) } + /// Returns plugin names that contain the exact stored scope, including opaque state versions. + /// + /// # Errors + /// Returns a typed error when the snapshot envelope version is unsupported. + pub fn plugin_names_with_scope( + &self, + scope: StoredStateScope, + ) -> Result, SessionStateError> { + let StoredSessionStateSnapshotInner::Supported(snapshot) = &self.inner else { + return Err(self.unsupported_schema_error()); + }; + Ok(snapshot + .plugins + .iter() + .filter_map(|(plugin, stored_scopes)| { + stored_scopes + .scopes + .as_ref() + .is_some_and(|scopes| scopes.contains_key(scope.as_str())) + .then_some(plugin.as_str()) + }) + .collect()) + } + + /// Enumerates well-formed entries matching a supported plugin state version. + /// + /// Malformed entries, unknown scopes, and other plugin state versions remain preserved but are + /// omitted from the result. + /// + /// # Errors + /// Returns a typed error when the snapshot envelope version is unsupported. + pub fn plugin_entries_for_apply( + &self, + schema_version: u32, + ) -> Result>, SessionStateError> { + let StoredSessionStateSnapshotInner::Supported(snapshot) = &self.inner else { + return Err(self.unsupported_schema_error()); + }; + let mut entries = Vec::new(); + for (plugin, stored_scopes) in &snapshot.plugins { + let Some(scopes) = stored_scopes.scopes.as_ref() else { + continue; + }; + for (scope_name, state) in scopes { + let Some(scope) = StoredStateScope::from_stored(scope_name) else { + continue; + }; + let (Some(found), Some(payload)) = (state.schema_version, state.payload()) else { + continue; + }; + if found == u64::from(schema_version) { + entries.push(StoredPluginStateEntry { + plugin, + scope, + payload, + }); + } + } + } + Ok(entries) + } + /// Validates persisted state before any plugin can apply it. /// /// These limits bound state accepted by plugins. The session reader separately caps each @@ -502,7 +668,8 @@ impl StoredSessionStateSnapshot { StoredSessionStateSnapshotInner::Supported(snapshot) => { validate_supported_snapshot(snapshot) } - StoredSessionStateSnapshotInner::Unsupported { .. } => { + StoredSessionStateSnapshotInner::Unsupported { .. } + | StoredSessionStateSnapshotInner::Malformed { .. } => { Err(self.unsupported_schema_error()) } } @@ -531,29 +698,10 @@ impl StoredSessionStateSnapshot { }); }; let Some(state) = scopes.get(scope.as_str()) else { - let Some(stored_scope) = scopes.keys().next() else { - return Err(SessionStateError::MissingPluginState { - plugin: plugin.to_owned(), - }); - }; - let Some(found) = StoredStateScope::from_stored(stored_scope) else { - return Err(SessionStateError::InvalidStoredScope { - plugin: plugin.to_owned(), - scope: stored_scope.clone(), - }); - }; - return Err(SessionStateError::ScopeMismatch { - plugin: plugin.to_owned(), - found, - expected: scope, - }); + return Ok(None); }; + validate_stored_plugin_state_size(plugin, state)?; let payload = state.payload(); - let measured = match payload { - Some(payload) => payload, - None => &state.raw, - }; - validate_plugin_state_size(plugin, measured)?; let (Some(found), Some(payload)) = (state.schema_version, payload) else { return Err(SessionStateError::InvalidPluginState { plugin: plugin.to_owned(), @@ -576,6 +724,9 @@ impl StoredSessionStateSnapshot { u64::from(snapshot.schema_version) } StoredSessionStateSnapshotInner::Unsupported { schema_version, .. } => *schema_version, + StoredSessionStateSnapshotInner::Malformed { .. } => { + return SessionStateError::InvalidEnvelope; + } }; SessionStateError::UnsupportedSchemaVersion { found, @@ -613,11 +764,7 @@ fn validate_supported_snapshot( continue; }; for state in scopes.values() { - let measured = match state.payload() { - Some(payload) => payload, - None => &state.raw, - }; - validate_plugin_state_size(plugin, measured)?; + validate_stored_plugin_state_size(plugin, state)?; } } let bytes = serde_json::to_vec(snapshot)?.len(); @@ -648,7 +795,37 @@ fn validate_plugin_state_size( plugin: &str, value: &serde_json::Value, ) -> Result<(), SessionStateError> { - let bytes = serde_json::to_vec(value)?.len(); + validate_plugin_state_bytes(plugin, serde_json::to_vec(value)?.len()) +} + +fn validate_stored_plugin_state_size( + plugin: &str, + state: &StoredPluginState, +) -> Result<(), SessionStateError> { + let serde_json::Value::Object(fields) = &state.raw else { + return validate_plugin_state_size(plugin, &state.raw); + }; + let Some(payload) = fields.get("payload") else { + return validate_plugin_state_size(plugin, &state.raw); + }; + let payload_bytes = serde_json::to_vec(payload)?.len(); + let opaque = fields + .iter() + .filter(|(name, _)| { + name.as_str() != "payload" + && (name.as_str() != "schema_version" || state.schema_version.is_none()) + }) + .map(|(name, value)| (name.clone(), value.clone())) + .collect::>(); + let opaque_bytes = if opaque.is_empty() { + 0 + } else { + serde_json::to_vec(&opaque)?.len() + }; + validate_plugin_state_bytes(plugin, payload_bytes.saturating_add(opaque_bytes)) +} + +fn validate_plugin_state_bytes(plugin: &str, bytes: usize) -> Result<(), SessionStateError> { if bytes > MAX_PLUGIN_STATE_BYTES { return Err(SessionStateError::PluginStateTooLarge { plugin: plugin.to_owned(), @@ -669,14 +846,27 @@ where let Some(raw) = raw else { return Ok(None); }; - match serde_json::from_value(raw) { + let bytes = serde_json::to_vec(&raw) + .map_err(serde::de::Error::custom)? + .len(); + if bytes > MAX_SESSION_STATE_BYTES { + return Err(serde::de::Error::custom( + SessionStateError::SnapshotTooLarge { + bytes, + maximum: MAX_SESSION_STATE_BYTES, + }, + )); + } + match serde_json::from_value(raw.clone()) { Ok(snapshot) => Ok(Some(snapshot)), Err(error) => { warn!( category = ?error.classify(), - "ignoring malformed session-state snapshot" + "quarantining malformed session-state snapshot" ); - Ok(None) + Ok(Some(StoredSessionStateSnapshot { + inner: StoredSessionStateSnapshotInner::Malformed { raw }, + })) } } } @@ -4565,7 +4755,7 @@ mod tests { fn session_state_snapshot_round_trips_unknown_plugin_entries() { let mut snapshot = super::StoredSessionStateSnapshot::new(7); snapshot - .insert_plugin_state( + .set_plugin_state( "future_plugin", 99, super::StoredStateScope::Root, @@ -4668,7 +4858,7 @@ mod tests { let mut session: TestSession = Session::new("model", "/project"); let mut snapshot = super::StoredSessionStateSnapshot::new(4); snapshot - .insert_plugin_state( + .set_plugin_state( "todo_write", 1, super::StoredStateScope::Root, @@ -4686,7 +4876,7 @@ mod tests { fn session_state_snapshot_supports_both_scopes_for_one_plugin() { let mut snapshot = super::StoredSessionStateSnapshot::new(1); snapshot - .insert_plugin_state( + .set_plugin_state( "plugin", 1, super::StoredStateScope::Root, @@ -4694,7 +4884,7 @@ mod tests { ) .unwrap(); snapshot - .insert_plugin_state( + .set_plugin_state( "plugin", 2, super::StoredStateScope::Session, @@ -4717,36 +4907,270 @@ mod tests { } #[test] - fn session_state_snapshot_rejects_scope_mismatch() { + fn session_state_snapshot_absent_requested_scope_is_none() { let mut snapshot = super::StoredSessionStateSnapshot::new(1); snapshot - .insert_plugin_state( + .set_plugin_state( "todo_write", 1, super::StoredStateScope::Root, serde_json::json!({ "todos": [] }), ) .unwrap(); + assert_eq!( + snapshot + .plugin_payload_for_apply("todo_write", 1, super::StoredStateScope::Session) + .unwrap(), + None + ); + } + + #[test] + fn session_state_snapshot_enumerates_only_valid_supported_entries() { + let snapshot: super::StoredSessionStateSnapshot = + serde_json::from_value(serde_json::json!({ + "schema_version": 1, + "state_revision": 3, + "plugins": { + "plugin": { + "root": { "schema_version": 1, "payload": { "valid": true } }, + "session": { "schema_version": 2, "payload": { "future": true } }, + "future_scope": { "schema_version": 1, "payload": { "opaque": true } } + }, + "malformed": { + "root": { "schema_version": "invalid", "payload": null } + }, + "malformed_container": null + } + })) + .unwrap(); + + let entries = snapshot.plugin_entries_for_apply(1).unwrap(); + assert_eq!(entries.len(), 1); + assert_eq!(entries[0].plugin, "plugin"); + assert_eq!(entries[0].scope, super::StoredStateScope::Root); + assert_eq!(entries[0].payload, &serde_json::json!({ "valid": true })); + } + + #[test] + fn session_state_snapshot_mutations_preserve_opaque_data() { + let raw = serde_json::json!({ + "schema_version": 1, + "state_revision": 3, + "future_envelope": { "kept": true }, + "plugins": { + "target": { + "root": { "schema_version": 1, "payload": "old", "state_extra": 7 }, + "future_scope": { "opaque": true }, + "session": { "malformed": true } + }, + "malformed_sibling": null, + "unknown_sibling": { "other": [1, 2, 3] } + } + }); + let mut snapshot: super::StoredSessionStateSnapshot = + serde_json::from_value(raw.clone()).unwrap(); + + snapshot.set_state_revision(4).unwrap(); + snapshot + .set_plugin_state( + "target", + 1, + super::StoredStateScope::Root, + serde_json::json!("new"), + ) + .unwrap(); + let after_set = serde_json::to_value(&snapshot).unwrap(); + assert_eq!(after_set["state_revision"], 4); + assert_eq!( + after_set["plugins"]["target"]["root"]["state_extra"], + raw["plugins"]["target"]["root"]["state_extra"] + ); + assert_eq!(after_set["future_envelope"], raw["future_envelope"]); + assert_eq!( + after_set["plugins"]["target"]["future_scope"], + raw["plugins"]["target"]["future_scope"] + ); + assert_eq!( + after_set["plugins"]["target"]["session"], + raw["plugins"]["target"]["session"] + ); + assert_eq!( + after_set["plugins"]["malformed_sibling"], + raw["plugins"]["malformed_sibling"] + ); + assert_eq!( + after_set["plugins"]["unknown_sibling"], + raw["plugins"]["unknown_sibling"] + ); + + snapshot + .remove_plugin_state("target", super::StoredStateScope::Root) + .unwrap(); + let after_remove = serde_json::to_value(snapshot).unwrap(); + assert!(after_remove["plugins"]["target"].get("root").is_none()); + assert_eq!( + after_remove["plugins"]["target"]["future_scope"], + raw["plugins"]["target"]["future_scope"] + ); + assert_eq!( + after_remove["plugins"]["target"]["session"], + raw["plugins"]["target"]["session"] + ); + } + + #[test] + fn session_state_snapshot_rejects_mutating_malformed_plugin_container() { + let raw = serde_json::json!({ + "schema_version": 1, + "state_revision": 3, + "plugins": {"target": null} + }); + let mut snapshot: super::StoredSessionStateSnapshot = + serde_json::from_value(raw.clone()).unwrap(); + + assert!(matches!( + snapshot.set_plugin_state( + "target", + 1, + super::StoredStateScope::Root, + serde_json::json!({"new": true}), + ), + Err(super::SessionStateError::InvalidPluginContainer { .. }) + )); + assert_eq!(serde_json::to_value(snapshot).unwrap(), raw); + } + + #[test] + fn session_state_snapshot_counts_opaque_state_fields_toward_entry_limit() { + let raw = serde_json::json!({ + "schema_version": 1, + "state_revision": 3, + "plugins": { + "target": { + "root": { + "schema_version": 1, + "payload": null, + "opaque": "x".repeat(super::MAX_PLUGIN_STATE_BYTES) + } + } + } + }); + + assert!(serde_json::from_value::(raw).is_err()); + } + + #[test] + fn session_state_snapshot_counts_malformed_version_as_opaque_state() { + let raw = serde_json::json!({ + "schema_version": 1, + "state_revision": 3, + "plugins": { + "target": { + "root": { + "schema_version": "x".repeat(super::MAX_PLUGIN_STATE_BYTES), + "payload": null + } + } + } + }); + + assert!(serde_json::from_value::(raw).is_err()); + } + + #[test] + fn session_state_snapshot_mutations_are_atomic() { + let mut snapshot = super::StoredSessionStateSnapshot::new(5); + snapshot + .set_plugin_state( + "kept", + 1, + super::StoredStateScope::Root, + serde_json::json!({ "value": true }), + ) + .unwrap(); + let before = serde_json::to_value(&snapshot).unwrap(); + assert!(matches!( - snapshot.plugin_payload_for_apply("todo_write", 1, super::StoredStateScope::Session), - Err(super::SessionStateError::ScopeMismatch { .. }) + snapshot.set_state_revision(4), + Err(super::SessionStateError::StateRevisionRegression { .. }) )); + assert_eq!(serde_json::to_value(&snapshot).unwrap(), before); + + assert!(matches!( + snapshot.set_plugin_state( + "oversized", + 1, + super::StoredStateScope::Session, + Value::String("x".repeat(super::MAX_PLUGIN_STATE_BYTES)), + ), + Err(super::SessionStateError::PluginStateTooLarge { .. }) + )); + assert_eq!(serde_json::to_value(&snapshot).unwrap(), before); + + assert!(matches!( + snapshot.remove_plugin_state("bad/name", super::StoredStateScope::Root), + Err(super::SessionStateError::InvalidPluginName { .. }) + )); + assert_eq!(serde_json::to_value(snapshot).unwrap(), before); + } + + #[test] + fn session_state_snapshot_unsupported_envelope_mutations_fail_unchanged() { + let raw = serde_json::json!({ + "schema_version": 2, + "state_revision": 9, + "opaque": { "kept": true } + }); + let mut snapshot: super::StoredSessionStateSnapshot = + serde_json::from_value(raw.clone()).unwrap(); + + assert!(matches!( + snapshot.plugin_entries_for_apply(1), + Err(super::SessionStateError::UnsupportedSchemaVersion { found: 2, .. }) + )); + assert!(matches!( + snapshot.set_state_revision(10), + Err(super::SessionStateError::UnsupportedSchemaVersion { found: 2, .. }) + )); + assert!(matches!( + snapshot.set_plugin_state( + "plugin", + 1, + super::StoredStateScope::Root, + serde_json::json!(null), + ), + Err(super::SessionStateError::UnsupportedSchemaVersion { found: 2, .. }) + )); + assert!(matches!( + snapshot.remove_plugin_state("plugin", super::StoredStateScope::Root), + Err(super::SessionStateError::UnsupportedSchemaVersion { found: 2, .. }) + )); + assert_eq!(serde_json::to_value(snapshot).unwrap(), raw); } #[test] fn malformed_session_state_does_not_invalidate_session_meta() { - let json = r#"{ - "mode":"plan", - "plan_written":true, - "state_snapshot":{ - "schema_version":1, - "plugins":{"todo_write":{"root":{"schema_version":1}}} + let raw = serde_json::json!({ + "mode": "plan", + "plan_written": true, + "state_snapshot": { + "schema_version": 1, + "plugins": {"todo_write": {"root": {"schema_version": 1}}} } - }"#; - let meta: super::SessionMeta = serde_json::from_str(json).unwrap(); + }); + let meta: super::SessionMeta = serde_json::from_value(raw.clone()).unwrap(); assert_eq!(meta.mode, Some(super::StoredMode::Plan)); assert!(meta.plan_written); - assert!(meta.state_snapshot.is_none()); + let snapshot = meta.state_snapshot.as_ref().unwrap(); + assert!(matches!( + snapshot.validate_for_apply(), + Err(super::SessionStateError::InvalidEnvelope) + )); + assert_eq!( + serde_json::to_value(meta).unwrap()["state_snapshot"], + raw["state_snapshot"] + ); } #[test] @@ -4767,13 +5191,15 @@ mod tests { } #[test] - fn session_state_snapshot_quarantines_missing_plugin_scope() { + fn session_state_snapshot_empty_plugin_scope_is_absent() { let json = r#"{"schema_version":1,"state_revision":1,"plugins":{"p":{}}}"#; let snapshot: super::StoredSessionStateSnapshot = serde_json::from_str(json).unwrap(); - assert!(matches!( - snapshot.plugin_payload_for_apply("p", 1, super::StoredStateScope::Root), - Err(super::SessionStateError::MissingPluginState { .. }) - )); + assert_eq!( + snapshot + .plugin_payload_for_apply("p", 1, super::StoredStateScope::Root) + .unwrap(), + None + ); } #[test] @@ -4781,7 +5207,7 @@ mod tests { let mut snapshot = super::StoredSessionStateSnapshot::new(1); let maximum_name = "x".repeat(super::MAX_PLUGIN_STATE_NAME_BYTES); snapshot - .insert_plugin_state( + .set_plugin_state( &maximum_name, 1, super::StoredStateScope::Root, @@ -4789,7 +5215,7 @@ mod tests { ) .unwrap(); assert!(matches!( - snapshot.insert_plugin_state( + snapshot.set_plugin_state( &"x".repeat(super::MAX_PLUGIN_STATE_NAME_BYTES + 1), 1, super::StoredStateScope::Root, @@ -4801,7 +5227,7 @@ mod tests { let mut entries = super::StoredSessionStateSnapshot::new(1); for index in 0..super::MAX_PLUGIN_STATE_ENTRIES { entries - .insert_plugin_state( + .set_plugin_state( &format!("plugin_{index}"), 1, super::StoredStateScope::Root, @@ -4810,7 +5236,7 @@ mod tests { .unwrap(); } assert!(matches!( - entries.insert_plugin_state( + entries.set_plugin_state( "one_too_many", 1, super::StoredStateScope::Root, @@ -4828,7 +5254,7 @@ mod tests { fn session_state_snapshot_rejects_unsafe_plugin_names(plugin: &str) { let mut snapshot = super::StoredSessionStateSnapshot::new(1); assert!(matches!( - snapshot.insert_plugin_state( + snapshot.set_plugin_state( plugin, 1, super::StoredStateScope::Root, @@ -4852,7 +5278,7 @@ mod tests { })) .unwrap(); snapshot - .insert_plugin_state( + .set_plugin_state( "usable", 1, super::StoredStateScope::Root, @@ -4871,7 +5297,7 @@ mod tests { fn session_state_snapshot_enforces_payload_and_aggregate_boundaries() { let mut payload = super::StoredSessionStateSnapshot::new(1); payload - .insert_plugin_state( + .set_plugin_state( "maximum", 1, super::StoredStateScope::Root, @@ -4879,7 +5305,7 @@ mod tests { ) .unwrap(); assert!(matches!( - payload.insert_plugin_state( + payload.set_plugin_state( "oversized", 1, super::StoredStateScope::Root, @@ -4891,7 +5317,7 @@ mod tests { let mut aggregate = super::StoredSessionStateSnapshot::new(1); for index in 0..4 { aggregate - .insert_plugin_state( + .set_plugin_state( &format!("large_{index}"), 1, super::StoredStateScope::Root, @@ -4900,7 +5326,7 @@ mod tests { .unwrap(); } assert!(matches!( - aggregate.insert_plugin_state( + aggregate.set_plugin_state( "aggregate_overflow", 1, super::StoredStateScope::Root, diff --git a/n00n-ui/src/agent/agent_loop.rs b/n00n-ui/src/agent/agent_loop.rs index b5b0f089e..022f5cccc 100644 --- a/n00n-ui/src/agent/agent_loop.rs +++ b/n00n-ui/src/agent/agent_loop.rs @@ -9,7 +9,8 @@ use n00n_agent::permissions::PermissionManager; use n00n_agent::template; use n00n_agent::template::Vars; use n00n_agent::tools::{ - DescriptionContext, FileReadTracker, ToolAudience, ToolFilter, ToolRegistry, ToolsSnapshot, + DescriptionContext, FileReadTracker, SessionIdentity, ToolAudience, ToolFilter, ToolRegistry, + ToolsSnapshot, }; use n00n_agent::{ Agent, AgentConfig, AgentEvent, AgentInput, AgentParams, AgentRunParams, CancelMap, @@ -55,7 +56,7 @@ pub(super) struct AgentLoop { agent_tx: flume::Sender, answer_rx: Arc>>, queue: Arc, - session_id: Option, + identity: Option, timeouts: n00n_providers::Timeouts, openai_options: OpenAiOptions, lua_handle: Option, @@ -144,7 +145,7 @@ impl AgentLoop { agent_tx, answer_rx: Arc::new(async_lock::Mutex::new(answer_rx)), queue, - session_id, + identity: session_id.map(SessionIdentity::root), timeouts, openai_options, lua_handle, @@ -342,7 +343,7 @@ impl AgentLoop { config: Arc::new(self.config.clone()), tool_output_lines: self.tool_output_lines, permissions: Arc::clone(&self.permissions), - session_id: self.session_id.clone(), + identity: self.identity.clone(), timeouts: self.timeouts, file_tracker: Arc::clone(&self.file_tracker), prompt_slots: Arc::new(prompt_slots),