feat: complete multi-tenant isolation — phases 2–4 (#1614)

* feat: complete multi-tenant isolation — per-user budgets, model selection, heartbeat cycling

Finishes the remaining isolation work from phases 2–4 of #59:

Phase 2 (DB scoping): Fix /status and /list commands to use _for_user
DB variants instead of global queries that leaked cross-user job data.

Phase 3 (Runtime isolation): Per-user workspace in routine engine's
spawn_fire so lightweight routines run in the correct user context.
Per-user daily cost tracking in CostGuard with configurable budget via
MAX_COST_PER_USER_PER_DAY_CENTS. Multi-user heartbeat that cycles
through all users with routines, auto-detected from GATEWAY_USER_TOKENS.

Phase 4 (Provider/tools): Per-user model selection via preferred_model
setting — looked up from SettingsStore on first iteration, threaded
through ReasoningContext.model_override to CompletionRequest. Works
with providers that support per-request model overrides (NearAI).

Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]>

* fix: use selected_model setting key to match /model command persistence

The dispatcher was reading "preferred_model" but the /model command
(merged from staging) persists to "selected_model". Since set_setting
is already per-user scoped, using the same key makes /model work as
the per-user model override in multi-tenant mode.

Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]>

* fix: heartbeat hygiene, /model multi-tenant guard, RigAdapter model override

Three follow-up fixes for multi-tenant isolation:

1. Multi-user heartbeat now runs memory hygiene per user before each
   heartbeat check, matching single-user heartbeat behavior.

2. /model command in multi-tenant mode only persists to per-user
   settings (selected_model) without calling set_model() on the shared
   LlmProvider. The per-request model_override in the dispatcher reads
   from the same setting. Added multi_tenant flag to AgentConfig
   (auto-detected from GATEWAY_USER_TOKENS).

3. RigAdapter now supports per-request model overrides by injecting the
   model name into rig-core's additional_params. OpenAI/Anthropic/Ollama
   API servers use last-key-wins for duplicate JSON keys, so the override
   takes effect via serde's flatten serialization order.

Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]>

* fix: address PR review — cost model attribution, heartbeat concurrency, pruning

Fixes from review comments on #1614:

- Cost tracking now uses the override model name (not active_model_name)
  when a per-user model override is active, for accurate attribution.
- Multi-user heartbeat runs per-user checks concurrently via JoinSet
  instead of sequentially, preventing one slow user from blocking others.
- Per-user failure counts tracked independently; users exceeding
  max_failures are skipped (matching single-user semantics).
- per_user_daily_cost HashMap pruned on day rollover to prevent
  unbounded growth in long-lived deployments.
- Doc comment fixed: says "routines" not "active routines".

Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]>

* fix: /status ownership, model persistence scoping, heartbeat robustness

Addresses second round of PR review on #1614:

- /status <job_id> DB path now validates job.user_id == requesting user
  before returning data (was missing ownership check, security fix).

- persist_selected_model takes user_id param instead of owner_id, and
  skips .env/TOML writes in multi-tenant mode (these are shared global
  files). handle_system_command now receives user_id from caller.

- JoinSet collection handles Err(JoinError) explicitly instead of
  silently dropping panicked tasks.

- Notification forwarder extracts owner_id from response metadata in
  multi-tenant mode for per-user routing instead of broadcasting to
  the agent owner.

Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]>

* fix: cost pricing, fire_manual workspace, heartbeat concurrency cap

Round 3 review fixes:

- Cost tracking passes None for cost_per_token when model override is
  active, letting CostGuard look up pricing by model name instead of
  using the default provider's rates (serrrfirat).

- fire_manual() now uses per-user workspace, matching spawn_fire()
  pattern (serrrfirat).

- Removed MULTI_TENANT env var — multi-tenant mode is auto-detected
  solely from GATEWAY_USER_TOKENS presence (serrrfirat + Copilot).

- Multi-user heartbeat capped at 8 concurrent tasks to avoid flooding
  the LLM provider (serrrfirat + Copilot).

- Fixed inject_model_override doc comment accuracy (Copilot).

- Added comment explaining multi-tenant notification routing priority
  (Copilot).

Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]>

* feat: user-scoped webhook endpoint for multi-tenant isolation

Adds POST /api/webhooks/u/{user_id}/{path} — a user-scoped webhook
endpoint that filters the routine lookup by user_id, preventing
cross-user webhook triggering when paths collide.

The existing /api/webhooks/{path} endpoint remains unchanged for
backward compatibility in single-user deployments.

Changes:
- get_webhook_routine_by_path gains user_id: Option<&str> param
- Both postgres and libsql implementations add AND user_id = ? filter
  when user_id is provided
- New webhook_trigger_user_scoped_handler extracts (user_id, path)
  from URL and passes to shared fire_webhook_inner logic
- Route registered on public router (webhooks are called by external
  services that can't send bearer tokens)

Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]>

* feat: add TenantCtx for compile-time tenant isolation

Implements zmanian's architectural proposal from #1614 review:
two-tier scoped database access (TenantScope/AdminScope) so handler
code cannot accidentally bypass tenant scoping.

TenantScope (default): wraps user_id + Arc<dyn Database>, auto-binds
user_id on every operation. ID-based lookups return None for cross-
tenant resources. No escape hatch — forgetting to scope is a compile
error.

AdminScope (explicit opt-in): cross-tenant access for system-level
components (heartbeat, routine engine, self-repair, scheduler, worker).

TenantCtx bundles TenantScope + workspace + cost guard + per-user
rate limiting. Constructed once per request in handle_message, threaded
through all command handlers and ChatDelegate.

Key changes:
- New src/tenant.rs (~920 lines): TenantScope, AdminScope, TenantCtx,
  TenantRateState, TenantRateRegistry
- All command handlers: user_id: &str → ctx: &TenantCtx
- ChatDelegate: cost check/record/settings via self.tenant
- System components: store field changed to AdminScope
- Config: TENANT_MAX_LLM_CONCURRENT, TENANT_MAX_JOBS_CONCURRENT env vars
- Fixes bug: /status <job_id> cross-tenant leak (now auto-filtered)

Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]>

---------

Co-authored-by: Claude Opus 4.6 (1M context) <[email protected]>
This commit is contained in:
Illia Polosukhin
2026-03-25 17:24:48 -07:00
committed by GitHub
co-authored by Claude Opus 4.6
parent 86d1143064
commit 4c043bf057
30 changed files with 1825 additions and 174 deletions
+136 -37
View File
@@ -13,7 +13,7 @@ use futures::StreamExt;
use uuid::Uuid; use uuid::Uuid;
use crate::agent::context_monitor::ContextMonitor; 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::routine_engine::{RoutineEngine, spawn_cron_ticker};
use crate::agent::self_repair::{DefaultSelfRepair, RepairResult, SelfRepair}; use crate::agent::self_repair::{DefaultSelfRepair, RepairResult, SelfRepair};
use crate::agent::session::ThreadState; use crate::agent::session::ThreadState;
@@ -182,6 +182,8 @@ pub struct AgentDeps {
/// Resolved LLM backend identifier (e.g., "nearai", "openai", "groq"). /// Resolved LLM backend identifier (e.g., "nearai", "openai", "groq").
/// Used by `/model` persistence to determine which env var to update. /// Used by `/model` persistence to determine which env var to update.
pub llm_backend: String, pub llm_backend: String,
/// Per-tenant rate limiting registry (lazily creates rate state per user).
pub tenant_rates: Arc<crate::tenant::TenantRateRegistry>,
} }
/// The main agent that coordinates all components. /// The main agent that coordinates all components.
@@ -244,7 +246,10 @@ impl Agent {
SchedulerDeps { SchedulerDeps {
tools: deps.tools.clone(), tools: deps.tools.clone(),
extension_manager: deps.extension_manager.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(), hooks: deps.hooks.clone(),
}, },
); );
@@ -325,6 +330,50 @@ impl Agent {
&self.deps.cost_guard &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<crate::tenant::AdminScope> {
self.deps
.store
.as_ref()
.map(|db| crate::tenant::AdminScope::new(Arc::clone(db)))
}
pub(super) fn skill_registry(&self) -> Option<&Arc<std::sync::RwLock<SkillRegistry>>> { pub(super) fn skill_registry(&self) -> Option<&Arc<std::sync::RwLock<SkillRegistry>>> {
self.deps.skill_registry.as_ref() self.deps.skill_registry.as_ref()
} }
@@ -410,8 +459,8 @@ impl Agent {
self.config.stuck_threshold, self.config.stuck_threshold,
self.config.max_repair_attempts, self.config.max_repair_attempts,
); );
if let Some(ref store) = self.deps.store { if let Some(admin) = self.admin_store() {
self_repair = self_repair.with_store(Arc::clone(store)); self_repair = self_repair.with_store(admin);
} }
if let Some(ref builder) = self.deps.builder { if let Some(ref builder) = self.deps.builder {
self_repair = self_repair.with_builder(Arc::clone(builder), Arc::clone(self.tools())); 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)); .with_interval(std::time::Duration::from_secs(hb_config.interval_secs));
config.quiet_hours_start = hb_config.quiet_hours_start; config.quiet_hours_start = hb_config.quiet_hours_start;
config.quiet_hours_end = hb_config.quiet_hours_end; config.quiet_hours_end = hb_config.quiet_hours_end;
config.multi_tenant = hb_config.multi_tenant;
config.timezone = hb_config config.timezone = hb_config
.timezone .timezone
.clone() .clone()
@@ -547,30 +597,52 @@ impl Agent {
.await; .await;
let notify_user = heartbeat_notify_user; let notify_user = heartbeat_notify_user;
let channels = self.channels.clone(); let channels = self.channels.clone();
let is_multi_tenant = hb_config.multi_tenant;
tokio::spawn(async move { tokio::spawn(async move {
while let Some(response) = notify_rx.recv().await { 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 // Try the configured channel first, fall back to
// broadcasting on all channels. // broadcasting on all channels.
let targeted_ok = if let Some(ref channel) = notify_channel let targeted_ok = if let Some(ref channel) = notify_channel {
&& let Some(ref user) = notify_target let target = effective_user.as_deref().or(notify_target.as_deref());
{ if let Some(user) = target {
channels channels
.broadcast(channel, user, response.clone()) .broadcast(channel, user, response.clone())
.await .await
.is_ok() .is_ok()
} else {
false
}
} else { } else {
false false
}; };
if !targeted_ok && let Some(ref user) = notify_user { if !targeted_ok {
let results = channels.broadcast_all(user, response).await; let fallback = effective_user.as_deref().or(notify_user.as_deref());
for (ch, result) in results { if let Some(user) = fallback {
if let Err(e) = result { let results = channels.broadcast_all(user, response).await;
tracing::warn!( for (ch, result) in results {
"Failed to broadcast heartbeat to {}: {}", if let Err(e) = result {
ch, tracing::warn!(
e "Failed to broadcast heartbeat to {}: {}",
); ch,
e
);
}
} }
} }
} }
@@ -583,14 +655,29 @@ impl Agent {
.map(|h| h.to_workspace_config()) .map(|h| h.to_workspace_config())
.unwrap_or_default(); .unwrap_or_default();
Some(spawn_heartbeat( if config.multi_tenant {
config, if let Some(admin) = self.admin_store() {
hygiene, Some(spawn_multi_user_heartbeat(
workspace.clone(), config,
self.cheap_llm().clone(), hygiene,
Some(notify_tx), self.cheap_llm().clone(),
self.store().map(Arc::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 { } else {
tracing::warn!("Heartbeat enabled but no workspace available"); tracing::warn!("Heartbeat enabled but no workspace available");
None None
@@ -612,7 +699,7 @@ impl Agent {
let engine = Arc::new(RoutineEngine::new( let engine = Arc::new(RoutineEngine::new(
rt_config.clone(), rt_config.clone(),
Arc::clone(store), crate::tenant::AdminScope::new(Arc::clone(store)),
self.llm().clone(), self.llm().clone(),
Arc::clone(workspace), Arc::clone(workspace),
notify_tx, 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); let session_for_empty_exit = Arc::clone(&session);
// Process based on submission type // Process based on submission type
let result = match submission { let result = match submission {
Submission::UserInput { content } => { Submission::UserInput { content } => {
let mut result = self 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; .await;
// Drain any messages queued during processing. // Drain any messages queued during processing.
@@ -1246,7 +1342,13 @@ impl Agent {
let mut queued_msg = message.clone(); let mut queued_msg = message.clone();
queued_msg.attachments.clear(); queued_msg.attachments.clear();
result = self 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; .await;
// If processing failed, re-queue the drained content so it // 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 // 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 .await
} }
Submission::Undo => self.process_undo(session, thread_id).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::Summarize => self.process_summarize(session, thread_id).await,
Submission::Suggest => self.process_suggest(session, thread_id).await, Submission::Suggest => self.process_suggest(session, thread_id).await,
Submission::JobStatus { job_id } => { Submission::JobStatus { job_id } => {
self.process_job_status(&message.user_id, job_id.as_deref()) self.process_job_status(&tenant, job_id.as_deref()).await
.await
}
Submission::JobCancel { job_id } => {
self.process_job_cancel(&message.user_id, &job_id).await
} }
Submission::JobCancel { job_id } => self.process_job_cancel(&tenant, &job_id).await,
Submission::Quit => return Ok(None), Submission::Quit => return Ok(None),
Submission::SwitchThread { thread_id: target } => { Submission::SwitchThread { thread_id: target } => {
self.process_switch_thread(message, target).await self.process_switch_thread(message, target).await
+88 -53
View File
@@ -33,6 +33,7 @@ impl Agent {
&self, &self,
intent: MessageIntent, intent: MessageIntent,
message: &IncomingMessage, message: &IncomingMessage,
tenant: &crate::tenant::TenantCtx,
) -> Result<SubmissionResult, Error> { ) -> Result<SubmissionResult, Error> {
// Send thinking status for non-trivial operations // Send thinking status for non-trivial operations
if let MessageIntent::CreateJob { .. } = &intent { if let MessageIntent::CreateJob { .. } = &intent {
@@ -52,24 +53,18 @@ impl Agent {
description, description,
category, category,
} => { } => {
self.handle_create_job(&message.user_id, title, description, category) self.handle_create_job(tenant, title, description, category)
.await? .await?
} }
MessageIntent::CheckJobStatus { job_id } => { MessageIntent::CheckJobStatus { job_id } => {
self.handle_check_status(&message.user_id, job_id).await? self.handle_check_status(tenant, 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?
} }
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 } => { MessageIntent::Command { command, args } => {
match self match self
.handle_command(&command, &args, &message.channel) .handle_command(&command, &args, &message.channel, tenant)
.await? .await?
{ {
Some(s) => s, Some(s) => s,
@@ -83,14 +78,14 @@ impl Agent {
async fn handle_create_job( async fn handle_create_job(
&self, &self,
user_id: &str, tenant: &crate::tenant::TenantCtx,
title: String, title: String,
description: String, description: String,
category: Option<String>, category: Option<String>,
) -> Result<String, Error> { ) -> Result<String, Error> {
let job_id = self let job_id = self
.scheduler .scheduler
.dispatch_job(user_id, &title, &description, None) .dispatch_job(tenant.user_id(), &title, &description, None)
.await?; .await?;
// Set the dedicated category field (not stored in metadata) // Set the dedicated category field (not stored in metadata)
@@ -113,7 +108,7 @@ impl Agent {
async fn handle_check_status( async fn handle_check_status(
&self, &self,
user_id: &str, tenant: &crate::tenant::TenantCtx,
job_id: Option<String>, job_id: Option<String>,
) -> Result<String, Error> { ) -> Result<String, Error> {
match job_id { match job_id {
@@ -122,7 +117,8 @@ impl Agent {
.map_err(|_| crate::error::JobError::NotFound { id: Uuid::nil() })?; .map_err(|_| crate::error::JobError::NotFound { id: Uuid::nil() })?;
// Try DB first for persistent state, fall back to ContextManager. // 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 && let Ok(Some(ctx)) = store.get_job(uuid).await
{ {
return Ok(format!( return Ok(format!(
@@ -138,7 +134,7 @@ impl Agent {
} }
let ctx = self.context_manager.get_context(uuid).await?; 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()); return Err(crate::error::JobError::NotFound { id: uuid }.into());
} }
@@ -155,7 +151,8 @@ impl Agent {
} }
None => { None => {
// Show summary from DB for consistency with Jobs tab. // 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 total = 0;
let mut in_progress = 0; let mut in_progress = 0;
let mut completed = 0; let mut completed = 0;
@@ -183,7 +180,7 @@ impl Agent {
} }
// Fallback to ContextManager if no DB. // 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!( Ok(format!(
"Jobs summary: Total: {} In Progress: {} Completed: {} Failed: {} Stuck: {}", "Jobs summary: Total: {} In Progress: {} Completed: {} Failed: {} Stuck: {}",
summary.total, summary.total,
@@ -196,19 +193,24 @@ impl Agent {
} }
} }
async fn handle_cancel_job(&self, user_id: &str, job_id: &str) -> Result<String, Error> { async fn handle_cancel_job(
&self,
tenant: &crate::tenant::TenantCtx,
job_id: &str,
) -> Result<String, Error> {
let uuid = Uuid::parse_str(job_id) let uuid = Uuid::parse_str(job_id)
.map_err(|_| crate::error::JobError::NotFound { id: Uuid::nil() })?; .map_err(|_| crate::error::JobError::NotFound { id: Uuid::nil() })?;
let ctx = self.context_manager.get_context(uuid).await?; 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()); return Err(crate::error::JobError::NotFound { id: uuid }.into());
} }
self.scheduler.stop(uuid).await?; self.scheduler.stop(uuid).await?;
// Also update DB so the Jobs tab reflects cancellation immediately. // 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 && let Err(e) = store
.update_job_status(uuid, JobState::Cancelled, Some("Cancelled by user")) .update_job_status(uuid, JobState::Cancelled, Some("Cancelled by user"))
.await .await
@@ -221,11 +223,12 @@ impl Agent {
async fn handle_list_jobs( async fn handle_list_jobs(
&self, &self,
user_id: &str, tenant: &crate::tenant::TenantCtx,
_filter: Option<String>, _filter: Option<String>,
) -> Result<String, Error> { ) -> Result<String, Error> {
// List from DB for consistency with Jobs tab. // 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 { let agent_jobs = match store.list_agent_jobs().await {
Ok(jobs) => jobs, Ok(jobs) => jobs,
Err(e) => { Err(e) => {
@@ -256,7 +259,7 @@ impl Agent {
} }
// Fallback to ContextManager if no DB. // 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() { if jobs.is_empty() {
return Ok("No jobs found.".to_string()); return Ok("No jobs found.".to_string());
} }
@@ -270,12 +273,16 @@ impl Agent {
Ok(output) Ok(output)
} }
async fn handle_help_job(&self, user_id: &str, job_id: &str) -> Result<String, Error> { async fn handle_help_job(
&self,
tenant: &crate::tenant::TenantCtx,
job_id: &str,
) -> Result<String, Error> {
let uuid = Uuid::parse_str(job_id) let uuid = Uuid::parse_str(job_id)
.map_err(|_| crate::error::JobError::NotFound { id: Uuid::nil() })?; .map_err(|_| crate::error::JobError::NotFound { id: Uuid::nil() })?;
let ctx = self.context_manager.get_context(uuid).await?; 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()); 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. /// Show job status inline — either all jobs (no id) or a specific job.
pub(super) async fn process_job_status( pub(super) async fn process_job_status(
&self, &self,
user_id: &str, tenant: &crate::tenant::TenantCtx,
job_id: Option<&str>, job_id: Option<&str>,
) -> Result<SubmissionResult, Error> { ) -> Result<SubmissionResult, Error> {
match self 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 .await
{ {
Ok(text) => Ok(SubmissionResult::response(text)), Ok(text) => Ok(SubmissionResult::response(text)),
@@ -323,10 +330,10 @@ impl Agent {
/// Cancel a job by ID. /// Cancel a job by ID.
pub(super) async fn process_job_cancel( pub(super) async fn process_job_cancel(
&self, &self,
user_id: &str, tenant: &crate::tenant::TenantCtx,
job_id: &str, job_id: &str,
) -> Result<SubmissionResult, Error> { ) -> Result<SubmissionResult, Error> {
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)), Ok(text) => Ok(SubmissionResult::response(text)),
Err(e) => Ok(SubmissionResult::error(format!("Cancel error: {}", e))), Err(e) => Ok(SubmissionResult::error(format!("Cancel error: {}", e))),
} }
@@ -559,6 +566,7 @@ impl Agent {
command: &str, command: &str,
args: &[String], args: &[String],
channel: &str, channel: &str,
tenant: &crate::tenant::TenantCtx,
) -> Result<SubmissionResult, Error> { ) -> Result<SubmissionResult, Error> {
match command { match command {
"help" => Ok(SubmissionResult::response(concat!( "help" => Ok(SubmissionResult::response(concat!(
@@ -752,19 +760,32 @@ impl Agent {
} }
} }
match self.llm().set_model(requested) { if self.config.multi_tenant {
Ok(()) => { // Multi-tenant: only persist to per-user DB settings.
// Persist the model choice so it survives restarts. // Do NOT call set_model() on the shared provider — that
self.persist_selected_model(requested).await; // would change the default for all users. The per-request
Ok(SubmissionResult::response(format!( // model_override in the dispatcher reads from the same
"Switched model to: {}", // "selected_model" setting and applies it per-user.
requested 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, command: &str,
args: &[String], args: &[String],
channel: &str, channel: &str,
tenant: &crate::tenant::TenantCtx,
) -> Result<Option<String>, Error> { ) -> Result<Option<String>, Error> {
// System commands are now handled directly via Submission::SystemCommand, // System commands are now handled directly via Submission::SystemCommand,
// but the router may still send us unknown /commands. // 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::Response { content } => Ok(Some(content)),
SubmissionResult::Ok { message } => Ok(message), SubmissionResult::Ok { message } => Ok(message),
SubmissionResult::Error { message } => Ok(Some(format!("Error: {}", 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, /// Best-effort: logs warnings on failure but does not propagate errors,
/// since the in-memory model switch already succeeded. /// since the in-memory model switch already succeeded.
async fn persist_selected_model(&self, model: &str) { ///
// 1. Persist to DB if available. /// In multi-tenant mode, only the per-user DB setting is written — global
if let Some(store) = self.store() { /// .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()); let value = serde_json::Value::String(model.to_string());
if let Err(e) = store if let Err(e) = store.set_setting("selected_model", &value).await {
.set_setting(self.owner_id(), "selected_model", &value)
.await
{
tracing::warn!("Failed to persist model to DB: {}", e); tracing::warn!("Failed to persist model to DB: {}", e);
} else { } else {
tracing::debug!("Persisted selected_model to DB: {}", model); tracing::debug!(
user_id = tenant.user_id(),
"Persisted selected_model to DB: {}",
model
);
} }
} else { } else {
tracing::warn!("No database store available — model choice will not persist to DB"); 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 model_owned = model.to_string();
let backend = self.deps.llm_backend.clone(); let backend = self.deps.llm_backend.clone();
if let Err(e) = tokio::task::spawn_blocking(move || { if let Err(e) = tokio::task::spawn_blocking(move || {
+236 -3
View File
@@ -21,6 +21,9 @@ pub struct CostGuardConfig {
pub max_cost_per_day_cents: Option<u64>, pub max_cost_per_day_cents: Option<u64>,
/// Maximum LLM calls per hour. None = unlimited. /// Maximum LLM calls per hour. None = unlimited.
pub max_actions_per_hour: Option<u64>, pub max_actions_per_hour: Option<u64>,
/// 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<u64>,
} }
/// Error returned when a cost limit is exceeded. /// Error returned when a cost limit is exceeded.
@@ -30,6 +33,12 @@ pub enum CostLimitExceeded {
DailyBudget { spent_cents: u64, limit_cents: u64 }, DailyBudget { spent_cents: u64, limit_cents: u64 },
/// Hourly action rate limit reached. /// Hourly action rate limit reached.
HourlyRate { actions: u64, limit: u64 }, 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 { 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", "Hourly action limit exceeded: {} actions of {} allowed per hour",
actions, limit 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. /// Per-model token usage since startup.
model_tokens: Mutex<HashMap<String, ModelTokens>>, model_tokens: Mutex<HashMap<String, ModelTokens>>,
/// Per-user daily cost tracking. Each entry resets independently at midnight UTC.
per_user_daily_cost: Mutex<HashMap<String, DailyCost>>,
} }
struct DailyCost { struct DailyCost {
@@ -97,6 +120,7 @@ impl CostGuard {
action_window: Mutex::new(VecDeque::new()), action_window: Mutex::new(VecDeque::new()),
budget_exceeded: AtomicBool::new(false), budget_exceeded: AtomicBool::new(false),
model_tokens: Mutex::new(HashMap::new()), model_tokens: Mutex::new(HashMap::new()),
per_user_daily_cost: Mutex::new(HashMap::new()),
} }
} }
@@ -203,6 +227,11 @@ impl CostGuard {
daily.reset_date = today; daily.reset_date = today;
self.budget_exceeded.store(false, Ordering::Relaxed); self.budget_exceeded.store(false, Ordering::Relaxed);
tracing::info!("Cost guard: daily counter reset for {}", today); 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; daily.total += cost;
@@ -248,6 +277,85 @@ impl CostGuard {
cost 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). /// Current daily spend in USD (as Decimal).
pub async fn daily_spend(&self) -> Decimal { pub async fn daily_spend(&self) -> Decimal {
let daily = self.daily_cost.lock().await; 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. /// Number of actions in the current hourly window.
pub async fn actions_this_hour(&self) -> u64 { pub async fn actions_this_hour(&self) -> u64 {
let mut window = self.action_window.lock().await; let mut window = self.action_window.lock().await;
@@ -314,7 +432,7 @@ mod tests {
async fn test_daily_budget_enforcement() { async fn test_daily_budget_enforcement() {
let guard = CostGuard::new(CostGuardConfig { let guard = CostGuard::new(CostGuardConfig {
max_cost_per_day_cents: Some(1), // $0.01 limit max_cost_per_day_cents: Some(1), // $0.01 limit
max_actions_per_hour: None, ..CostGuardConfig::default()
}); });
// First call allowed // First call allowed
@@ -350,8 +468,8 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn test_hourly_rate_enforcement() { async fn test_hourly_rate_enforcement() {
let guard = CostGuard::new(CostGuardConfig { let guard = CostGuard::new(CostGuardConfig {
max_cost_per_day_cents: None,
max_actions_per_hour: Some(3), max_actions_per_hour: Some(3),
..CostGuardConfig::default()
}); });
// First 3 actions allowed // First 3 actions allowed
@@ -633,8 +751,8 @@ mod tests {
// A fresh CostGuard with rate limits should not panic even if // A fresh CostGuard with rate limits should not panic even if
// checked_sub returns None (simulating short uptime). // checked_sub returns None (simulating short uptime).
let guard = CostGuard::new(CostGuardConfig { let guard = CostGuard::new(CostGuardConfig {
max_cost_per_day_cents: None,
max_actions_per_hour: Some(100), max_actions_per_hour: Some(100),
..CostGuardConfig::default()
}); });
// These must not panic regardless of system uptime // These must not panic regardless of system uptime
@@ -656,4 +774,119 @@ mod tests {
let result = Instant::now().checked_sub(std::time::Duration::MAX); let result = Instant::now().checked_sub(std::time::Duration::MAX);
assert!(result.is_none()); 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"));
}
} }
+53 -9
View File
@@ -42,6 +42,7 @@ impl Agent {
pub(super) async fn run_agentic_loop( pub(super) async fn run_agentic_loop(
&self, &self,
message: &IncomingMessage, message: &IncomingMessage,
tenant: crate::tenant::TenantCtx,
session: Arc<Mutex<Session>>, session: Arc<Mutex<Session>>,
thread_id: Uuid, thread_id: Uuid,
initial_messages: Vec<ChatMessage>, initial_messages: Vec<ChatMessage>,
@@ -168,6 +169,7 @@ impl Agent {
let delegate = ChatDelegate { let delegate = ChatDelegate {
agent: self, agent: self,
tenant,
session: session.clone(), session: session.clone(),
thread_id, thread_id,
message, message,
@@ -240,6 +242,7 @@ impl Agent {
/// auth intercept, and cost tracking. /// auth intercept, and cost tracking.
struct ChatDelegate<'a> { struct ChatDelegate<'a> {
agent: &'a Agent, agent: &'a Agent,
tenant: crate::tenant::TenantCtx,
session: Arc<Mutex<Session>>, session: Arc<Mutex<Session>>,
thread_id: Uuid, thread_id: Uuid,
message: &'a IncomingMessage, message: &'a IncomingMessage,
@@ -336,8 +339,8 @@ impl<'a> LoopDelegate for ChatDelegate<'a> {
reason_ctx: &mut ReasoningContext, reason_ctx: &mut ReasoningContext,
iteration: usize, iteration: usize,
) -> Result<crate::llm::RespondOutput, Error> { ) -> Result<crate::llm::RespondOutput, Error> {
// Enforce cost guardrails before the LLM call // Enforce cost guardrails before the LLM call (global + per-user)
if let Err(limit) = self.agent.cost_guard().check_allowed().await { if let Err(limit) = self.tenant.check_cost_allowed().await {
return Err(crate::error::LlmError::InvalidResponse { return Err(crate::error::LlmError::InvalidResponse {
provider: "agent".to_string(), provider: "agent".to_string(),
reason: limit.to_string(), reason: limit.to_string(),
@@ -345,6 +348,21 @@ impl<'a> LoopDelegate for ChatDelegate<'a> {
.into()); .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 { let output = match reasoning.respond_with_tools(reason_ctx).await {
Ok(output) => output, Ok(output) => output,
Err(crate::error::LlmError::ContextLengthExceeded { used, limit }) => { Err(crate::error::LlmError::ContextLengthExceeded { used, limit }) => {
@@ -379,13 +397,22 @@ impl<'a> LoopDelegate for ChatDelegate<'a> {
Err(e) => return Err(e.into()), Err(e) => return Err(e.into()),
}; };
// Record cost and track token usage // Record cost and track token usage (global + per-user).
let model_name = self.agent.llm().active_model_name(); // 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 read_discount = self.agent.llm().cache_read_discount();
let write_multiplier = self.agent.llm().cache_write_multiplier(); let write_multiplier = self.agent.llm().cache_write_multiplier();
let call_cost = self let call_cost = self
.agent .tenant
.cost_guard()
.record_llm_call( .record_llm_call(
&model_name, &model_name,
output.usage.input_tokens, output.usage.input_tokens,
@@ -394,7 +421,7 @@ impl<'a> LoopDelegate for ChatDelegate<'a> {
output.usage.cache_creation_input_tokens, output.usage.cache_creation_input_tokens,
read_discount, read_discount,
write_multiplier, write_multiplier,
Some(self.agent.llm().cost_per_token()), cost_per_token,
) )
.await; .await;
tracing::debug!( tracing::debug!(
@@ -1305,6 +1332,7 @@ mod tests {
sandbox_readiness: crate::agent::routine_engine::SandboxReadiness::DisabledByConfig, sandbox_readiness: crate::agent::routine_engine::SandboxReadiness::DisabledByConfig,
builder: None, builder: None,
llm_backend: "nearai".to_string(), llm_backend: "nearai".to_string(),
tenant_rates: Arc::new(crate::tenant::TenantRateRegistry::new(4, 3)),
}; };
Agent::new( Agent::new(
@@ -1320,10 +1348,14 @@ mod tests {
allow_local_tools: false, allow_local_tools: false,
max_cost_per_day_cents: None, max_cost_per_day_cents: None,
max_actions_per_hour: None, max_actions_per_hour: None,
max_cost_per_user_per_day_cents: None,
max_tool_iterations: 50, max_tool_iterations: 50,
auto_approve_tools: false, auto_approve_tools: false,
default_timezone: "UTC".to_string(), default_timezone: "UTC".to_string(),
max_tokens_per_job: 0, max_tokens_per_job: 0,
multi_tenant: false,
max_llm_concurrent_per_user: None,
max_jobs_concurrent_per_user: None,
}, },
deps, deps,
Arc::new(ChannelManager::new()), Arc::new(ChannelManager::new()),
@@ -2181,6 +2213,7 @@ mod tests {
sandbox_readiness: crate::agent::routine_engine::SandboxReadiness::DisabledByConfig, sandbox_readiness: crate::agent::routine_engine::SandboxReadiness::DisabledByConfig,
builder: None, builder: None,
llm_backend: "nearai".to_string(), llm_backend: "nearai".to_string(),
tenant_rates: Arc::new(crate::tenant::TenantRateRegistry::new(4, 3)),
}; };
Agent::new( Agent::new(
@@ -2196,10 +2229,14 @@ mod tests {
allow_local_tools: false, allow_local_tools: false,
max_cost_per_day_cents: None, max_cost_per_day_cents: None,
max_actions_per_hour: None, max_actions_per_hour: None,
max_cost_per_user_per_day_cents: None,
max_tool_iterations, max_tool_iterations,
auto_approve_tools: true, auto_approve_tools: true,
default_timezone: "UTC".to_string(), default_timezone: "UTC".to_string(),
max_tokens_per_job: 0, max_tokens_per_job: 0,
multi_tenant: false,
max_llm_concurrent_per_user: None,
max_jobs_concurrent_per_user: None,
}, },
deps, deps,
Arc::new(ChannelManager::new()), Arc::new(ChannelManager::new()),
@@ -2234,13 +2271,14 @@ mod tests {
let message = IncomingMessage::new("test", "test-user", "do something"); let message = IncomingMessage::new("test", "test-user", "do something");
let initial_messages = vec![ChatMessage::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 // The dispatcher must terminate within 5 seconds. If there is an
// infinite loop bug (e.g., index not advancing on tool failure), the // infinite loop bug (e.g., index not advancing on tool failure), the
// timeout will fire and the test will fail. // timeout will fire and the test will fail.
let result = tokio::time::timeout( let result = tokio::time::timeout(
Duration::from_secs(5), 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; .await;
@@ -2302,6 +2340,7 @@ mod tests {
sandbox_readiness: crate::agent::routine_engine::SandboxReadiness::DisabledByConfig, sandbox_readiness: crate::agent::routine_engine::SandboxReadiness::DisabledByConfig,
builder: None, builder: None,
llm_backend: "nearai".to_string(), llm_backend: "nearai".to_string(),
tenant_rates: Arc::new(crate::tenant::TenantRateRegistry::new(4, 3)),
}; };
Agent::new( Agent::new(
@@ -2317,10 +2356,14 @@ mod tests {
allow_local_tools: false, allow_local_tools: false,
max_cost_per_day_cents: None, max_cost_per_day_cents: None,
max_actions_per_hour: None, max_actions_per_hour: None,
max_cost_per_user_per_day_cents: None,
max_tool_iterations: max_iter, max_tool_iterations: max_iter,
auto_approve_tools: true, auto_approve_tools: true,
default_timezone: "UTC".to_string(), default_timezone: "UTC".to_string(),
max_tokens_per_job: 0, max_tokens_per_job: 0,
multi_tenant: false,
max_llm_concurrent_per_user: None,
max_jobs_concurrent_per_user: None,
}, },
deps, deps,
Arc::new(ChannelManager::new()), Arc::new(ChannelManager::new()),
@@ -2340,13 +2383,14 @@ mod tests {
let message = IncomingMessage::new("test", "test-user", "keep calling tools"); let message = IncomingMessage::new("test", "test-user", "keep calling tools");
let initial_messages = vec![ChatMessage::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 // Even with an LLM that always wants to call tools, the dispatcher
// must terminate within the timeout thanks to force_text at // must terminate within the timeout thanks to force_text at
// max_tool_iterations. // max_tool_iterations.
let result = tokio::time::timeout( let result = tokio::time::timeout(
Duration::from_secs(5), 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; .await;
+184 -7
View File
@@ -31,8 +31,8 @@ use chrono_tz::Tz;
use tokio::sync::mpsc; use tokio::sync::mpsc;
use crate::channels::OutgoingResponse; use crate::channels::OutgoingResponse;
use crate::db::Database;
use crate::llm::{ChatMessage, CompletionRequest, LlmProvider, Reasoning}; use crate::llm::{ChatMessage, CompletionRequest, LlmProvider, Reasoning};
use crate::tenant::AdminScope;
use crate::workspace::Workspace; use crate::workspace::Workspace;
use crate::workspace::hygiene::HygieneConfig; use crate::workspace::hygiene::HygieneConfig;
@@ -57,6 +57,9 @@ pub struct HeartbeatConfig {
pub quiet_hours_end: Option<u32>, pub quiet_hours_end: Option<u32>,
/// Timezone for fire_at and quiet hours evaluation (IANA name). /// Timezone for fire_at and quiet hours evaluation (IANA name).
pub timezone: Option<String>, pub timezone: Option<String>,
/// 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 { impl Default for HeartbeatConfig {
@@ -71,6 +74,7 @@ impl Default for HeartbeatConfig {
quiet_hours_start: None, quiet_hours_start: None,
quiet_hours_end: None, quiet_hours_end: None,
timezone: None, timezone: None,
multi_tenant: false,
} }
} }
} }
@@ -178,7 +182,7 @@ pub struct HeartbeatRunner {
workspace: Arc<Workspace>, workspace: Arc<Workspace>,
llm: Arc<dyn LlmProvider>, llm: Arc<dyn LlmProvider>,
response_tx: Option<mpsc::Sender<OutgoingResponse>>, response_tx: Option<mpsc::Sender<OutgoingResponse>>,
store: Option<Arc<dyn Database>>, store: Option<AdminScope>,
consecutive_failures: u32, consecutive_failures: u32,
} }
@@ -207,8 +211,8 @@ impl HeartbeatRunner {
self self
} }
/// Set the database store for persistent heartbeat conversations. /// Set the admin-scoped database store for persistent heartbeat conversations.
pub fn with_store(mut self, store: Arc<dyn Database>) -> Self { pub fn with_store(mut self, store: AdminScope) -> Self {
self.store = Some(store); self.store = Some(store);
self self
} }
@@ -396,7 +400,7 @@ impl HeartbeatRunner {
} }
/// Send a notification about heartbeat findings. /// 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 { let Some(ref tx) = self.response_tx else {
tracing::debug!("No response channel configured for heartbeat notifications"); tracing::debug!("No response channel configured for heartbeat notifications");
return; return;
@@ -493,7 +497,7 @@ pub fn spawn_heartbeat(
workspace: Arc<Workspace>, workspace: Arc<Workspace>,
llm: Arc<dyn LlmProvider>, llm: Arc<dyn LlmProvider>,
response_tx: Option<mpsc::Sender<OutgoingResponse>>, response_tx: Option<mpsc::Sender<OutgoingResponse>>,
store: Option<Arc<dyn Database>>, store: Option<AdminScope>,
) -> tokio::task::JoinHandle<()> { ) -> tokio::task::JoinHandle<()> {
let mut runner = HeartbeatRunner::new(config, hygiene_config, workspace, llm); let mut runner = HeartbeatRunner::new(config, hygiene_config, workspace, llm);
if let Some(tx) = response_tx { 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<dyn LlmProvider>,
response_tx: Option<mpsc::Sender<OutgoingResponse>>,
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<String, u32> =
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<String> = routines
.iter()
.map(|r| r.user_id.clone())
.collect::<std::collections::HashSet<_>>()
.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<String, u32>,
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)] #[cfg(test)]
mod tests { mod tests {
use super::*; use super::*;
@@ -726,7 +903,7 @@ mod tests {
Arc<crate::workspace::Workspace>, Arc<crate::workspace::Workspace>,
Arc<dyn crate::llm::LlmProvider>, Arc<dyn crate::llm::LlmProvider>,
Option<tokio::sync::mpsc::Sender<crate::channels::OutgoingResponse>>, Option<tokio::sync::mpsc::Sender<crate::channels::OutgoingResponse>>,
Option<Arc<dyn crate::db::Database>>, Option<AdminScope>,
) -> tokio::task::JoinHandle<()> = spawn_heartbeat; ) -> tokio::task::JoinHandle<()> = spawn_heartbeat;
let _ = _fn_ptr; let _ = _fn_ptr;
} }
+3 -1
View File
@@ -36,7 +36,9 @@ pub(crate) use agent_loop::truncate_for_preview;
pub use agent_loop::{Agent, AgentDeps}; pub use agent_loop::{Agent, AgentDeps};
pub use compaction::{CompactionResult, ContextCompactor}; pub use compaction::{CompactionResult, ContextCompactor};
pub use context_monitor::{CompactionStrategy, ContextBreakdown, ContextMonitor}; 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 router::{MessageIntent, Router};
pub use routine::{Routine, RoutineAction, RoutineRun, Trigger}; pub use routine::{Routine, RoutineAction, RoutineRun, Trigger};
pub use routine_engine::{RoutineEngine, SandboxReadiness}; pub use routine_engine::{RoutineEngine, SandboxReadiness};
+30 -8
View File
@@ -28,12 +28,12 @@ use crate::agent::routine::{
use crate::channels::{IncomingMessage, OutgoingResponse}; use crate::channels::{IncomingMessage, OutgoingResponse};
use crate::config::RoutineConfig; use crate::config::RoutineConfig;
use crate::context::{JobContext, JobState}; use crate::context::{JobContext, JobState};
use crate::db::Database;
use crate::error::RoutineError; use crate::error::RoutineError;
use crate::extensions::ExtensionManager; use crate::extensions::ExtensionManager;
use crate::llm::{ use crate::llm::{
ChatMessage, CompletionRequest, FinishReason, LlmProvider, ToolCall, ToolCompletionRequest, ChatMessage, CompletionRequest, FinishReason, LlmProvider, ToolCall, ToolCompletionRequest,
}; };
use crate::tenant::AdminScope;
use crate::tools::{ use crate::tools::{
ToolError, ToolRegistry, autonomous_allowed_tool_names, autonomous_unavailable_message, ToolError, ToolRegistry, autonomous_allowed_tool_names, autonomous_unavailable_message,
prepare_tool_params, prepare_tool_params,
@@ -99,7 +99,7 @@ pub(crate) fn routine_matches_message(routine: &Routine, message: &IncomingMessa
/// The routine execution engine. /// The routine execution engine.
pub struct RoutineEngine { pub struct RoutineEngine {
config: RoutineConfig, config: RoutineConfig,
store: Arc<dyn Database>, store: AdminScope,
llm: Arc<dyn LlmProvider>, llm: Arc<dyn LlmProvider>,
workspace: Arc<Workspace>, workspace: Arc<Workspace>,
/// Sender for notifications (routed to channel manager). /// Sender for notifications (routed to channel manager).
@@ -128,7 +128,7 @@ impl RoutineEngine {
#[allow(clippy::too_many_arguments)] #[allow(clippy::too_many_arguments)]
pub fn new( pub fn new(
config: RoutineConfig, config: RoutineConfig,
store: Arc<dyn Database>, store: AdminScope,
llm: Arc<dyn LlmProvider>, llm: Arc<dyn LlmProvider>,
workspace: Arc<Workspace>, workspace: Arc<Workspace>,
notify_tx: mpsc::Sender<OutgoingResponse>, notify_tx: mpsc::Sender<OutgoingResponse>,
@@ -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) // Execute inline for manual triggers (caller wants to wait)
let engine = EngineContext { let engine = EngineContext {
config: self.config.clone(), config: self.config.clone(),
store: self.store.clone(), store: self.store.clone(),
llm: self.llm.clone(), llm: self.llm.clone(),
workspace: self.workspace.clone(), workspace: routine_workspace,
notify_tx: self.notify_tx.clone(), notify_tx: self.notify_tx.clone(),
running_count: self.running_count.clone(), running_count: self.running_count.clone(),
scheduler: self.scheduler.clone(), scheduler: self.scheduler.clone(),
@@ -910,11 +920,23 @@ impl RoutineEngine {
created_at: Utc::now(), 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 { let engine = EngineContext {
config: self.config.clone(), config: self.config.clone(),
store: self.store.clone(), store: self.store.clone(),
llm: self.llm.clone(), llm: self.llm.clone(),
workspace: self.workspace.clone(), workspace: routine_workspace,
notify_tx: self.notify_tx.clone(), notify_tx: self.notify_tx.clone(),
running_count: self.running_count.clone(), running_count: self.running_count.clone(),
scheduler: self.scheduler.clone(), scheduler: self.scheduler.clone(),
@@ -967,7 +989,7 @@ impl RoutineEngine {
/// an active state (Pending/InProgress/Stuck). Maps the final `JobState` to /// an active state (Pending/InProgress/Stuck). Maps the final `JobState` to
/// a `RunStatus` for the routine run. /// a `RunStatus` for the routine run.
struct FullJobWatcher { struct FullJobWatcher {
store: Arc<dyn Database>, store: AdminScope,
job_id: Uuid, job_id: Uuid,
routine_name: String, routine_name: String,
} }
@@ -978,7 +1000,7 @@ impl FullJobWatcher {
/// Safety ceiling: 24 hours, derived from POLL_INTERVAL. /// Safety ceiling: 24 hours, derived from POLL_INTERVAL.
const MAX_POLLS: u32 = (24 * 60 * 60) / Self::POLL_INTERVAL.as_secs() as u32; const MAX_POLLS: u32 = (24 * 60 * 60) / Self::POLL_INTERVAL.as_secs() as u32;
fn new(store: Arc<dyn Database>, job_id: Uuid, routine_name: String) -> Self { fn new(store: AdminScope, job_id: Uuid, routine_name: String) -> Self {
Self { Self {
store, store,
job_id, job_id,
@@ -1050,7 +1072,7 @@ impl FullJobWatcher {
/// Shared context passed to the execution function. /// Shared context passed to the execution function.
struct EngineContext { struct EngineContext {
config: RoutineConfig, config: RoutineConfig,
store: Arc<dyn Database>, store: AdminScope,
llm: Arc<dyn LlmProvider>, llm: Arc<dyn LlmProvider>,
workspace: Arc<Workspace>, workspace: Arc<Workspace>,
notify_tx: mpsc::Sender<OutgoingResponse>, notify_tx: mpsc::Sender<OutgoingResponse>,
+7 -3
View File
@@ -11,12 +11,12 @@ use uuid::Uuid;
use crate::agent::task::{Task, TaskContext, TaskOutput}; use crate::agent::task::{Task, TaskContext, TaskOutput};
use crate::config::AgentConfig; use crate::config::AgentConfig;
use crate::context::{ContextManager, JobContext, JobState}; use crate::context::{ContextManager, JobContext, JobState};
use crate::db::Database;
use crate::error::{Error, JobError}; use crate::error::{Error, JobError};
use crate::extensions::ExtensionManager; use crate::extensions::ExtensionManager;
use crate::hooks::HookRegistry; use crate::hooks::HookRegistry;
use crate::llm::LlmProvider; use crate::llm::LlmProvider;
use crate::safety::SafetyLayer; use crate::safety::SafetyLayer;
use crate::tenant::AdminScope;
use crate::tools::{ use crate::tools::{
ApprovalContext, ToolRegistry, autonomous_allowed_tool_names, autonomous_unavailable_error, ApprovalContext, ToolRegistry, autonomous_allowed_tool_names, autonomous_unavailable_error,
prepare_tool_params, prepare_tool_params,
@@ -52,7 +52,7 @@ struct ScheduledSubtask {
pub struct SchedulerDeps { pub struct SchedulerDeps {
pub tools: Arc<ToolRegistry>, pub tools: Arc<ToolRegistry>,
pub extension_manager: Option<Arc<ExtensionManager>>, pub extension_manager: Option<Arc<ExtensionManager>>,
pub store: Option<Arc<dyn Database>>, pub store: Option<AdminScope>,
pub hooks: Arc<HookRegistry>, pub hooks: Arc<HookRegistry>,
} }
@@ -64,7 +64,7 @@ pub struct Scheduler {
safety: Arc<SafetyLayer>, safety: Arc<SafetyLayer>,
tools: Arc<ToolRegistry>, tools: Arc<ToolRegistry>,
extension_manager: Option<Arc<ExtensionManager>>, extension_manager: Option<Arc<ExtensionManager>>,
store: Option<Arc<dyn Database>>, store: Option<AdminScope>,
hooks: Arc<HookRegistry>, hooks: Arc<HookRegistry>,
/// SSE manager for live job event streaming. /// SSE manager for live job event streaming.
sse_tx: Option<Arc<crate::channels::web::sse::SseManager>>, sse_tx: Option<Arc<crate::channels::web::sse::SseManager>>,
@@ -780,10 +780,14 @@ mod tests {
allow_local_tools: true, allow_local_tools: true,
max_cost_per_day_cents: None, max_cost_per_day_cents: None,
max_actions_per_hour: None, max_actions_per_hour: None,
max_cost_per_user_per_day_cents: None,
max_tool_iterations: 10, max_tool_iterations: 10,
auto_approve_tools: true, auto_approve_tools: true,
default_timezone: "UTC".to_string(), default_timezone: "UTC".to_string(),
max_tokens_per_job, 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 cm = Arc::new(ContextManager::new(5));
let llm: Arc<dyn LlmProvider> = Arc::new(StubLlm); let llm: Arc<dyn LlmProvider> = Arc::new(StubLlm);
+5 -5
View File
@@ -8,8 +8,8 @@ use chrono::{DateTime, Utc};
use uuid::Uuid; use uuid::Uuid;
use crate::context::{ContextManager, JobState}; use crate::context::{ContextManager, JobState};
use crate::db::Database;
use crate::error::RepairError; use crate::error::RepairError;
use crate::tenant::AdminScope;
use crate::tools::{BuildRequirement, Language, SoftwareBuilder, SoftwareType, ToolRegistry}; use crate::tools::{BuildRequirement, Language, SoftwareBuilder, SoftwareType, ToolRegistry};
/// A job that has been detected as stuck. /// 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. /// Jobs in `InProgress` longer than this are treated as stuck.
stuck_threshold: Duration, stuck_threshold: Duration,
max_repair_attempts: u32, max_repair_attempts: u32,
store: Option<Arc<dyn Database>>, store: Option<AdminScope>,
builder: Option<Arc<dyn SoftwareBuilder>>, builder: Option<Arc<dyn SoftwareBuilder>>,
tools: Option<Arc<ToolRegistry>>, tools: Option<Arc<ToolRegistry>>,
} }
@@ -91,8 +91,8 @@ impl DefaultSelfRepair {
} }
} }
/// Add a Store for tool failure tracking. /// Add an admin-scoped store for tool failure tracking.
pub fn with_store(mut self, store: Arc<dyn Database>) -> Self { pub fn with_store(mut self, store: AdminScope) -> Self {
self.store = Some(store); self.store = Some(store);
self self
} }
@@ -806,7 +806,7 @@ mod tests {
// Create self-repair with zero threshold (detect immediately), // Create self-repair with zero threshold (detect immediately),
// wired with store, builder, and tools. // wired with store, builder, and tools.
let repair = DefaultSelfRepair::new(Arc::clone(&cm), Duration::from_secs(0), 3) 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( .with_builder(
Arc::clone(&builder) as Arc<dyn crate::tools::SoftwareBuilder>, Arc::clone(&builder) as Arc<dyn crate::tools::SoftwareBuilder>,
tools, tools,
+10 -3
View File
@@ -175,6 +175,7 @@ impl Agent {
pub(super) async fn process_user_input( pub(super) async fn process_user_input(
&self, &self,
message: &IncomingMessage, message: &IncomingMessage,
tenant: crate::tenant::TenantCtx,
session: Arc<Mutex<Session>>, session: Arc<Mutex<Session>>,
thread_id: Uuid, thread_id: Uuid,
content: &str, content: &str,
@@ -351,7 +352,7 @@ impl Agent {
if let Some(intent) = self.router.route_command(&temp_message) { if let Some(intent) = self.router.route_command(&temp_message) {
// Explicit command like /status, /job, /list - handle directly // 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 // Natural language goes through the agentic loop
@@ -462,7 +463,7 @@ impl Agent {
// Run the agentic tool execution loop // Run the agentic tool execution loop
let result = self 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; .await;
// Re-acquire lock and check if interrupted // Re-acquire lock and check if interrupted
@@ -1473,7 +1474,13 @@ impl Agent {
// Continue the agentic loop (a tool was already executed this turn) // Continue the agentic loop (a tool was already executed this turn)
let result = self 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; .await;
// Handle the result // Handle the result
+1
View File
@@ -880,6 +880,7 @@ impl AppBuilder {
crate::agent::cost_guard::CostGuardConfig { crate::agent::cost_guard::CostGuardConfig {
max_cost_per_day_cents: self.config.agent.max_cost_per_day_cents, 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_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,
}, },
)); ));
+30 -3
View File
@@ -54,10 +54,37 @@ fn validate_webhook_secret(
/// ///
/// This endpoint is **public** (no gateway auth token required) but protected /// This endpoint is **public** (no gateway auth token required) but protected
/// by the per-routine webhook secret sent via the `X-Webhook-Secret` header. /// 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( pub async fn webhook_trigger_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
Path(path): Path<String>, Path(path): Path<String>,
headers: HeaderMap, headers: HeaderMap,
) -> Result<Json<serde_json::Value>, (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<Arc<GatewayState>>,
Path((user_id, path)): Path<(String, String)>,
headers: HeaderMap,
) -> Result<Json<serde_json::Value>, (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<GatewayState>,
path: &str,
user_id: Option<&str>,
headers: &HeaderMap,
) -> Result<Json<serde_json::Value>, (StatusCode, String)> { ) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
// Rate limit check // Rate limit check
if !state.webhook_rate_limiter.check() { if !state.webhook_rate_limiter.check() {
@@ -72,9 +99,9 @@ pub async fn webhook_trigger_handler(
"Database not available".to_string(), "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 let routine = store
.get_webhook_routine_by_path(&path) .get_webhook_routine_by_path(path, user_id)
.await .await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))? .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?
.ok_or(( .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 { let status = match &e {
crate::error::RoutineError::NotFound { .. } => StatusCode::NOT_FOUND, crate::error::RoutineError::NotFound { .. } => StatusCode::NOT_FOUND,
crate::error::RoutineError::Disabled { .. } crate::error::RoutineError::Disabled { .. }
+5
View File
@@ -414,6 +414,11 @@ pub async fn start_server(
.route( .route(
"/api/webhooks/{path}", "/api/webhooks/{path}",
post(crate::channels::web::handlers::webhooks::webhook_trigger_handler), 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) // Protected routes (require auth)
+20 -1
View File
@@ -1,6 +1,6 @@
use std::time::Duration; 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::error::ConfigError;
use crate::settings::Settings; use crate::settings::Settings;
@@ -23,6 +23,8 @@ pub struct AgentConfig {
pub max_cost_per_day_cents: Option<u64>, pub max_cost_per_day_cents: Option<u64>,
/// Maximum LLM/tool actions per hour. None = unlimited. /// Maximum LLM/tool actions per hour. None = unlimited.
pub max_actions_per_hour: Option<u64>, pub max_actions_per_hour: Option<u64>,
/// Maximum daily LLM spend per user in cents. None = unlimited.
pub max_cost_per_user_per_day_cents: Option<u64>,
/// Maximum tool-call iterations per agentic loop invocation. Default 50. /// Maximum tool-call iterations per agentic loop invocation. Default 50.
pub max_tool_iterations: usize, pub max_tool_iterations: usize,
/// When true, skip tool approval checks entirely. For benchmarks/CI. /// When true, skip tool approval checks entirely. For benchmarks/CI.
@@ -31,6 +33,13 @@ pub struct AgentConfig {
pub default_timezone: String, pub default_timezone: String,
/// Maximum tokens per job (0 = unlimited). /// Maximum tokens per job (0 = unlimited).
pub max_tokens_per_job: u64, 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<usize>,
/// Maximum concurrent jobs per user. None = use default (3).
pub max_jobs_concurrent_per_user: Option<usize>,
} }
impl AgentConfig { impl AgentConfig {
@@ -49,10 +58,14 @@ impl AgentConfig {
allow_local_tools: true, allow_local_tools: true,
max_cost_per_day_cents: None, max_cost_per_day_cents: None,
max_actions_per_hour: None, max_actions_per_hour: None,
max_cost_per_user_per_day_cents: None,
max_tool_iterations: 10, max_tool_iterations: 10,
auto_approve_tools: true, auto_approve_tools: true,
default_timezone: "UTC".to_string(), default_timezone: "UTC".to_string(),
max_tokens_per_job: 0, 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)?, allow_local_tools: parse_bool_env("ALLOW_LOCAL_TOOLS", false)?,
max_cost_per_day_cents: parse_option_env("MAX_COST_PER_DAY_CENTS")?, 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_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( max_tool_iterations: parse_optional_env(
"AGENT_MAX_TOOL_ITERATIONS", "AGENT_MAX_TOOL_ITERATIONS",
settings.agent.max_tool_iterations, settings.agent.max_tool_iterations,
@@ -112,6 +126,11 @@ impl AgentConfig {
"AGENT_MAX_TOKENS_PER_JOB", "AGENT_MAX_TOKENS_PER_JOB",
settings.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")?,
}) })
} }
} }
+10
View File
@@ -21,6 +21,9 @@ pub struct HeartbeatConfig {
pub quiet_hours_end: Option<u32>, pub quiet_hours_end: Option<u32>,
/// Timezone for fire_at and quiet hours evaluation (IANA name). /// Timezone for fire_at and quiet hours evaluation (IANA name).
pub timezone: Option<String>, pub timezone: Option<String>,
/// 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 { impl Default for HeartbeatConfig {
@@ -34,6 +37,7 @@ impl Default for HeartbeatConfig {
quiet_hours_start: None, quiet_hours_start: None,
quiet_hours_end: None, quiet_hours_end: None,
timezone: None, timezone: None,
multi_tenant: false,
} }
} }
} }
@@ -101,6 +105,12 @@ impl HeartbeatConfig {
} }
tz 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(),
)?,
}) })
} }
} }
+18 -3
View File
@@ -530,10 +530,24 @@ impl RoutineStore for LibSqlBackend {
async fn get_webhook_routine_by_path( async fn get_webhook_routine_by_path(
&self, &self,
path: &str, path: &str,
user_id: Option<&str>,
) -> Result<Option<Routine>, DatabaseError> { ) -> Result<Option<Routine>, DatabaseError> {
let conn = self.connect().await?; let conn = self.connect().await?;
let mut rows = conn let mut rows = if let Some(uid) = user_id {
.query( 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!( &format!(
"SELECT {} FROM routines WHERE enabled = 1 AND trigger_type = 'webhook' \ "SELECT {} FROM routines WHERE enabled = 1 AND trigger_type = 'webhook' \
AND (json_extract(trigger_config, '$.path') = ?1 \ AND (json_extract(trigger_config, '$.path') = ?1 \
@@ -543,7 +557,8 @@ impl RoutineStore for LibSqlBackend {
params![path], params![path],
) )
.await .await
.map_err(|e| DatabaseError::Query(e.to_string()))?; .map_err(|e| DatabaseError::Query(e.to_string()))?
};
match rows match rows
.next() .next()
+1
View File
@@ -545,6 +545,7 @@ pub trait RoutineStore: Send + Sync {
async fn get_webhook_routine_by_path( async fn get_webhook_routine_by_path(
&self, &self,
path: &str, path: &str,
user_id: Option<&str>,
) -> Result<Option<Routine>, DatabaseError>; ) -> Result<Option<Routine>, DatabaseError>;
/// List routine runs that were dispatched as full_job but have not yet /// List routine runs that were dispatched as full_job but have not yet
+2 -1
View File
@@ -529,8 +529,9 @@ impl RoutineStore for PgBackend {
async fn get_webhook_routine_by_path( async fn get_webhook_routine_by_path(
&self, &self,
path: &str, path: &str,
user_id: Option<&str>,
) -> Result<Option<Routine>, DatabaseError> { ) -> Result<Option<Routine>, 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<Vec<RoutineRun>, DatabaseError> { async fn list_dispatched_routine_runs(&self) -> Result<Vec<RoutineRun>, DatabaseError> {
+13 -3
View File
@@ -1162,15 +1162,25 @@ impl Store {
pub async fn get_webhook_routine_by_path( pub async fn get_webhook_routine_by_path(
&self, &self,
path: &str, path: &str,
user_id: Option<&str>,
) -> Result<Option<Routine>, DatabaseError> { ) -> Result<Option<Routine>, DatabaseError> {
let conn = self.conn().await?; let conn = self.conn().await?;
let row = conn let row = if let Some(uid) = user_id {
.query_opt( 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' \ "SELECT * FROM routines WHERE enabled AND trigger_type = 'webhook' \
AND (trigger_config->>'path' = $1 OR (trigger_config->>'path' IS NULL AND id::text = $1))", AND (trigger_config->>'path' = $1 OR (trigger_config->>'path' IS NULL AND id::text = $1))",
&[&path], &[&path],
) )
.await?; .await?
};
row.as_ref().map(row_to_routine).transpose() row.as_ref().map(row_to_routine).transpose()
} }
+1
View File
@@ -69,6 +69,7 @@ pub mod service;
pub mod settings; pub mod settings;
pub mod setup; pub mod setup;
pub mod skills; pub mod skills;
pub mod tenant;
pub mod timezone; pub mod timezone;
pub mod tools; pub mod tools;
pub mod tracing_fmt; pub mod tracing_fmt;
+11
View File
@@ -199,6 +199,10 @@ pub struct ReasoningContext {
/// instead of calling `build_system_prompt_with_tools`. Allows callers to build /// instead of calling `build_system_prompt_with_tools`. Allows callers to build
/// the prompt once and reuse it across iterations. /// the prompt once and reuse it across iterations.
pub system_prompt: Option<String>, pub system_prompt: Option<String>,
/// 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<String>,
} }
impl ReasoningContext { impl ReasoningContext {
@@ -212,6 +216,7 @@ impl ReasoningContext {
metadata: std::collections::HashMap::new(), metadata: std::collections::HashMap::new(),
force_text: false, force_text: false,
system_prompt: None, system_prompt: None,
model_override: None,
} }
} }
@@ -671,6 +676,9 @@ Respond in JSON format:
.with_temperature(0.7) .with_temperature(0.7)
.with_tool_choice("auto"); .with_tool_choice("auto");
request.metadata = context.metadata.clone(); 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 response = self.llm.complete_with_tools(request).await?;
let usage = TokenUsage { let usage = TokenUsage {
@@ -773,6 +781,9 @@ Respond in JSON format:
.with_max_tokens(4096) .with_max_tokens(4096)
.with_temperature(0.7); .with_temperature(0.7);
request.metadata = context.metadata.clone(); 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 response = self.llm.complete(request).await?;
let pre_truncated = truncate_at_tool_tags(&response.content); let pre_truncated = truncate_at_tool_tags(&response.content);
+32 -20
View File
@@ -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] #[async_trait]
impl<M> LlmProvider for RigAdapter<M> impl<M> LlmProvider for RigAdapter<M>
where where
@@ -632,15 +656,7 @@ where
&self, &self,
mut request: CompletionRequest, mut request: CompletionRequest,
) -> Result<CompletionResponse, LlmError> { ) -> Result<CompletionResponse, LlmError> {
if let Some(requested_model) = request.model.as_deref() let model_override = request.model.take();
&& 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"
);
}
self.strip_unsupported_completion_params(&mut request); self.strip_unsupported_completion_params(&mut request);
@@ -648,7 +664,7 @@ where
crate::llm::provider::sanitize_tool_messages(&mut messages); crate::llm::provider::sanitize_tool_messages(&mut messages);
let (preamble, history) = convert_messages(&messages); let (preamble, history) = convert_messages(&messages);
let rig_req = build_rig_request( let mut rig_req = build_rig_request(
preamble, preamble,
history, history,
Vec::new(), Vec::new(),
@@ -658,6 +674,8 @@ where
self.cache_retention, self.cache_retention,
)?; )?;
inject_model_override(&mut rig_req, model_override.as_deref());
let response = let response =
self.model self.model
.completion(rig_req) .completion(rig_req)
@@ -695,15 +713,7 @@ where
&self, &self,
mut request: ToolCompletionRequest, mut request: ToolCompletionRequest,
) -> Result<ToolCompletionResponse, LlmError> { ) -> Result<ToolCompletionResponse, LlmError> {
if let Some(requested_model) = request.model.as_deref() let model_override = request.model.take();
&& 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"
);
}
self.strip_unsupported_tool_params(&mut request); self.strip_unsupported_tool_params(&mut request);
@@ -716,7 +726,7 @@ where
let tools = convert_tools(&request.tools); let tools = convert_tools(&request.tools);
let tool_choice = convert_tool_choice(request.tool_choice.as_deref()); 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, preamble,
history, history,
tools, tools,
@@ -726,6 +736,8 @@ where
self.cache_retention, self.cache_retention,
)?; )?;
inject_model_override(&mut rig_req, model_override.as_deref());
let response = let response =
self.model self.model
.completion(rig_req) .completion(rig_req)
+4
View File
@@ -914,6 +914,10 @@ async fn async_main() -> anyhow::Result<()> {
}, },
builder: components.builder, builder: components.builder,
llm_backend: config.llm.backend.clone(), 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); let channels_for_warnings = Arc::clone(&channels);
+906
View File
@@ -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<dyn Database>,
}
impl TenantScope {
pub fn new(user_id: impl Into<String>, db: Arc<dyn Database>) -> 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<Vec<AgentJobRecord>, DatabaseError> {
self.inner.list_agent_jobs_for_user(&self.user_id).await
}
pub async fn agent_job_summary(&self) -> Result<AgentJobSummary, DatabaseError> {
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<Option<JobContext>, 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<Option<String>, 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<Vec<SandboxJobRecord>, DatabaseError> {
self.inner.list_sandbox_jobs_for_user(&self.user_id).await
}
pub async fn sandbox_job_summary(&self) -> Result<SandboxJobSummary, DatabaseError> {
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<Option<SandboxJobRecord>, 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<bool, DatabaseError> {
self.inner
.sandbox_job_belongs_to_user(job_id, &self.user_id)
.await
}
// === Routines ===
pub async fn list_routines(&self) -> Result<Vec<Routine>, DatabaseError> {
self.inner.list_routines(&self.user_id).await
}
pub async fn get_routine_by_name(&self, name: &str) -> Result<Option<Routine>, 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<Option<Routine>, 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<bool, DatabaseError> {
// 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<Vec<RoutineRun>, 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<Option<Routine>, DatabaseError> {
self.inner
.get_webhook_routine_by_path(path, Some(&self.user_id))
.await
}
// === Settings ===
pub async fn get_setting(&self, key: &str) -> Result<Option<serde_json::Value>, DatabaseError> {
self.inner.get_setting(&self.user_id, key).await
}
pub async fn get_setting_full(&self, key: &str) -> Result<Option<SettingRow>, 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<bool, DatabaseError> {
self.inner.delete_setting(&self.user_id, key).await
}
pub async fn list_settings(&self) -> Result<Vec<SettingRow>, DatabaseError> {
self.inner.list_settings(&self.user_id).await
}
pub async fn get_all_settings(
&self,
) -> Result<HashMap<String, serde_json::Value>, DatabaseError> {
self.inner.get_all_settings(&self.user_id).await
}
pub async fn set_all_settings(
&self,
settings: &HashMap<String, serde_json::Value>,
) -> Result<(), DatabaseError> {
self.inner.set_all_settings(&self.user_id, settings).await
}
pub async fn has_settings(&self) -> Result<bool, DatabaseError> {
self.inner.has_settings(&self.user_id).await
}
// === Conversations ===
pub async fn create_conversation(
&self,
channel: &str,
thread_id: Option<&str>,
) -> Result<Uuid, DatabaseError> {
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<bool, DatabaseError> {
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<Vec<ConversationSummary>, DatabaseError> {
self.inner
.list_conversations_with_preview(&self.user_id, channel, limit)
.await
}
pub async fn list_conversations_all_channels(
&self,
limit: i64,
) -> Result<Vec<ConversationSummary>, 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<Uuid, DatabaseError> {
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<Uuid, DatabaseError> {
self.inner
.get_or_create_heartbeat_conversation(&self.user_id)
.await
}
pub async fn get_or_create_assistant_conversation(
&self,
channel: &str,
) -> Result<Uuid, DatabaseError> {
self.inner
.get_or_create_assistant_conversation(&self.user_id, channel)
.await
}
pub async fn conversation_belongs_to_user(
&self,
conversation_id: Uuid,
) -> Result<bool, DatabaseError> {
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<Uuid, DatabaseError> {
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<Vec<ConversationMessage>, DatabaseError> {
self.inner.list_conversation_messages(conversation_id).await
}
pub async fn list_conversation_messages_paginated(
&self,
conversation_id: Uuid,
before: Option<DateTime<Utc>>,
limit: i64,
) -> Result<(Vec<ConversationMessage>, 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<Uuid, DatabaseError> {
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<Option<serde_json::Value>, 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<dyn Database>,
}
impl AdminScope {
pub fn new(db: Arc<dyn Database>) -> 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<dyn Database> {
&self.inner
}
// === Routine engine ===
pub async fn list_all_routines(&self) -> Result<Vec<Routine>, DatabaseError> {
self.inner.list_all_routines().await
}
pub async fn list_event_routines(&self) -> Result<Vec<Routine>, DatabaseError> {
self.inner.list_event_routines().await
}
pub async fn list_due_cron_routines(&self) -> Result<Vec<Routine>, DatabaseError> {
self.inner.list_due_cron_routines().await
}
pub async fn list_dispatched_routine_runs(&self) -> Result<Vec<RoutineRun>, DatabaseError> {
self.inner.list_dispatched_routine_runs().await
}
pub async fn count_running_routine_runs_batch(
&self,
routine_ids: &[Uuid],
) -> Result<HashMap<Uuid, i64>, DatabaseError> {
self.inner
.count_running_routine_runs_batch(routine_ids)
.await
}
pub async fn batch_get_last_run_status(
&self,
routine_ids: &[Uuid],
) -> Result<HashMap<Uuid, RunStatus>, DatabaseError> {
self.inner.batch_get_last_run_status(routine_ids).await
}
pub async fn count_running_routine_runs(&self, routine_id: Uuid) -> Result<i64, DatabaseError> {
self.inner.count_running_routine_runs(routine_id).await
}
pub async fn update_routine_runtime(
&self,
id: Uuid,
last_run_at: DateTime<Utc>,
next_fire_at: Option<DateTime<Utc>>,
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<i32>,
) -> 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<Option<Routine>, 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<Vec<Uuid>, DatabaseError> {
self.inner.get_stuck_jobs().await
}
pub async fn get_broken_tools(&self, threshold: i32) -> Result<Vec<BrokenTool>, 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<u64, DatabaseError> {
self.inner.cleanup_stale_sandbox_jobs().await
}
pub async fn get_sandbox_job(
&self,
id: Uuid,
) -> Result<Option<SandboxJobRecord>, 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<bool>,
message: Option<&str>,
started_at: Option<DateTime<Utc>>,
completed_at: Option<DateTime<Utc>>,
) -> 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<Option<String>, 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<i64>,
) -> Result<Vec<crate::history::JobEventRecord>, DatabaseError> {
self.inner.list_job_events(job_id, limit).await
}
// === Job persistence (scheduler, worker) ===
pub async fn get_job(&self, id: Uuid) -> Result<Option<JobContext>, 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<Vec<AgentJobRecord>, DatabaseError> {
self.inner.list_agent_jobs().await
}
pub async fn get_agent_job_failure_reason(
&self,
id: Uuid,
) -> Result<Option<String>, DatabaseError> {
self.inner.get_agent_job_failure_reason(id).await
}
// === LLM call recording ===
pub async fn record_llm_call(&self, record: &LlmCallRecord<'_>) -> Result<Uuid, DatabaseError> {
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<Vec<ActionRecord>, 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<Uuid, DatabaseError> {
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<Decimal>,
) -> 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<Uuid, DatabaseError> {
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<Uuid, DatabaseError> {
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<Uuid, DatabaseError> {
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<Semaphore>,
/// Limits concurrent jobs for this user.
pub job_semaphore: Arc<Semaphore>,
}
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<HashMap>` (consistent with the rest of the
/// codebase — no DashMap dependency).
pub struct TenantRateRegistry {
state: tokio::sync::RwLock<HashMap<String, Arc<TenantRateState>>>,
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<TenantRateState> {
// 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<TenantScope>,
workspace: Option<Arc<Workspace>>,
cost_guard: Arc<CostGuard>,
rate: Arc<TenantRateState>,
}
impl TenantCtx {
pub fn new(
user_id: impl Into<String>,
store: Option<TenantScope>,
workspace: Option<Arc<Workspace>>,
cost_guard: Arc<CostGuard>,
rate: Arc<TenantRateState>,
) -> 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<Workspace>> {
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<SemaphorePermit<'_>, 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));
}
}
+2
View File
@@ -532,6 +532,7 @@ impl TestHarnessBuilder {
let cost_guard = Arc::new(CostGuard::new(CostGuardConfig { let cost_guard = Arc::new(CostGuard::new(CostGuardConfig {
max_cost_per_day_cents: None, max_cost_per_day_cents: None,
max_actions_per_hour: None, max_actions_per_hour: None,
max_cost_per_user_per_day_cents: None,
})); }));
let channel = if self.stub_channel { let channel = if self.stub_channel {
@@ -564,6 +565,7 @@ impl TestHarnessBuilder {
sandbox_readiness: crate::agent::routine_engine::SandboxReadiness::DisabledByConfig, sandbox_readiness: crate::agent::routine_engine::SandboxReadiness::DisabledByConfig,
builder: None, builder: None,
llm_backend: "nearai".to_string(), llm_backend: "nearai".to_string(),
tenant_rates: std::sync::Arc::new(crate::tenant::TenantRateRegistry::new(4, 3)),
}; };
TestHarness { TestHarness {
+3 -3
View File
@@ -20,7 +20,6 @@ use crate::agent::scheduler::WorkerMessage;
use crate::agent::task::TaskOutput; use crate::agent::task::TaskOutput;
use crate::channels::web::types::ToolDecisionDto; use crate::channels::web::types::ToolDecisionDto;
use crate::context::{ContextManager, JobState}; use crate::context::{ContextManager, JobState};
use crate::db::Database;
use crate::error::Error; use crate::error::Error;
use crate::hooks::HookRegistry; use crate::hooks::HookRegistry;
use crate::llm::{ use crate::llm::{
@@ -28,6 +27,7 @@ use crate::llm::{
ToolSelection, ToolSelection,
}; };
use crate::safety::SafetyLayer; use crate::safety::SafetyLayer;
use crate::tenant::AdminScope;
use crate::tools::execute::process_tool_result; use crate::tools::execute::process_tool_result;
use crate::tools::rate_limiter::RateLimitResult; use crate::tools::rate_limiter::RateLimitResult;
use crate::tools::{ use crate::tools::{
@@ -45,7 +45,7 @@ pub struct WorkerDeps {
pub llm: Arc<dyn LlmProvider>, pub llm: Arc<dyn LlmProvider>,
pub safety: Arc<SafetyLayer>, pub safety: Arc<SafetyLayer>,
pub tools: Arc<ToolRegistry>, pub tools: Arc<ToolRegistry>,
pub store: Option<Arc<dyn Database>>, pub store: Option<AdminScope>,
pub hooks: Arc<HookRegistry>, pub hooks: Arc<HookRegistry>,
pub timeout: Duration, pub timeout: Duration,
pub use_planning: bool, pub use_planning: bool,
@@ -94,7 +94,7 @@ impl Worker {
&self.deps.tools &self.deps.tools
} }
fn store(&self) -> Option<&Arc<dyn Database>> { fn store(&self) -> Option<&AdminScope> {
self.deps.store.as_ref() self.deps.store.as_ref()
} }
+10 -10
View File
@@ -337,14 +337,14 @@ mod tests {
SchedulerDeps { SchedulerDeps {
tools: registry.clone(), tools: registry.clone(),
extension_manager: extension_manager.clone(), extension_manager: extension_manager.clone(),
store: Some(db.clone()), store: Some(ironclaw::tenant::AdminScope::new(db.clone())),
hooks: Arc::new(HookRegistry::new()), hooks: Arc::new(HookRegistry::new()),
}, },
)); ));
Arc::new(RoutineEngine::new( Arc::new(RoutineEngine::new(
RoutineConfig::default(), RoutineConfig::default(),
db, ironclaw::tenant::AdminScope::new(db),
llm, llm,
ws, ws,
notify_tx, notify_tx,
@@ -448,7 +448,7 @@ mod tests {
let engine = Arc::new(RoutineEngine::new( let engine = Arc::new(RoutineEngine::new(
RoutineConfig::default(), RoutineConfig::default(),
db.clone(), ironclaw::tenant::AdminScope::new(db.clone()),
llm, llm,
ws, ws,
notify_tx, notify_tx,
@@ -527,7 +527,7 @@ mod tests {
let engine = Arc::new(RoutineEngine::new( let engine = Arc::new(RoutineEngine::new(
RoutineConfig::default(), RoutineConfig::default(),
db.clone(), ironclaw::tenant::AdminScope::new(db.clone()),
llm, llm,
ws, ws,
notify_tx, notify_tx,
@@ -614,7 +614,7 @@ mod tests {
let engine = Arc::new(RoutineEngine::new( let engine = Arc::new(RoutineEngine::new(
RoutineConfig::default(), RoutineConfig::default(),
db.clone(), ironclaw::tenant::AdminScope::new(db.clone()),
llm, llm,
ws, ws,
notify_tx, notify_tx,
@@ -723,7 +723,7 @@ mod tests {
let engine = Arc::new(RoutineEngine::new( let engine = Arc::new(RoutineEngine::new(
RoutineConfig::default(), RoutineConfig::default(),
db.clone(), ironclaw::tenant::AdminScope::new(db.clone()),
llm, llm,
ws, ws,
notify_tx, notify_tx,
@@ -866,7 +866,7 @@ mod tests {
let engine = Arc::new(RoutineEngine::new( let engine = Arc::new(RoutineEngine::new(
RoutineConfig::default(), RoutineConfig::default(),
db.clone(), ironclaw::tenant::AdminScope::new(db.clone()),
llm, llm,
ws, ws,
notify_tx, notify_tx,
@@ -1049,7 +1049,7 @@ mod tests {
let engine = Arc::new(RoutineEngine::new( let engine = Arc::new(RoutineEngine::new(
RoutineConfig::default(), RoutineConfig::default(),
Arc::clone(&db), ironclaw::tenant::AdminScope::new(Arc::clone(&db)),
llm, llm,
ws, ws,
notify_tx, notify_tx,
@@ -1171,7 +1171,7 @@ mod tests {
let engine = Arc::new(RoutineEngine::new( let engine = Arc::new(RoutineEngine::new(
RoutineConfig::default(), RoutineConfig::default(),
db.clone(), ironclaw::tenant::AdminScope::new(db.clone()),
llm, llm,
ws, ws,
notify_tx, notify_tx,
@@ -1279,7 +1279,7 @@ mod tests {
let engine = Arc::new(RoutineEngine::new( let engine = Arc::new(RoutineEngine::new(
config, config,
db.clone(), ironclaw::tenant::AdminScope::new(db.clone()),
llm, llm,
ws, ws,
notify_tx, notify_tx,
+1
View File
@@ -201,6 +201,7 @@ mod tests {
sandbox_readiness: ironclaw::agent::SandboxReadiness::DisabledByConfig, sandbox_readiness: ironclaw::agent::SandboxReadiness::DisabledByConfig,
builder: None, builder: None,
llm_backend: "nearai".to_string(), llm_backend: "nearai".to_string(),
tenant_rates: std::sync::Arc::new(ironclaw::tenant::TenantRateRegistry::new(4, 3)),
}; };
let gateway = Arc::new(TestChannel::new()); let gateway = Arc::new(TestChannel::new());
@@ -266,6 +266,7 @@ impl GatewayWorkflowHarness {
sandbox_readiness: ironclaw::agent::SandboxReadiness::DisabledByConfig, sandbox_readiness: ironclaw::agent::SandboxReadiness::DisabledByConfig,
builder: None, builder: None,
llm_backend: "nearai".to_string(), llm_backend: "nearai".to_string(),
tenant_rates: std::sync::Arc::new(ironclaw::tenant::TenantRateRegistry::new(4, 3)),
}, },
channels, channels,
None, None,
+2 -1
View File
@@ -642,7 +642,7 @@ impl TestRigBuilder {
let (notify_tx, _notify_rx) = tokio::sync::mpsc::channel(16); let (notify_tx, _notify_rx) = tokio::sync::mpsc::channel(16);
let engine = Arc::new(RoutineEngine::new( let engine = Arc::new(RoutineEngine::new(
routine_config, routine_config,
Arc::clone(db_arc), ironclaw::tenant::AdminScope::new(Arc::clone(db_arc)),
components.llm.clone(), components.llm.clone(),
Arc::clone(ws), Arc::clone(ws),
notify_tx, notify_tx,
@@ -762,6 +762,7 @@ impl TestRigBuilder {
sandbox_readiness: ironclaw::agent::SandboxReadiness::Available, // tests don't use real Docker sandbox_readiness: ironclaw::agent::SandboxReadiness::Available, // tests don't use real Docker
builder: None, builder: None,
llm_backend: "nearai".to_string(), llm_backend: "nearai".to_string(),
tenant_rates: std::sync::Arc::new(ironclaw::tenant::TenantRateRegistry::new(4, 3)),
}; };
// 7. Create TestChannel and ChannelManager. // 7. Create TestChannel and ChannelManager.