diff --git a/src/app.rs b/src/app.rs index 94d949be..4e1d506f 100644 --- a/src/app.rs +++ b/src/app.rs @@ -327,7 +327,7 @@ impl AppBuilder { .with_search_config(&self.config.search); if let Some(ref emb) = embeddings { - ws = ws.with_embeddings_cached(emb.clone(), emb_cache_config); + ws = ws.with_embeddings_cached(emb.clone(), emb_cache_config.clone()); } // Wire workspace-level settings (read scopes, memory layers) @@ -341,7 +341,34 @@ impl AppBuilder { } ws = ws.with_memory_layers(self.config.workspace.memory_layers.clone()); let ws = Arc::new(ws); - tools.register_memory_tools(Arc::clone(&ws)); + + // Detect multi-tenant mode: when GATEWAY_USER_TOKENS is configured, + // each authenticated user needs their own workspace scope. Use + // PerUserWorkspaceResolver to create per-user workspaces on demand + // instead of sharing the startup workspace across all users. + let is_multi_tenant = self + .config + .channels + .gateway + .as_ref() + .is_some_and(|gw| gw.user_tokens.is_some()); + + if is_multi_tenant { + let resolver = Arc::new( + crate::tools::builtin::memory::PerUserWorkspaceResolver::new( + self.db.as_ref().unwrap().clone(), + embeddings.clone(), + emb_cache_config, + self.config.search.clone(), + self.config.workspace.clone(), + ), + ); + tools.register_memory_tools_with_resolver(resolver); + tracing::info!("Memory tools configured with per-user workspace resolver (multi-tenant mode)"); + } else { + tools.register_memory_tools(Arc::clone(&ws)); + } + Some(ws) } else { None diff --git a/src/extensions/manager.rs b/src/extensions/manager.rs index c61c5a2d..3c1ebbdb 100644 --- a/src/extensions/manager.rs +++ b/src/extensions/manager.rs @@ -3403,9 +3403,10 @@ impl ExtensionManager { // No auth section — setup_is_complete was already checked above, // so if we reach here the setup requirements are satisfied. - if cap_file.setup.is_none() { - return ToolAuthState::NoAuth; - } + let setup = match &cap_file.setup { + Some(s) => s, + None => return ToolAuthState::NoAuth, + }; let all_provided = futures::future::join_all( setup diff --git a/src/tools/builtin/memory.rs b/src/tools/builtin/memory.rs index edbc4f1c..5dd777c9 100644 --- a/src/tools/builtin/memory.rs +++ b/src/tools/builtin/memory.rs @@ -12,15 +12,119 @@ //! Use `memory_write` to persist important facts that should be remembered //! across sessions. +use std::collections::HashMap; use std::path::Path; use std::sync::Arc; use async_trait::async_trait; +use tokio::sync::RwLock; use crate::context::JobContext; use crate::tools::tool::{Tool, ToolError, ToolOutput, require_str}; use crate::workspace::{Workspace, paths}; +// ── WorkspaceResolver ────────────────────────────────────────────── + +/// Resolves a workspace for a given user ID. +/// +/// In single-user mode, always returns the same workspace. +/// In multi-tenant mode, creates per-user workspaces on demand. +#[async_trait] +pub trait WorkspaceResolver: Send + Sync { + async fn resolve(&self, user_id: &str) -> Arc; +} + +/// Returns a fixed workspace regardless of user ID (single-user mode). +pub struct FixedWorkspaceResolver { + workspace: Arc, +} + +impl FixedWorkspaceResolver { + pub fn new(workspace: Arc) -> Self { + Self { workspace } + } +} + +#[async_trait] +impl WorkspaceResolver for FixedWorkspaceResolver { + async fn resolve(&self, _user_id: &str) -> Arc { + Arc::clone(&self.workspace) + } +} + +/// Creates per-user workspaces on demand, caching them for reuse. +/// +/// Used in multi-tenant mode where each authenticated user gets their own +/// workspace scope. The workspace is constructed with the same configuration +/// (embeddings, search config, memory layers) as the startup workspace. +pub struct PerUserWorkspaceResolver { + db: Arc, + embeddings: Option>, + embedding_cache_config: crate::workspace::EmbeddingCacheConfig, + search_config: crate::config::WorkspaceSearchConfig, + workspace_config: crate::config::WorkspaceConfig, + cache: RwLock>>, +} + +impl PerUserWorkspaceResolver { + pub fn new( + db: Arc, + embeddings: Option>, + embedding_cache_config: crate::workspace::EmbeddingCacheConfig, + search_config: crate::config::WorkspaceSearchConfig, + workspace_config: crate::config::WorkspaceConfig, + ) -> Self { + Self { + db, + embeddings, + embedding_cache_config, + search_config, + workspace_config, + cache: RwLock::new(HashMap::new()), + } + } + + fn build_workspace(&self, user_id: &str) -> Arc { + let mut ws = Workspace::new_with_db(user_id, Arc::clone(&self.db)) + .with_search_config(&self.search_config); + + if let Some(ref emb) = self.embeddings { + ws = ws.with_embeddings_cached(Arc::clone(emb), self.embedding_cache_config.clone()); + } + + if !self.workspace_config.read_scopes.is_empty() { + ws = ws.with_additional_read_scopes(self.workspace_config.read_scopes.clone()); + } + ws = ws.with_memory_layers(self.workspace_config.memory_layers.clone()); + + Arc::new(ws) + } +} + +#[async_trait] +impl WorkspaceResolver for PerUserWorkspaceResolver { + async fn resolve(&self, user_id: &str) -> Arc { + // Fast path: read lock + { + let cache = self.cache.read().await; + if let Some(ws) = cache.get(user_id) { + return Arc::clone(ws); + } + } + + // Slow path: write lock, double-check + let mut cache = self.cache.write().await; + if let Some(ws) = cache.get(user_id) { + return Arc::clone(ws); + } + + let ws = self.build_workspace(user_id); + cache.insert(user_id.to_string(), Arc::clone(&ws)); + tracing::debug!(user_id = user_id, "Created per-user workspace"); + ws + } +} + /// Detect paths that are clearly local filesystem references, not workspace-memory docs. /// /// Examples: @@ -62,13 +166,20 @@ fn map_write_err(e: crate::error::WorkspaceError) -> ToolError { /// The agent should call this tool before answering questions about /// prior work, decisions, preferences, or any historical context. pub struct MemorySearchTool { - workspace: Arc, + resolver: Arc, } impl MemorySearchTool { - /// Create a new memory search tool. - pub fn new(workspace: Arc) -> Self { - Self { workspace } + /// Create a new memory search tool with a workspace resolver. + pub fn new(resolver: Arc) -> Self { + Self { resolver } + } + + /// Create from a fixed workspace (backward compatibility). + pub fn from_workspace(workspace: Arc) -> Self { + Self { + resolver: Arc::new(FixedWorkspaceResolver::new(workspace)), + } } } @@ -107,7 +218,7 @@ impl Tool for MemorySearchTool { async fn execute( &self, params: serde_json::Value, - _ctx: &JobContext, + ctx: &JobContext, ) -> Result { let start = std::time::Instant::now(); @@ -119,8 +230,8 @@ impl Tool for MemorySearchTool { .unwrap_or(5) .min(20) as usize; - let results = self - .workspace + let workspace = self.resolver.resolve(&ctx.user_id).await; + let results = workspace .search(query, limit) .await .map_err(|e| ToolError::ExecutionFailed(format!("Search failed: {}", e)))?; @@ -151,13 +262,20 @@ impl Tool for MemorySearchTool { /// Use this to persist important information that should be remembered /// across sessions: decisions, preferences, facts, lessons learned. pub struct MemoryWriteTool { - workspace: Arc, + resolver: Arc, } impl MemoryWriteTool { - /// Create a new memory write tool. - pub fn new(workspace: Arc) -> Self { - Self { workspace } + /// Create a new memory write tool with a workspace resolver. + pub fn new(resolver: Arc) -> Self { + Self { resolver } + } + + /// Create from a fixed workspace (backward compatibility). + pub fn from_workspace(workspace: Arc) -> Self { + Self { + resolver: Arc::new(FixedWorkspaceResolver::new(workspace)), + } } } @@ -231,19 +349,21 @@ impl Tool for MemoryWriteTool { ))); } + let workspace = self.resolver.resolve(&ctx.user_id).await; + // Bootstrap target: clear BOOTSTRAP.md to mark first-run ritual complete. // Handled early because it accepts empty content (unlike other targets). if target == "bootstrap" { // Write empty content to effectively disable the bootstrap injection. // system_prompt_for_context() skips empty files. - self.workspace + workspace .write(paths::BOOTSTRAP, "") .await .map_err(map_write_err)?; // Also set the in-memory flag so BOOTSTRAP.md injection stops // immediately without waiting for a restart. - self.workspace.mark_bootstrap_completed(); + workspace.mark_bootstrap_completed(); let output = serde_json::json!({ "status": "cleared", @@ -289,12 +409,12 @@ impl Tool for MemoryWriteTool { // Otherwise, use default workspace methods (which include injection scanning). let layer_result = if let Some(layer_name) = layer { let result = if append { - self.workspace + workspace .append_to_layer(layer_name, &resolved_path, content, force) .await .map_err(map_write_err)? } else { - self.workspace + workspace .write_to_layer(layer_name, &resolved_path, content, force) .await .map_err(map_write_err)? @@ -307,31 +427,33 @@ impl Tool for MemoryWriteTool { match target { "memory" => { if append { - self.workspace + workspace .append_memory(content) .await .map_err(map_write_err)?; } else { - self.workspace + workspace .write(paths::MEMORY, content) .await .map_err(map_write_err)?; } } "daily_log" => { - self.workspace + let tz = crate::timezone::parse_timezone(&ctx.user_timezone) + .unwrap_or(chrono_tz::Tz::UTC); + workspace .append_daily_log_tz(content, tz) .await .map_err(map_write_err)?; } _ => { if append { - self.workspace + workspace .append(&resolved_path, content) .await .map_err(map_write_err)?; } else { - self.workspace + workspace .write(&resolved_path, content) .await .map_err(map_write_err)?; @@ -361,12 +483,12 @@ impl Tool for MemoryWriteTool { }; let mut synced_docs: Vec<&str> = Vec::new(); if normalized_path == paths::PROFILE { - match self.workspace.sync_profile_documents().await { + match workspace.sync_profile_documents().await { Ok(true) => { tracing::info!("profile write: synced USER.md + assistant-directives.md"); synced_docs.extend_from_slice(&[paths::USER, paths::ASSISTANT_DIRECTIVES]); - self.workspace.mark_bootstrap_completed(); + workspace.mark_bootstrap_completed(); let toml_path = crate::settings::Settings::default_toml_path(); if let Ok(Some(mut settings)) = crate::settings::Settings::load_toml(&toml_path) && !settings.profile_onboarding_completed @@ -416,13 +538,20 @@ impl Tool for MemoryWriteTool { /// /// Use this to read the full content of any file in the workspace. pub struct MemoryReadTool { - workspace: Arc, + resolver: Arc, } impl MemoryReadTool { - /// Create a new memory read tool. - pub fn new(workspace: Arc) -> Self { - Self { workspace } + /// Create a new memory read tool with a workspace resolver. + pub fn new(resolver: Arc) -> Self { + Self { resolver } + } + + /// Create from a fixed workspace (backward compatibility). + pub fn from_workspace(workspace: Arc) -> Self { + Self { + resolver: Arc::new(FixedWorkspaceResolver::new(workspace)), + } } } @@ -456,7 +585,7 @@ impl Tool for MemoryReadTool { async fn execute( &self, params: serde_json::Value, - _ctx: &JobContext, + ctx: &JobContext, ) -> Result { let start = std::time::Instant::now(); @@ -470,8 +599,8 @@ impl Tool for MemoryReadTool { ))); } - let doc = self - .workspace + let workspace = self.resolver.resolve(&ctx.user_id).await; + let doc = workspace .read(path) .await .map_err(|e| ToolError::ExecutionFailed(format!("Read failed: {}", e)))?; @@ -495,20 +624,27 @@ impl Tool for MemoryReadTool { /// /// Returns a hierarchical view of files and directories with configurable depth. pub struct MemoryTreeTool { - workspace: Arc, + resolver: Arc, } impl MemoryTreeTool { - /// Create a new memory tree tool. - pub fn new(workspace: Arc) -> Self { - Self { workspace } + /// Create a new memory tree tool with a workspace resolver. + pub fn new(resolver: Arc) -> Self { + Self { resolver } + } + + /// Create from a fixed workspace (backward compatibility). + pub fn from_workspace(workspace: Arc) -> Self { + Self { + resolver: Arc::new(FixedWorkspaceResolver::new(workspace)), + } } /// Recursively build tree structure. /// /// Returns a compact format where directories end with `/` and may have children. async fn build_tree( - &self, + workspace: &Arc, path: &str, current_depth: usize, max_depth: usize, @@ -517,8 +653,7 @@ impl MemoryTreeTool { return Ok(Vec::new()); } - let entries = self - .workspace + let entries = workspace .list(path) .await .map_err(|e| ToolError::ExecutionFailed(format!("Tree failed: {}", e)))?; @@ -534,7 +669,7 @@ impl MemoryTreeTool { if entry.is_directory && current_depth < max_depth { let children = - Box::pin(self.build_tree(&entry.path, current_depth + 1, max_depth)).await?; + Box::pin(Self::build_tree(workspace, &entry.path, current_depth + 1, max_depth)).await?; if children.is_empty() { result.push(serde_json::Value::String(display_path)); } else { @@ -584,7 +719,7 @@ impl Tool for MemoryTreeTool { async fn execute( &self, params: serde_json::Value, - _ctx: &JobContext, + ctx: &JobContext, ) -> Result { let start = std::time::Instant::now(); @@ -596,7 +731,8 @@ impl Tool for MemoryTreeTool { .unwrap_or(1) .clamp(1, 10) as usize; - let tree = self.build_tree(path, 1, depth).await?; + let workspace = self.resolver.resolve(&ctx.user_id).await; + let tree = Self::build_tree(&workspace, path, 1, depth).await?; // Compact output: just the tree array Ok(ToolOutput::success( @@ -650,7 +786,7 @@ mod tests { #[test] fn test_memory_search_schema() { let workspace = make_test_workspace(); - let tool = MemorySearchTool::new(workspace); + let tool = MemorySearchTool::from_workspace(workspace); assert_eq!(tool.name(), "memory_search"); assert!(!tool.requires_sanitization()); @@ -668,7 +804,7 @@ mod tests { #[test] fn test_memory_write_schema() { let workspace = make_test_workspace(); - let tool = MemoryWriteTool::new(workspace); + let tool = MemoryWriteTool::from_workspace(workspace); assert_eq!(tool.name(), "memory_write"); @@ -681,7 +817,7 @@ mod tests { #[test] fn test_memory_read_schema() { let workspace = make_test_workspace(); - let tool = MemoryReadTool::new(workspace); + let tool = MemoryReadTool::from_workspace(workspace); assert_eq!(tool.name(), "memory_read"); @@ -698,7 +834,7 @@ mod tests { #[test] fn test_memory_tree_schema() { let workspace = make_test_workspace(); - let tool = MemoryTreeTool::new(workspace); + let tool = MemoryTreeTool::from_workspace(workspace); assert_eq!(tool.name(), "memory_tree"); @@ -711,7 +847,7 @@ mod tests { #[tokio::test] async fn test_memory_write_rejects_injection_to_identity_file() { let workspace = make_test_workspace(); - let tool = MemoryWriteTool::new(workspace); + let tool = MemoryWriteTool::from_workspace(workspace); let ctx = JobContext::default(); let params = serde_json::json!({ @@ -733,4 +869,176 @@ mod tests { } } } + + // Regression tests for per-user workspace scoping (multi-tenant mode). + // See: https://github.com/nearai/ironclaw/pull/1118 + // Bug: memory tools used a single startup workspace regardless of which + // user was chatting. Fix: resolve workspace per-request via JobContext.user_id. + + #[cfg(feature = "postgres")] + mod resolver_tests { + use super::*; + + fn make_test_workspace_for_user(user_id: &str) -> Arc { + Arc::new(Workspace::new( + user_id, + deadpool_postgres::Pool::builder(deadpool_postgres::Manager::new( + tokio_postgres::Config::new(), + tokio_postgres::NoTls, + )) + .build() + .unwrap(), + )) + } + + #[tokio::test] + async fn test_fixed_workspace_resolver_ignores_user_id() { + let ws = make_test_workspace_for_user("alice"); + let resolver = FixedWorkspaceResolver::new(Arc::clone(&ws)); + + let ws_alice = resolver.resolve("alice").await; + let ws_bob = resolver.resolve("bob").await; + + // Both should return the exact same Arc (pointer equality) + assert!(Arc::ptr_eq(&ws_alice, &ws_bob)); + assert_eq!(ws_alice.user_id(), "alice"); + } + + /// Tracking resolver that records which user_ids were requested. + struct TrackingWorkspaceResolver { + inner: FixedWorkspaceResolver, + resolved_users: std::sync::Mutex>, + } + + impl TrackingWorkspaceResolver { + fn new(workspace: Arc) -> Self { + Self { + inner: FixedWorkspaceResolver::new(workspace), + resolved_users: std::sync::Mutex::new(Vec::new()), + } + } + + fn resolved_users(&self) -> Vec { + self.resolved_users.lock().unwrap().clone() + } + } + + #[async_trait] + impl WorkspaceResolver for TrackingWorkspaceResolver { + async fn resolve(&self, user_id: &str) -> Arc { + self.resolved_users + .lock() + .unwrap() + .push(user_id.to_string()); + self.inner.resolve(user_id).await + } + } + + #[tokio::test] + async fn test_memory_search_uses_job_context_user_id() { + let ws = make_test_workspace_for_user("default"); + let tracker = Arc::new(TrackingWorkspaceResolver::new(ws)); + let tool = MemorySearchTool::new(tracker.clone() as Arc); + + // Execute with user_id "alice" + let ctx_alice = JobContext::with_user("alice", "test", "test"); + let params = serde_json::json!({"query": "test"}); + // The search will fail (no real DB) but we only care about resolver call + let _ = tool.execute(params, &ctx_alice).await; + + // Execute with user_id "bob" + let ctx_bob = JobContext::with_user("bob", "test", "test"); + let params = serde_json::json!({"query": "test"}); + let _ = tool.execute(params, &ctx_bob).await; + + let resolved = tracker.resolved_users(); + assert_eq!(resolved, vec!["alice", "bob"]); + } + + #[tokio::test] + async fn test_memory_write_uses_job_context_user_id() { + let ws = make_test_workspace_for_user("default"); + let tracker = Arc::new(TrackingWorkspaceResolver::new(ws)); + let tool = MemoryWriteTool::new(tracker.clone() as Arc); + + // Execute with user_id "alice" + let ctx_alice = JobContext::with_user("alice", "test", "test"); + let params = serde_json::json!({ + "content": "remember this", + "target": "daily_log", + }); + let _ = tool.execute(params, &ctx_alice).await; + + // Execute with user_id "bob" + let ctx_bob = JobContext::with_user("bob", "test", "test"); + let params = serde_json::json!({ + "content": "remember that", + "target": "daily_log", + }); + let _ = tool.execute(params, &ctx_bob).await; + + let resolved = tracker.resolved_users(); + assert_eq!(resolved, vec!["alice", "bob"]); + } + } + + #[cfg(feature = "libsql")] + mod per_user_resolver_tests { + use super::*; + + async fn make_test_db() -> Arc { + use crate::db::libsql::LibSqlBackend; + let temp_dir = tempfile::tempdir().expect("tempdir"); + let db_path = temp_dir.path().join("resolver_test.db"); + let backend = LibSqlBackend::new_local(&db_path) + .await + .expect("LibSqlBackend"); + ::run_migrations(&backend) + .await + .expect("migrations"); + // Leak the tempdir so it outlives the test (cleaned up on process exit). + std::mem::forget(temp_dir); + Arc::new(backend) + } + + #[tokio::test] + async fn test_per_user_workspace_resolver_returns_different_workspaces() { + let db = make_test_db().await; + + let resolver = PerUserWorkspaceResolver::new( + db, + None, + crate::workspace::EmbeddingCacheConfig::default(), + crate::config::WorkspaceSearchConfig::default(), + crate::config::WorkspaceConfig::default(), + ); + + let ws_alice = resolver.resolve("alice").await; + let ws_bob = resolver.resolve("bob").await; + + // Different user IDs should get different workspaces + assert_eq!(ws_alice.user_id(), "alice"); + assert_eq!(ws_bob.user_id(), "bob"); + assert!(!Arc::ptr_eq(&ws_alice, &ws_bob)); + } + + #[tokio::test] + async fn test_per_user_workspace_resolver_caches_workspace() { + let db = make_test_db().await; + + let resolver = PerUserWorkspaceResolver::new( + db, + None, + crate::workspace::EmbeddingCacheConfig::default(), + crate::config::WorkspaceSearchConfig::default(), + crate::config::WorkspaceConfig::default(), + ); + + let ws1 = resolver.resolve("alice").await; + let ws2 = resolver.resolve("alice").await; + + // Same user_id should return the same cached Arc (pointer equality) + assert!(Arc::ptr_eq(&ws1, &ws2)); + } + } } diff --git a/src/tools/builtin/mod.rs b/src/tools/builtin/mod.rs index 8ba8e57b..d196b12c 100644 --- a/src/tools/builtin/mod.rs +++ b/src/tools/builtin/mod.rs @@ -6,7 +6,7 @@ mod file; mod http; mod job; mod json; -mod memory; +pub mod memory; mod message; pub mod path_utils; mod restart; diff --git a/src/tools/registry.rs b/src/tools/registry.rs index 9dfee0a3..dd72f33b 100644 --- a/src/tools/registry.rs +++ b/src/tools/registry.rs @@ -334,15 +334,31 @@ impl ToolRegistry { tracing::debug!("Registered 5 development tools"); } - /// Register memory tools with a workspace. + /// Register memory tools with a workspace resolver. + /// + /// Memory tools require a workspace resolver for persistence. Call this after + /// `register_builtin_tools()` if you have a workspace available. + pub fn register_memory_tools_with_resolver( + &self, + resolver: Arc, + ) { + self.register_sync(Arc::new(MemorySearchTool::new(Arc::clone(&resolver)))); + self.register_sync(Arc::new(MemoryWriteTool::new(Arc::clone(&resolver)))); + self.register_sync(Arc::new(MemoryReadTool::new(Arc::clone(&resolver)))); + self.register_sync(Arc::new(MemoryTreeTool::new(resolver))); + + tracing::debug!("Registered 4 memory tools"); + } + + /// Register memory tools with a fixed workspace (backward compatibility). /// /// Memory tools require a workspace for persistence. Call this after /// `register_builtin_tools()` if you have a workspace available. pub fn register_memory_tools(&self, workspace: Arc) { - self.register_sync(Arc::new(MemorySearchTool::new(Arc::clone(&workspace)))); - self.register_sync(Arc::new(MemoryWriteTool::new(Arc::clone(&workspace)))); - self.register_sync(Arc::new(MemoryReadTool::new(Arc::clone(&workspace)))); - self.register_sync(Arc::new(MemoryTreeTool::new(workspace))); + self.register_sync(Arc::new(MemorySearchTool::from_workspace(Arc::clone(&workspace)))); + self.register_sync(Arc::new(MemoryWriteTool::from_workspace(Arc::clone(&workspace)))); + self.register_sync(Arc::new(MemoryReadTool::from_workspace(Arc::clone(&workspace)))); + self.register_sync(Arc::new(MemoryTreeTool::from_workspace(workspace))); tracing::debug!("Registered 4 memory tools"); }