mirror of
https://github.com/outbackdingo/optimclaw.git
synced 2026-08-27 08:00:17 +00:00
feat: multi-tenant auth with per-user workspace isolation (#1118)
* feat: multi-tenant auth with per-user scoping Multi-user authentication and authorization for IronClaw gateway: - Token-based auth mapping tokens to user IDs via GATEWAY_USER_TOKENS - Per-user SSE broadcast scoping - Per-user rate limiting with poisoned lock recovery - Handler auth and ownership checks for jobs, settings, routines - Extension secrets scoped per-user - Chat handlers use authenticated identity - Reverse proxy deployment documentation - Comprehensive integration tests for auth, SSE, rate limiting, and job isolation * fix: scope memory tools per-user in multi-tenant mode Memory tools (search, write, read, tree) held a single workspace created at startup with GATEWAY_USER_ID. In multi-tenant mode, all users' tool calls searched the default user's scope. Add WorkspaceResolver trait that resolves workspaces per-request using JobContext.user_id. In single-user mode, returns the startup workspace. In multi-tenant mode (GATEWAY_USER_TOKENS configured), creates and caches per-user workspaces on demand. Includes regression tests for workspace resolution and user isolation. Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]> * fix: comprehensive multi-tenant isolation audit Address all review findings from @serrrfirat plus 7 additional gaps found via full security audit: Reviewer findings (5): - WorkspacePool now applies search config, memory layers, embedding cache, identity read scopes, and global config scopes (was bare) - jobs_summary_handler uses per-user queries instead of global counters - jobs_prompt_handler restructured to not 404 agent jobs + ownership check - jobs_restart_handler agent branch now verifies user ownership - agent_job_summary_for_user added to Database trait + both backends Audit findings (7): - Delete dead handlers/memory.rs (stale copies with no auth) - Add AuthenticatedUser to logs_events, logs_level_get, logs_level_set - Add AuthenticatedUser to extensions_tools_handler, gateway_status_handler - Add auth + ownership checks to all 6 routines handlers - Add auth to all 4 skills handlers with audit logging on mutations - Scope extension setup SSE broadcast to user (broadcast_for_user) - Fix pre-existing test compilation errors in extensions/manager.rs 17 new multi-tenant isolation tests covering: - WorkspacePool config propagation and scope merging - Jobs handler per-user isolation (summary, restart, prompt, cancel) - Routines handler auth enforcement and cross-user rejection - Auth middleware enforcement on logs, skills, status endpoints Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]> * fix: second-pass multi-tenant audit — scope SSE broadcasts, DB queries, dead handlers Second audit pass applying learned patterns across the codebase: - OAuth callback SSE broadcasts now use broadcast_for_user (lines 773, 912) - jobs_list_handler uses list_agent_jobs_for_user instead of fetching all users' jobs and filtering in Rust - list_agent_jobs_for_user added to Database trait + postgres + libsql - Dead handler files (extensions.rs, static_files.rs) hardened with AuthenticatedUser to prevent auth regression if migrated Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]> * fix: address review findings — token hashing, broadcast scoping, error handling Security fixes: - Hash tokens with SHA-256 at construction time so authentication compares fixed-size 32-byte digests, eliminating length-oracle timing leaks - Scope auth SSE broadcasts per-user in chat_auth_token_handler — AuthRequired/AuthCompleted events were leaking across tenants - Propagate DB errors in restart handlers instead of silently swallowing via `if let Ok(Some(...))` pattern Code quality: - Log SSE serialization failures instead of silently producing empty strings via unwrap_or_default() - Remove dead `pub type AuthState = MultiAuthState` alias - Replace `.unwrap()` with `Arc::clone(db)` in app.rs multi-tenant workspace setup (db is guaranteed Some in context, but unwrap violates project convention) - Fix telegram setup test to inject UserIdentity into request extensions (handler now requires AuthenticatedUser) - Add safety comments on test-only expect/unwrap calls for CI - Apply cargo fmt to fix pre-existing formatting Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]> * fix: address review findings — unify workspace pool, fix SSE regression, cache job owners - Unify WorkspacePool and PerUserWorkspaceResolver: WorkspacePool now implements WorkspaceResolver, eliminating duplicate per-user workspace construction logic. app.rs uses WorkspacePool directly. - Fix sse_tx: None scheduler regression: change scheduler/worker SSE broadcasting from broadcast::Sender<SseEvent> to Arc<SseManager>, restoring SSE event delivery for scheduled agent jobs. - Cache job owner in orchestrator: add job_owner_cache to OrchestratorState so job_event_handler avoids a DB round-trip on every event after the first per job. - Deduplicate ext_user_id computation in main.rs. - Remove unused _gateway_state variable. - Fix pre-existing test: first_token() returns None in multi-user mode by design; align test assertion. Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]> * style: fix formatting in app.rs Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]> * refactor: extract memory handlers back into handlers/memory.rs Move memory API handlers out of server.rs into their own module, consistent with how jobs, routines, and skills handlers are organized. The resolve_workspace() helper moves with them since it is only used by memory handlers. Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]> --------- Co-authored-by: Claude Opus 4.6 (1M context) <[email protected]> Co-authored-by: [email protected] <[email protected]>
This commit is contained in:
@@ -130,7 +130,7 @@ impl Tool for ToolInstallTool {
|
||||
async fn execute(
|
||||
&self,
|
||||
params: serde_json::Value,
|
||||
_ctx: &JobContext,
|
||||
ctx: &JobContext,
|
||||
) -> Result<ToolOutput, ToolError> {
|
||||
let start = std::time::Instant::now();
|
||||
|
||||
@@ -150,7 +150,7 @@ impl Tool for ToolInstallTool {
|
||||
|
||||
let result = self
|
||||
.manager
|
||||
.install(name, url, kind_hint)
|
||||
.install(name, url, kind_hint, &ctx.user_id)
|
||||
.await
|
||||
.map_err(|e| ToolError::ExecutionFailed(e.to_string()))?;
|
||||
|
||||
@@ -205,7 +205,7 @@ impl Tool for ToolAuthTool {
|
||||
async fn execute(
|
||||
&self,
|
||||
params: serde_json::Value,
|
||||
_ctx: &JobContext,
|
||||
ctx: &JobContext,
|
||||
) -> Result<ToolOutput, ToolError> {
|
||||
let start = std::time::Instant::now();
|
||||
|
||||
@@ -213,13 +213,13 @@ impl Tool for ToolAuthTool {
|
||||
|
||||
let result = self
|
||||
.manager
|
||||
.auth(name)
|
||||
.auth(name, &ctx.user_id)
|
||||
.await
|
||||
.map_err(|e| ToolError::ExecutionFailed(e.to_string()))?;
|
||||
|
||||
// Auto-activate after successful auth so tools are available immediately
|
||||
if result.is_authenticated() {
|
||||
match self.manager.activate(name).await {
|
||||
match self.manager.activate(name, &ctx.user_id).await {
|
||||
Ok(activate_result) => {
|
||||
let output = serde_json::json!({
|
||||
"status": "authenticated_and_activated",
|
||||
@@ -304,13 +304,13 @@ impl Tool for ToolActivateTool {
|
||||
async fn execute(
|
||||
&self,
|
||||
params: serde_json::Value,
|
||||
_ctx: &JobContext,
|
||||
ctx: &JobContext,
|
||||
) -> Result<ToolOutput, ToolError> {
|
||||
let start = std::time::Instant::now();
|
||||
|
||||
let name = require_str(¶ms, "name")?;
|
||||
|
||||
match self.manager.activate(name).await {
|
||||
match self.manager.activate(name, &ctx.user_id).await {
|
||||
Ok(result) => {
|
||||
let output = serde_json::to_value(&result)
|
||||
.unwrap_or_else(|_| serde_json::json!({"error": "serialization failed"}));
|
||||
@@ -329,12 +329,12 @@ impl Tool for ToolActivateTool {
|
||||
|
||||
// Activation failed due to missing auth; initiate auth flow
|
||||
// so the agent loop can show the auth card.
|
||||
match self.manager.auth(name).await {
|
||||
match self.manager.auth(name, &ctx.user_id).await {
|
||||
Ok(auth_result) if auth_result.is_authenticated() => {
|
||||
// Auth succeeded (e.g. env var was set); retry activation.
|
||||
let result = self
|
||||
.manager
|
||||
.activate(name)
|
||||
.activate(name, &ctx.user_id)
|
||||
.await
|
||||
.map_err(|e| ToolError::ExecutionFailed(e.to_string()))?;
|
||||
let output = serde_json::to_value(&result).unwrap_or_else(
|
||||
@@ -404,7 +404,7 @@ impl Tool for ToolListTool {
|
||||
async fn execute(
|
||||
&self,
|
||||
params: serde_json::Value,
|
||||
_ctx: &JobContext,
|
||||
ctx: &JobContext,
|
||||
) -> Result<ToolOutput, ToolError> {
|
||||
let start = std::time::Instant::now();
|
||||
|
||||
@@ -425,7 +425,7 @@ impl Tool for ToolListTool {
|
||||
|
||||
let extensions = self
|
||||
.manager
|
||||
.list(kind_filter, include_available)
|
||||
.list(kind_filter, include_available, &ctx.user_id)
|
||||
.await
|
||||
.map_err(|e| ToolError::ExecutionFailed(e.to_string()))?;
|
||||
|
||||
@@ -477,7 +477,7 @@ impl Tool for ToolRemoveTool {
|
||||
async fn execute(
|
||||
&self,
|
||||
params: serde_json::Value,
|
||||
_ctx: &JobContext,
|
||||
ctx: &JobContext,
|
||||
) -> Result<ToolOutput, ToolError> {
|
||||
let start = std::time::Instant::now();
|
||||
|
||||
@@ -485,7 +485,7 @@ impl Tool for ToolRemoveTool {
|
||||
|
||||
let message = self
|
||||
.manager
|
||||
.remove(name)
|
||||
.remove(name, &ctx.user_id)
|
||||
.await
|
||||
.map_err(|e| ToolError::ExecutionFailed(e.to_string()))?;
|
||||
|
||||
@@ -541,7 +541,7 @@ impl Tool for ToolUpgradeTool {
|
||||
async fn execute(
|
||||
&self,
|
||||
params: serde_json::Value,
|
||||
_ctx: &JobContext,
|
||||
ctx: &JobContext,
|
||||
) -> Result<ToolOutput, ToolError> {
|
||||
let start = std::time::Instant::now();
|
||||
|
||||
@@ -549,7 +549,7 @@ impl Tool for ToolUpgradeTool {
|
||||
|
||||
let result = self
|
||||
.manager
|
||||
.upgrade(name)
|
||||
.upgrade(name, &ctx.user_id)
|
||||
.await
|
||||
.map_err(|e| ToolError::ExecutionFailed(e.to_string()))?;
|
||||
|
||||
@@ -603,7 +603,7 @@ impl Tool for ExtensionInfoTool {
|
||||
async fn execute(
|
||||
&self,
|
||||
params: serde_json::Value,
|
||||
_ctx: &JobContext,
|
||||
ctx: &JobContext,
|
||||
) -> Result<ToolOutput, ToolError> {
|
||||
let start = std::time::Instant::now();
|
||||
|
||||
@@ -611,7 +611,7 @@ impl Tool for ExtensionInfoTool {
|
||||
|
||||
let info = self
|
||||
.manager
|
||||
.extension_info(name)
|
||||
.extension_info(name, &ctx.user_id)
|
||||
.await
|
||||
.map_err(|e| ToolError::ExecutionFailed(e.to_string()))?;
|
||||
|
||||
|
||||
@@ -85,7 +85,7 @@ pub struct CreateJobTool {
|
||||
job_manager: Option<Arc<ContainerJobManager>>,
|
||||
store: Option<Arc<dyn Database>>,
|
||||
/// Broadcast sender for job events (used to subscribe a monitor).
|
||||
event_tx: Option<tokio::sync::broadcast::Sender<(Uuid, SseEvent)>>,
|
||||
event_tx: Option<tokio::sync::broadcast::Sender<(Uuid, String, SseEvent)>>,
|
||||
/// Injection channel for pushing messages into the agent loop.
|
||||
inject_tx: Option<tokio::sync::mpsc::Sender<IncomingMessage>>,
|
||||
/// Encrypted secrets store for validating credential grants.
|
||||
@@ -120,7 +120,7 @@ impl CreateJobTool {
|
||||
/// monitor that forwards Claude Code output to the main agent loop.
|
||||
pub fn with_monitor_deps(
|
||||
mut self,
|
||||
event_tx: tokio::sync::broadcast::Sender<(Uuid, SseEvent)>,
|
||||
event_tx: tokio::sync::broadcast::Sender<(Uuid, String, SseEvent)>,
|
||||
inject_tx: tokio::sync::mpsc::Sender<IncomingMessage>,
|
||||
) -> Self {
|
||||
self.event_tx = Some(event_tx);
|
||||
|
||||
+283
-45
@@ -21,6 +21,35 @@ 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<Workspace>;
|
||||
}
|
||||
|
||||
/// Returns a fixed workspace regardless of user ID (single-user mode).
|
||||
pub struct FixedWorkspaceResolver {
|
||||
workspace: Arc<Workspace>,
|
||||
}
|
||||
|
||||
impl FixedWorkspaceResolver {
|
||||
pub fn new(workspace: Arc<Workspace>) -> Self {
|
||||
Self { workspace }
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl WorkspaceResolver for FixedWorkspaceResolver {
|
||||
async fn resolve(&self, _user_id: &str) -> Arc<Workspace> {
|
||||
Arc::clone(&self.workspace)
|
||||
}
|
||||
}
|
||||
|
||||
/// Detect paths that are clearly local filesystem references, not workspace-memory docs.
|
||||
///
|
||||
/// Examples:
|
||||
@@ -62,13 +91,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<Workspace>,
|
||||
resolver: Arc<dyn WorkspaceResolver>,
|
||||
}
|
||||
|
||||
impl MemorySearchTool {
|
||||
/// Create a new memory search tool.
|
||||
pub fn new(workspace: Arc<Workspace>) -> Self {
|
||||
Self { workspace }
|
||||
/// Create a new memory search tool with a workspace resolver.
|
||||
pub fn new(resolver: Arc<dyn WorkspaceResolver>) -> Self {
|
||||
Self { resolver }
|
||||
}
|
||||
|
||||
/// Create from a fixed workspace (backward compatibility).
|
||||
pub fn from_workspace(workspace: Arc<Workspace>) -> Self {
|
||||
Self {
|
||||
resolver: Arc::new(FixedWorkspaceResolver::new(workspace)),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -107,7 +143,7 @@ impl Tool for MemorySearchTool {
|
||||
async fn execute(
|
||||
&self,
|
||||
params: serde_json::Value,
|
||||
_ctx: &JobContext,
|
||||
ctx: &JobContext,
|
||||
) -> Result<ToolOutput, ToolError> {
|
||||
let start = std::time::Instant::now();
|
||||
|
||||
@@ -119,8 +155,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 +187,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<Workspace>,
|
||||
resolver: Arc<dyn WorkspaceResolver>,
|
||||
}
|
||||
|
||||
impl MemoryWriteTool {
|
||||
/// Create a new memory write tool.
|
||||
pub fn new(workspace: Arc<Workspace>) -> Self {
|
||||
Self { workspace }
|
||||
/// Create a new memory write tool with a workspace resolver.
|
||||
pub fn new(resolver: Arc<dyn WorkspaceResolver>) -> Self {
|
||||
Self { resolver }
|
||||
}
|
||||
|
||||
/// Create from a fixed workspace (backward compatibility).
|
||||
pub fn from_workspace(workspace: Arc<Workspace>) -> Self {
|
||||
Self {
|
||||
resolver: Arc::new(FixedWorkspaceResolver::new(workspace)),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -231,19 +274,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 +334,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 +352,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 +408,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 +463,20 @@ impl Tool for MemoryWriteTool {
|
||||
///
|
||||
/// Use this to read the full content of any file in the workspace.
|
||||
pub struct MemoryReadTool {
|
||||
workspace: Arc<Workspace>,
|
||||
resolver: Arc<dyn WorkspaceResolver>,
|
||||
}
|
||||
|
||||
impl MemoryReadTool {
|
||||
/// Create a new memory read tool.
|
||||
pub fn new(workspace: Arc<Workspace>) -> Self {
|
||||
Self { workspace }
|
||||
/// Create a new memory read tool with a workspace resolver.
|
||||
pub fn new(resolver: Arc<dyn WorkspaceResolver>) -> Self {
|
||||
Self { resolver }
|
||||
}
|
||||
|
||||
/// Create from a fixed workspace (backward compatibility).
|
||||
pub fn from_workspace(workspace: Arc<Workspace>) -> Self {
|
||||
Self {
|
||||
resolver: Arc::new(FixedWorkspaceResolver::new(workspace)),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -456,7 +510,7 @@ impl Tool for MemoryReadTool {
|
||||
async fn execute(
|
||||
&self,
|
||||
params: serde_json::Value,
|
||||
_ctx: &JobContext,
|
||||
ctx: &JobContext,
|
||||
) -> Result<ToolOutput, ToolError> {
|
||||
let start = std::time::Instant::now();
|
||||
|
||||
@@ -470,8 +524,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 +549,27 @@ impl Tool for MemoryReadTool {
|
||||
///
|
||||
/// Returns a hierarchical view of files and directories with configurable depth.
|
||||
pub struct MemoryTreeTool {
|
||||
workspace: Arc<Workspace>,
|
||||
resolver: Arc<dyn WorkspaceResolver>,
|
||||
}
|
||||
|
||||
impl MemoryTreeTool {
|
||||
/// Create a new memory tree tool.
|
||||
pub fn new(workspace: Arc<Workspace>) -> Self {
|
||||
Self { workspace }
|
||||
/// Create a new memory tree tool with a workspace resolver.
|
||||
pub fn new(resolver: Arc<dyn WorkspaceResolver>) -> Self {
|
||||
Self { resolver }
|
||||
}
|
||||
|
||||
/// Create from a fixed workspace (backward compatibility).
|
||||
pub fn from_workspace(workspace: Arc<Workspace>) -> 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<Workspace>,
|
||||
path: &str,
|
||||
current_depth: usize,
|
||||
max_depth: usize,
|
||||
@@ -517,8 +578,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)))?;
|
||||
@@ -533,8 +593,13 @@ 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?;
|
||||
let children = 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 +649,7 @@ impl Tool for MemoryTreeTool {
|
||||
async fn execute(
|
||||
&self,
|
||||
params: serde_json::Value,
|
||||
_ctx: &JobContext,
|
||||
ctx: &JobContext,
|
||||
) -> Result<ToolOutput, ToolError> {
|
||||
let start = std::time::Instant::now();
|
||||
|
||||
@@ -596,7 +661,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 +716,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 +734,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 +747,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 +764,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 +777,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 +799,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<Workspace> {
|
||||
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<Vec<String>>,
|
||||
}
|
||||
|
||||
impl TrackingWorkspaceResolver {
|
||||
fn new(workspace: Arc<Workspace>) -> Self {
|
||||
Self {
|
||||
inner: FixedWorkspaceResolver::new(workspace),
|
||||
resolved_users: std::sync::Mutex::new(Vec::new()),
|
||||
}
|
||||
}
|
||||
|
||||
fn resolved_users(&self) -> Vec<String> {
|
||||
self.resolved_users.lock().unwrap().clone()
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl WorkspaceResolver for TrackingWorkspaceResolver {
|
||||
async fn resolve(&self, user_id: &str) -> Arc<Workspace> {
|
||||
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<dyn WorkspaceResolver>);
|
||||
|
||||
// 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<dyn WorkspaceResolver>);
|
||||
|
||||
// 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<dyn crate::db::Database> {
|
||||
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");
|
||||
<LibSqlBackend as crate::db::Database>::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_workspace_pool_resolver_returns_different_workspaces() {
|
||||
let db = make_test_db().await;
|
||||
|
||||
let pool = crate::channels::web::server::WorkspacePool::new(
|
||||
db,
|
||||
None,
|
||||
crate::workspace::EmbeddingCacheConfig::default(),
|
||||
crate::config::WorkspaceSearchConfig::default(),
|
||||
crate::config::WorkspaceConfig::default(),
|
||||
);
|
||||
|
||||
let ws_alice = pool.resolve("alice").await;
|
||||
let ws_bob = pool.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_workspace_pool_resolver_caches_workspace() {
|
||||
let db = make_test_db().await;
|
||||
|
||||
let pool = crate::channels::web::server::WorkspacePool::new(
|
||||
db,
|
||||
None,
|
||||
crate::workspace::EmbeddingCacheConfig::default(),
|
||||
crate::config::WorkspaceSearchConfig::default(),
|
||||
crate::config::WorkspaceConfig::default(),
|
||||
);
|
||||
|
||||
let ws1 = pool.resolve("alice").await;
|
||||
let ws2 = pool.resolve("alice").await;
|
||||
|
||||
// Same user_id should return the same cached Arc (pointer equality)
|
||||
assert!(Arc::ptr_eq(&ws1, &ws2));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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;
|
||||
|
||||
+32
-6
@@ -334,15 +334,37 @@ 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<dyn crate::tools::builtin::memory::WorkspaceResolver>,
|
||||
) {
|
||||
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<Workspace>) {
|
||||
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");
|
||||
}
|
||||
@@ -361,7 +383,11 @@ impl ToolRegistry {
|
||||
job_manager: Option<Arc<ContainerJobManager>>,
|
||||
store: Option<Arc<dyn Database>>,
|
||||
job_event_tx: Option<
|
||||
tokio::sync::broadcast::Sender<(uuid::Uuid, crate::channels::web::types::SseEvent)>,
|
||||
tokio::sync::broadcast::Sender<(
|
||||
uuid::Uuid,
|
||||
String,
|
||||
crate::channels::web::types::SseEvent,
|
||||
)>,
|
||||
>,
|
||||
inject_tx: Option<tokio::sync::mpsc::Sender<crate::channels::IncomingMessage>>,
|
||||
prompt_queue: Option<PromptQueue>,
|
||||
|
||||
Reference in New Issue
Block a user