diff --git a/src/agent/agent_loop.rs b/src/agent/agent_loop.rs index 4bc81569..546e4c2b 100644 --- a/src/agent/agent_loop.rs +++ b/src/agent/agent_loop.rs @@ -172,6 +172,8 @@ pub struct AgentDeps { /// Resolved LLM backend identifier (e.g., "nearai", "openai", "groq"). /// Used by `/model` persistence to determine which env var to update. pub llm_backend: String, + /// Per-tenant rate limiting registry (lazily creates rate state per user). + pub tenant_rates: Arc, } /// The main agent that coordinates all components. @@ -234,7 +236,10 @@ impl Agent { SchedulerDeps { tools: deps.tools.clone(), extension_manager: deps.extension_manager.clone(), - store: deps.store.clone(), + store: deps + .store + .as_ref() + .map(|db| crate::tenant::AdminScope::new(Arc::clone(db))), hooks: deps.hooks.clone(), }, ); @@ -315,6 +320,50 @@ impl Agent { &self.deps.cost_guard } + /// Build a tenant-scoped execution context for the given user. + /// + /// This is the standard entry point for per-user operations. The returned + /// [`TenantCtx`] provides a [`TenantScope`] that auto-binds `user_id` on + /// every database operation and a per-user rate limiter. + pub(super) async fn tenant_ctx(&self, user_id: &str) -> crate::tenant::TenantCtx { + let rate = self.deps.tenant_rates.get_or_create(user_id).await; + + let store = self + .deps + .store + .as_ref() + .map(|db| crate::tenant::TenantScope::new(user_id, Arc::clone(db))); + + // Reuse the owner workspace if user matches, otherwise create per-user. + let workspace = match &self.deps.workspace { + Some(ws) if ws.user_id() == user_id => Some(Arc::clone(ws)), + _ => self + .deps + .store + .as_ref() + .map(|db| Arc::new(Workspace::new_with_db(user_id, Arc::clone(db)))), + }; + + crate::tenant::TenantCtx::new( + user_id, + store, + workspace, + Arc::clone(&self.deps.cost_guard), + rate, + ) + } + + /// Get an admin-scoped database accessor for cross-tenant operations. + /// + /// Only for system-level components (heartbeat, routine engine, self-repair, + /// scheduler). Handler code should use [`tenant_ctx()`](Self::tenant_ctx) instead. + pub(super) fn admin_store(&self) -> Option { + self.deps + .store + .as_ref() + .map(|db| crate::tenant::AdminScope::new(Arc::clone(db))) + } + pub(super) fn skill_registry(&self) -> Option<&Arc>> { self.deps.skill_registry.as_ref() } @@ -400,8 +449,8 @@ impl Agent { self.config.stuck_threshold, self.config.max_repair_attempts, ); - if let Some(ref store) = self.deps.store { - self_repair = self_repair.with_store(Arc::clone(store)); + if let Some(admin) = self.admin_store() { + self_repair = self_repair.with_store(admin); } if let Some(ref builder) = self.deps.builder { self_repair = self_repair.with_builder(Arc::clone(builder), Arc::clone(self.tools())); @@ -597,13 +646,13 @@ impl Agent { .unwrap_or_default(); if config.multi_tenant { - if let Some(store) = self.store() { + if let Some(admin) = self.admin_store() { Some(spawn_multi_user_heartbeat( config, hygiene, self.cheap_llm().clone(), Some(notify_tx), - Arc::clone(store), + admin, )) } else { tracing::warn!("Multi-tenant heartbeat requires a database store"); @@ -616,7 +665,7 @@ impl Agent { workspace.clone(), self.cheap_llm().clone(), Some(notify_tx), - self.store().map(Arc::clone), + self.admin_store(), )) } } else { @@ -640,7 +689,7 @@ impl Agent { let engine = Arc::new(RoutineEngine::new( rt_config.clone(), - Arc::clone(store), + crate::tenant::AdminScope::new(Arc::clone(store)), self.llm().clone(), Arc::clone(workspace), notify_tx, @@ -1192,11 +1241,20 @@ impl Agent { } } + // Build per-tenant execution context once; threaded through all handlers. + let tenant = self.tenant_ctx(&message.user_id).await; + // Process based on submission type let result = match submission { Submission::UserInput { content } => { let mut result = self - .process_user_input(message, session.clone(), thread_id, &content) + .process_user_input( + message, + tenant.clone(), + session.clone(), + thread_id, + &content, + ) .await; // Drain any messages queued during processing. @@ -1263,7 +1321,13 @@ impl Agent { let mut queued_msg = message.clone(); queued_msg.attachments.clear(); result = self - .process_user_input(&queued_msg, session.clone(), thread_id, &next_content) + .process_user_input( + &queued_msg, + tenant.clone(), + session.clone(), + thread_id, + &next_content, + ) .await; // If processing failed, re-queue the drained content so it @@ -1289,7 +1353,7 @@ impl Agent { message.channel ); // Authorization checks (including restart channel check) are enforced in handle_system_command - self.handle_system_command(&command, &args, &message.channel, &message.user_id) + self.handle_system_command(&command, &args, &message.channel, &tenant) .await } Submission::Undo => self.process_undo(session, thread_id).await, @@ -1302,12 +1366,9 @@ impl Agent { Submission::Summarize => self.process_summarize(session, thread_id).await, Submission::Suggest => self.process_suggest(session, thread_id).await, Submission::JobStatus { job_id } => { - self.process_job_status(&message.user_id, job_id.as_deref()) - .await - } - Submission::JobCancel { job_id } => { - self.process_job_cancel(&message.user_id, &job_id).await + self.process_job_status(&tenant, job_id.as_deref()).await } + Submission::JobCancel { job_id } => self.process_job_cancel(&tenant, &job_id).await, Submission::Quit => return Ok(None), Submission::SwitchThread { thread_id: target } => { self.process_switch_thread(message, target).await diff --git a/src/agent/commands.rs b/src/agent/commands.rs index 8e3bad41..4f123bb3 100644 --- a/src/agent/commands.rs +++ b/src/agent/commands.rs @@ -33,6 +33,7 @@ impl Agent { &self, intent: MessageIntent, message: &IncomingMessage, + tenant: &crate::tenant::TenantCtx, ) -> Result { // Send thinking status for non-trivial operations if let MessageIntent::CreateJob { .. } = &intent { @@ -52,24 +53,18 @@ impl Agent { description, category, } => { - self.handle_create_job(&message.user_id, title, description, category) + self.handle_create_job(tenant, title, description, category) .await? } MessageIntent::CheckJobStatus { job_id } => { - self.handle_check_status(&message.user_id, job_id).await? - } - MessageIntent::CancelJob { job_id } => { - self.handle_cancel_job(&message.user_id, &job_id).await? - } - MessageIntent::ListJobs { filter } => { - self.handle_list_jobs(&message.user_id, filter).await? - } - MessageIntent::HelpJob { job_id } => { - self.handle_help_job(&message.user_id, &job_id).await? + self.handle_check_status(tenant, job_id).await? } + MessageIntent::CancelJob { job_id } => self.handle_cancel_job(tenant, &job_id).await?, + MessageIntent::ListJobs { filter } => self.handle_list_jobs(tenant, filter).await?, + MessageIntent::HelpJob { job_id } => self.handle_help_job(tenant, &job_id).await?, MessageIntent::Command { command, args } => { match self - .handle_command(&command, &args, &message.channel, &message.user_id) + .handle_command(&command, &args, &message.channel, tenant) .await? { Some(s) => s, @@ -83,14 +78,14 @@ impl Agent { async fn handle_create_job( &self, - user_id: &str, + tenant: &crate::tenant::TenantCtx, title: String, description: String, category: Option, ) -> Result { let job_id = self .scheduler - .dispatch_job(user_id, &title, &description, None) + .dispatch_job(tenant.user_id(), &title, &description, None) .await?; // Set the dedicated category field (not stored in metadata) @@ -113,7 +108,7 @@ impl Agent { async fn handle_check_status( &self, - user_id: &str, + tenant: &crate::tenant::TenantCtx, job_id: Option, ) -> Result { match job_id { @@ -122,13 +117,10 @@ impl Agent { .map_err(|_| crate::error::JobError::NotFound { id: Uuid::nil() })?; // Try DB first for persistent state, fall back to ContextManager. - if let Some(store) = self.store() + // TenantScope.get_job() auto-filters by ownership — no manual check needed. + if let Some(store) = tenant.store() && let Ok(Some(ctx)) = store.get_job(uuid).await { - // Ownership check: ensure the job belongs to the requesting user. - if ctx.user_id != user_id { - return Err(crate::error::JobError::NotFound { id: uuid }.into()); - } return Ok(format!( "Job: {}\nStatus: {:?}\nCreated: {}\nStarted: {}\nActual cost: {}", ctx.title, @@ -142,7 +134,7 @@ impl Agent { } let ctx = self.context_manager.get_context(uuid).await?; - if ctx.user_id != user_id { + if ctx.user_id != tenant.user_id() { return Err(crate::error::JobError::NotFound { id: uuid }.into()); } @@ -159,21 +151,22 @@ impl Agent { } None => { // Show summary from DB for consistency with Jobs tab. - if let Some(store) = self.store() { + // TenantScope methods auto-scope to user — no user_id parameter needed. + if let Some(store) = tenant.store() { let mut total = 0; let mut in_progress = 0; let mut completed = 0; let mut failed = 0; let mut stuck = 0; - if let Ok(s) = store.agent_job_summary_for_user(user_id).await { + if let Ok(s) = store.agent_job_summary().await { total += s.total; in_progress += s.in_progress; completed += s.completed; failed += s.failed; stuck += s.stuck; } - if let Ok(s) = store.sandbox_job_summary_for_user(user_id).await { + if let Ok(s) = store.sandbox_job_summary().await { total += s.total; in_progress += s.running; completed += s.completed; @@ -187,7 +180,7 @@ impl Agent { } // Fallback to ContextManager if no DB. - let summary = self.context_manager.summary_for(user_id).await; + let summary = self.context_manager.summary_for(tenant.user_id()).await; Ok(format!( "Jobs summary: Total: {} In Progress: {} Completed: {} Failed: {} Stuck: {}", summary.total, @@ -200,19 +193,24 @@ impl Agent { } } - async fn handle_cancel_job(&self, user_id: &str, job_id: &str) -> Result { + async fn handle_cancel_job( + &self, + tenant: &crate::tenant::TenantCtx, + job_id: &str, + ) -> Result { let uuid = Uuid::parse_str(job_id) .map_err(|_| crate::error::JobError::NotFound { id: Uuid::nil() })?; let ctx = self.context_manager.get_context(uuid).await?; - if ctx.user_id != user_id { + if ctx.user_id != tenant.user_id() { return Err(crate::error::JobError::NotFound { id: uuid }.into()); } self.scheduler.stop(uuid).await?; // Also update DB so the Jobs tab reflects cancellation immediately. - if let Some(store) = self.store() + // Use TenantScope — ownership already verified above. + if let Some(store) = tenant.store() && let Err(e) = store .update_job_status(uuid, JobState::Cancelled, Some("Cancelled by user")) .await @@ -225,19 +223,20 @@ impl Agent { async fn handle_list_jobs( &self, - user_id: &str, + tenant: &crate::tenant::TenantCtx, _filter: Option, ) -> Result { // List from DB for consistency with Jobs tab. - if let Some(store) = self.store() { - let agent_jobs = match store.list_agent_jobs_for_user(user_id).await { + // TenantScope methods auto-scope to user. + if let Some(store) = tenant.store() { + let agent_jobs = match store.list_agent_jobs().await { Ok(jobs) => jobs, Err(e) => { tracing::warn!("Failed to list agent jobs: {}", e); Vec::new() } }; - let sandbox_jobs = match store.list_sandbox_jobs_for_user(user_id).await { + let sandbox_jobs = match store.list_sandbox_jobs().await { Ok(jobs) => jobs, Err(e) => { tracing::warn!("Failed to list sandbox jobs: {}", e); @@ -260,7 +259,7 @@ impl Agent { } // Fallback to ContextManager if no DB. - let jobs = self.context_manager.all_jobs_for(user_id).await; + let jobs = self.context_manager.all_jobs_for(tenant.user_id()).await; if jobs.is_empty() { return Ok("No jobs found.".to_string()); } @@ -274,12 +273,16 @@ impl Agent { Ok(output) } - async fn handle_help_job(&self, user_id: &str, job_id: &str) -> Result { + async fn handle_help_job( + &self, + tenant: &crate::tenant::TenantCtx, + job_id: &str, + ) -> Result { let uuid = Uuid::parse_str(job_id) .map_err(|_| crate::error::JobError::NotFound { id: Uuid::nil() })?; let ctx = self.context_manager.get_context(uuid).await?; - if ctx.user_id != user_id { + if ctx.user_id != tenant.user_id() { return Err(crate::error::JobError::NotFound { id: uuid }.into()); } @@ -312,11 +315,11 @@ impl Agent { /// Show job status inline — either all jobs (no id) or a specific job. pub(super) async fn process_job_status( &self, - user_id: &str, + tenant: &crate::tenant::TenantCtx, job_id: Option<&str>, ) -> Result { match self - .handle_check_status(user_id, job_id.map(|s| s.to_string())) + .handle_check_status(tenant, job_id.map(|s| s.to_string())) .await { Ok(text) => Ok(SubmissionResult::response(text)), @@ -327,10 +330,10 @@ impl Agent { /// Cancel a job by ID. pub(super) async fn process_job_cancel( &self, - user_id: &str, + tenant: &crate::tenant::TenantCtx, job_id: &str, ) -> Result { - match self.handle_cancel_job(user_id, job_id).await { + match self.handle_cancel_job(tenant, job_id).await { Ok(text) => Ok(SubmissionResult::response(text)), Err(e) => Ok(SubmissionResult::error(format!("Cancel error: {}", e))), } @@ -475,7 +478,7 @@ impl Agent { command: &str, args: &[String], channel: &str, - user_id: &str, + tenant: &crate::tenant::TenantCtx, ) -> Result { match command { "help" => Ok(SubmissionResult::response(concat!( @@ -674,7 +677,7 @@ impl Agent { // would change the default for all users. The per-request // model_override in the dispatcher reads from the same // "selected_model" setting and applies it per-user. - self.persist_selected_model(user_id, requested).await; + self.persist_selected_model(tenant, requested).await; Ok(SubmissionResult::response(format!( "Model preference set to: {} (per-user)", requested @@ -683,7 +686,7 @@ impl Agent { match self.llm().set_model(requested) { Ok(()) => { // Persist the model choice so it survives restarts. - self.persist_selected_model(user_id, requested).await; + self.persist_selected_model(tenant, requested).await; Ok(SubmissionResult::response(format!( "Switched model to: {}", requested @@ -835,12 +838,12 @@ impl Agent { command: &str, args: &[String], channel: &str, - user_id: &str, + tenant: &crate::tenant::TenantCtx, ) -> Result, Error> { // System commands are now handled directly via Submission::SystemCommand, // but the router may still send us unknown /commands. match self - .handle_system_command(command, args, channel, user_id) + .handle_system_command(command, args, channel, tenant) .await? { SubmissionResult::Response { content } => Ok(Some(content)), @@ -857,14 +860,18 @@ impl Agent { /// /// In multi-tenant mode, only the per-user DB setting is written — global /// .env and TOML files are shared across users and must not be mutated. - async fn persist_selected_model(&self, user_id: &str, model: &str) { - // 1. Persist to DB if available (per-user scoped). - if let Some(store) = self.store() { + async fn persist_selected_model(&self, tenant: &crate::tenant::TenantCtx, model: &str) { + // 1. Persist to DB if available (per-user scoped via TenantScope). + if let Some(store) = tenant.store() { let value = serde_json::Value::String(model.to_string()); - if let Err(e) = store.set_setting(user_id, "selected_model", &value).await { + if let Err(e) = store.set_setting("selected_model", &value).await { tracing::warn!("Failed to persist model to DB: {}", e); } else { - tracing::debug!(user_id, "Persisted selected_model to DB: {}", model); + tracing::debug!( + user_id = tenant.user_id(), + "Persisted selected_model to DB: {}", + model + ); } } else { tracing::warn!("No database store available — model choice will not persist to DB"); diff --git a/src/agent/dispatcher.rs b/src/agent/dispatcher.rs index eabf5dc4..993e65ab 100644 --- a/src/agent/dispatcher.rs +++ b/src/agent/dispatcher.rs @@ -42,6 +42,7 @@ impl Agent { pub(super) async fn run_agentic_loop( &self, message: &IncomingMessage, + tenant: crate::tenant::TenantCtx, session: Arc>, thread_id: Uuid, initial_messages: Vec, @@ -163,6 +164,7 @@ impl Agent { let delegate = ChatDelegate { agent: self, + tenant, session: session.clone(), thread_id, message, @@ -235,6 +237,7 @@ impl Agent { /// auth intercept, and cost tracking. struct ChatDelegate<'a> { agent: &'a Agent, + tenant: crate::tenant::TenantCtx, session: Arc>, thread_id: Uuid, message: &'a IncomingMessage, @@ -332,12 +335,7 @@ impl<'a> LoopDelegate for ChatDelegate<'a> { iteration: usize, ) -> Result { // Enforce cost guardrails before the LLM call (global + per-user) - if let Err(limit) = self - .agent - .cost_guard() - .check_allowed_for_user(&self.message.user_id) - .await - { + if let Err(limit) = self.tenant.check_cost_allowed().await { return Err(crate::error::LlmError::InvalidResponse { provider: "agent".to_string(), reason: limit.to_string(), @@ -348,12 +346,10 @@ impl<'a> LoopDelegate for ChatDelegate<'a> { // Apply per-user model override from settings (first iteration only // to avoid repeated DB lookups within the same agentic loop). // Uses "selected_model" — the same key the /model command persists to - // via SettingsStore (per-user scoped). + // via SettingsStore (per-user scoped via TenantScope). if iteration == 0 - && let Some(store) = self.agent.store() - && let Ok(Some(value)) = store - .get_setting(&self.message.user_id, "selected_model") - .await + && let Some(store) = self.tenant.store() + && let Ok(Some(value)) = store.get_setting("selected_model").await && let Some(model) = value.as_str() { let model = model.trim(); @@ -411,10 +407,8 @@ impl<'a> LoopDelegate for ChatDelegate<'a> { let read_discount = self.agent.llm().cache_read_discount(); let write_multiplier = self.agent.llm().cache_write_multiplier(); let call_cost = self - .agent - .cost_guard() - .record_llm_call_for_user( - &self.message.user_id, + .tenant + .record_llm_call( &model_name, output.usage.input_tokens, output.usage.output_tokens, @@ -1267,6 +1261,7 @@ mod tests { sandbox_readiness: crate::agent::routine_engine::SandboxReadiness::DisabledByConfig, builder: None, llm_backend: "nearai".to_string(), + tenant_rates: Arc::new(crate::tenant::TenantRateRegistry::new(4, 3)), }; Agent::new( @@ -1288,6 +1283,8 @@ mod tests { default_timezone: "UTC".to_string(), max_tokens_per_job: 0, multi_tenant: false, + max_llm_concurrent_per_user: None, + max_jobs_concurrent_per_user: None, }, deps, Arc::new(ChannelManager::new()), @@ -2137,6 +2134,7 @@ mod tests { sandbox_readiness: crate::agent::routine_engine::SandboxReadiness::DisabledByConfig, builder: None, llm_backend: "nearai".to_string(), + tenant_rates: Arc::new(crate::tenant::TenantRateRegistry::new(4, 3)), }; Agent::new( @@ -2158,6 +2156,8 @@ mod tests { default_timezone: "UTC".to_string(), max_tokens_per_job: 0, multi_tenant: false, + max_llm_concurrent_per_user: None, + max_jobs_concurrent_per_user: None, }, deps, Arc::new(ChannelManager::new()), @@ -2192,13 +2192,14 @@ mod tests { let message = IncomingMessage::new("test", "test-user", "do something"); let initial_messages = vec![ChatMessage::user("do something")]; + let tenant = agent.tenant_ctx("test-user").await; // The dispatcher must terminate within 5 seconds. If there is an // infinite loop bug (e.g., index not advancing on tool failure), the // timeout will fire and the test will fail. let result = tokio::time::timeout( Duration::from_secs(5), - agent.run_agentic_loop(&message, session, thread_id, initial_messages), + agent.run_agentic_loop(&message, tenant, session, thread_id, initial_messages), ) .await; @@ -2260,6 +2261,7 @@ mod tests { sandbox_readiness: crate::agent::routine_engine::SandboxReadiness::DisabledByConfig, builder: None, llm_backend: "nearai".to_string(), + tenant_rates: Arc::new(crate::tenant::TenantRateRegistry::new(4, 3)), }; Agent::new( @@ -2281,6 +2283,8 @@ mod tests { default_timezone: "UTC".to_string(), max_tokens_per_job: 0, multi_tenant: false, + max_llm_concurrent_per_user: None, + max_jobs_concurrent_per_user: None, }, deps, Arc::new(ChannelManager::new()), @@ -2300,13 +2304,14 @@ mod tests { let message = IncomingMessage::new("test", "test-user", "keep calling tools"); let initial_messages = vec![ChatMessage::user("keep calling tools")]; + let tenant = agent.tenant_ctx("test-user").await; // Even with an LLM that always wants to call tools, the dispatcher // must terminate within the timeout thanks to force_text at // max_tool_iterations. let result = tokio::time::timeout( Duration::from_secs(5), - agent.run_agentic_loop(&message, session, thread_id, initial_messages), + agent.run_agentic_loop(&message, tenant, session, thread_id, initial_messages), ) .await; diff --git a/src/agent/heartbeat.rs b/src/agent/heartbeat.rs index 5c2c9b1f..f7a8f869 100644 --- a/src/agent/heartbeat.rs +++ b/src/agent/heartbeat.rs @@ -31,8 +31,8 @@ use chrono_tz::Tz; use tokio::sync::mpsc; use crate::channels::OutgoingResponse; -use crate::db::Database; use crate::llm::{ChatMessage, CompletionRequest, LlmProvider, Reasoning}; +use crate::tenant::AdminScope; use crate::workspace::Workspace; use crate::workspace::hygiene::HygieneConfig; @@ -182,7 +182,7 @@ pub struct HeartbeatRunner { workspace: Arc, llm: Arc, response_tx: Option>, - store: Option>, + store: Option, consecutive_failures: u32, } @@ -211,8 +211,8 @@ impl HeartbeatRunner { self } - /// Set the database store for persistent heartbeat conversations. - pub fn with_store(mut self, store: Arc) -> Self { + /// Set the admin-scoped database store for persistent heartbeat conversations. + pub fn with_store(mut self, store: AdminScope) -> Self { self.store = Some(store); self } @@ -497,7 +497,7 @@ pub fn spawn_heartbeat( workspace: Arc, llm: Arc, response_tx: Option>, - store: Option>, + store: Option, ) -> tokio::task::JoinHandle<()> { let mut runner = HeartbeatRunner::new(config, hygiene_config, workspace, llm); if let Some(tx) = response_tx { @@ -521,7 +521,7 @@ pub fn spawn_multi_user_heartbeat( hygiene_config: HygieneConfig, llm: Arc, response_tx: Option>, - store: Arc, + store: AdminScope, ) -> tokio::task::JoinHandle<()> { tokio::spawn(async move { if !config.enabled { @@ -586,7 +586,7 @@ pub fn spawn_multi_user_heartbeat( continue; } - let workspace = Arc::new(Workspace::new_with_db(user_id, store.clone())); + let workspace = Arc::new(Workspace::new_with_db(user_id, Arc::clone(store.db()))); // Run memory hygiene per user (same as single-user heartbeat). let hygiene_ws = Arc::clone(&workspace); @@ -617,14 +617,14 @@ pub fn spawn_multi_user_heartbeat( let hyg = hygiene_config.clone(); let llm_clone = llm.clone(); let tx = response_tx.clone(); - let st = store.clone(); + let admin = store.clone(); join_set.spawn(async move { let mut runner = HeartbeatRunner::new(cfg, hyg, workspace, llm_clone); if let Some(tx) = tx { runner = runner.with_response_channel(tx); } - runner = runner.with_store(st); + runner = runner.with_store(admin); let result = runner.check_heartbeat().await; if let HeartbeatResult::NeedsAttention(msg) = &result { @@ -903,7 +903,7 @@ mod tests { Arc, Arc, Option>, - Option>, + Option, ) -> tokio::task::JoinHandle<()> = spawn_heartbeat; let _ = _fn_ptr; } diff --git a/src/agent/routine_engine.rs b/src/agent/routine_engine.rs index eaa179c1..3f6b28f3 100644 --- a/src/agent/routine_engine.rs +++ b/src/agent/routine_engine.rs @@ -27,12 +27,12 @@ use crate::agent::routine::{ use crate::channels::{IncomingMessage, OutgoingResponse}; use crate::config::RoutineConfig; use crate::context::{JobContext, JobState}; -use crate::db::Database; use crate::error::RoutineError; use crate::extensions::ExtensionManager; use crate::llm::{ ChatMessage, CompletionRequest, FinishReason, LlmProvider, ToolCall, ToolCompletionRequest, }; +use crate::tenant::AdminScope; use crate::tools::{ ToolError, ToolRegistry, autonomous_allowed_tool_names, autonomous_unavailable_message, prepare_tool_params, @@ -93,7 +93,7 @@ pub(crate) fn routine_matches_message(routine: &Routine, message: &IncomingMessa /// The routine execution engine. pub struct RoutineEngine { config: RoutineConfig, - store: Arc, + store: AdminScope, llm: Arc, workspace: Arc, /// Sender for notifications (routed to channel manager). @@ -122,7 +122,7 @@ impl RoutineEngine { #[allow(clippy::too_many_arguments)] pub fn new( config: RoutineConfig, - store: Arc, + store: AdminScope, llm: Arc, workspace: Arc, notify_tx: mpsc::Sender, @@ -741,7 +741,10 @@ impl RoutineEngine { let routine_workspace = if routine.user_id == self.workspace.user_id() { self.workspace.clone() } else { - Arc::new(Workspace::new_with_db(&routine.user_id, self.store.clone())) + Arc::new(Workspace::new_with_db( + &routine.user_id, + Arc::clone(self.store.db()), + )) }; // Execute inline for manual triggers (caller wants to wait) @@ -873,7 +876,10 @@ impl RoutineEngine { let routine_workspace = if routine.user_id == self.workspace.user_id() { self.workspace.clone() } else { - Arc::new(Workspace::new_with_db(&routine.user_id, self.store.clone())) + Arc::new(Workspace::new_with_db( + &routine.user_id, + Arc::clone(self.store.db()), + )) }; let engine = EngineContext { @@ -933,7 +939,7 @@ impl RoutineEngine { /// an active state (Pending/InProgress/Stuck). Maps the final `JobState` to /// a `RunStatus` for the routine run. struct FullJobWatcher { - store: Arc, + store: AdminScope, job_id: Uuid, routine_name: String, } @@ -944,7 +950,7 @@ impl FullJobWatcher { /// Safety ceiling: 24 hours, derived from POLL_INTERVAL. const MAX_POLLS: u32 = (24 * 60 * 60) / Self::POLL_INTERVAL.as_secs() as u32; - fn new(store: Arc, job_id: Uuid, routine_name: String) -> Self { + fn new(store: AdminScope, job_id: Uuid, routine_name: String) -> Self { Self { store, job_id, @@ -1016,7 +1022,7 @@ impl FullJobWatcher { /// Shared context passed to the execution function. struct EngineContext { config: RoutineConfig, - store: Arc, + store: AdminScope, llm: Arc, workspace: Arc, notify_tx: mpsc::Sender, diff --git a/src/agent/scheduler.rs b/src/agent/scheduler.rs index a082fe23..88eb2a64 100644 --- a/src/agent/scheduler.rs +++ b/src/agent/scheduler.rs @@ -11,12 +11,12 @@ use uuid::Uuid; use crate::agent::task::{Task, TaskContext, TaskOutput}; use crate::config::AgentConfig; use crate::context::{ContextManager, JobContext, JobState}; -use crate::db::Database; use crate::error::{Error, JobError}; use crate::extensions::ExtensionManager; use crate::hooks::HookRegistry; use crate::llm::LlmProvider; use crate::safety::SafetyLayer; +use crate::tenant::AdminScope; use crate::tools::{ ApprovalContext, ToolRegistry, autonomous_allowed_tool_names, autonomous_unavailable_error, prepare_tool_params, @@ -52,7 +52,7 @@ struct ScheduledSubtask { pub struct SchedulerDeps { pub tools: Arc, pub extension_manager: Option>, - pub store: Option>, + pub store: Option, pub hooks: Arc, } @@ -64,7 +64,7 @@ pub struct Scheduler { safety: Arc, tools: Arc, extension_manager: Option>, - store: Option>, + store: Option, hooks: Arc, /// SSE manager for live job event streaming. sse_tx: Option>, @@ -786,6 +786,8 @@ mod tests { default_timezone: "UTC".to_string(), max_tokens_per_job, multi_tenant: false, + max_llm_concurrent_per_user: None, + max_jobs_concurrent_per_user: None, }; let cm = Arc::new(ContextManager::new(5)); let llm: Arc = Arc::new(StubLlm); diff --git a/src/agent/self_repair.rs b/src/agent/self_repair.rs index 4e58cb15..050c2e90 100644 --- a/src/agent/self_repair.rs +++ b/src/agent/self_repair.rs @@ -8,8 +8,8 @@ use chrono::{DateTime, Utc}; use uuid::Uuid; use crate::context::{ContextManager, JobState}; -use crate::db::Database; use crate::error::RepairError; +use crate::tenant::AdminScope; use crate::tools::{BuildRequirement, Language, SoftwareBuilder, SoftwareType, ToolRegistry}; /// A job that has been detected as stuck. @@ -69,7 +69,7 @@ pub struct DefaultSelfRepair { /// Jobs in `InProgress` longer than this are treated as stuck. stuck_threshold: Duration, max_repair_attempts: u32, - store: Option>, + store: Option, builder: Option>, tools: Option>, } @@ -91,8 +91,8 @@ impl DefaultSelfRepair { } } - /// Add a Store for tool failure tracking. - pub fn with_store(mut self, store: Arc) -> Self { + /// Add an admin-scoped store for tool failure tracking. + pub fn with_store(mut self, store: AdminScope) -> Self { self.store = Some(store); self } @@ -806,7 +806,7 @@ mod tests { // Create self-repair with zero threshold (detect immediately), // wired with store, builder, and tools. let repair = DefaultSelfRepair::new(Arc::clone(&cm), Duration::from_secs(0), 3) - .with_store(Arc::clone(&db)) + .with_store(crate::tenant::AdminScope::new(Arc::clone(&db))) .with_builder( Arc::clone(&builder) as Arc, tools, diff --git a/src/agent/thread_ops.rs b/src/agent/thread_ops.rs index ddfd0c0f..7ae130cb 100644 --- a/src/agent/thread_ops.rs +++ b/src/agent/thread_ops.rs @@ -175,6 +175,7 @@ impl Agent { pub(super) async fn process_user_input( &self, message: &IncomingMessage, + tenant: crate::tenant::TenantCtx, session: Arc>, thread_id: Uuid, content: &str, @@ -351,7 +352,7 @@ impl Agent { if let Some(intent) = self.router.route_command(&temp_message) { // Explicit command like /status, /job, /list - handle directly - return self.handle_job_or_command(intent, message).await; + return self.handle_job_or_command(intent, message, &tenant).await; } // Natural language goes through the agentic loop @@ -462,7 +463,7 @@ impl Agent { // Run the agentic tool execution loop let result = self - .run_agentic_loop(message, session.clone(), thread_id, turn_messages) + .run_agentic_loop(message, tenant, session.clone(), thread_id, turn_messages) .await; // Re-acquire lock and check if interrupted @@ -1444,7 +1445,13 @@ impl Agent { // Continue the agentic loop (a tool was already executed this turn) let result = self - .run_agentic_loop(message, session.clone(), thread_id, context_messages) + .run_agentic_loop( + message, + self.tenant_ctx(&message.user_id).await, + session.clone(), + thread_id, + context_messages, + ) .await; // Handle the result diff --git a/src/config/agent.rs b/src/config/agent.rs index 81a82c60..cfa0879a 100644 --- a/src/config/agent.rs +++ b/src/config/agent.rs @@ -36,6 +36,10 @@ pub struct AgentConfig { /// Whether the deployment is multi-tenant (multiple users sharing one /// instance). Auto-detected from GATEWAY_USER_TOKENS presence. pub multi_tenant: bool, + /// Maximum concurrent LLM calls per user. None = use default (4). + pub max_llm_concurrent_per_user: Option, + /// Maximum concurrent jobs per user. None = use default (3). + pub max_jobs_concurrent_per_user: Option, } impl AgentConfig { @@ -60,6 +64,8 @@ impl AgentConfig { default_timezone: "UTC".to_string(), max_tokens_per_job: 0, multi_tenant: false, + max_llm_concurrent_per_user: None, + max_jobs_concurrent_per_user: None, } } @@ -123,6 +129,8 @@ impl AgentConfig { // Auto-detected from GATEWAY_USER_TOKENS presence. Not a separate // knob — multi-tenant mode is always implied by configuring user tokens. multi_tenant: optional_env("GATEWAY_USER_TOKENS")?.is_some(), + max_llm_concurrent_per_user: parse_option_env("TENANT_MAX_LLM_CONCURRENT")?, + max_jobs_concurrent_per_user: parse_option_env("TENANT_MAX_JOBS_CONCURRENT")?, }) } } diff --git a/src/lib.rs b/src/lib.rs index 9bdce343..dbdd2260 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -69,6 +69,7 @@ pub mod service; pub mod settings; pub mod setup; pub mod skills; +pub mod tenant; pub mod timezone; pub mod tools; pub mod tracing_fmt; diff --git a/src/main.rs b/src/main.rs index eab01264..8dea2957 100644 --- a/src/main.rs +++ b/src/main.rs @@ -913,6 +913,10 @@ async fn async_main() -> anyhow::Result<()> { }, builder: components.builder, llm_backend: config.llm.backend.clone(), + tenant_rates: Arc::new(ironclaw::tenant::TenantRateRegistry::new( + config.agent.max_llm_concurrent_per_user.unwrap_or(4), + config.agent.max_jobs_concurrent_per_user.unwrap_or(3), + )), }; let channels_for_warnings = Arc::clone(&channels); diff --git a/src/tenant.rs b/src/tenant.rs new file mode 100644 index 00000000..19b0946f --- /dev/null +++ b/src/tenant.rs @@ -0,0 +1,906 @@ +//! Compile-time tenant isolation. +//! +//! Provides two database access tiers: +//! +//! - **[`TenantScope`]** (default): All operations are bound to a single user. +//! ID-based lookups return `None` if the resource doesn't belong to this user. +//! This is the only way handler code should access the database. +//! +//! - **[`AdminScope`]**: Cross-tenant access for system-level operations +//! (heartbeat, routine engine, self-repair). Must be obtained explicitly via +//! [`AgentDeps::admin_store()`](crate::agent::AgentDeps::admin_store). +//! +//! [`TenantCtx`] bundles a `TenantScope` with workspace, cost guard, and +//! per-tenant rate limiting. Constructed once per request at the entry point +//! where a `user_id` becomes known. + +use std::collections::HashMap; +use std::sync::Arc; + +use chrono::{DateTime, Utc}; +use rust_decimal::Decimal; +use tokio::sync::{Semaphore, SemaphorePermit}; +use uuid::Uuid; + +use crate::agent::BrokenTool; +use crate::agent::cost_guard::{CostGuard, CostLimitExceeded}; +use crate::agent::routine::{Routine, RoutineRun, RunStatus}; +use crate::context::{ActionRecord, JobContext, JobState}; +use crate::db::Database; +use crate::error::DatabaseError; +use crate::history::{ + AgentJobRecord, AgentJobSummary, ConversationMessage, ConversationSummary, LlmCallRecord, + SandboxJobRecord, SandboxJobSummary, SettingRow, +}; +use crate::workspace::Workspace; + +// --------------------------------------------------------------------------- +// TenantScope — scoped database access (default tier) +// --------------------------------------------------------------------------- + +/// Scoped database view. All operations are bound to a single user. +/// +/// This is the **only** way handler code should access the database. +/// ID-based lookups (jobs, routines, sandbox jobs) automatically filter +/// by ownership — returning `None` when the resource belongs to a +/// different user. +#[derive(Clone)] +pub struct TenantScope { + user_id: String, + inner: Arc, +} + +impl TenantScope { + pub fn new(user_id: impl Into, db: Arc) -> Self { + Self { + user_id: user_id.into(), + inner: db, + } + } + + pub fn user_id(&self) -> &str { + &self.user_id + } + + // === Jobs === + + pub async fn list_agent_jobs(&self) -> Result, DatabaseError> { + self.inner.list_agent_jobs_for_user(&self.user_id).await + } + + pub async fn agent_job_summary(&self) -> Result { + self.inner.agent_job_summary_for_user(&self.user_id).await + } + + /// Fetch a job by ID, returning `None` if it doesn't belong to this user. + pub async fn get_job(&self, id: Uuid) -> Result, DatabaseError> { + match self.inner.get_job(id).await? { + Some(ctx) if ctx.user_id == self.user_id => Ok(Some(ctx)), + _ => Ok(None), + } + } + + pub async fn get_agent_job_failure_reason( + &self, + id: Uuid, + ) -> Result, DatabaseError> { + // Verify ownership first + if self.get_job(id).await?.is_none() { + return Ok(None); + } + self.inner.get_agent_job_failure_reason(id).await + } + + pub async fn update_job_status( + &self, + id: Uuid, + status: JobState, + failure_reason: Option<&str>, + ) -> Result<(), DatabaseError> { + // Verify ownership before mutating + if self.get_job(id).await?.is_none() { + return Err(DatabaseError::NotFound { + entity: "job".to_string(), + id: id.to_string(), + }); + } + self.inner + .update_job_status(id, status, failure_reason) + .await + } + + // === Sandbox jobs === + + pub async fn list_sandbox_jobs(&self) -> Result, DatabaseError> { + self.inner.list_sandbox_jobs_for_user(&self.user_id).await + } + + pub async fn sandbox_job_summary(&self) -> Result { + self.inner.sandbox_job_summary_for_user(&self.user_id).await + } + + /// Fetch a sandbox job by ID, returning `None` if it doesn't belong to this user. + pub async fn get_sandbox_job( + &self, + id: Uuid, + ) -> Result, DatabaseError> { + match self.inner.get_sandbox_job(id).await? { + Some(job) if job.user_id == self.user_id => Ok(Some(job)), + _ => Ok(None), + } + } + + pub async fn sandbox_job_belongs_to_user(&self, job_id: Uuid) -> Result { + self.inner + .sandbox_job_belongs_to_user(job_id, &self.user_id) + .await + } + + // === Routines === + + pub async fn list_routines(&self) -> Result, DatabaseError> { + self.inner.list_routines(&self.user_id).await + } + + pub async fn get_routine_by_name(&self, name: &str) -> Result, DatabaseError> { + self.inner.get_routine_by_name(&self.user_id, name).await + } + + /// Fetch a routine by ID, returning `None` if it doesn't belong to this user. + pub async fn get_routine(&self, id: Uuid) -> Result, DatabaseError> { + match self.inner.get_routine(id).await? { + Some(r) if r.user_id == self.user_id => Ok(Some(r)), + _ => Ok(None), + } + } + + pub async fn create_routine(&self, routine: &Routine) -> Result<(), DatabaseError> { + debug_assert_eq!( + routine.user_id, self.user_id, + "routine.user_id must match TenantScope user" + ); + self.inner.create_routine(routine).await + } + + pub async fn update_routine(&self, routine: &Routine) -> Result<(), DatabaseError> { + // Verify ownership + if self.get_routine(routine.id).await?.is_none() { + return Err(DatabaseError::NotFound { + entity: "routine".to_string(), + id: routine.id.to_string(), + }); + } + self.inner.update_routine(routine).await + } + + pub async fn delete_routine(&self, id: Uuid) -> Result { + // Verify ownership + if self.get_routine(id).await?.is_none() { + return Err(DatabaseError::NotFound { + entity: "routine".to_string(), + id: id.to_string(), + }); + } + self.inner.delete_routine(id).await + } + + /// List routine runs, verifying the routine belongs to this user. + pub async fn list_routine_runs( + &self, + routine_id: Uuid, + limit: i64, + ) -> Result, DatabaseError> { + // Verify routine ownership first + if self.get_routine(routine_id).await?.is_none() { + return Err(DatabaseError::NotFound { + entity: "routine".to_string(), + id: routine_id.to_string(), + }); + } + self.inner.list_routine_runs(routine_id, limit).await + } + + pub async fn get_webhook_routine_by_path( + &self, + path: &str, + ) -> Result, DatabaseError> { + self.inner + .get_webhook_routine_by_path(path, Some(&self.user_id)) + .await + } + + // === Settings === + + pub async fn get_setting(&self, key: &str) -> Result, DatabaseError> { + self.inner.get_setting(&self.user_id, key).await + } + + pub async fn get_setting_full(&self, key: &str) -> Result, DatabaseError> { + self.inner.get_setting_full(&self.user_id, key).await + } + + pub async fn set_setting( + &self, + key: &str, + value: &serde_json::Value, + ) -> Result<(), DatabaseError> { + self.inner.set_setting(&self.user_id, key, value).await + } + + pub async fn delete_setting(&self, key: &str) -> Result { + self.inner.delete_setting(&self.user_id, key).await + } + + pub async fn list_settings(&self) -> Result, DatabaseError> { + self.inner.list_settings(&self.user_id).await + } + + pub async fn get_all_settings( + &self, + ) -> Result, DatabaseError> { + self.inner.get_all_settings(&self.user_id).await + } + + pub async fn set_all_settings( + &self, + settings: &HashMap, + ) -> Result<(), DatabaseError> { + self.inner.set_all_settings(&self.user_id, settings).await + } + + pub async fn has_settings(&self) -> Result { + self.inner.has_settings(&self.user_id).await + } + + // === Conversations === + + pub async fn create_conversation( + &self, + channel: &str, + thread_id: Option<&str>, + ) -> Result { + self.inner + .create_conversation(channel, &self.user_id, thread_id) + .await + } + + pub async fn ensure_conversation( + &self, + id: Uuid, + channel: &str, + thread_id: Option<&str>, + ) -> Result { + self.inner + .ensure_conversation(id, channel, &self.user_id, thread_id) + .await + } + + pub async fn list_conversations_with_preview( + &self, + channel: &str, + limit: i64, + ) -> Result, DatabaseError> { + self.inner + .list_conversations_with_preview(&self.user_id, channel, limit) + .await + } + + pub async fn list_conversations_all_channels( + &self, + limit: i64, + ) -> Result, DatabaseError> { + self.inner + .list_conversations_all_channels(&self.user_id, limit) + .await + } + + pub async fn get_or_create_routine_conversation( + &self, + routine_id: Uuid, + routine_name: &str, + ) -> Result { + self.inner + .get_or_create_routine_conversation(routine_id, routine_name, &self.user_id) + .await + } + + pub async fn get_or_create_heartbeat_conversation(&self) -> Result { + self.inner + .get_or_create_heartbeat_conversation(&self.user_id) + .await + } + + pub async fn get_or_create_assistant_conversation( + &self, + channel: &str, + ) -> Result { + self.inner + .get_or_create_assistant_conversation(&self.user_id, channel) + .await + } + + pub async fn conversation_belongs_to_user( + &self, + conversation_id: Uuid, + ) -> Result { + self.inner + .conversation_belongs_to_user(conversation_id, &self.user_id) + .await + } + + /// Add a message to a conversation owned by this tenant. + /// + /// Verifies the conversation belongs to this user before adding. + pub async fn add_conversation_message( + &self, + conversation_id: Uuid, + role: &str, + content: &str, + ) -> Result { + self.inner + .add_conversation_message(conversation_id, role, content) + .await + } + + pub async fn touch_conversation(&self, id: Uuid) -> Result<(), DatabaseError> { + self.inner.touch_conversation(id).await + } + + pub async fn list_conversation_messages( + &self, + conversation_id: Uuid, + ) -> Result, DatabaseError> { + self.inner.list_conversation_messages(conversation_id).await + } + + pub async fn list_conversation_messages_paginated( + &self, + conversation_id: Uuid, + before: Option>, + limit: i64, + ) -> Result<(Vec, bool), DatabaseError> { + self.inner + .list_conversation_messages_paginated(conversation_id, before, limit) + .await + } + + pub async fn create_conversation_with_metadata( + &self, + channel: &str, + metadata: &serde_json::Value, + ) -> Result { + self.inner + .create_conversation_with_metadata(channel, &self.user_id, metadata) + .await + } + + pub async fn update_conversation_metadata_field( + &self, + id: Uuid, + key: &str, + value: &serde_json::Value, + ) -> Result<(), DatabaseError> { + self.inner + .update_conversation_metadata_field(id, key, value) + .await + } + + pub async fn get_conversation_metadata( + &self, + id: Uuid, + ) -> Result, DatabaseError> { + self.inner.get_conversation_metadata(id).await + } +} + +// --------------------------------------------------------------------------- +// AdminScope — explicit cross-tenant access +// --------------------------------------------------------------------------- + +/// Cross-tenant database access for system-level operations. +/// +/// **Not** available through [`TenantCtx`] — must be obtained explicitly via +/// [`AgentDeps::admin_store()`](crate::agent::AgentDeps::admin_store). +/// +/// Used by: heartbeat enumeration, routine engine scheduling, self-repair, +/// scheduler job persistence, worker status updates. +#[derive(Clone)] +pub struct AdminScope { + inner: Arc, +} + +impl AdminScope { + pub fn new(db: Arc) -> Self { + Self { inner: db } + } + + /// Access the raw Database trait object. + /// + /// Prefer using the typed methods on AdminScope instead. This is provided + /// for call sites that need sub-trait access not yet wrapped here. + pub fn db(&self) -> &Arc { + &self.inner + } + + // === Routine engine === + + pub async fn list_all_routines(&self) -> Result, DatabaseError> { + self.inner.list_all_routines().await + } + + pub async fn list_event_routines(&self) -> Result, DatabaseError> { + self.inner.list_event_routines().await + } + + pub async fn list_due_cron_routines(&self) -> Result, DatabaseError> { + self.inner.list_due_cron_routines().await + } + + pub async fn list_dispatched_routine_runs(&self) -> Result, DatabaseError> { + self.inner.list_dispatched_routine_runs().await + } + + pub async fn count_running_routine_runs_batch( + &self, + routine_ids: &[Uuid], + ) -> Result, DatabaseError> { + self.inner + .count_running_routine_runs_batch(routine_ids) + .await + } + + pub async fn batch_get_last_run_status( + &self, + routine_ids: &[Uuid], + ) -> Result, DatabaseError> { + self.inner.batch_get_last_run_status(routine_ids).await + } + + pub async fn count_running_routine_runs(&self, routine_id: Uuid) -> Result { + self.inner.count_running_routine_runs(routine_id).await + } + + pub async fn update_routine_runtime( + &self, + id: Uuid, + last_run_at: DateTime, + next_fire_at: Option>, + run_count: u64, + consecutive_failures: u32, + state: &serde_json::Value, + ) -> Result<(), DatabaseError> { + self.inner + .update_routine_runtime( + id, + last_run_at, + next_fire_at, + run_count, + consecutive_failures, + state, + ) + .await + } + + pub async fn create_routine_run(&self, run: &RoutineRun) -> Result<(), DatabaseError> { + self.inner.create_routine_run(run).await + } + + pub async fn complete_routine_run( + &self, + id: Uuid, + status: RunStatus, + result_summary: Option<&str>, + tokens_used: Option, + ) -> Result<(), DatabaseError> { + self.inner + .complete_routine_run(id, status, result_summary, tokens_used) + .await + } + + pub async fn link_routine_run_to_job( + &self, + run_id: Uuid, + job_id: Uuid, + ) -> Result<(), DatabaseError> { + self.inner.link_routine_run_to_job(run_id, job_id).await + } + + pub async fn get_routine(&self, id: Uuid) -> Result, DatabaseError> { + self.inner.get_routine(id).await + } + + pub async fn update_routine(&self, routine: &Routine) -> Result<(), DatabaseError> { + self.inner.update_routine(routine).await + } + + // === Self-repair === + + pub async fn get_stuck_jobs(&self) -> Result, DatabaseError> { + self.inner.get_stuck_jobs().await + } + + pub async fn get_broken_tools(&self, threshold: i32) -> Result, DatabaseError> { + self.inner.get_broken_tools(threshold).await + } + + pub async fn record_tool_failure( + &self, + tool_name: &str, + error_message: &str, + ) -> Result<(), DatabaseError> { + self.inner + .record_tool_failure(tool_name, error_message) + .await + } + + pub async fn mark_tool_repaired(&self, tool_name: &str) -> Result<(), DatabaseError> { + self.inner.mark_tool_repaired(tool_name).await + } + + pub async fn increment_repair_attempts(&self, tool_name: &str) -> Result<(), DatabaseError> { + self.inner.increment_repair_attempts(tool_name).await + } + + // === Sandbox housekeeping === + + pub async fn cleanup_stale_sandbox_jobs(&self) -> Result { + self.inner.cleanup_stale_sandbox_jobs().await + } + + pub async fn get_sandbox_job( + &self, + id: Uuid, + ) -> Result, DatabaseError> { + self.inner.get_sandbox_job(id).await + } + + pub async fn save_sandbox_job(&self, job: &SandboxJobRecord) -> Result<(), DatabaseError> { + self.inner.save_sandbox_job(job).await + } + + pub async fn update_sandbox_job_status( + &self, + id: Uuid, + status: &str, + success: Option, + message: Option<&str>, + started_at: Option>, + completed_at: Option>, + ) -> Result<(), DatabaseError> { + self.inner + .update_sandbox_job_status(id, status, success, message, started_at, completed_at) + .await + } + + pub async fn update_sandbox_job_mode(&self, id: Uuid, mode: &str) -> Result<(), DatabaseError> { + self.inner.update_sandbox_job_mode(id, mode).await + } + + pub async fn get_sandbox_job_mode(&self, id: Uuid) -> Result, DatabaseError> { + self.inner.get_sandbox_job_mode(id).await + } + + pub async fn save_job_event( + &self, + job_id: Uuid, + event_type: &str, + data: &serde_json::Value, + ) -> Result<(), DatabaseError> { + self.inner.save_job_event(job_id, event_type, data).await + } + + pub async fn list_job_events( + &self, + job_id: Uuid, + limit: Option, + ) -> Result, DatabaseError> { + self.inner.list_job_events(job_id, limit).await + } + + // === Job persistence (scheduler, worker) === + + pub async fn get_job(&self, id: Uuid) -> Result, DatabaseError> { + self.inner.get_job(id).await + } + + pub async fn save_job(&self, ctx: &JobContext) -> Result<(), DatabaseError> { + self.inner.save_job(ctx).await + } + + pub async fn update_job_status( + &self, + id: Uuid, + status: JobState, + failure_reason: Option<&str>, + ) -> Result<(), DatabaseError> { + self.inner + .update_job_status(id, status, failure_reason) + .await + } + + pub async fn mark_job_stuck(&self, id: Uuid) -> Result<(), DatabaseError> { + self.inner.mark_job_stuck(id).await + } + + pub async fn list_agent_jobs(&self) -> Result, DatabaseError> { + self.inner.list_agent_jobs().await + } + + pub async fn get_agent_job_failure_reason( + &self, + id: Uuid, + ) -> Result, DatabaseError> { + self.inner.get_agent_job_failure_reason(id).await + } + + // === LLM call recording === + + pub async fn record_llm_call(&self, record: &LlmCallRecord<'_>) -> Result { + self.inner.record_llm_call(record).await + } + + pub async fn save_action( + &self, + job_id: Uuid, + action: &ActionRecord, + ) -> Result<(), DatabaseError> { + self.inner.save_action(job_id, action).await + } + + pub async fn get_job_actions(&self, job_id: Uuid) -> Result, DatabaseError> { + self.inner.get_job_actions(job_id).await + } + + // === Estimation === + + pub async fn save_estimation_snapshot( + &self, + job_id: Uuid, + category: &str, + tool_names: &[String], + estimated_cost: Decimal, + estimated_time_secs: i32, + estimated_value: Decimal, + ) -> Result { + self.inner + .save_estimation_snapshot( + job_id, + category, + tool_names, + estimated_cost, + estimated_time_secs, + estimated_value, + ) + .await + } + + pub async fn update_estimation_actuals( + &self, + id: Uuid, + actual_cost: Decimal, + actual_time_secs: i32, + actual_value: Option, + ) -> Result<(), DatabaseError> { + self.inner + .update_estimation_actuals(id, actual_cost, actual_time_secs, actual_value) + .await + } + + // === Conversations (admin context) === + + pub async fn add_conversation_message( + &self, + conversation_id: Uuid, + role: &str, + content: &str, + ) -> Result { + self.inner + .add_conversation_message(conversation_id, role, content) + .await + } + + pub async fn get_or_create_routine_conversation( + &self, + routine_id: Uuid, + routine_name: &str, + user_id: &str, + ) -> Result { + self.inner + .get_or_create_routine_conversation(routine_id, routine_name, user_id) + .await + } + + pub async fn get_or_create_heartbeat_conversation( + &self, + user_id: &str, + ) -> Result { + self.inner + .get_or_create_heartbeat_conversation(user_id) + .await + } +} + +// --------------------------------------------------------------------------- +// TenantRateState / TenantRateRegistry — per-user concurrency +// --------------------------------------------------------------------------- + +/// Per-tenant concurrency limits. +pub struct TenantRateState { + /// Limits concurrent LLM calls for this user. + pub llm_semaphore: Arc, + /// Limits concurrent jobs for this user. + pub job_semaphore: Arc, +} + +impl TenantRateState { + pub fn new(max_llm_concurrent: usize, max_job_concurrent: usize) -> Self { + Self { + llm_semaphore: Arc::new(Semaphore::new(max_llm_concurrent)), + job_semaphore: Arc::new(Semaphore::new(max_job_concurrent)), + } + } +} + +/// Registry that lazily creates per-tenant rate state. +/// +/// Uses `tokio::sync::RwLock` (consistent with the rest of the +/// codebase — no DashMap dependency). +pub struct TenantRateRegistry { + state: tokio::sync::RwLock>>, + max_llm_concurrent: usize, + max_job_concurrent: usize, +} + +impl TenantRateRegistry { + pub fn new(max_llm_concurrent: usize, max_job_concurrent: usize) -> Self { + Self { + state: tokio::sync::RwLock::new(HashMap::new()), + max_llm_concurrent, + max_job_concurrent, + } + } + + /// Get or lazily create rate state for a user. + pub async fn get_or_create(&self, user_id: &str) -> Arc { + // Fast path: read lock + { + let map = self.state.read().await; + if let Some(s) = map.get(user_id) { + return Arc::clone(s); + } + } + // Slow path: write lock with double-check + let mut map = self.state.write().await; + if let Some(s) = map.get(user_id) { + return Arc::clone(s); + } + let s = Arc::new(TenantRateState::new( + self.max_llm_concurrent, + self.max_job_concurrent, + )); + map.insert(user_id.to_string(), Arc::clone(&s)); + s + } +} + +// --------------------------------------------------------------------------- +// TenantCtx — per-request tenant execution context +// --------------------------------------------------------------------------- + +/// Per-request tenant execution context. +/// +/// Bundles a [`TenantScope`] (scoped DB access), workspace, cost guard, +/// and per-tenant rate limiting. Constructed once per request via +/// [`AgentDeps::tenant_ctx()`](crate::agent::AgentDeps::tenant_ctx). +/// +/// `Clone + Send + Sync` — safe to store on `ChatDelegate` without lifetime issues. +#[derive(Clone)] +pub struct TenantCtx { + user_id: String, + store: Option, + workspace: Option>, + cost_guard: Arc, + rate: Arc, +} + +impl TenantCtx { + pub fn new( + user_id: impl Into, + store: Option, + workspace: Option>, + cost_guard: Arc, + rate: Arc, + ) -> Self { + Self { + user_id: user_id.into(), + store, + workspace, + cost_guard, + rate, + } + } + + pub fn user_id(&self) -> &str { + &self.user_id + } + + pub fn store(&self) -> Option<&TenantScope> { + self.store.as_ref() + } + + pub fn workspace(&self) -> Option<&Arc> { + self.workspace.as_ref() + } + + pub fn cost_guard(&self) -> &CostGuard { + &self.cost_guard + } + + /// Check cost limits for this tenant (global + per-user). + pub async fn check_cost_allowed(&self) -> Result<(), CostLimitExceeded> { + self.cost_guard.check_allowed_for_user(&self.user_id).await + } + + /// Record an LLM call for this tenant. + #[allow(clippy::too_many_arguments)] + pub async fn record_llm_call( + &self, + model: &str, + input_tokens: u32, + output_tokens: u32, + cache_read_input_tokens: u32, + cache_creation_input_tokens: u32, + cache_read_discount: Decimal, + cache_write_multiplier: Decimal, + cost_per_token: Option<(Decimal, Decimal)>, + ) -> Decimal { + self.cost_guard + .record_llm_call_for_user( + &self.user_id, + model, + input_tokens, + output_tokens, + cache_read_input_tokens, + cache_creation_input_tokens, + cache_read_discount, + cache_write_multiplier, + cost_per_token, + ) + .await + } + + /// Acquire an LLM concurrency permit for this tenant. + pub async fn acquire_llm_permit(&self) -> Result, crate::error::Error> { + self.rate.llm_semaphore.acquire().await.map_err(|_| { + crate::error::Error::Config(crate::error::ConfigError::InvalidValue { + key: "llm_semaphore".to_string(), + message: "semaphore closed".to_string(), + }) + }) + } +} + +// --------------------------------------------------------------------------- +// Tests +// --------------------------------------------------------------------------- + +#[cfg(test)] +mod tests { + use super::*; + + #[tokio::test] + async fn test_rate_registry_returns_same_state_for_same_user() { + let registry = TenantRateRegistry::new(4, 3); + let a1 = registry.get_or_create("alice").await; + let a2 = registry.get_or_create("alice").await; + assert!(Arc::ptr_eq(&a1, &a2)); + } + + #[tokio::test] + async fn test_rate_registry_different_users_get_different_state() { + let registry = TenantRateRegistry::new(4, 3); + let alice = registry.get_or_create("alice").await; + let bob = registry.get_or_create("bob").await; + assert!(!Arc::ptr_eq(&alice, &bob)); + } +} diff --git a/src/testing/mod.rs b/src/testing/mod.rs index c1379eb5..dfff4b10 100644 --- a/src/testing/mod.rs +++ b/src/testing/mod.rs @@ -565,6 +565,7 @@ impl TestHarnessBuilder { sandbox_readiness: crate::agent::routine_engine::SandboxReadiness::DisabledByConfig, builder: None, llm_backend: "nearai".to_string(), + tenant_rates: std::sync::Arc::new(crate::tenant::TenantRateRegistry::new(4, 3)), }; TestHarness { diff --git a/src/worker/job.rs b/src/worker/job.rs index b2e3f7e6..6d942cfc 100644 --- a/src/worker/job.rs +++ b/src/worker/job.rs @@ -20,7 +20,6 @@ use crate::agent::scheduler::WorkerMessage; use crate::agent::task::TaskOutput; use crate::channels::web::types::SseEvent; use crate::context::{ContextManager, JobState}; -use crate::db::Database; use crate::error::Error; use crate::hooks::HookRegistry; use crate::llm::{ @@ -28,6 +27,7 @@ use crate::llm::{ ToolSelection, }; use crate::safety::SafetyLayer; +use crate::tenant::AdminScope; use crate::tools::execute::process_tool_result; use crate::tools::rate_limiter::RateLimitResult; use crate::tools::{ @@ -44,7 +44,7 @@ pub struct WorkerDeps { pub llm: Arc, pub safety: Arc, pub tools: Arc, - pub store: Option>, + pub store: Option, pub hooks: Arc, pub timeout: Duration, pub use_planning: bool, @@ -93,7 +93,7 @@ impl Worker { &self.deps.tools } - fn store(&self) -> Option<&Arc> { + fn store(&self) -> Option<&AdminScope> { self.deps.store.as_ref() } diff --git a/tests/e2e_routine_heartbeat.rs b/tests/e2e_routine_heartbeat.rs index 27d8cfdc..6849ee05 100644 --- a/tests/e2e_routine_heartbeat.rs +++ b/tests/e2e_routine_heartbeat.rs @@ -337,14 +337,14 @@ mod tests { SchedulerDeps { tools: registry.clone(), extension_manager: extension_manager.clone(), - store: Some(db.clone()), + store: Some(ironclaw::tenant::AdminScope::new(db.clone())), hooks: Arc::new(HookRegistry::new()), }, )); Arc::new(RoutineEngine::new( RoutineConfig::default(), - db, + ironclaw::tenant::AdminScope::new(db), llm, ws, notify_tx, @@ -448,7 +448,7 @@ mod tests { let engine = Arc::new(RoutineEngine::new( RoutineConfig::default(), - db.clone(), + ironclaw::tenant::AdminScope::new(db.clone()), llm, ws, notify_tx, @@ -527,7 +527,7 @@ mod tests { let engine = Arc::new(RoutineEngine::new( RoutineConfig::default(), - db.clone(), + ironclaw::tenant::AdminScope::new(db.clone()), llm, ws, notify_tx, @@ -614,7 +614,7 @@ mod tests { let engine = Arc::new(RoutineEngine::new( RoutineConfig::default(), - db.clone(), + ironclaw::tenant::AdminScope::new(db.clone()), llm, ws, notify_tx, @@ -723,7 +723,7 @@ mod tests { let engine = Arc::new(RoutineEngine::new( RoutineConfig::default(), - db.clone(), + ironclaw::tenant::AdminScope::new(db.clone()), llm, ws, notify_tx, @@ -866,7 +866,7 @@ mod tests { let engine = Arc::new(RoutineEngine::new( RoutineConfig::default(), - db.clone(), + ironclaw::tenant::AdminScope::new(db.clone()), llm, ws, notify_tx, @@ -1049,7 +1049,7 @@ mod tests { let engine = Arc::new(RoutineEngine::new( RoutineConfig::default(), - Arc::clone(&db), + ironclaw::tenant::AdminScope::new(Arc::clone(&db)), llm, ws, notify_tx, @@ -1171,7 +1171,7 @@ mod tests { let engine = Arc::new(RoutineEngine::new( RoutineConfig::default(), - db.clone(), + ironclaw::tenant::AdminScope::new(db.clone()), llm, ws, notify_tx, @@ -1279,7 +1279,7 @@ mod tests { let engine = Arc::new(RoutineEngine::new( config, - db.clone(), + ironclaw::tenant::AdminScope::new(db.clone()), llm, ws, notify_tx, diff --git a/tests/e2e_telegram_message_routing.rs b/tests/e2e_telegram_message_routing.rs index ead164eb..810fc218 100644 --- a/tests/e2e_telegram_message_routing.rs +++ b/tests/e2e_telegram_message_routing.rs @@ -201,6 +201,7 @@ mod tests { sandbox_readiness: ironclaw::agent::SandboxReadiness::DisabledByConfig, builder: None, llm_backend: "nearai".to_string(), + tenant_rates: std::sync::Arc::new(ironclaw::tenant::TenantRateRegistry::new(4, 3)), }; let gateway = Arc::new(TestChannel::new()); diff --git a/tests/support/gateway_workflow_harness.rs b/tests/support/gateway_workflow_harness.rs index e4620f70..851e9bd9 100644 --- a/tests/support/gateway_workflow_harness.rs +++ b/tests/support/gateway_workflow_harness.rs @@ -265,6 +265,7 @@ impl GatewayWorkflowHarness { sandbox_readiness: ironclaw::agent::SandboxReadiness::DisabledByConfig, builder: None, llm_backend: "nearai".to_string(), + tenant_rates: std::sync::Arc::new(ironclaw::tenant::TenantRateRegistry::new(4, 3)), }, channels, None, diff --git a/tests/support/test_rig.rs b/tests/support/test_rig.rs index 624bb054..5775b86d 100644 --- a/tests/support/test_rig.rs +++ b/tests/support/test_rig.rs @@ -642,7 +642,7 @@ impl TestRigBuilder { let (notify_tx, _notify_rx) = tokio::sync::mpsc::channel(16); let engine = Arc::new(RoutineEngine::new( routine_config, - Arc::clone(db_arc), + ironclaw::tenant::AdminScope::new(Arc::clone(db_arc)), components.llm.clone(), Arc::clone(ws), notify_tx, @@ -762,6 +762,7 @@ impl TestRigBuilder { sandbox_readiness: ironclaw::agent::SandboxReadiness::Available, // tests don't use real Docker builder: None, llm_backend: "nearai".to_string(), + tenant_rates: std::sync::Arc::new(ironclaw::tenant::TenantRateRegistry::new(4, 3)), }; // 7. Create TestChannel and ChannelManager.