diff --git a/src/agent/agent_loop.rs b/src/agent/agent_loop.rs index e28f11d0..4ee846f7 100644 --- a/src/agent/agent_loop.rs +++ b/src/agent/agent_loop.rs @@ -13,7 +13,7 @@ use futures::StreamExt; use uuid::Uuid; use crate::agent::context_monitor::ContextMonitor; -use crate::agent::heartbeat::spawn_heartbeat; +use crate::agent::heartbeat::{spawn_heartbeat, spawn_multi_user_heartbeat}; use crate::agent::routine_engine::{RoutineEngine, spawn_cron_ticker}; use crate::agent::self_repair::{DefaultSelfRepair, RepairResult, SelfRepair}; use crate::agent::session::ThreadState; @@ -182,6 +182,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. @@ -244,7 +246,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(), }, ); @@ -325,6 +330,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() } @@ -410,8 +459,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())); @@ -518,6 +567,7 @@ impl Agent { .with_interval(std::time::Duration::from_secs(hb_config.interval_secs)); config.quiet_hours_start = hb_config.quiet_hours_start; config.quiet_hours_end = hb_config.quiet_hours_end; + config.multi_tenant = hb_config.multi_tenant; config.timezone = hb_config .timezone .clone() @@ -547,30 +597,52 @@ impl Agent { .await; let notify_user = heartbeat_notify_user; let channels = self.channels.clone(); + let is_multi_tenant = hb_config.multi_tenant; tokio::spawn(async move { while let Some(response) = notify_rx.recv().await { + // In multi-tenant mode, extract the owning user_id from + // the response metadata so notifications reach the + // correct user rather than the agent's owner. + // This intentionally overrides the configured notify_target + // because each user's heartbeat should notify that user. + let effective_user = if is_multi_tenant { + response + .metadata + .get("owner_id") + .and_then(|v| v.as_str()) + .map(String::from) + } else { + None + }; + // Try the configured channel first, fall back to // broadcasting on all channels. - let targeted_ok = if let Some(ref channel) = notify_channel - && let Some(ref user) = notify_target - { - channels - .broadcast(channel, user, response.clone()) - .await - .is_ok() + let targeted_ok = if let Some(ref channel) = notify_channel { + let target = effective_user.as_deref().or(notify_target.as_deref()); + if let Some(user) = target { + channels + .broadcast(channel, user, response.clone()) + .await + .is_ok() + } else { + false + } } else { false }; - if !targeted_ok && let Some(ref user) = notify_user { - let results = channels.broadcast_all(user, response).await; - for (ch, result) in results { - if let Err(e) = result { - tracing::warn!( - "Failed to broadcast heartbeat to {}: {}", - ch, - e - ); + if !targeted_ok { + let fallback = effective_user.as_deref().or(notify_user.as_deref()); + if let Some(user) = fallback { + let results = channels.broadcast_all(user, response).await; + for (ch, result) in results { + if let Err(e) = result { + tracing::warn!( + "Failed to broadcast heartbeat to {}: {}", + ch, + e + ); + } } } } @@ -583,14 +655,29 @@ impl Agent { .map(|h| h.to_workspace_config()) .unwrap_or_default(); - Some(spawn_heartbeat( - config, - hygiene, - workspace.clone(), - self.cheap_llm().clone(), - Some(notify_tx), - self.store().map(Arc::clone), - )) + if config.multi_tenant { + if let Some(admin) = self.admin_store() { + Some(spawn_multi_user_heartbeat( + config, + hygiene, + self.cheap_llm().clone(), + Some(notify_tx), + admin, + )) + } else { + tracing::warn!("Multi-tenant heartbeat requires a database store"); + None + } + } else { + Some(spawn_heartbeat( + config, + hygiene, + workspace.clone(), + self.cheap_llm().clone(), + Some(notify_tx), + self.admin_store(), + )) + } } else { tracing::warn!("Heartbeat enabled but no workspace available"); None @@ -612,7 +699,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, @@ -1173,13 +1260,22 @@ impl Agent { } } + // Build per-tenant execution context once; threaded through all handlers. + let tenant = self.tenant_ctx(&message.user_id).await; + let session_for_empty_exit = Arc::clone(&session); // 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. @@ -1246,7 +1342,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 @@ -1294,7 +1396,7 @@ impl Agent { }; } // Authorization checks (including restart channel check) are enforced in handle_system_command - self.handle_system_command(&command, &args, &message.channel) + self.handle_system_command(&command, &args, &message.channel, &tenant) .await } Submission::Undo => self.process_undo(session, thread_id).await, @@ -1307,12 +1409,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/agentic_loop.rs b/src/agent/agentic_loop.rs index e61856dc..27c2ab72 100644 --- a/src/agent/agentic_loop.rs +++ b/src/agent/agentic_loop.rs @@ -10,7 +10,7 @@ use std::borrow::Cow; use crate::agent::session::PendingApproval; use crate::error::Error; -use crate::llm::{ChatMessage, Reasoning, ReasoningContext, RespondResult}; +use crate::llm::{ChatMessage, FinishReason, Reasoning, ReasoningContext, RespondResult}; /// Signal from the delegate indicating how the loop should proceed. pub enum LoopSignal { @@ -134,6 +134,9 @@ pub async fn run_agentic_loop( config: &AgenticLoopConfig, ) -> Result { let mut consecutive_tool_intent_nudges: u32 = 0; + // Accumulates across all iterations (not reset by text responses) so + // non-consecutive truncations still escalate to force_text. + let mut truncation_count: u32 = 0; for iteration in 1..=config.max_iterations { // Check for external signals (stop, cancellation, user messages) @@ -215,7 +218,35 @@ pub async fn run_agentic_loop( tool_calls, content, } => { + // If the response was truncated, tool call parameters are likely + // incomplete. Discard them and tell the LLM to try a different + // approach rather than executing malformed tool calls. + if output.finish_reason == FinishReason::Length { + truncation_count += 1; + let names: Vec<&str> = tool_calls.iter().map(|tc| tc.name.as_str()).collect(); + tracing::warn!( + iteration, + tools = ?names, + truncation_count, + "Discarding truncated tool calls (finish_reason=Length)" + ); + if let Some(ref text) = content { + reason_ctx.messages.push(ChatMessage::assistant(text)); + } + reason_ctx + .messages + .push(ChatMessage::user(crate::llm::TRUNCATED_TOOL_CALL_NOTICE)); + // After repeated truncations, force text-only mode so the LLM + // stops attempting tool calls it can't fit in the output budget. + if truncation_count >= 3 { + reason_ctx.force_text = true; + } + delegate.after_iteration(iteration).await; + continue; + } + consecutive_tool_intent_nudges = 0; + truncation_count = 0; if let Some(outcome) = delegate .execute_tool_calls(tool_calls, content, reason_ctx) @@ -271,6 +302,7 @@ mod tests { RespondOutput { result: RespondResult::Text(text.to_string()), usage: zero_usage(), + finish_reason: FinishReason::Stop, } } @@ -281,6 +313,7 @@ mod tests { content: None, }, usage: zero_usage(), + finish_reason: FinishReason::ToolUse, } } @@ -622,4 +655,95 @@ mod tests { let result = truncate_for_preview("café", 4); assert_eq!(result, "caf..."); } + + #[tokio::test] + async fn test_truncated_tool_calls_discarded_on_length() { + let truncated_tool_call = ToolCall { + id: "call_1".to_string(), + name: "memory_write".to_string(), + arguments: serde_json::json!({}), // empty — truncated + reasoning: None, + }; + let truncated_output = RespondOutput { + result: RespondResult::ToolCalls { + tool_calls: vec![truncated_tool_call], + content: Some("I'll write the report.".to_string()), + }, + usage: zero_usage(), + finish_reason: FinishReason::Length, // response was truncated + }; + let delegate = MockDelegate::new(vec![truncated_output, text_output("Summarized it.")]); + let reasoning = stub_reasoning(); + let mut ctx = ReasoningContext::new(); + let config = AgenticLoopConfig { + max_iterations: 5, + ..Default::default() + }; + + let outcome = run_agentic_loop(&delegate, &reasoning, &mut ctx, &config) + .await + .unwrap(); + + // Tool calls should NOT have been executed + assert_eq!(delegate.tool_exec_count.load(Ordering::SeqCst), 0); + // The loop should have continued and returned the text response + assert!(matches!(outcome, LoopOutcome::Response(ref t) if t == "Summarized it.")); + // A truncation notice should have been injected into context + assert!( + ctx.messages + .iter() + .any(|m| m.role == crate::llm::Role::User && m.content.contains("truncated")), + "Should inject truncation notice into context" + ); + // The partial assistant content should have been preserved + assert!( + ctx.messages + .iter() + .any(|m| m.role == crate::llm::Role::Assistant + && m.content.contains("write the report")), + "Should preserve partial assistant content" + ); + } + + #[tokio::test] + async fn test_repeated_truncations_force_text_mode() { + let make_truncated = || RespondOutput { + result: RespondResult::ToolCalls { + tool_calls: vec![ToolCall { + id: "call_1".to_string(), + name: "memory_write".to_string(), + arguments: serde_json::json!({}), + reasoning: None, + }], + content: None, + }, + usage: zero_usage(), + finish_reason: FinishReason::Length, + }; + // Three truncated responses, then a text response + let delegate = MockDelegate::new(vec![ + make_truncated(), + make_truncated(), + make_truncated(), + text_output("Gave up on tool calls."), + ]); + let reasoning = stub_reasoning(); + let mut ctx = ReasoningContext::new(); + let config = AgenticLoopConfig { + max_iterations: 5, + ..Default::default() + }; + + let outcome = run_agentic_loop(&delegate, &reasoning, &mut ctx, &config) + .await + .unwrap(); + + assert!(matches!(outcome, LoopOutcome::Response(_))); + assert_eq!(delegate.tool_exec_count.load(Ordering::SeqCst), 0); + // After 3 truncations, force_text should be set + assert!( + ctx.force_text, + "Should escalate to force_text after repeated truncations" + ); + } } diff --git a/src/agent/commands.rs b/src/agent/commands.rs index e02b33db..643d8c7c 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) + .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,7 +117,8 @@ 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 { return Ok(format!( @@ -138,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()); } @@ -155,7 +151,8 @@ 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; @@ -183,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, @@ -196,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 @@ -221,11 +223,12 @@ 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() { + // 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) => { @@ -256,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()); } @@ -270,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()); } @@ -308,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)), @@ -323,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))), } @@ -559,6 +566,7 @@ impl Agent { command: &str, args: &[String], channel: &str, + tenant: &crate::tenant::TenantCtx, ) -> Result { match command { "help" => Ok(SubmissionResult::response(concat!( @@ -752,19 +760,32 @@ impl Agent { } } - match self.llm().set_model(requested) { - Ok(()) => { - // Persist the model choice so it survives restarts. - self.persist_selected_model(requested).await; - Ok(SubmissionResult::response(format!( - "Switched model to: {}", - requested - ))) + if self.config.multi_tenant { + // Multi-tenant: only persist to per-user DB settings. + // Do NOT call set_model() on the shared provider — that + // 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(tenant, requested).await; + Ok(SubmissionResult::response(format!( + "Model preference set to: {} (per-user)", + requested + ))) + } else { + match self.llm().set_model(requested) { + Ok(()) => { + // Persist the model choice so it survives restarts. + self.persist_selected_model(tenant, requested).await; + Ok(SubmissionResult::response(format!( + "Switched model to: {}", + requested + ))) + } + Err(e) => Ok(SubmissionResult::error(format!( + "Failed to switch model: {}", + e + ))), } - Err(e) => Ok(SubmissionResult::error(format!( - "Failed to switch model: {}", - e - ))), } } } @@ -906,10 +927,14 @@ impl Agent { command: &str, args: &[String], channel: &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).await? { + match self + .handle_system_command(command, args, channel, tenant) + .await? + { SubmissionResult::Response { content } => Ok(Some(content)), SubmissionResult::Ok { message } => Ok(message), SubmissionResult::Error { message } => Ok(Some(format!("Error: {}", message))), @@ -921,23 +946,33 @@ impl Agent { /// /// Best-effort: logs warnings on failure but does not propagate errors, /// since the in-memory model switch already succeeded. - async fn persist_selected_model(&self, model: &str) { - // 1. Persist to DB if available. - if let Some(store) = self.store() { + /// + /// 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, 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(self.owner_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!("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"); } - // 2. Update .env and TOML config file (sync I/O in spawn_blocking). + // 2. In multi-tenant mode, skip .env/TOML writes — these are global + // files shared by all users. The per-user DB setting is sufficient. + if self.config.multi_tenant { + return; + } + + // 3. Update .env and TOML config file (sync I/O in spawn_blocking). let model_owned = model.to_string(); let backend = self.deps.llm_backend.clone(); if let Err(e) = tokio::task::spawn_blocking(move || { diff --git a/src/agent/cost_guard.rs b/src/agent/cost_guard.rs index 4563bbbe..4885364b 100644 --- a/src/agent/cost_guard.rs +++ b/src/agent/cost_guard.rs @@ -21,6 +21,9 @@ pub struct CostGuardConfig { pub max_cost_per_day_cents: Option, /// Maximum LLM calls per hour. None = unlimited. pub max_actions_per_hour: Option, + /// Maximum spend per user per day in cents. None = unlimited. + /// Applied independently per user alongside the global budget. + pub max_cost_per_user_per_day_cents: Option, } /// Error returned when a cost limit is exceeded. @@ -30,6 +33,12 @@ pub enum CostLimitExceeded { DailyBudget { spent_cents: u64, limit_cents: u64 }, /// Hourly action rate limit reached. HourlyRate { actions: u64, limit: u64 }, + /// Per-user daily spending cap reached. + UserDailyBudget { + user_id: String, + spent_cents: u64, + limit_cents: u64, + }, } impl std::fmt::Display for CostLimitExceeded { @@ -49,6 +58,17 @@ impl std::fmt::Display for CostLimitExceeded { "Hourly action limit exceeded: {} actions of {} allowed per hour", actions, limit ), + Self::UserDailyBudget { + user_id, + spent_cents, + limit_cents, + } => write!( + f, + "User '{}' daily cost limit exceeded: spent ${:.2} of ${:.2} allowed", + user_id, + *spent_cents as f64 / 100.0, + *limit_cents as f64 / 100.0 + ), } } } @@ -78,6 +98,9 @@ pub struct CostGuard { /// Per-model token usage since startup. model_tokens: Mutex>, + + /// Per-user daily cost tracking. Each entry resets independently at midnight UTC. + per_user_daily_cost: Mutex>, } struct DailyCost { @@ -97,6 +120,7 @@ impl CostGuard { action_window: Mutex::new(VecDeque::new()), budget_exceeded: AtomicBool::new(false), model_tokens: Mutex::new(HashMap::new()), + per_user_daily_cost: Mutex::new(HashMap::new()), } } @@ -203,6 +227,11 @@ impl CostGuard { daily.reset_date = today; self.budget_exceeded.store(false, Ordering::Relaxed); tracing::info!("Cost guard: daily counter reset for {}", today); + + // Prune per-user entries from previous days to prevent + // unbounded HashMap growth in long-lived deployments. + let mut per_user = self.per_user_daily_cost.lock().await; + per_user.retain(|_, entry| entry.reset_date == today); } daily.total += cost; @@ -248,6 +277,85 @@ impl CostGuard { cost } + /// Record an LLM call with per-user attribution. + /// + /// Delegates to `record_llm_call` for global tracking, then additionally + /// records the cost against the user's daily budget. + #[allow(clippy::too_many_arguments)] + pub async fn record_llm_call_for_user( + &self, + user_id: &str, + 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 { + let cost = self + .record_llm_call( + model, + input_tokens, + output_tokens, + cache_read_input_tokens, + cache_creation_input_tokens, + cache_read_discount, + cache_write_multiplier, + cost_per_token, + ) + .await; + + // Track per-user daily cost + { + let today = chrono::Utc::now().date_naive(); + let mut per_user = self.per_user_daily_cost.lock().await; + let entry = per_user + .entry(user_id.to_string()) + .or_insert_with(|| DailyCost { + total: Decimal::ZERO, + reset_date: today, + }); + if today != entry.reset_date { + entry.total = Decimal::ZERO; + entry.reset_date = today; + } + entry.total += cost; + } + + cost + } + + /// Check whether the next action is allowed for a specific user. + /// + /// Checks the global limits first (via `check_allowed`), then additionally + /// checks the per-user daily budget if configured. + pub async fn check_allowed_for_user(&self, user_id: &str) -> Result<(), CostLimitExceeded> { + // Check global limits first + self.check_allowed().await?; + + // Check per-user daily budget + if let Some(limit_cents) = self.config.max_cost_per_user_per_day_cents { + let today = chrono::Utc::now().date_naive(); + let per_user = self.per_user_daily_cost.lock().await; + if let Some(entry) = per_user.get(user_id) + && entry.reset_date == today + { + let spent_cents = to_cents(entry.total); + if spent_cents >= limit_cents { + return Err(CostLimitExceeded::UserDailyBudget { + user_id: user_id.to_string(), + spent_cents, + limit_cents, + }); + } + } + } + + Ok(()) + } + /// Current daily spend in USD (as Decimal). pub async fn daily_spend(&self) -> Decimal { let daily = self.daily_cost.lock().await; @@ -259,6 +367,16 @@ impl CostGuard { } } + /// Current daily spend for a specific user in USD (as Decimal). + pub async fn daily_spend_for_user(&self, user_id: &str) -> Decimal { + let today = chrono::Utc::now().date_naive(); + let per_user = self.per_user_daily_cost.lock().await; + match per_user.get(user_id) { + Some(entry) if entry.reset_date == today => entry.total, + _ => Decimal::ZERO, + } + } + /// Number of actions in the current hourly window. pub async fn actions_this_hour(&self) -> u64 { let mut window = self.action_window.lock().await; @@ -314,7 +432,7 @@ mod tests { async fn test_daily_budget_enforcement() { let guard = CostGuard::new(CostGuardConfig { max_cost_per_day_cents: Some(1), // $0.01 limit - max_actions_per_hour: None, + ..CostGuardConfig::default() }); // First call allowed @@ -350,8 +468,8 @@ mod tests { #[tokio::test] async fn test_hourly_rate_enforcement() { let guard = CostGuard::new(CostGuardConfig { - max_cost_per_day_cents: None, max_actions_per_hour: Some(3), + ..CostGuardConfig::default() }); // First 3 actions allowed @@ -633,8 +751,8 @@ mod tests { // A fresh CostGuard with rate limits should not panic even if // checked_sub returns None (simulating short uptime). let guard = CostGuard::new(CostGuardConfig { - max_cost_per_day_cents: None, max_actions_per_hour: Some(100), + ..CostGuardConfig::default() }); // These must not panic regardless of system uptime @@ -656,4 +774,119 @@ mod tests { let result = Instant::now().checked_sub(std::time::Duration::MAX); assert!(result.is_none()); } + + #[tokio::test] + async fn test_per_user_daily_budget_enforcement() { + let guard = CostGuard::new(CostGuardConfig { + max_cost_per_day_cents: None, + max_actions_per_hour: None, + max_cost_per_user_per_day_cents: Some(1), // $0.01 per user + }); + + // Both users initially allowed + assert!(guard.check_allowed_for_user("alice").await.is_ok()); + assert!(guard.check_allowed_for_user("bob").await.is_ok()); + + // Alice makes an expensive call + guard + .record_llm_call_for_user( + "alice", + "gpt-4o", + 10_000, + 10_000, + 0, + 0, + Decimal::ONE, + Decimal::ONE, + None, + ) + .await; + + // Alice should be blocked, Bob should still be allowed + let result = guard.check_allowed_for_user("alice").await; + assert!(result.is_err()); + match result.unwrap_err() { + CostLimitExceeded::UserDailyBudget { + user_id, + limit_cents, + .. + } => { + assert_eq!(user_id, "alice"); + assert_eq!(limit_cents, 1); + } + other => panic!("Expected UserDailyBudget, got {:?}", other), + } + assert!(guard.check_allowed_for_user("bob").await.is_ok()); + } + + #[tokio::test] + async fn test_per_user_daily_spend_tracking() { + let guard = CostGuard::new(CostGuardConfig::default()); + + assert_eq!(guard.daily_spend_for_user("alice").await, Decimal::ZERO); + assert_eq!(guard.daily_spend_for_user("bob").await, Decimal::ZERO); + + let cost = guard + .record_llm_call_for_user( + "alice", + "gpt-4o", + 1000, + 500, + 0, + 0, + Decimal::ONE, + Decimal::ONE, + None, + ) + .await; + + assert_eq!(guard.daily_spend_for_user("alice").await, cost); + assert_eq!(guard.daily_spend_for_user("bob").await, Decimal::ZERO); + // Global spend should also be tracked + assert_eq!(guard.daily_spend().await, cost); + } + + #[tokio::test] + async fn test_per_user_budget_independent_of_global() { + let guard = CostGuard::new(CostGuardConfig { + max_cost_per_day_cents: Some(100_000), // $1000 global limit + max_actions_per_hour: None, + max_cost_per_user_per_day_cents: Some(1), // $0.01 per user + }); + + // User hits their personal limit + guard + .record_llm_call_for_user( + "alice", + "gpt-4o", + 10_000, + 10_000, + 0, + 0, + Decimal::ONE, + Decimal::ONE, + None, + ) + .await; + + // Alice blocked by per-user limit, not global + assert!(guard.check_allowed_for_user("alice").await.is_err()); + // Global limit is far from reached + assert!(guard.check_allowed().await.is_ok()); + // Bob is unaffected + assert!(guard.check_allowed_for_user("bob").await.is_ok()); + } + + #[test] + fn test_user_cost_limit_display() { + let limit = CostLimitExceeded::UserDailyBudget { + user_id: "alice".to_string(), + spent_cents: 150, + limit_cents: 100, + }; + let msg = limit.to_string(); + assert!(msg.contains("alice")); + assert!(msg.contains("$1.50")); + assert!(msg.contains("$1.00")); + } } diff --git a/src/agent/dispatcher.rs b/src/agent/dispatcher.rs index fe208c1b..a5f9cd6f 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, @@ -168,6 +169,7 @@ impl Agent { let delegate = ChatDelegate { agent: self, + tenant, session: session.clone(), thread_id, message, @@ -240,6 +242,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, @@ -303,6 +306,8 @@ impl<'a> LoopDelegate for ChatDelegate<'a> { // Update context for this iteration reason_ctx.available_tools = tool_defs; + // Preserve force_text if already set (e.g. by truncation escalation). + let force_text = force_text || reason_ctx.force_text; reason_ctx.system_prompt = Some(if force_text { self.cached_prompt_no_tools.clone() } else { @@ -336,8 +341,8 @@ impl<'a> LoopDelegate for ChatDelegate<'a> { reason_ctx: &mut ReasoningContext, iteration: usize, ) -> Result { - // Enforce cost guardrails before the LLM call - if let Err(limit) = self.agent.cost_guard().check_allowed().await { + // Enforce cost guardrails before the LLM call (global + per-user) + if let Err(limit) = self.tenant.check_cost_allowed().await { return Err(crate::error::LlmError::InvalidResponse { provider: "agent".to_string(), reason: limit.to_string(), @@ -345,6 +350,21 @@ impl<'a> LoopDelegate for ChatDelegate<'a> { .into()); } + // 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 TenantScope). + if iteration == 0 + && 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(); + if !model.is_empty() { + reason_ctx.model_override = Some(model.to_string()); + } + } + let output = match reasoning.respond_with_tools(reason_ctx).await { Ok(output) => output, Err(crate::error::LlmError::ContextLengthExceeded { used, limit }) => { @@ -379,13 +399,22 @@ impl<'a> LoopDelegate for ChatDelegate<'a> { Err(e) => return Err(e.into()), }; - // Record cost and track token usage - let model_name = self.agent.llm().active_model_name(); + // Record cost and track token usage (global + per-user). + // When a model override is active, use the override name for attribution + // and let CostGuard look up pricing via costs::model_cost() instead of + // using the default provider's cost_per_token (which reflects the wrong model). + let (model_name, cost_per_token) = if let Some(ref ovr) = reason_ctx.model_override { + (ovr.clone(), None) + } else { + ( + self.agent.llm().active_model_name(), + Some(self.agent.llm().cost_per_token()), + ) + }; 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() + .tenant .record_llm_call( &model_name, output.usage.input_tokens, @@ -394,7 +423,7 @@ impl<'a> LoopDelegate for ChatDelegate<'a> { output.usage.cache_creation_input_tokens, read_discount, write_multiplier, - Some(self.agent.llm().cost_per_token()), + cost_per_token, ) .await; tracing::debug!( @@ -1305,6 +1334,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( @@ -1320,10 +1350,14 @@ mod tests { allow_local_tools: false, max_cost_per_day_cents: None, max_actions_per_hour: None, + max_cost_per_user_per_day_cents: None, max_tool_iterations: 50, auto_approve_tools: false, 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()), @@ -2181,6 +2215,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( @@ -2196,10 +2231,14 @@ mod tests { allow_local_tools: false, max_cost_per_day_cents: None, max_actions_per_hour: None, + max_cost_per_user_per_day_cents: None, max_tool_iterations, auto_approve_tools: true, 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()), @@ -2234,13 +2273,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; @@ -2302,6 +2342,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( @@ -2317,10 +2358,14 @@ mod tests { allow_local_tools: false, max_cost_per_day_cents: None, max_actions_per_hour: None, + max_cost_per_user_per_day_cents: None, max_tool_iterations: max_iter, auto_approve_tools: true, 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()), @@ -2340,13 +2385,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 ec4cd5e9..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; @@ -57,6 +57,9 @@ pub struct HeartbeatConfig { pub quiet_hours_end: Option, /// Timezone for fire_at and quiet hours evaluation (IANA name). pub timezone: Option, + /// When true, cycle through all users with routines instead of + /// running heartbeat for a single user. Requires a database store. + pub multi_tenant: bool, } impl Default for HeartbeatConfig { @@ -71,6 +74,7 @@ impl Default for HeartbeatConfig { quiet_hours_start: None, quiet_hours_end: None, timezone: None, + multi_tenant: false, } } } @@ -178,7 +182,7 @@ pub struct HeartbeatRunner { workspace: Arc, llm: Arc, response_tx: Option>, - store: Option>, + store: Option, consecutive_failures: u32, } @@ -207,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 } @@ -396,7 +400,7 @@ impl HeartbeatRunner { } /// Send a notification about heartbeat findings. - async fn send_notification(&self, message: &str) { + pub(crate) async fn send_notification(&self, message: &str) { let Some(ref tx) = self.response_tx else { tracing::debug!("No response channel configured for heartbeat notifications"); return; @@ -493,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 { @@ -508,6 +512,179 @@ pub fn spawn_heartbeat( }) } +/// Spawn a multi-user heartbeat runner that cycles through all users that +/// own routines (enabled or not). Each tick, it queries the DB for distinct +/// user_ids, creates a per-user workspace, and runs a heartbeat check for +/// each user concurrently. Per-user failure counts are tracked independently. +pub fn spawn_multi_user_heartbeat( + config: HeartbeatConfig, + hygiene_config: HygieneConfig, + llm: Arc, + response_tx: Option>, + store: AdminScope, +) -> tokio::task::JoinHandle<()> { + tokio::spawn(async move { + if !config.enabled { + tracing::info!("Multi-user heartbeat is disabled"); + return; + } + + let mut tick_interval = if config.fire_at.is_none() { + let mut iv = tokio::time::interval(config.interval); + iv.tick().await; // skip immediate tick + Some(iv) + } else { + None + }; + + // Track consecutive failures per user so we can disable heartbeat + // for persistently-failing users (same semantics as single-user mode). + let mut user_failures: std::collections::HashMap = + std::collections::HashMap::new(); + + tracing::info!("Starting multi-user heartbeat loop"); + + loop { + if let Some(fire_at) = config.fire_at { + let sleep_dur = duration_until_next_fire(fire_at, config.resolved_tz()); + tokio::time::sleep(sleep_dur).await; + } else if let Some(ref mut iv) = tick_interval { + iv.tick().await; + } + + if config.is_quiet_hours() { + continue; + } + + // Get distinct user_ids from routines + let user_ids = match store.list_all_routines().await { + Ok(routines) => { + let mut ids: Vec = routines + .iter() + .map(|r| r.user_id.clone()) + .collect::>() + .into_iter() + .collect(); + ids.sort(); + ids + } + Err(e) => { + tracing::error!("Multi-user heartbeat: failed to list routines: {}", e); + continue; + } + }; + + // Run user heartbeats concurrently so one slow LLM call doesn't + // block others. Cap concurrency to avoid flooding the LLM provider. + const MAX_CONCURRENT_HEARTBEATS: usize = 8; + let mut join_set = tokio::task::JoinSet::new(); + + for user_id in &user_ids { + // Skip users that have exceeded max_failures + let failures = user_failures.get(user_id).copied().unwrap_or(0); + if failures >= config.max_failures { + continue; + } + + 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); + let hygiene_cfg = hygiene_config.clone(); + let hygiene_user = user_id.clone(); + tokio::spawn(async move { + let report = + crate::workspace::hygiene::run_if_due(&hygiene_ws, &hygiene_cfg).await; + if report.had_work() { + tracing::info!( + user_id = hygiene_user, + daily_logs_deleted = report.daily_logs_deleted, + conversation_docs_deleted = report.conversation_docs_deleted, + "multi-user heartbeat: memory hygiene deleted stale documents" + ); + } + }); + + // Drain completed tasks to stay within the concurrency cap. + while join_set.len() >= MAX_CONCURRENT_HEARTBEATS { + if let Some(join_result) = join_set.join_next().await { + collect_heartbeat_result(join_result, &mut user_failures, &config); + } + } + + let uid = user_id.clone(); + let cfg = config.clone(); + let hyg = hygiene_config.clone(); + let llm_clone = llm.clone(); + let tx = response_tx.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(admin); + + let result = runner.check_heartbeat().await; + if let HeartbeatResult::NeedsAttention(msg) = &result { + runner.send_notification(msg).await; + } + (uid, result) + }); + } + + // Collect remaining results and update failure counts + while let Some(join_result) = join_set.join_next().await { + collect_heartbeat_result(join_result, &mut user_failures, &config); + } + } + }) +} + +/// Process a single JoinSet result from the multi-user heartbeat loop. +fn collect_heartbeat_result( + join_result: Result<(String, HeartbeatResult), tokio::task::JoinError>, + user_failures: &mut std::collections::HashMap, + config: &HeartbeatConfig, +) { + let (uid, result) = match join_result { + Ok(pair) => pair, + Err(e) => { + tracing::error!("Multi-user heartbeat task panicked: {}", e); + return; + } + }; + match result { + HeartbeatResult::Ok => { + tracing::trace!(user_id = uid, "Multi-user heartbeat OK"); + user_failures.remove(&uid); + } + HeartbeatResult::NeedsAttention(_) => { + tracing::info!(user_id = uid, "Multi-user heartbeat needs attention"); + user_failures.remove(&uid); + } + HeartbeatResult::Skipped => {} + HeartbeatResult::Failed(err) => { + let count = user_failures.entry(uid.clone()).or_insert(0); + *count += 1; + tracing::error!( + user_id = uid, + consecutive_failures = *count, + "Multi-user heartbeat failed: {}", + err + ); + if *count >= config.max_failures { + tracing::error!( + user_id = uid, + "Multi-user heartbeat disabled for user after {} consecutive failures", + count + ); + } + } + } +} + #[cfg(test)] mod tests { use super::*; @@ -726,7 +903,7 @@ mod tests { Arc, Arc, Option>, - Option>, + Option, ) -> tokio::task::JoinHandle<()> = spawn_heartbeat; let _ = _fn_ptr; } diff --git a/src/agent/mod.rs b/src/agent/mod.rs index 84155666..e7242845 100644 --- a/src/agent/mod.rs +++ b/src/agent/mod.rs @@ -36,7 +36,9 @@ pub(crate) use agent_loop::truncate_for_preview; pub use agent_loop::{Agent, AgentDeps}; pub use compaction::{CompactionResult, ContextCompactor}; pub use context_monitor::{CompactionStrategy, ContextBreakdown, ContextMonitor}; -pub use heartbeat::{HeartbeatConfig, HeartbeatResult, HeartbeatRunner, spawn_heartbeat}; +pub use heartbeat::{ + HeartbeatConfig, HeartbeatResult, HeartbeatRunner, spawn_heartbeat, spawn_multi_user_heartbeat, +}; pub use router::{MessageIntent, Router}; pub use routine::{Routine, RoutineAction, RoutineRun, Trigger}; pub use routine_engine::{RoutineEngine, SandboxReadiness}; diff --git a/src/agent/routine_engine.rs b/src/agent/routine_engine.rs index a3cdb6cd..64c3b94c 100644 --- a/src/agent/routine_engine.rs +++ b/src/agent/routine_engine.rs @@ -28,12 +28,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, @@ -99,7 +99,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). @@ -128,7 +128,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, @@ -782,12 +782,22 @@ impl RoutineEngine { }); } + // Per-user workspace (same pattern as spawn_fire). + let routine_workspace = if routine.user_id == self.workspace.user_id() { + self.workspace.clone() + } else { + Arc::new(Workspace::new_with_db( + &routine.user_id, + Arc::clone(self.store.db()), + )) + }; + // Execute inline for manual triggers (caller wants to wait) let engine = EngineContext { config: self.config.clone(), store: self.store.clone(), llm: self.llm.clone(), - workspace: self.workspace.clone(), + workspace: routine_workspace, notify_tx: self.notify_tx.clone(), running_count: self.running_count.clone(), scheduler: self.scheduler.clone(), @@ -910,11 +920,23 @@ impl RoutineEngine { created_at: Utc::now(), }; + // Use per-user workspace so each routine executes in the correct + // user's context. Fall back to the engine-wide workspace when the + // routine belongs to the same user (avoids unnecessary allocation). + let routine_workspace = if routine.user_id == self.workspace.user_id() { + self.workspace.clone() + } else { + Arc::new(Workspace::new_with_db( + &routine.user_id, + Arc::clone(self.store.db()), + )) + }; + let engine = EngineContext { config: self.config.clone(), store: self.store.clone(), llm: self.llm.clone(), - workspace: self.workspace.clone(), + workspace: routine_workspace, notify_tx: self.notify_tx.clone(), running_count: self.running_count.clone(), scheduler: self.scheduler.clone(), @@ -967,7 +989,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, } @@ -978,7 +1000,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, @@ -1050,7 +1072,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 02953a4b..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>, @@ -780,10 +780,14 @@ mod tests { allow_local_tools: true, max_cost_per_day_cents: None, max_actions_per_hour: None, + max_cost_per_user_per_day_cents: None, max_tool_iterations: 10, auto_approve_tools: true, 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 11f211f9..a5288f68 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 @@ -1473,7 +1474,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/app.rs b/src/app.rs index 074e9479..8fb950fb 100644 --- a/src/app.rs +++ b/src/app.rs @@ -880,6 +880,7 @@ impl AppBuilder { crate::agent::cost_guard::CostGuardConfig { max_cost_per_day_cents: self.config.agent.max_cost_per_day_cents, max_actions_per_hour: self.config.agent.max_actions_per_hour, + max_cost_per_user_per_day_cents: self.config.agent.max_cost_per_user_per_day_cents, }, )); diff --git a/src/channels/relay/client.rs b/src/channels/relay/client.rs index 81fbb56c..b67f2c5e 100644 --- a/src/channels/relay/client.rs +++ b/src/channels/relay/client.rs @@ -122,18 +122,32 @@ impl RelayClient { /// instance_url in chat-api. IronClaw only passes an optional CSRF nonce /// for validating the callback — no URLs. pub async fn initiate_oauth(&self, state_nonce: Option<&str>) -> Result { + let url = format!("{}/oauth/slack/auth", self.base_url); + tracing::debug!(relay_url = %url, "RelayClient::initiate_oauth: sending request"); let mut query: Vec<(&str, &str)> = vec![]; if let Some(nonce) = state_nonce { query.push(("state_nonce", nonce)); } let resp = self .http - .get(format!("{}/oauth/slack/auth", self.base_url)) + .get(&url) .bearer_auth(self.api_key.expose_secret()) .query(&query) .send() .await - .map_err(|e| RelayError::Network(e.to_string()))?; + .map_err(|e| { + tracing::warn!( + relay_url = %url, + error = %e, + "RelayClient::initiate_oauth: network request failed" + ); + RelayError::Network(e.to_string()) + })?; + tracing::debug!( + relay_url = %url, + status = %resp.status(), + "RelayClient::initiate_oauth: received response" + ); let status = resp.status(); if status.is_redirection() { @@ -224,20 +238,39 @@ impl RelayClient { method: &str, body: serde_json::Value, ) -> Result { + let url = format!("{}/proxy/{}/{}", self.base_url, provider, method); + tracing::debug!( + relay_url = %url, + provider = %provider, + method = %method, + "RelayClient::proxy_provider: sending request" + ); let query: Vec<(&str, &str)> = vec![("team_id", team_id)]; let resp = self .http - .post(format!("{}/proxy/{}/{}", self.base_url, provider, method)) + .post(&url) .bearer_auth(self.api_key.expose_secret()) .query(&query) .json(&body) .send() .await - .map_err(|e| RelayError::Network(e.to_string()))?; + .map_err(|e| { + tracing::warn!( + relay_url = %url, + error = %e, + "RelayClient::proxy_provider: network request failed" + ); + RelayError::Network(e.to_string()) + })?; if !resp.status().is_success() { let status = resp.status().as_u16(); let body = resp.text().await.unwrap_or_default(); + tracing::warn!( + relay_url = %url, + status = status, + "RelayClient::proxy_provider: channel-relay returned error" + ); return Err(RelayError::Api { status, message: body, @@ -255,23 +288,45 @@ impl RelayClient { /// 32-byte secret. Called once at activation time; the result is cached in the /// extension manager so subsequent calls to `relay_signing_secret()` use it. pub async fn get_signing_secret(&self, team_id: &str) -> Result, RelayError> { + let url = format!("{}/relay/signing-secret", self.base_url); + tracing::debug!( + relay_url = %url, + "RelayClient::get_signing_secret: fetching signing secret" + ); let resp = self .http - .get(format!("{}/relay/signing-secret", self.base_url)) + .get(&url) .bearer_auth(self.api_key.expose_secret()) .query(&[("team_id", team_id)]) .send() .await - .map_err(|e| RelayError::Network(e.to_string()))?; + .map_err(|e| { + tracing::warn!( + relay_url = %url, + error = %e, + "RelayClient::get_signing_secret: network request failed" + ); + RelayError::Network(e.to_string()) + })?; if !resp.status().is_success() { let status = resp.status().as_u16(); let body = resp.text().await.unwrap_or_default(); + tracing::warn!( + relay_url = %url, + status = status, + body = %body, + "RelayClient::get_signing_secret: channel-relay returned error" + ); return Err(RelayError::Api { status, message: body, }); } + tracing::debug!( + relay_url = %url, + "RelayClient::get_signing_secret: received successful response" + ); let body: serde_json::Value = resp .json() diff --git a/src/channels/web/handlers/webhooks.rs b/src/channels/web/handlers/webhooks.rs index 7b041a06..1fd78c66 100644 --- a/src/channels/web/handlers/webhooks.rs +++ b/src/channels/web/handlers/webhooks.rs @@ -54,10 +54,37 @@ fn validate_webhook_secret( /// /// This endpoint is **public** (no gateway auth token required) but protected /// by the per-routine webhook secret sent via the `X-Webhook-Secret` header. +/// +/// **Single-user/backward-compatible**: looks up routines by path across all +/// users. For multi-tenant isolation, use the user-scoped endpoint at +/// `/api/webhooks/u/{user_id}/{path}` instead. pub async fn webhook_trigger_handler( State(state): State>, Path(path): Path, headers: HeaderMap, +) -> Result, (StatusCode, String)> { + fire_webhook_inner(state, &path, None, &headers).await +} + +/// Handle incoming webhook POST to `/api/webhooks/u/{user_id}/{path}`. +/// +/// User-scoped variant for multi-tenant deployments. The `user_id` in the URL +/// restricts the routine lookup to that user only, preventing cross-user +/// webhook triggering even when paths collide. +pub async fn webhook_trigger_user_scoped_handler( + State(state): State>, + Path((user_id, path)): Path<(String, String)>, + headers: HeaderMap, +) -> Result, (StatusCode, String)> { + fire_webhook_inner(state, &path, Some(&user_id), &headers).await +} + +/// Shared webhook logic for both scoped and unscoped endpoints. +async fn fire_webhook_inner( + state: Arc, + path: &str, + user_id: Option<&str>, + headers: &HeaderMap, ) -> Result, (StatusCode, String)> { // Rate limit check if !state.webhook_rate_limiter.check() { @@ -72,9 +99,9 @@ pub async fn webhook_trigger_handler( "Database not available".to_string(), ))?; - // Targeted query instead of loading all routines + // Targeted query — when user_id is provided, restrict to that user's routines let routine = store - .get_webhook_routine_by_path(&path) + .get_webhook_routine_by_path(path, user_id) .await .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))? .ok_or(( @@ -99,7 +126,7 @@ pub async fn webhook_trigger_handler( ))? }; - let run_id = engine.fire_webhook(routine.id, &path).await.map_err(|e| { + let run_id = engine.fire_webhook(routine.id, path).await.map_err(|e| { let status = match &e { crate::error::RoutineError::NotFound { .. } => StatusCode::NOT_FOUND, crate::error::RoutineError::Disabled { .. } diff --git a/src/channels/web/server.rs b/src/channels/web/server.rs index c24ceb16..26c005d4 100644 --- a/src/channels/web/server.rs +++ b/src/channels/web/server.rs @@ -414,6 +414,11 @@ pub async fn start_server( .route( "/api/webhooks/{path}", post(crate::channels::web::handlers::webhooks::webhook_trigger_handler), + ) + // User-scoped webhook endpoint for multi-tenant isolation + .route( + "/api/webhooks/u/{user_id}/{path}", + post(crate::channels::web::handlers::webhooks::webhook_trigger_user_scoped_handler), ); // Protected routes (require auth) @@ -1172,11 +1177,31 @@ async fn slack_relay_oauth_callback_handler( // Store team_id in settings let team_id_key = format!("relay:{}:team_id", DEFAULT_RELAY_NAME); - let _ = store + tracing::info!( + relay = DEFAULT_RELAY_NAME, + owner_id = %state.owner_id, + team_id_key = %team_id_key, + "relay OAuth callback: storing team_id in settings" + ); + store .set_setting(&state.owner_id, &team_id_key, &serde_json::json!(team_id)) - .await; + .await + .map_err(|e| { + tracing::error!( + relay = DEFAULT_RELAY_NAME, + owner_id = %state.owner_id, + error = %e, + "relay OAuth callback: failed to persist team_id to settings store" + ); + format!("Failed to persist relay team_id: {e}") + })?; // Activate the relay channel + tracing::info!( + relay = DEFAULT_RELAY_NAME, + owner_id = %state.owner_id, + "relay OAuth callback: activating relay channel" + ); ext_mgr .activate_stored_relay(DEFAULT_RELAY_NAME, &state.owner_id) .await @@ -2176,6 +2201,11 @@ async fn extensions_activate_handler( AuthenticatedUser(user): AuthenticatedUser, Path(name): Path, ) -> Result, (StatusCode, String)> { + tracing::debug!( + extension = %name, + user_id = %user.user_id, + "extensions_activate_handler: received activate request" + ); let ext_mgr = state.extension_manager.as_ref().ok_or(( StatusCode::NOT_IMPLEMENTED, "Extension manager not available (secrets store required)".to_string(), @@ -2183,6 +2213,10 @@ async fn extensions_activate_handler( match ext_mgr.activate(&name, &user.user_id).await { Ok(result) => { + tracing::info!( + extension = %name, + "extensions_activate_handler: activation succeeded" + ); // Activation loaded the WASM module. Check if the tool needs // OAuth scope expansion (e.g., adding google-docs when gmail // already has a token but missing the documents scope). @@ -2201,6 +2235,13 @@ async fn extensions_activate_handler( crate::extensions::ExtensionError::AuthRequired ); + tracing::debug!( + extension = %name, + error = %activate_err, + needs_auth = needs_auth, + "extensions_activate_handler: activation failed, attempting auth fallback" + ); + if !needs_auth { return Ok(Json(ActionResponse::fail(activate_err.to_string()))); } @@ -2208,10 +2249,21 @@ async fn extensions_activate_handler( // Activation failed due to auth; try authenticating first. match ext_mgr.auth(&name, &user.user_id).await { Ok(auth_result) if auth_result.is_authenticated() => { + tracing::debug!( + extension = %name, + "extensions_activate_handler: auth reports authenticated, retrying activate" + ); // Auth succeeded, retry activation. match ext_mgr.activate(&name, &user.user_id).await { Ok(result) => Ok(Json(ActionResponse::ok(result.message))), - Err(e) => Ok(Json(ActionResponse::fail(e.to_string()))), + Err(e) => { + tracing::warn!( + extension = %name, + error = %e, + "extensions_activate_handler: retry after auth still failed" + ); + Ok(Json(ActionResponse::fail(e.to_string()))) + } } } Ok(auth_result) => { diff --git a/src/config/agent.rs b/src/config/agent.rs index cb09707d..cfa0879a 100644 --- a/src/config/agent.rs +++ b/src/config/agent.rs @@ -1,6 +1,6 @@ use std::time::Duration; -use crate::config::helpers::{parse_bool_env, parse_option_env, parse_optional_env}; +use crate::config::helpers::{optional_env, parse_bool_env, parse_option_env, parse_optional_env}; use crate::error::ConfigError; use crate::settings::Settings; @@ -23,6 +23,8 @@ pub struct AgentConfig { pub max_cost_per_day_cents: Option, /// Maximum LLM/tool actions per hour. None = unlimited. pub max_actions_per_hour: Option, + /// Maximum daily LLM spend per user in cents. None = unlimited. + pub max_cost_per_user_per_day_cents: Option, /// Maximum tool-call iterations per agentic loop invocation. Default 50. pub max_tool_iterations: usize, /// When true, skip tool approval checks entirely. For benchmarks/CI. @@ -31,6 +33,13 @@ pub struct AgentConfig { pub default_timezone: String, /// Maximum tokens per job (0 = unlimited). pub max_tokens_per_job: u64, + /// 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 { @@ -49,10 +58,14 @@ impl AgentConfig { allow_local_tools: true, max_cost_per_day_cents: None, max_actions_per_hour: None, + max_cost_per_user_per_day_cents: None, max_tool_iterations: 10, auto_approve_tools: true, 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, } } @@ -87,6 +100,7 @@ impl AgentConfig { allow_local_tools: parse_bool_env("ALLOW_LOCAL_TOOLS", false)?, max_cost_per_day_cents: parse_option_env("MAX_COST_PER_DAY_CENTS")?, max_actions_per_hour: parse_option_env("MAX_ACTIONS_PER_HOUR")?, + max_cost_per_user_per_day_cents: parse_option_env("MAX_COST_PER_USER_PER_DAY_CENTS")?, max_tool_iterations: parse_optional_env( "AGENT_MAX_TOOL_ITERATIONS", settings.agent.max_tool_iterations, @@ -112,6 +126,11 @@ impl AgentConfig { "AGENT_MAX_TOKENS_PER_JOB", settings.agent.max_tokens_per_job, )?, + // 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/config/heartbeat.rs b/src/config/heartbeat.rs index 1dd456d7..09b8f0cd 100644 --- a/src/config/heartbeat.rs +++ b/src/config/heartbeat.rs @@ -21,6 +21,9 @@ pub struct HeartbeatConfig { pub quiet_hours_end: Option, /// Timezone for fire_at and quiet hours evaluation (IANA name). pub timezone: Option, + /// When true, cycle through all users with routines. Auto-detected from + /// GATEWAY_USER_TOKENS or set explicitly via HEARTBEAT_MULTI_TENANT. + pub multi_tenant: bool, } impl Default for HeartbeatConfig { @@ -34,6 +37,7 @@ impl Default for HeartbeatConfig { quiet_hours_start: None, quiet_hours_end: None, timezone: None, + multi_tenant: false, } } } @@ -101,6 +105,12 @@ impl HeartbeatConfig { } tz }, + // Auto-detect multi-tenant mode from GATEWAY_USER_TOKENS presence, + // or allow explicit override via HEARTBEAT_MULTI_TENANT. + multi_tenant: parse_bool_env( + "HEARTBEAT_MULTI_TENANT", + optional_env("GATEWAY_USER_TOKENS")?.is_some(), + )?, }) } } diff --git a/src/db/libsql/routines.rs b/src/db/libsql/routines.rs index 69c9f5c0..504d77dc 100644 --- a/src/db/libsql/routines.rs +++ b/src/db/libsql/routines.rs @@ -530,10 +530,24 @@ impl RoutineStore for LibSqlBackend { async fn get_webhook_routine_by_path( &self, path: &str, + user_id: Option<&str>, ) -> Result, DatabaseError> { let conn = self.connect().await?; - let mut rows = conn - .query( + let mut rows = if let Some(uid) = user_id { + conn.query( + &format!( + "SELECT {} FROM routines WHERE enabled = 1 AND trigger_type = 'webhook' \ + AND user_id = ?2 \ + AND (json_extract(trigger_config, '$.path') = ?1 \ + OR (json_extract(trigger_config, '$.path') IS NULL AND CAST(id AS TEXT) = ?1))", + ROUTINE_COLUMNS + ), + params![path, uid], + ) + .await + .map_err(|e| DatabaseError::Query(e.to_string()))? + } else { + conn.query( &format!( "SELECT {} FROM routines WHERE enabled = 1 AND trigger_type = 'webhook' \ AND (json_extract(trigger_config, '$.path') = ?1 \ @@ -543,7 +557,8 @@ impl RoutineStore for LibSqlBackend { params![path], ) .await - .map_err(|e| DatabaseError::Query(e.to_string()))?; + .map_err(|e| DatabaseError::Query(e.to_string()))? + }; match rows .next() diff --git a/src/db/mod.rs b/src/db/mod.rs index 6d984fed..d89b976e 100644 --- a/src/db/mod.rs +++ b/src/db/mod.rs @@ -545,6 +545,7 @@ pub trait RoutineStore: Send + Sync { async fn get_webhook_routine_by_path( &self, path: &str, + user_id: Option<&str>, ) -> Result, DatabaseError>; /// List routine runs that were dispatched as full_job but have not yet diff --git a/src/db/postgres.rs b/src/db/postgres.rs index 7bf76001..9e5ea9ce 100644 --- a/src/db/postgres.rs +++ b/src/db/postgres.rs @@ -529,8 +529,9 @@ impl RoutineStore for PgBackend { async fn get_webhook_routine_by_path( &self, path: &str, + user_id: Option<&str>, ) -> Result, DatabaseError> { - self.store.get_webhook_routine_by_path(path).await + self.store.get_webhook_routine_by_path(path, user_id).await } async fn list_dispatched_routine_runs(&self) -> Result, DatabaseError> { diff --git a/src/extensions/manager.rs b/src/extensions/manager.rs index 90920767..47b45a0f 100644 --- a/src/extensions/manager.rs +++ b/src/extensions/manager.rs @@ -659,6 +659,66 @@ impl ExtensionManager { }) } + /// Resolve the relay URL override for an extension from settings. + /// + /// Returns `Some(url)` if a non-empty per-extension `relay_url` override is + /// set for the given extension; otherwise returns `None` and callers should + /// fall back to the env-level `RelayConfig`. + /// + /// Uses `self.user_id` (owner scope) for consistency with `configure()`, + /// which also writes setting_path fields under the owner scope. + /// + /// The override is validated: only `http` / `https` schemes are accepted + /// and the URL must not contain userinfo (embedded credentials). This + /// prevents a malicious override from exfiltrating the instance-wide relay + /// API key to an attacker-controlled host. + async fn effective_relay_url(&self, name: &str) -> Option { + if let Some(ref store) = self.store { + let key = format!("extensions.{name}.relay_url"); + if let Ok(Some(v)) = store.get_setting(&self.user_id, &key).await { + let url = v + .as_str() + .map(|s| s.trim().to_string()) + .filter(|s| !s.is_empty()); + if let Some(ref u) = url { + // Validate the override to prevent API-key exfiltration: + // only allow http(s) with no embedded credentials. + match url::Url::parse(u) { + Ok(parsed) + if (parsed.scheme() == "http" || parsed.scheme() == "https") + && parsed.username().is_empty() + && parsed.password().is_none() => + { + tracing::debug!( + extension = %name, + relay_url_host = %parsed.host_str().unwrap_or("unknown"), + "effective_relay_url: using per-extension override from settings" + ); + return url; + } + Ok(parsed) => { + tracing::warn!( + extension = %name, + scheme = %parsed.scheme(), + has_userinfo = !parsed.username().is_empty() || parsed.password().is_some(), + "effective_relay_url: rejecting override — \ + only http/https without embedded credentials is allowed" + ); + } + Err(e) => { + tracing::warn!( + extension = %name, + error = %e, + "effective_relay_url: rejecting override — invalid URL" + ); + } + } + } + } + } + None + } + /// Get the shared relay event sender for the webhook endpoint. pub fn relay_event_tx( &self, @@ -892,6 +952,46 @@ impl ExtensionManager { false } + /// Check whether a stored `team_id` setting exists for the given relay extension. + /// + /// Unlike [`is_relay_channel`], this does **not** consult the in-memory + /// `installed_relay_extensions` set — it only looks at the persistent settings + /// store. This distinction matters for `auth_channel_relay`: an extension can + /// be *installed* (present in the in-memory set) but not yet *authenticated* + /// (no OAuth completed, no team_id stored). + async fn has_stored_team_id(&self, name: &str, _user_id: &str) -> bool { + if let Some(ref store) = self.store { + let key = format!("relay:{}:team_id", name); + // Use owner scope (self.user_id) for consistency: the OAuth callback + // stores team_id under state.owner_id which maps to self.user_id. + match store.get_setting(&self.user_id, &key).await { + Ok(Some(v)) => { + let has_id = v.as_str().is_some_and(|s| !s.is_empty()); + tracing::debug!( + extension = %name, + has_team_id = has_id, + "has_stored_team_id: checked store" + ); + return has_id; + } + Ok(None) => { + tracing::debug!( + extension = %name, + "has_stored_team_id: no team_id setting found" + ); + } + Err(e) => { + tracing::warn!( + extension = %name, + error = %e, + "has_stored_team_id: failed to read from settings store" + ); + } + } + } + false + } + /// Restore persisted relay channels after startup. /// /// Loads the persisted active channel list, filters to relay types (those with @@ -1418,7 +1518,7 @@ impl ExtensionManager { let errors = self.activation_errors.read().await; for name in installed.iter() { let active = active_names.contains(name); - let authenticated = self.is_relay_channel(name, user_id).await; + let authenticated = self.has_stored_team_id(name, user_id).await; let activation_error = errors.get(name).cloned(); let registry_entry = self .registry @@ -4191,20 +4291,69 @@ impl ExtensionManager { name: &str, user_id: &str, ) -> Result { - // Check if already authenticated (team_id setting exists) - if self.is_relay_channel(name, user_id).await { + tracing::debug!( + extension = %name, + user_id = %user_id, + "auth_channel_relay: starting" + ); + + // Check if already authenticated by looking for a stored team_id. + // We intentionally skip the `installed_relay_extensions` in-memory set + // here because that set only tracks *installed* extensions — an extension + // can be installed (via registry) but not yet authenticated (no OAuth + // completed). Checking just `is_relay_channel()` would short-circuit + // to "authenticated" even when no team_id exists, preventing the OAuth + // flow from being offered to the user. + if self.has_stored_team_id(name, user_id).await { + tracing::debug!( + extension = %name, + "auth_channel_relay: already authenticated (team_id in store)" + ); return Ok(AuthResult::authenticated(name, ExtensionKind::ChannelRelay)); } + tracing::debug!( + extension = %name, + "auth_channel_relay: no stored team_id, initiating OAuth" + ); + // Use relay config captured at startup - let relay_config = self.relay_config()?; + let relay_config = self.relay_config().map_err(|e| { + tracing::warn!( + extension = %name, + error = %e, + "auth_channel_relay: relay config not available — \ + CHANNEL_RELAY_URL and CHANNEL_RELAY_API_KEY must be set" + ); + e + })?; + + // Allow per-extension URL override from settings + let effective_url = self + .effective_relay_url(name) + .await + .unwrap_or_else(|| relay_config.url.clone()); + + tracing::debug!( + extension = %name, + relay_url = %effective_url, + "auth_channel_relay: creating relay client for OAuth" + ); let client = crate::channels::relay::RelayClient::new( - relay_config.url.clone(), + effective_url.clone(), relay_config.api_key.clone(), relay_config.request_timeout_secs, ) - .map_err(|e| ExtensionError::Config(e.to_string()))?; + .map_err(|e| { + tracing::warn!( + extension = %name, + relay_url = %effective_url, + error = %e, + "auth_channel_relay: failed to create relay HTTP client" + ); + ExtensionError::Config(e.to_string()) + })?; // Generate CSRF nonce — IronClaw validates this on the callback to ensure // the OAuth completion is legitimate. Channel-relay embeds it in the signed @@ -4216,18 +4365,44 @@ impl ExtensionManager { self.secrets .create(user_id, CreateSecretParams::new(&state_key, &state_nonce)) .await - .map_err(|e| ExtensionError::AuthFailed(format!("Failed to store OAuth state: {e}")))?; + .map_err(|e| { + tracing::warn!( + extension = %name, + error = %e, + "auth_channel_relay: failed to store OAuth state nonce" + ); + ExtensionError::AuthFailed(format!("Failed to store OAuth state: {e}")) + })?; // Channel-relay derives all URLs from trusted instance_url in chat-api. // We only pass the nonce for CSRF validation on the callback. + tracing::debug!( + extension = %name, + relay_url = %effective_url, + "auth_channel_relay: calling initiate_oauth on channel-relay" + ); match client.initiate_oauth(Some(&state_nonce)).await { - Ok(auth_url) => Ok(AuthResult::awaiting_authorization( - name, - ExtensionKind::ChannelRelay, - auth_url, - "redirect".to_string(), - )), - Err(e) => Err(ExtensionError::AuthFailed(e.to_string())), + Ok(auth_url) => { + tracing::info!( + extension = %name, + "auth_channel_relay: OAuth URL obtained, awaiting user authorization" + ); + Ok(AuthResult::awaiting_authorization( + name, + ExtensionKind::ChannelRelay, + auth_url, + "redirect".to_string(), + )) + } + Err(e) => { + tracing::warn!( + extension = %name, + relay_url = %effective_url, + error = %e, + "auth_channel_relay: initiate_oauth call to channel-relay failed" + ); + Err(ExtensionError::AuthFailed(e.to_string())) + } } } @@ -4237,40 +4412,112 @@ impl ExtensionManager { name: &str, user_id: &str, ) -> Result { + tracing::debug!( + extension = %name, + user_id = %user_id, + "activate_channel_relay: starting" + ); + let team_id_key = format!("relay:{}:team_id", name); // Get team_id from settings (stored by the OAuth callback) let team_id = if let Some(ref store) = self.store { - store - .get_setting(user_id, &team_id_key) - .await - .ok() - .flatten() - .and_then(|v| v.as_str().map(|s| s.to_string())) - .unwrap_or_default() + match store.get_setting(user_id, &team_id_key).await { + Ok(Some(v)) => { + let id = v.as_str().map(|s| s.to_string()).unwrap_or_default(); + tracing::debug!( + extension = %name, + team_id_empty = id.is_empty(), + "activate_channel_relay: loaded team_id from store" + ); + id + } + Ok(None) => { + tracing::debug!( + extension = %name, + setting_key = %team_id_key, + "activate_channel_relay: no team_id in settings store" + ); + String::new() + } + Err(e) => { + tracing::warn!( + extension = %name, + error = %e, + "activate_channel_relay: failed to read team_id from settings store" + ); + String::new() + } + } } else { + tracing::debug!( + extension = %name, + "activate_channel_relay: no settings store available" + ); String::new() }; if team_id.is_empty() { + tracing::debug!( + extension = %name, + "activate_channel_relay: team_id is empty, returning AuthRequired" + ); return Err(ExtensionError::AuthRequired); } // Use relay config captured at startup - let relay_config = self.relay_config()?; + let relay_config = self.relay_config().map_err(|e| { + tracing::warn!( + extension = %name, + error = %e, + "activate_channel_relay: relay config not available" + ); + e + })?; + + // Allow per-extension URL override from settings + let effective_url = self + .effective_relay_url(name) + .await + .unwrap_or_else(|| relay_config.url.clone()); + + tracing::debug!( + extension = %name, + relay_url = %effective_url, + "activate_channel_relay: relay config loaded" + ); let instance_id = self.relay_instance_id(relay_config, user_id); let client = crate::channels::relay::RelayClient::new( - relay_config.url.clone(), + effective_url.clone(), relay_config.api_key.clone(), relay_config.request_timeout_secs, ) - .map_err(|e| ExtensionError::ActivationFailed(e.to_string()))?; + .map_err(|e| { + tracing::warn!( + extension = %name, + relay_url = %effective_url, + error = %e, + "activate_channel_relay: failed to create relay HTTP client" + ); + ExtensionError::ActivationFailed(e.to_string()) + })?; // Fetch the per-instance signing secret from channel-relay. // This must succeed — there is no fallback. + tracing::debug!( + extension = %name, + relay_url = %effective_url, + "activate_channel_relay: fetching signing secret from channel-relay" + ); let signing_secret = client.get_signing_secret(&team_id).await.map_err(|e| { + tracing::warn!( + extension = %name, + relay_url = %effective_url, + error = %e, + "activate_channel_relay: failed to fetch signing secret from channel-relay" + ); ExtensionError::Config(format!("Failed to fetch relay signing secret: {e}")) })?; @@ -4289,16 +4536,29 @@ impl ExtensionManager { // Hot-add to channel manager let cm_guard = self.relay_channel_manager.read().await; let channel_mgr = cm_guard.as_ref().ok_or_else(|| { + tracing::warn!( + extension = %name, + "activate_channel_relay: channel manager not initialized" + ); ExtensionError::ActivationFailed("Channel manager not initialized".to_string()) })?; - channel_mgr - .hot_add(Box::new(channel)) - .await - .map_err(|e| ExtensionError::ActivationFailed(e.to_string()))?; + channel_mgr.hot_add(Box::new(channel)).await.map_err(|e| { + tracing::warn!( + extension = %name, + error = %e, + "activate_channel_relay: hot_add to channel manager failed" + ); + ExtensionError::ActivationFailed(e.to_string()) + })?; if let Ok(mut cache) = self.relay_signing_secret_cache.lock() { *cache = Some(signing_secret); + } else { + tracing::warn!( + extension = %name, + "activate_channel_relay: failed to cache signing secret (mutex poisoned)" + ); } // Store the event sender so the web gateway's relay webhook endpoint can push events @@ -4316,6 +4576,12 @@ impl ExtensionManager { self.broadcast_extension_status(name, "active", Some(&status_msg)) .await; + tracing::info!( + extension = %name, + instance_id = %instance_id, + "activate_channel_relay: relay channel activated successfully" + ); + Ok(ActivateResult { name: name.to_string(), kind: ExtensionKind::ChannelRelay, @@ -4595,6 +4861,41 @@ impl ExtensionManager { } Ok(ExtensionSetupSchema { secrets, fields }) } + ExtensionKind::ChannelRelay => { + let relay_url_key = format!("extensions.{name}.relay_url"); + let current_url = if let Some(ref store) = self.store { + match store.get_setting(&self.user_id, &relay_url_key).await { + Ok(value_opt) => value_opt + .and_then(|v| v.as_str().map(|s| s.to_string())) + .filter(|s| !s.is_empty()), + Err(e) => { + tracing::warn!( + extension = %name, + setting_key = %relay_url_key, + error = %e, + "get_setup_schema: failed to read relay_url from settings" + ); + None + } + } + } else { + None + }; + let env_url = self.relay_config.as_ref().map(|c| c.url.as_str()); + Ok(ExtensionSetupSchema { + secrets: Vec::new(), + fields: vec![crate::channels::web::types::SetupFieldInfo { + name: "relay_url".to_string(), + prompt: format!( + "Channel-relay service URL (leave empty to use env default{})", + env_url.map(|u| format!(": {u}")).unwrap_or_default() + ), + optional: true, + provided: current_url.is_some(), + input_type: crate::tools::wasm::ToolSetupFieldInputType::Text, + }], + }) + } _ => Ok(ExtensionSetupSchema { secrets: Vec::new(), fields: Vec::new(), @@ -4997,7 +5298,17 @@ impl ExtensionManager { names.insert(server.token_secret_name()); (names, Vec::new()) } - ExtensionKind::ChannelRelay => (std::collections::HashSet::new(), Vec::new()), + ExtensionKind::ChannelRelay => { + let relay_fields = vec![crate::tools::wasm::ToolFieldSetupSchema { + name: "relay_url".to_string(), + prompt: "Channel-relay service URL override".to_string(), + optional: true, + setting_path: Some(format!("extensions.{name}.relay_url")), + input_type: crate::tools::wasm::ToolSetupFieldInputType::Text, + restart_required: false, + }]; + (std::collections::HashSet::new(), relay_fields) + } }; let allowed_fields: std::collections::HashSet = @@ -5088,13 +5399,28 @@ impl ExtensionManager { ))); } let trimmed = field_value.trim(); + let field_def = setup_field_defs.get(field_name); + + // Empty value on an optional field with a setting_path: clear the + // stored override so the system reverts to the env/default value. if trimmed.is_empty() { + if let Some(def) = field_def + && def.optional + { + stored_fields.remove(field_name); + if let Some(setting_path) = &def.setting_path { + Self::validate_setup_setting_path(name, setting_path)?; + if let Some(store) = self.store.as_ref() { + let _ = store.delete_setting(&self.user_id, setting_path).await; + } + } + } continue; } stored_fields.insert(field_name.clone(), trimmed.to_string()); - if let Some(field_def) = setup_field_defs.get(field_name) { + if let Some(field_def) = field_def { if field_def.restart_required { restart_required = true; } @@ -7058,6 +7384,39 @@ mod tests { ); } + /// Regression: installed-but-not-authenticated relay must NOT short-circuit + /// `auth_channel_relay()` to "authenticated". Previously, `auth_channel_relay` + /// called `is_relay_channel()` which checked the in-memory + /// `installed_relay_extensions` set; that returned `true` even when no team_id + /// existed in the store, so the OAuth URL was never offered. + #[tokio::test] + async fn test_auth_channel_relay_installed_without_team_id_is_not_authenticated() { + let dir = tempfile::tempdir().expect("temp dir"); + let mgr = make_test_manager(None, dir.path().to_path_buf()); + + // Mark as installed (simulates clicking Install in the UI) + mgr.installed_relay_extensions + .write() + .await + .insert("slack-relay".to_string()); + + // Without a stored team_id, auth should NOT return authenticated. + // It should fail because relay config is missing (no CHANNEL_RELAY_URL), + // but the key assertion is that it does NOT return Ok(authenticated). + let result = mgr.auth_channel_relay("slack-relay", "test").await; + match result { + Ok(ref auth_result) if auth_result.is_authenticated() => { + panic!( + "auth_channel_relay returned authenticated for installed-but-no-team-id relay; \ + expected either an OAuth URL or a config error" + ); + } + _ => { + // Config error (no relay URL) or awaiting_authorization — both are correct + } + } + } + #[tokio::test] async fn test_remove_relay_shuts_down_via_relay_channel_manager() { // Regression: remove() only checked channel_runtime for shutdown, missing diff --git a/src/history/store.rs b/src/history/store.rs index 1e4cdd82..625e8b1e 100644 --- a/src/history/store.rs +++ b/src/history/store.rs @@ -1162,15 +1162,25 @@ impl Store { pub async fn get_webhook_routine_by_path( &self, path: &str, + user_id: Option<&str>, ) -> Result, DatabaseError> { let conn = self.conn().await?; - let row = conn - .query_opt( + let row = if let Some(uid) = user_id { + conn.query_opt( + "SELECT * FROM routines WHERE enabled AND trigger_type = 'webhook' \ + AND user_id = $2 \ + AND (trigger_config->>'path' = $1 OR (trigger_config->>'path' IS NULL AND id::text = $1))", + &[&path, &uid], + ) + .await? + } else { + conn.query_opt( "SELECT * FROM routines WHERE enabled AND trigger_type = 'webhook' \ AND (trigger_config->>'path' = $1 OR (trigger_config->>'path' IS NULL AND id::text = $1))", &[&path], ) - .await?; + .await? + }; row.as_ref().map(row_to_routine).transpose() } 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/llm/mod.rs b/src/llm/mod.rs index 308b3983..d681547d 100644 --- a/src/llm/mod.rs +++ b/src/llm/mod.rs @@ -63,7 +63,8 @@ pub use provider::{ }; pub use reasoning::{ ActionPlan, Reasoning, ReasoningContext, RespondOutput, RespondResult, SILENT_REPLY_TOKEN, - TOOL_INTENT_NUDGE, TokenUsage, ToolSelection, is_silent_reply, llm_signals_tool_intent, + TOOL_INTENT_NUDGE, TRUNCATED_TOOL_CALL_NOTICE, TokenUsage, ToolSelection, is_silent_reply, + llm_signals_tool_intent, }; pub use recording::RecordingLlm; pub use registry::{ProviderDefinition, ProviderProtocol, ProviderRegistry}; diff --git a/src/llm/reasoning.rs b/src/llm/reasoning.rs index 77905f95..6e078ac7 100644 --- a/src/llm/reasoning.rs +++ b/src/llm/reasoning.rs @@ -8,8 +8,8 @@ use serde::{Deserialize, Serialize}; use crate::llm::error::LlmError; use crate::llm::{ - ChatMessage, CompletionRequest, LlmProvider, Role, ToolCall, ToolCompletionRequest, - ToolDefinition, + ChatMessage, CompletionRequest, FinishReason, LlmProvider, Role, ToolCall, + ToolCompletionRequest, ToolDefinition, }; /// Token the agent returns when it has nothing to say (e.g. in group chats). @@ -23,6 +23,13 @@ You said you would perform an action, but you did not include any tool calls.\n\ Do NOT describe what you intend to do — actually call the tool now.\n\ Use the tool_calls mechanism to invoke the appropriate tool."; +/// Notice injected when the LLM's response was truncated mid-tool-call, +/// causing incomplete parameters. Tells the LLM to try a different approach. +pub const TRUNCATED_TOOL_CALL_NOTICE: &str = "\ +Your previous response was truncated while generating tool call parameters. \ +The tool calls were discarded. Please try a different approach — \ +summarize or transform the data instead of echoing it verbatim in a tool call."; + /// Seed value used as the second argument to `generate_tool_call_id` when /// recovering tool calls from malformed LLM text responses. This must differ /// from the `0` seed used in `rig_adapter::normalized_tool_call_id` to avoid @@ -194,11 +201,17 @@ pub struct ReasoningContext { pub metadata: std::collections::HashMap, /// When true, force a text-only response (ignore available tools). /// Used by the agentic loop to guarantee termination near the iteration limit. + /// Sticky: once set, never cleared within a loop invocation. Callers must + /// create a fresh `ReasoningContext` per `run_agentic_loop()` call. pub force_text: bool, /// Pre-built system prompt. When set, `respond_with_tools` uses this directly /// instead of calling `build_system_prompt_with_tools`. Allows callers to build /// the prompt once and reuse it across iterations. pub system_prompt: Option, + /// Per-user model override. When set, completion requests use this model + /// instead of the provider's default. Only effective with providers that + /// support per-request model overrides (e.g. NearAI). + pub model_override: Option, } impl ReasoningContext { @@ -212,6 +225,7 @@ impl ReasoningContext { metadata: std::collections::HashMap::new(), force_text: false, system_prompt: None, + model_override: None, } } @@ -344,6 +358,7 @@ pub enum RespondResult { pub struct RespondOutput { pub result: RespondResult, pub usage: TokenUsage, + pub finish_reason: FinishReason, } /// Reasoning engine for the agent. @@ -525,6 +540,17 @@ impl Reasoning { let response = self.llm.complete_with_tools(request).await?; + // If the response was truncated, tool call parameters are likely incomplete. + // Return empty so the caller can fall through to respond_with_tools() which + // has a larger output token budget. + if response.finish_reason == FinishReason::Length { + tracing::warn!( + "select_tools response truncated (finish_reason=Length), \ + discarding potentially incomplete tool selections" + ); + return Ok(vec![]); + } + let shared_reasoning = response .content .map(|c| { @@ -671,6 +697,9 @@ Respond in JSON format: .with_temperature(0.7) .with_tool_choice("auto"); request.metadata = context.metadata.clone(); + if let Some(ref model) = context.model_override { + request.model = Some(model.clone()); + } let response = self.llm.complete_with_tools(request).await?; let usage = TokenUsage { @@ -714,6 +743,7 @@ Respond in JSON format: content: narrative, }, usage, + finish_reason: response.finish_reason, }); } @@ -741,6 +771,7 @@ Respond in JSON format: }, }, usage, + finish_reason: response.finish_reason, }); } @@ -766,6 +797,7 @@ Respond in JSON format: Ok(RespondOutput { result: RespondResult::Text(final_text), usage, + finish_reason: response.finish_reason, }) } else { // No tools, use simple completion @@ -773,6 +805,9 @@ Respond in JSON format: .with_max_tokens(4096) .with_temperature(0.7); request.metadata = context.metadata.clone(); + if let Some(ref model) = context.model_override { + request.model = Some(model.clone()); + } let response = self.llm.complete(request).await?; let pre_truncated = truncate_at_tool_tags(&response.content); @@ -794,6 +829,7 @@ Respond in JSON format: cache_read_input_tokens: response.cache_read_input_tokens, cache_creation_input_tokens: response.cache_creation_input_tokens, }, + finish_reason: response.finish_reason, }) } } @@ -1334,6 +1370,49 @@ fn is_inside_code(pos: usize, regions: &[CodeRegion]) -> bool { regions.iter().any(|r| pos >= r.start && pos < r.end) } +/// Check whether a byte range overlaps any code region. +fn overlaps_code_region(start: usize, end: usize, regions: &[CodeRegion]) -> bool { + regions.iter().any(|r| start < r.end && end > r.start) +} + +/// Return the byte bounds of the line containing `pos`, excluding the trailing newline. +fn line_bounds(text: &str, pos: usize) -> (usize, usize) { + let start = text[..pos].rfind('\n').map_or(0, |idx| idx + 1); + let end = text[pos..].find('\n').map_or(text.len(), |idx| pos + idx); + (start, end) +} + +/// Only recover XML-style tool calls when they are isolated content outside +/// markdown code and quote contexts. This avoids converting code examples or +/// quoted snippets into executable tool calls. +fn is_recoverable_tool_call_segment( + text: &str, + start: usize, + end: usize, + code_regions: &[CodeRegion], +) -> bool { + if overlaps_code_region(start, end, code_regions) { + return false; + } + + let (first_line_start, first_line_end) = line_bounds(text, start); + let first_line = &text[first_line_start..first_line_end]; + + if first_line.trim_start().starts_with('>') { + return false; + } + + let (_, last_line_end) = line_bounds(text, end.saturating_sub(1)); + let first_line_prefix = &text[first_line_start..start]; + let last_line_suffix = &text[end..last_line_end]; + + if !first_line_prefix.trim().is_empty() || !last_line_suffix.trim().is_empty() { + return false; + } + + true +} + /// Clean up LLM response by stripping model-internal tags and reasoning patterns. /// /// Some models (GLM-4.7, etc.) emit XML-tagged internal state like @@ -1353,6 +1432,7 @@ fn recover_tool_calls_from_content( ) -> Vec { let tool_names: std::collections::HashSet<&str> = available_tools.iter().map(|t| t.name.as_str()).collect(); + let code_regions = find_code_regions(content); let mut calls = Vec::new(); for (open, close) in &[ @@ -1361,15 +1441,23 @@ fn recover_tool_calls_from_content( ("", ""), ("<|function_call|>", "<|/function_call|>"), ] { - let mut remaining = content; - while let Some(start) = remaining.find(open) { + let mut search_from = 0; + while let Some(offset) = content[search_from..].find(open) { + let start = search_from + offset; let inner_start = start + open.len(); - let after = &remaining[inner_start..]; - let Some(end) = after.find(close) else { + let after = &content[inner_start..]; + let Some(end_offset) = after.find(close) else { break; }; - let inner = after[..end].trim(); - remaining = &after[end + close.len()..]; + let end = inner_start + end_offset; + let segment_end = end + close.len(); + search_from = segment_end; + + if !is_recoverable_tool_call_segment(content, start, segment_end, &code_regions) { + continue; + } + + let inner = content[inner_start..end].trim(); if inner.is_empty() { continue; @@ -2302,6 +2390,40 @@ That's my plan."#; assert_eq!(calls[0].name, "tool_list"); } + #[test] + fn test_recover_tool_call_in_fenced_code_block_ignored() { + let tools = make_tools(&["tool_list"]); + let content = "Here is the XML format:\n\n```xml\ntool_list\n```"; + let calls = recover_tool_calls_from_content(content, &tools); + assert!(calls.is_empty()); + } + + #[test] + fn test_recover_tool_call_in_inline_code_ignored() { + let tools = make_tools(&["tool_list"]); + let content = "Use `tool_list` to illustrate the syntax."; + let calls = recover_tool_calls_from_content(content, &tools); + assert!(calls.is_empty()); + } + + #[test] + fn test_recover_tool_call_in_blockquote_ignored() { + let tools = make_tools(&["tool_list"]); + let content = "The page replied:\n> tool_list"; + let calls = recover_tool_calls_from_content(content, &tools); + assert!(calls.is_empty()); + } + + #[test] + fn test_recover_multiline_json_tool_call_on_own_line() { + let tools = make_tools(&["memory_search"]); + let content = "Let me check.\n\n\n{\"name\": \"memory_search\", \"arguments\": {\"query\": \"test\"}}\n\n\nDone."; + let calls = recover_tool_calls_from_content(content, &tools); + assert_eq!(calls.len(), 1); + assert_eq!(calls[0].name, "memory_search"); + assert_eq!(calls[0].arguments, serde_json::json!({"query": "test"})); + } + // ---- System prompt building tests (issue #565) ---- fn make_test_reasoning() -> Reasoning { @@ -3218,4 +3340,85 @@ That's my plan."#; let cleaned = clean_response(&pre_truncated); assert!(cleaned.trim().is_empty()); } + + // ---- select_tools truncation guard ---- + + /// Mock provider that returns tool calls with a configurable finish_reason. + struct TruncatingLlm { + finish_reason: crate::llm::FinishReason, + } + + #[async_trait::async_trait] + impl crate::llm::LlmProvider for TruncatingLlm { + fn model_name(&self) -> &str { + "truncating-stub" + } + fn cost_per_token(&self) -> (rust_decimal::Decimal, rust_decimal::Decimal) { + (rust_decimal::Decimal::ZERO, rust_decimal::Decimal::ZERO) + } + async fn complete( + &self, + _request: crate::llm::CompletionRequest, + ) -> Result { + unimplemented!() + } + async fn complete_with_tools( + &self, + _request: crate::llm::ToolCompletionRequest, + ) -> Result { + Ok(crate::llm::ToolCompletionResponse { + content: Some("I'll write the report.".to_string()), + tool_calls: vec![ToolCall { + id: "call_1".to_string(), + name: "memory_write".to_string(), + arguments: serde_json::json!({}), + reasoning: None, + }], + input_tokens: 5000, + output_tokens: 1024, + finish_reason: self.finish_reason, + cache_read_input_tokens: 0, + cache_creation_input_tokens: 0, + }) + } + } + + #[tokio::test] + async fn test_select_tools_returns_empty_on_truncation() { + let llm = Arc::new(TruncatingLlm { + finish_reason: FinishReason::Length, + }); + let reasoning = Reasoning::new(llm); + let mut ctx = ReasoningContext::new().with_message(ChatMessage::user("Write a report")); + ctx.available_tools.push(ToolDefinition { + name: "memory_write".to_string(), + description: "Write to memory".to_string(), + parameters: serde_json::json!({"type": "object"}), + }); + + let selections = reasoning.select_tools(&ctx).await.unwrap(); + assert!( + selections.is_empty(), + "Truncated tool selections should be discarded (got {} selections)", + selections.len() + ); + } + + #[tokio::test] + async fn test_select_tools_returns_selections_when_not_truncated() { + let llm = Arc::new(TruncatingLlm { + finish_reason: FinishReason::ToolUse, + }); + let reasoning = Reasoning::new(llm); + let mut ctx = ReasoningContext::new().with_message(ChatMessage::user("Write a report")); + ctx.available_tools.push(ToolDefinition { + name: "memory_write".to_string(), + description: "Write to memory".to_string(), + parameters: serde_json::json!({"type": "object"}), + }); + + let selections = reasoning.select_tools(&ctx).await.unwrap(); + assert_eq!(selections.len(), 1); + assert_eq!(selections[0].tool_name, "memory_write"); + } } diff --git a/src/llm/rig_adapter.rs b/src/llm/rig_adapter.rs index 7a6b2ae8..038236fd 100644 --- a/src/llm/rig_adapter.rs +++ b/src/llm/rig_adapter.rs @@ -598,6 +598,30 @@ fn build_rig_request( }) } +/// Inject a per-request model override into the rig request's `additional_params`. +/// +/// Rig-core bakes the model name at construction time inside each provider's +/// `CompletionModel` implementation. The actual HTTP request body includes a +/// `model` field set by the provider. Rig-core's `#[serde(flatten)]` on +/// `additional_params` emits these fields AFTER the provider's own fields. +/// Most API servers (Python, Go) use last-key-wins when deserializing +/// duplicate JSON keys, so the injected `model` value takes effect. +fn inject_model_override(rig_req: &mut RigRequest, model_override: Option<&str>) { + let Some(model) = model_override else { + return; + }; + match rig_req.additional_params { + Some(ref mut params) => { + if let Some(obj) = params.as_object_mut() { + obj.insert("model".to_string(), serde_json::json!(model)); + } + } + None => { + rig_req.additional_params = Some(serde_json::json!({ "model": model })); + } + } +} + #[async_trait] impl LlmProvider for RigAdapter where @@ -632,15 +656,7 @@ where &self, mut request: CompletionRequest, ) -> Result { - if let Some(requested_model) = request.model.as_deref() - && requested_model != self.model_name.as_str() - { - tracing::warn!( - requested_model = requested_model, - active_model = %self.model_name, - "Per-request model override is not supported for this provider; using configured model" - ); - } + let model_override = request.model.take(); self.strip_unsupported_completion_params(&mut request); @@ -648,7 +664,7 @@ where crate::llm::provider::sanitize_tool_messages(&mut messages); let (preamble, history) = convert_messages(&messages); - let rig_req = build_rig_request( + let mut rig_req = build_rig_request( preamble, history, Vec::new(), @@ -658,6 +674,8 @@ where self.cache_retention, )?; + inject_model_override(&mut rig_req, model_override.as_deref()); + let response = self.model .completion(rig_req) @@ -695,15 +713,7 @@ where &self, mut request: ToolCompletionRequest, ) -> Result { - if let Some(requested_model) = request.model.as_deref() - && requested_model != self.model_name.as_str() - { - tracing::warn!( - requested_model = requested_model, - active_model = %self.model_name, - "Per-request model override is not supported for this provider; using configured model" - ); - } + let model_override = request.model.take(); self.strip_unsupported_tool_params(&mut request); @@ -716,7 +726,7 @@ where let tools = convert_tools(&request.tools); let tool_choice = convert_tool_choice(request.tool_choice.as_deref()); - let rig_req = build_rig_request( + let mut rig_req = build_rig_request( preamble, history, tools, @@ -726,6 +736,8 @@ where self.cache_retention, )?; + inject_model_override(&mut rig_req, model_override.as_deref()); + let response = self.model .completion(rig_req) diff --git a/src/main.rs b/src/main.rs index e885cb7d..3a43ce0d 100644 --- a/src/main.rs +++ b/src/main.rs @@ -914,6 +914,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 e580b169..dfff4b10 100644 --- a/src/testing/mod.rs +++ b/src/testing/mod.rs @@ -532,6 +532,7 @@ impl TestHarnessBuilder { let cost_guard = Arc::new(CostGuard::new(CostGuardConfig { max_cost_per_day_cents: None, max_actions_per_hour: None, + max_cost_per_user_per_day_cents: None, })); let channel = if self.stub_channel { @@ -564,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 669c69f0..f74d4ec8 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::ToolDecisionDto; 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::{ @@ -45,7 +45,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, @@ -94,7 +94,7 @@ impl Worker { &self.deps.tools } - fn store(&self) -> Option<&Arc> { + fn store(&self) -> Option<&AdminScope> { self.deps.store.as_ref() } @@ -1158,6 +1158,7 @@ impl<'a> JobDelegate<'a> { Ok(crate::llm::RespondOutput { result: RespondResult::Text(String::new()), usage: crate::llm::TokenUsage::default(), + finish_reason: crate::llm::FinishReason::Stop, }) } } @@ -1283,6 +1284,7 @@ impl<'a> LoopDelegate for JobDelegate<'a> { content: reasoning_text, }, usage: crate::llm::TokenUsage::default(), + finish_reason: crate::llm::FinishReason::ToolUse, }); } Ok(_) => {} // empty selections, fall through 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 5f477de0..ac35b160 100644 --- a/tests/support/gateway_workflow_harness.rs +++ b/tests/support/gateway_workflow_harness.rs @@ -266,6 +266,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.