mirror of
https://github.com/outbackdingo/optimclaw.git
synced 2026-08-27 08:00:17 +00:00
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]>
This commit is contained in:
+76
-15
@@ -172,6 +172,8 @@ pub struct AgentDeps {
|
||||
/// Resolved LLM backend identifier (e.g., "nearai", "openai", "groq").
|
||||
/// Used by `/model` persistence to determine which env var to update.
|
||||
pub llm_backend: String,
|
||||
/// Per-tenant rate limiting registry (lazily creates rate state per user).
|
||||
pub tenant_rates: Arc<crate::tenant::TenantRateRegistry>,
|
||||
}
|
||||
|
||||
/// The main agent that coordinates all components.
|
||||
@@ -234,7 +236,10 @@ impl Agent {
|
||||
SchedulerDeps {
|
||||
tools: deps.tools.clone(),
|
||||
extension_manager: deps.extension_manager.clone(),
|
||||
store: deps.store.clone(),
|
||||
store: deps
|
||||
.store
|
||||
.as_ref()
|
||||
.map(|db| crate::tenant::AdminScope::new(Arc::clone(db))),
|
||||
hooks: deps.hooks.clone(),
|
||||
},
|
||||
);
|
||||
@@ -315,6 +320,50 @@ impl Agent {
|
||||
&self.deps.cost_guard
|
||||
}
|
||||
|
||||
/// Build a tenant-scoped execution context for the given user.
|
||||
///
|
||||
/// This is the standard entry point for per-user operations. The returned
|
||||
/// [`TenantCtx`] provides a [`TenantScope`] that auto-binds `user_id` on
|
||||
/// every database operation and a per-user rate limiter.
|
||||
pub(super) async fn tenant_ctx(&self, user_id: &str) -> crate::tenant::TenantCtx {
|
||||
let rate = self.deps.tenant_rates.get_or_create(user_id).await;
|
||||
|
||||
let store = self
|
||||
.deps
|
||||
.store
|
||||
.as_ref()
|
||||
.map(|db| crate::tenant::TenantScope::new(user_id, Arc::clone(db)));
|
||||
|
||||
// Reuse the owner workspace if user matches, otherwise create per-user.
|
||||
let workspace = match &self.deps.workspace {
|
||||
Some(ws) if ws.user_id() == user_id => Some(Arc::clone(ws)),
|
||||
_ => self
|
||||
.deps
|
||||
.store
|
||||
.as_ref()
|
||||
.map(|db| Arc::new(Workspace::new_with_db(user_id, Arc::clone(db)))),
|
||||
};
|
||||
|
||||
crate::tenant::TenantCtx::new(
|
||||
user_id,
|
||||
store,
|
||||
workspace,
|
||||
Arc::clone(&self.deps.cost_guard),
|
||||
rate,
|
||||
)
|
||||
}
|
||||
|
||||
/// Get an admin-scoped database accessor for cross-tenant operations.
|
||||
///
|
||||
/// Only for system-level components (heartbeat, routine engine, self-repair,
|
||||
/// scheduler). Handler code should use [`tenant_ctx()`](Self::tenant_ctx) instead.
|
||||
pub(super) fn admin_store(&self) -> Option<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>>> {
|
||||
self.deps.skill_registry.as_ref()
|
||||
}
|
||||
@@ -400,8 +449,8 @@ impl Agent {
|
||||
self.config.stuck_threshold,
|
||||
self.config.max_repair_attempts,
|
||||
);
|
||||
if let Some(ref store) = self.deps.store {
|
||||
self_repair = self_repair.with_store(Arc::clone(store));
|
||||
if let Some(admin) = self.admin_store() {
|
||||
self_repair = self_repair.with_store(admin);
|
||||
}
|
||||
if let Some(ref builder) = self.deps.builder {
|
||||
self_repair = self_repair.with_builder(Arc::clone(builder), Arc::clone(self.tools()));
|
||||
@@ -597,13 +646,13 @@ impl Agent {
|
||||
.unwrap_or_default();
|
||||
|
||||
if config.multi_tenant {
|
||||
if let Some(store) = self.store() {
|
||||
if let Some(admin) = self.admin_store() {
|
||||
Some(spawn_multi_user_heartbeat(
|
||||
config,
|
||||
hygiene,
|
||||
self.cheap_llm().clone(),
|
||||
Some(notify_tx),
|
||||
Arc::clone(store),
|
||||
admin,
|
||||
))
|
||||
} else {
|
||||
tracing::warn!("Multi-tenant heartbeat requires a database store");
|
||||
@@ -616,7 +665,7 @@ impl Agent {
|
||||
workspace.clone(),
|
||||
self.cheap_llm().clone(),
|
||||
Some(notify_tx),
|
||||
self.store().map(Arc::clone),
|
||||
self.admin_store(),
|
||||
))
|
||||
}
|
||||
} else {
|
||||
@@ -640,7 +689,7 @@ impl Agent {
|
||||
|
||||
let engine = Arc::new(RoutineEngine::new(
|
||||
rt_config.clone(),
|
||||
Arc::clone(store),
|
||||
crate::tenant::AdminScope::new(Arc::clone(store)),
|
||||
self.llm().clone(),
|
||||
Arc::clone(workspace),
|
||||
notify_tx,
|
||||
@@ -1192,11 +1241,20 @@ impl Agent {
|
||||
}
|
||||
}
|
||||
|
||||
// Build per-tenant execution context once; threaded through all handlers.
|
||||
let tenant = self.tenant_ctx(&message.user_id).await;
|
||||
|
||||
// Process based on submission type
|
||||
let result = match submission {
|
||||
Submission::UserInput { content } => {
|
||||
let mut result = self
|
||||
.process_user_input(message, session.clone(), thread_id, &content)
|
||||
.process_user_input(
|
||||
message,
|
||||
tenant.clone(),
|
||||
session.clone(),
|
||||
thread_id,
|
||||
&content,
|
||||
)
|
||||
.await;
|
||||
|
||||
// Drain any messages queued during processing.
|
||||
@@ -1263,7 +1321,13 @@ impl Agent {
|
||||
let mut queued_msg = message.clone();
|
||||
queued_msg.attachments.clear();
|
||||
result = self
|
||||
.process_user_input(&queued_msg, session.clone(), thread_id, &next_content)
|
||||
.process_user_input(
|
||||
&queued_msg,
|
||||
tenant.clone(),
|
||||
session.clone(),
|
||||
thread_id,
|
||||
&next_content,
|
||||
)
|
||||
.await;
|
||||
|
||||
// If processing failed, re-queue the drained content so it
|
||||
@@ -1289,7 +1353,7 @@ impl Agent {
|
||||
message.channel
|
||||
);
|
||||
// Authorization checks (including restart channel check) are enforced in handle_system_command
|
||||
self.handle_system_command(&command, &args, &message.channel, &message.user_id)
|
||||
self.handle_system_command(&command, &args, &message.channel, &tenant)
|
||||
.await
|
||||
}
|
||||
Submission::Undo => self.process_undo(session, thread_id).await,
|
||||
@@ -1302,12 +1366,9 @@ impl Agent {
|
||||
Submission::Summarize => self.process_summarize(session, thread_id).await,
|
||||
Submission::Suggest => self.process_suggest(session, thread_id).await,
|
||||
Submission::JobStatus { job_id } => {
|
||||
self.process_job_status(&message.user_id, job_id.as_deref())
|
||||
.await
|
||||
}
|
||||
Submission::JobCancel { job_id } => {
|
||||
self.process_job_cancel(&message.user_id, &job_id).await
|
||||
self.process_job_status(&tenant, job_id.as_deref()).await
|
||||
}
|
||||
Submission::JobCancel { job_id } => self.process_job_cancel(&tenant, &job_id).await,
|
||||
Submission::Quit => return Ok(None),
|
||||
Submission::SwitchThread { thread_id: target } => {
|
||||
self.process_switch_thread(message, target).await
|
||||
|
||||
+56
-49
@@ -33,6 +33,7 @@ impl Agent {
|
||||
&self,
|
||||
intent: MessageIntent,
|
||||
message: &IncomingMessage,
|
||||
tenant: &crate::tenant::TenantCtx,
|
||||
) -> Result<SubmissionResult, Error> {
|
||||
// Send thinking status for non-trivial operations
|
||||
if let MessageIntent::CreateJob { .. } = &intent {
|
||||
@@ -52,24 +53,18 @@ impl Agent {
|
||||
description,
|
||||
category,
|
||||
} => {
|
||||
self.handle_create_job(&message.user_id, title, description, category)
|
||||
self.handle_create_job(tenant, title, description, category)
|
||||
.await?
|
||||
}
|
||||
MessageIntent::CheckJobStatus { job_id } => {
|
||||
self.handle_check_status(&message.user_id, job_id).await?
|
||||
}
|
||||
MessageIntent::CancelJob { job_id } => {
|
||||
self.handle_cancel_job(&message.user_id, &job_id).await?
|
||||
}
|
||||
MessageIntent::ListJobs { filter } => {
|
||||
self.handle_list_jobs(&message.user_id, filter).await?
|
||||
}
|
||||
MessageIntent::HelpJob { job_id } => {
|
||||
self.handle_help_job(&message.user_id, &job_id).await?
|
||||
self.handle_check_status(tenant, job_id).await?
|
||||
}
|
||||
MessageIntent::CancelJob { job_id } => self.handle_cancel_job(tenant, &job_id).await?,
|
||||
MessageIntent::ListJobs { filter } => self.handle_list_jobs(tenant, filter).await?,
|
||||
MessageIntent::HelpJob { job_id } => self.handle_help_job(tenant, &job_id).await?,
|
||||
MessageIntent::Command { command, args } => {
|
||||
match self
|
||||
.handle_command(&command, &args, &message.channel, &message.user_id)
|
||||
.handle_command(&command, &args, &message.channel, tenant)
|
||||
.await?
|
||||
{
|
||||
Some(s) => s,
|
||||
@@ -83,14 +78,14 @@ impl Agent {
|
||||
|
||||
async fn handle_create_job(
|
||||
&self,
|
||||
user_id: &str,
|
||||
tenant: &crate::tenant::TenantCtx,
|
||||
title: String,
|
||||
description: String,
|
||||
category: Option<String>,
|
||||
) -> Result<String, Error> {
|
||||
let job_id = self
|
||||
.scheduler
|
||||
.dispatch_job(user_id, &title, &description, None)
|
||||
.dispatch_job(tenant.user_id(), &title, &description, None)
|
||||
.await?;
|
||||
|
||||
// Set the dedicated category field (not stored in metadata)
|
||||
@@ -113,7 +108,7 @@ impl Agent {
|
||||
|
||||
async fn handle_check_status(
|
||||
&self,
|
||||
user_id: &str,
|
||||
tenant: &crate::tenant::TenantCtx,
|
||||
job_id: Option<String>,
|
||||
) -> Result<String, Error> {
|
||||
match job_id {
|
||||
@@ -122,13 +117,10 @@ impl Agent {
|
||||
.map_err(|_| crate::error::JobError::NotFound { id: Uuid::nil() })?;
|
||||
|
||||
// Try DB first for persistent state, fall back to ContextManager.
|
||||
if let Some(store) = self.store()
|
||||
// TenantScope.get_job() auto-filters by ownership — no manual check needed.
|
||||
if let Some(store) = tenant.store()
|
||||
&& let Ok(Some(ctx)) = store.get_job(uuid).await
|
||||
{
|
||||
// Ownership check: ensure the job belongs to the requesting user.
|
||||
if ctx.user_id != user_id {
|
||||
return Err(crate::error::JobError::NotFound { id: uuid }.into());
|
||||
}
|
||||
return Ok(format!(
|
||||
"Job: {}\nStatus: {:?}\nCreated: {}\nStarted: {}\nActual cost: {}",
|
||||
ctx.title,
|
||||
@@ -142,7 +134,7 @@ impl Agent {
|
||||
}
|
||||
|
||||
let ctx = self.context_manager.get_context(uuid).await?;
|
||||
if ctx.user_id != user_id {
|
||||
if ctx.user_id != tenant.user_id() {
|
||||
return Err(crate::error::JobError::NotFound { id: uuid }.into());
|
||||
}
|
||||
|
||||
@@ -159,21 +151,22 @@ impl Agent {
|
||||
}
|
||||
None => {
|
||||
// Show summary from DB for consistency with Jobs tab.
|
||||
if let Some(store) = self.store() {
|
||||
// TenantScope methods auto-scope to user — no user_id parameter needed.
|
||||
if let Some(store) = tenant.store() {
|
||||
let mut total = 0;
|
||||
let mut in_progress = 0;
|
||||
let mut completed = 0;
|
||||
let mut failed = 0;
|
||||
let mut stuck = 0;
|
||||
|
||||
if let Ok(s) = store.agent_job_summary_for_user(user_id).await {
|
||||
if let Ok(s) = store.agent_job_summary().await {
|
||||
total += s.total;
|
||||
in_progress += s.in_progress;
|
||||
completed += s.completed;
|
||||
failed += s.failed;
|
||||
stuck += s.stuck;
|
||||
}
|
||||
if let Ok(s) = store.sandbox_job_summary_for_user(user_id).await {
|
||||
if let Ok(s) = store.sandbox_job_summary().await {
|
||||
total += s.total;
|
||||
in_progress += s.running;
|
||||
completed += s.completed;
|
||||
@@ -187,7 +180,7 @@ impl Agent {
|
||||
}
|
||||
|
||||
// Fallback to ContextManager if no DB.
|
||||
let summary = self.context_manager.summary_for(user_id).await;
|
||||
let summary = self.context_manager.summary_for(tenant.user_id()).await;
|
||||
Ok(format!(
|
||||
"Jobs summary: Total: {} In Progress: {} Completed: {} Failed: {} Stuck: {}",
|
||||
summary.total,
|
||||
@@ -200,19 +193,24 @@ impl Agent {
|
||||
}
|
||||
}
|
||||
|
||||
async fn handle_cancel_job(&self, user_id: &str, job_id: &str) -> Result<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)
|
||||
.map_err(|_| crate::error::JobError::NotFound { id: Uuid::nil() })?;
|
||||
|
||||
let ctx = self.context_manager.get_context(uuid).await?;
|
||||
if ctx.user_id != user_id {
|
||||
if ctx.user_id != tenant.user_id() {
|
||||
return Err(crate::error::JobError::NotFound { id: uuid }.into());
|
||||
}
|
||||
|
||||
self.scheduler.stop(uuid).await?;
|
||||
|
||||
// Also update DB so the Jobs tab reflects cancellation immediately.
|
||||
if let Some(store) = self.store()
|
||||
// Use TenantScope — ownership already verified above.
|
||||
if let Some(store) = tenant.store()
|
||||
&& let Err(e) = store
|
||||
.update_job_status(uuid, JobState::Cancelled, Some("Cancelled by user"))
|
||||
.await
|
||||
@@ -225,19 +223,20 @@ impl Agent {
|
||||
|
||||
async fn handle_list_jobs(
|
||||
&self,
|
||||
user_id: &str,
|
||||
tenant: &crate::tenant::TenantCtx,
|
||||
_filter: Option<String>,
|
||||
) -> Result<String, Error> {
|
||||
// List from DB for consistency with Jobs tab.
|
||||
if let Some(store) = self.store() {
|
||||
let agent_jobs = match store.list_agent_jobs_for_user(user_id).await {
|
||||
// TenantScope methods auto-scope to user.
|
||||
if let Some(store) = tenant.store() {
|
||||
let agent_jobs = match store.list_agent_jobs().await {
|
||||
Ok(jobs) => jobs,
|
||||
Err(e) => {
|
||||
tracing::warn!("Failed to list agent jobs: {}", e);
|
||||
Vec::new()
|
||||
}
|
||||
};
|
||||
let sandbox_jobs = match store.list_sandbox_jobs_for_user(user_id).await {
|
||||
let sandbox_jobs = match store.list_sandbox_jobs().await {
|
||||
Ok(jobs) => jobs,
|
||||
Err(e) => {
|
||||
tracing::warn!("Failed to list sandbox jobs: {}", e);
|
||||
@@ -260,7 +259,7 @@ impl Agent {
|
||||
}
|
||||
|
||||
// Fallback to ContextManager if no DB.
|
||||
let jobs = self.context_manager.all_jobs_for(user_id).await;
|
||||
let jobs = self.context_manager.all_jobs_for(tenant.user_id()).await;
|
||||
if jobs.is_empty() {
|
||||
return Ok("No jobs found.".to_string());
|
||||
}
|
||||
@@ -274,12 +273,16 @@ impl Agent {
|
||||
Ok(output)
|
||||
}
|
||||
|
||||
async fn handle_help_job(&self, user_id: &str, job_id: &str) -> Result<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)
|
||||
.map_err(|_| crate::error::JobError::NotFound { id: Uuid::nil() })?;
|
||||
|
||||
let ctx = self.context_manager.get_context(uuid).await?;
|
||||
if ctx.user_id != user_id {
|
||||
if ctx.user_id != tenant.user_id() {
|
||||
return Err(crate::error::JobError::NotFound { id: uuid }.into());
|
||||
}
|
||||
|
||||
@@ -312,11 +315,11 @@ impl Agent {
|
||||
/// Show job status inline — either all jobs (no id) or a specific job.
|
||||
pub(super) async fn process_job_status(
|
||||
&self,
|
||||
user_id: &str,
|
||||
tenant: &crate::tenant::TenantCtx,
|
||||
job_id: Option<&str>,
|
||||
) -> Result<SubmissionResult, Error> {
|
||||
match self
|
||||
.handle_check_status(user_id, job_id.map(|s| s.to_string()))
|
||||
.handle_check_status(tenant, job_id.map(|s| s.to_string()))
|
||||
.await
|
||||
{
|
||||
Ok(text) => Ok(SubmissionResult::response(text)),
|
||||
@@ -327,10 +330,10 @@ impl Agent {
|
||||
/// Cancel a job by ID.
|
||||
pub(super) async fn process_job_cancel(
|
||||
&self,
|
||||
user_id: &str,
|
||||
tenant: &crate::tenant::TenantCtx,
|
||||
job_id: &str,
|
||||
) -> Result<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)),
|
||||
Err(e) => Ok(SubmissionResult::error(format!("Cancel error: {}", e))),
|
||||
}
|
||||
@@ -475,7 +478,7 @@ impl Agent {
|
||||
command: &str,
|
||||
args: &[String],
|
||||
channel: &str,
|
||||
user_id: &str,
|
||||
tenant: &crate::tenant::TenantCtx,
|
||||
) -> Result<SubmissionResult, Error> {
|
||||
match command {
|
||||
"help" => Ok(SubmissionResult::response(concat!(
|
||||
@@ -674,7 +677,7 @@ impl Agent {
|
||||
// would change the default for all users. The per-request
|
||||
// model_override in the dispatcher reads from the same
|
||||
// "selected_model" setting and applies it per-user.
|
||||
self.persist_selected_model(user_id, requested).await;
|
||||
self.persist_selected_model(tenant, requested).await;
|
||||
Ok(SubmissionResult::response(format!(
|
||||
"Model preference set to: {} (per-user)",
|
||||
requested
|
||||
@@ -683,7 +686,7 @@ impl Agent {
|
||||
match self.llm().set_model(requested) {
|
||||
Ok(()) => {
|
||||
// Persist the model choice so it survives restarts.
|
||||
self.persist_selected_model(user_id, requested).await;
|
||||
self.persist_selected_model(tenant, requested).await;
|
||||
Ok(SubmissionResult::response(format!(
|
||||
"Switched model to: {}",
|
||||
requested
|
||||
@@ -835,12 +838,12 @@ impl Agent {
|
||||
command: &str,
|
||||
args: &[String],
|
||||
channel: &str,
|
||||
user_id: &str,
|
||||
tenant: &crate::tenant::TenantCtx,
|
||||
) -> Result<Option<String>, Error> {
|
||||
// System commands are now handled directly via Submission::SystemCommand,
|
||||
// but the router may still send us unknown /commands.
|
||||
match self
|
||||
.handle_system_command(command, args, channel, user_id)
|
||||
.handle_system_command(command, args, channel, tenant)
|
||||
.await?
|
||||
{
|
||||
SubmissionResult::Response { content } => Ok(Some(content)),
|
||||
@@ -857,14 +860,18 @@ impl Agent {
|
||||
///
|
||||
/// In multi-tenant mode, only the per-user DB setting is written — global
|
||||
/// .env and TOML files are shared across users and must not be mutated.
|
||||
async fn persist_selected_model(&self, user_id: &str, model: &str) {
|
||||
// 1. Persist to DB if available (per-user scoped).
|
||||
if let Some(store) = self.store() {
|
||||
async fn persist_selected_model(&self, tenant: &crate::tenant::TenantCtx, model: &str) {
|
||||
// 1. Persist to DB if available (per-user scoped via TenantScope).
|
||||
if let Some(store) = tenant.store() {
|
||||
let value = serde_json::Value::String(model.to_string());
|
||||
if let Err(e) = store.set_setting(user_id, "selected_model", &value).await {
|
||||
if let Err(e) = store.set_setting("selected_model", &value).await {
|
||||
tracing::warn!("Failed to persist model to DB: {}", e);
|
||||
} else {
|
||||
tracing::debug!(user_id, "Persisted selected_model to DB: {}", model);
|
||||
tracing::debug!(
|
||||
user_id = tenant.user_id(),
|
||||
"Persisted selected_model to DB: {}",
|
||||
model
|
||||
);
|
||||
}
|
||||
} else {
|
||||
tracing::warn!("No database store available — model choice will not persist to DB");
|
||||
|
||||
+22
-17
@@ -42,6 +42,7 @@ impl Agent {
|
||||
pub(super) async fn run_agentic_loop(
|
||||
&self,
|
||||
message: &IncomingMessage,
|
||||
tenant: crate::tenant::TenantCtx,
|
||||
session: Arc<Mutex<Session>>,
|
||||
thread_id: Uuid,
|
||||
initial_messages: Vec<ChatMessage>,
|
||||
@@ -163,6 +164,7 @@ impl Agent {
|
||||
|
||||
let delegate = ChatDelegate {
|
||||
agent: self,
|
||||
tenant,
|
||||
session: session.clone(),
|
||||
thread_id,
|
||||
message,
|
||||
@@ -235,6 +237,7 @@ impl Agent {
|
||||
/// auth intercept, and cost tracking.
|
||||
struct ChatDelegate<'a> {
|
||||
agent: &'a Agent,
|
||||
tenant: crate::tenant::TenantCtx,
|
||||
session: Arc<Mutex<Session>>,
|
||||
thread_id: Uuid,
|
||||
message: &'a IncomingMessage,
|
||||
@@ -332,12 +335,7 @@ impl<'a> LoopDelegate for ChatDelegate<'a> {
|
||||
iteration: usize,
|
||||
) -> Result<crate::llm::RespondOutput, Error> {
|
||||
// Enforce cost guardrails before the LLM call (global + per-user)
|
||||
if let Err(limit) = self
|
||||
.agent
|
||||
.cost_guard()
|
||||
.check_allowed_for_user(&self.message.user_id)
|
||||
.await
|
||||
{
|
||||
if let Err(limit) = self.tenant.check_cost_allowed().await {
|
||||
return Err(crate::error::LlmError::InvalidResponse {
|
||||
provider: "agent".to_string(),
|
||||
reason: limit.to_string(),
|
||||
@@ -348,12 +346,10 @@ impl<'a> LoopDelegate for ChatDelegate<'a> {
|
||||
// Apply per-user model override from settings (first iteration only
|
||||
// to avoid repeated DB lookups within the same agentic loop).
|
||||
// Uses "selected_model" — the same key the /model command persists to
|
||||
// via SettingsStore (per-user scoped).
|
||||
// via SettingsStore (per-user scoped via TenantScope).
|
||||
if iteration == 0
|
||||
&& let Some(store) = self.agent.store()
|
||||
&& let Ok(Some(value)) = store
|
||||
.get_setting(&self.message.user_id, "selected_model")
|
||||
.await
|
||||
&& let Some(store) = self.tenant.store()
|
||||
&& let Ok(Some(value)) = store.get_setting("selected_model").await
|
||||
&& let Some(model) = value.as_str()
|
||||
{
|
||||
let model = model.trim();
|
||||
@@ -411,10 +407,8 @@ impl<'a> LoopDelegate for ChatDelegate<'a> {
|
||||
let read_discount = self.agent.llm().cache_read_discount();
|
||||
let write_multiplier = self.agent.llm().cache_write_multiplier();
|
||||
let call_cost = self
|
||||
.agent
|
||||
.cost_guard()
|
||||
.record_llm_call_for_user(
|
||||
&self.message.user_id,
|
||||
.tenant
|
||||
.record_llm_call(
|
||||
&model_name,
|
||||
output.usage.input_tokens,
|
||||
output.usage.output_tokens,
|
||||
@@ -1267,6 +1261,7 @@ mod tests {
|
||||
sandbox_readiness: crate::agent::routine_engine::SandboxReadiness::DisabledByConfig,
|
||||
builder: None,
|
||||
llm_backend: "nearai".to_string(),
|
||||
tenant_rates: Arc::new(crate::tenant::TenantRateRegistry::new(4, 3)),
|
||||
};
|
||||
|
||||
Agent::new(
|
||||
@@ -1288,6 +1283,8 @@ mod tests {
|
||||
default_timezone: "UTC".to_string(),
|
||||
max_tokens_per_job: 0,
|
||||
multi_tenant: false,
|
||||
max_llm_concurrent_per_user: None,
|
||||
max_jobs_concurrent_per_user: None,
|
||||
},
|
||||
deps,
|
||||
Arc::new(ChannelManager::new()),
|
||||
@@ -2137,6 +2134,7 @@ mod tests {
|
||||
sandbox_readiness: crate::agent::routine_engine::SandboxReadiness::DisabledByConfig,
|
||||
builder: None,
|
||||
llm_backend: "nearai".to_string(),
|
||||
tenant_rates: Arc::new(crate::tenant::TenantRateRegistry::new(4, 3)),
|
||||
};
|
||||
|
||||
Agent::new(
|
||||
@@ -2158,6 +2156,8 @@ mod tests {
|
||||
default_timezone: "UTC".to_string(),
|
||||
max_tokens_per_job: 0,
|
||||
multi_tenant: false,
|
||||
max_llm_concurrent_per_user: None,
|
||||
max_jobs_concurrent_per_user: None,
|
||||
},
|
||||
deps,
|
||||
Arc::new(ChannelManager::new()),
|
||||
@@ -2192,13 +2192,14 @@ mod tests {
|
||||
|
||||
let message = IncomingMessage::new("test", "test-user", "do something");
|
||||
let initial_messages = vec![ChatMessage::user("do something")];
|
||||
let tenant = agent.tenant_ctx("test-user").await;
|
||||
|
||||
// The dispatcher must terminate within 5 seconds. If there is an
|
||||
// infinite loop bug (e.g., index not advancing on tool failure), the
|
||||
// timeout will fire and the test will fail.
|
||||
let result = tokio::time::timeout(
|
||||
Duration::from_secs(5),
|
||||
agent.run_agentic_loop(&message, session, thread_id, initial_messages),
|
||||
agent.run_agentic_loop(&message, tenant, session, thread_id, initial_messages),
|
||||
)
|
||||
.await;
|
||||
|
||||
@@ -2260,6 +2261,7 @@ mod tests {
|
||||
sandbox_readiness: crate::agent::routine_engine::SandboxReadiness::DisabledByConfig,
|
||||
builder: None,
|
||||
llm_backend: "nearai".to_string(),
|
||||
tenant_rates: Arc::new(crate::tenant::TenantRateRegistry::new(4, 3)),
|
||||
};
|
||||
|
||||
Agent::new(
|
||||
@@ -2281,6 +2283,8 @@ mod tests {
|
||||
default_timezone: "UTC".to_string(),
|
||||
max_tokens_per_job: 0,
|
||||
multi_tenant: false,
|
||||
max_llm_concurrent_per_user: None,
|
||||
max_jobs_concurrent_per_user: None,
|
||||
},
|
||||
deps,
|
||||
Arc::new(ChannelManager::new()),
|
||||
@@ -2300,13 +2304,14 @@ mod tests {
|
||||
|
||||
let message = IncomingMessage::new("test", "test-user", "keep calling tools");
|
||||
let initial_messages = vec![ChatMessage::user("keep calling tools")];
|
||||
let tenant = agent.tenant_ctx("test-user").await;
|
||||
|
||||
// Even with an LLM that always wants to call tools, the dispatcher
|
||||
// must terminate within the timeout thanks to force_text at
|
||||
// max_tool_iterations.
|
||||
let result = tokio::time::timeout(
|
||||
Duration::from_secs(5),
|
||||
agent.run_agentic_loop(&message, session, thread_id, initial_messages),
|
||||
agent.run_agentic_loop(&message, tenant, session, thread_id, initial_messages),
|
||||
)
|
||||
.await;
|
||||
|
||||
|
||||
+10
-10
@@ -31,8 +31,8 @@ use chrono_tz::Tz;
|
||||
use tokio::sync::mpsc;
|
||||
|
||||
use crate::channels::OutgoingResponse;
|
||||
use crate::db::Database;
|
||||
use crate::llm::{ChatMessage, CompletionRequest, LlmProvider, Reasoning};
|
||||
use crate::tenant::AdminScope;
|
||||
use crate::workspace::Workspace;
|
||||
use crate::workspace::hygiene::HygieneConfig;
|
||||
|
||||
@@ -182,7 +182,7 @@ pub struct HeartbeatRunner {
|
||||
workspace: Arc<Workspace>,
|
||||
llm: Arc<dyn LlmProvider>,
|
||||
response_tx: Option<mpsc::Sender<OutgoingResponse>>,
|
||||
store: Option<Arc<dyn Database>>,
|
||||
store: Option<AdminScope>,
|
||||
consecutive_failures: u32,
|
||||
}
|
||||
|
||||
@@ -211,8 +211,8 @@ impl HeartbeatRunner {
|
||||
self
|
||||
}
|
||||
|
||||
/// Set the database store for persistent heartbeat conversations.
|
||||
pub fn with_store(mut self, store: Arc<dyn Database>) -> Self {
|
||||
/// Set the admin-scoped database store for persistent heartbeat conversations.
|
||||
pub fn with_store(mut self, store: AdminScope) -> Self {
|
||||
self.store = Some(store);
|
||||
self
|
||||
}
|
||||
@@ -497,7 +497,7 @@ pub fn spawn_heartbeat(
|
||||
workspace: Arc<Workspace>,
|
||||
llm: Arc<dyn LlmProvider>,
|
||||
response_tx: Option<mpsc::Sender<OutgoingResponse>>,
|
||||
store: Option<Arc<dyn Database>>,
|
||||
store: Option<AdminScope>,
|
||||
) -> tokio::task::JoinHandle<()> {
|
||||
let mut runner = HeartbeatRunner::new(config, hygiene_config, workspace, llm);
|
||||
if let Some(tx) = response_tx {
|
||||
@@ -521,7 +521,7 @@ pub fn spawn_multi_user_heartbeat(
|
||||
hygiene_config: HygieneConfig,
|
||||
llm: Arc<dyn LlmProvider>,
|
||||
response_tx: Option<mpsc::Sender<OutgoingResponse>>,
|
||||
store: Arc<dyn Database>,
|
||||
store: AdminScope,
|
||||
) -> tokio::task::JoinHandle<()> {
|
||||
tokio::spawn(async move {
|
||||
if !config.enabled {
|
||||
@@ -586,7 +586,7 @@ pub fn spawn_multi_user_heartbeat(
|
||||
continue;
|
||||
}
|
||||
|
||||
let workspace = Arc::new(Workspace::new_with_db(user_id, store.clone()));
|
||||
let workspace = Arc::new(Workspace::new_with_db(user_id, Arc::clone(store.db())));
|
||||
|
||||
// Run memory hygiene per user (same as single-user heartbeat).
|
||||
let hygiene_ws = Arc::clone(&workspace);
|
||||
@@ -617,14 +617,14 @@ pub fn spawn_multi_user_heartbeat(
|
||||
let hyg = hygiene_config.clone();
|
||||
let llm_clone = llm.clone();
|
||||
let tx = response_tx.clone();
|
||||
let st = store.clone();
|
||||
let admin = store.clone();
|
||||
|
||||
join_set.spawn(async move {
|
||||
let mut runner = HeartbeatRunner::new(cfg, hyg, workspace, llm_clone);
|
||||
if let Some(tx) = tx {
|
||||
runner = runner.with_response_channel(tx);
|
||||
}
|
||||
runner = runner.with_store(st);
|
||||
runner = runner.with_store(admin);
|
||||
|
||||
let result = runner.check_heartbeat().await;
|
||||
if let HeartbeatResult::NeedsAttention(msg) = &result {
|
||||
@@ -903,7 +903,7 @@ mod tests {
|
||||
Arc<crate::workspace::Workspace>,
|
||||
Arc<dyn crate::llm::LlmProvider>,
|
||||
Option<tokio::sync::mpsc::Sender<crate::channels::OutgoingResponse>>,
|
||||
Option<Arc<dyn crate::db::Database>>,
|
||||
Option<AdminScope>,
|
||||
) -> tokio::task::JoinHandle<()> = spawn_heartbeat;
|
||||
let _ = _fn_ptr;
|
||||
}
|
||||
|
||||
@@ -27,12 +27,12 @@ use crate::agent::routine::{
|
||||
use crate::channels::{IncomingMessage, OutgoingResponse};
|
||||
use crate::config::RoutineConfig;
|
||||
use crate::context::{JobContext, JobState};
|
||||
use crate::db::Database;
|
||||
use crate::error::RoutineError;
|
||||
use crate::extensions::ExtensionManager;
|
||||
use crate::llm::{
|
||||
ChatMessage, CompletionRequest, FinishReason, LlmProvider, ToolCall, ToolCompletionRequest,
|
||||
};
|
||||
use crate::tenant::AdminScope;
|
||||
use crate::tools::{
|
||||
ToolError, ToolRegistry, autonomous_allowed_tool_names, autonomous_unavailable_message,
|
||||
prepare_tool_params,
|
||||
@@ -93,7 +93,7 @@ pub(crate) fn routine_matches_message(routine: &Routine, message: &IncomingMessa
|
||||
/// The routine execution engine.
|
||||
pub struct RoutineEngine {
|
||||
config: RoutineConfig,
|
||||
store: Arc<dyn Database>,
|
||||
store: AdminScope,
|
||||
llm: Arc<dyn LlmProvider>,
|
||||
workspace: Arc<Workspace>,
|
||||
/// Sender for notifications (routed to channel manager).
|
||||
@@ -122,7 +122,7 @@ impl RoutineEngine {
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
pub fn new(
|
||||
config: RoutineConfig,
|
||||
store: Arc<dyn Database>,
|
||||
store: AdminScope,
|
||||
llm: Arc<dyn LlmProvider>,
|
||||
workspace: Arc<Workspace>,
|
||||
notify_tx: mpsc::Sender<OutgoingResponse>,
|
||||
@@ -741,7 +741,10 @@ impl RoutineEngine {
|
||||
let routine_workspace = if routine.user_id == self.workspace.user_id() {
|
||||
self.workspace.clone()
|
||||
} else {
|
||||
Arc::new(Workspace::new_with_db(&routine.user_id, self.store.clone()))
|
||||
Arc::new(Workspace::new_with_db(
|
||||
&routine.user_id,
|
||||
Arc::clone(self.store.db()),
|
||||
))
|
||||
};
|
||||
|
||||
// Execute inline for manual triggers (caller wants to wait)
|
||||
@@ -873,7 +876,10 @@ impl RoutineEngine {
|
||||
let routine_workspace = if routine.user_id == self.workspace.user_id() {
|
||||
self.workspace.clone()
|
||||
} else {
|
||||
Arc::new(Workspace::new_with_db(&routine.user_id, self.store.clone()))
|
||||
Arc::new(Workspace::new_with_db(
|
||||
&routine.user_id,
|
||||
Arc::clone(self.store.db()),
|
||||
))
|
||||
};
|
||||
|
||||
let engine = EngineContext {
|
||||
@@ -933,7 +939,7 @@ impl RoutineEngine {
|
||||
/// an active state (Pending/InProgress/Stuck). Maps the final `JobState` to
|
||||
/// a `RunStatus` for the routine run.
|
||||
struct FullJobWatcher {
|
||||
store: Arc<dyn Database>,
|
||||
store: AdminScope,
|
||||
job_id: Uuid,
|
||||
routine_name: String,
|
||||
}
|
||||
@@ -944,7 +950,7 @@ impl FullJobWatcher {
|
||||
/// Safety ceiling: 24 hours, derived from POLL_INTERVAL.
|
||||
const MAX_POLLS: u32 = (24 * 60 * 60) / Self::POLL_INTERVAL.as_secs() as u32;
|
||||
|
||||
fn new(store: Arc<dyn Database>, job_id: Uuid, routine_name: String) -> Self {
|
||||
fn new(store: AdminScope, job_id: Uuid, routine_name: String) -> Self {
|
||||
Self {
|
||||
store,
|
||||
job_id,
|
||||
@@ -1016,7 +1022,7 @@ impl FullJobWatcher {
|
||||
/// Shared context passed to the execution function.
|
||||
struct EngineContext {
|
||||
config: RoutineConfig,
|
||||
store: Arc<dyn Database>,
|
||||
store: AdminScope,
|
||||
llm: Arc<dyn LlmProvider>,
|
||||
workspace: Arc<Workspace>,
|
||||
notify_tx: mpsc::Sender<OutgoingResponse>,
|
||||
|
||||
@@ -11,12 +11,12 @@ use uuid::Uuid;
|
||||
use crate::agent::task::{Task, TaskContext, TaskOutput};
|
||||
use crate::config::AgentConfig;
|
||||
use crate::context::{ContextManager, JobContext, JobState};
|
||||
use crate::db::Database;
|
||||
use crate::error::{Error, JobError};
|
||||
use crate::extensions::ExtensionManager;
|
||||
use crate::hooks::HookRegistry;
|
||||
use crate::llm::LlmProvider;
|
||||
use crate::safety::SafetyLayer;
|
||||
use crate::tenant::AdminScope;
|
||||
use crate::tools::{
|
||||
ApprovalContext, ToolRegistry, autonomous_allowed_tool_names, autonomous_unavailable_error,
|
||||
prepare_tool_params,
|
||||
@@ -52,7 +52,7 @@ struct ScheduledSubtask {
|
||||
pub struct SchedulerDeps {
|
||||
pub tools: Arc<ToolRegistry>,
|
||||
pub extension_manager: Option<Arc<ExtensionManager>>,
|
||||
pub store: Option<Arc<dyn Database>>,
|
||||
pub store: Option<AdminScope>,
|
||||
pub hooks: Arc<HookRegistry>,
|
||||
}
|
||||
|
||||
@@ -64,7 +64,7 @@ pub struct Scheduler {
|
||||
safety: Arc<SafetyLayer>,
|
||||
tools: Arc<ToolRegistry>,
|
||||
extension_manager: Option<Arc<ExtensionManager>>,
|
||||
store: Option<Arc<dyn Database>>,
|
||||
store: Option<AdminScope>,
|
||||
hooks: Arc<HookRegistry>,
|
||||
/// SSE manager for live job event streaming.
|
||||
sse_tx: Option<Arc<crate::channels::web::sse::SseManager>>,
|
||||
@@ -786,6 +786,8 @@ mod tests {
|
||||
default_timezone: "UTC".to_string(),
|
||||
max_tokens_per_job,
|
||||
multi_tenant: false,
|
||||
max_llm_concurrent_per_user: None,
|
||||
max_jobs_concurrent_per_user: None,
|
||||
};
|
||||
let cm = Arc::new(ContextManager::new(5));
|
||||
let llm: Arc<dyn LlmProvider> = Arc::new(StubLlm);
|
||||
|
||||
@@ -8,8 +8,8 @@ use chrono::{DateTime, Utc};
|
||||
use uuid::Uuid;
|
||||
|
||||
use crate::context::{ContextManager, JobState};
|
||||
use crate::db::Database;
|
||||
use crate::error::RepairError;
|
||||
use crate::tenant::AdminScope;
|
||||
use crate::tools::{BuildRequirement, Language, SoftwareBuilder, SoftwareType, ToolRegistry};
|
||||
|
||||
/// A job that has been detected as stuck.
|
||||
@@ -69,7 +69,7 @@ pub struct DefaultSelfRepair {
|
||||
/// Jobs in `InProgress` longer than this are treated as stuck.
|
||||
stuck_threshold: Duration,
|
||||
max_repair_attempts: u32,
|
||||
store: Option<Arc<dyn Database>>,
|
||||
store: Option<AdminScope>,
|
||||
builder: Option<Arc<dyn SoftwareBuilder>>,
|
||||
tools: Option<Arc<ToolRegistry>>,
|
||||
}
|
||||
@@ -91,8 +91,8 @@ impl DefaultSelfRepair {
|
||||
}
|
||||
}
|
||||
|
||||
/// Add a Store for tool failure tracking.
|
||||
pub fn with_store(mut self, store: Arc<dyn Database>) -> Self {
|
||||
/// Add an admin-scoped store for tool failure tracking.
|
||||
pub fn with_store(mut self, store: AdminScope) -> Self {
|
||||
self.store = Some(store);
|
||||
self
|
||||
}
|
||||
@@ -806,7 +806,7 @@ mod tests {
|
||||
// Create self-repair with zero threshold (detect immediately),
|
||||
// wired with store, builder, and tools.
|
||||
let repair = DefaultSelfRepair::new(Arc::clone(&cm), Duration::from_secs(0), 3)
|
||||
.with_store(Arc::clone(&db))
|
||||
.with_store(crate::tenant::AdminScope::new(Arc::clone(&db)))
|
||||
.with_builder(
|
||||
Arc::clone(&builder) as Arc<dyn crate::tools::SoftwareBuilder>,
|
||||
tools,
|
||||
|
||||
+10
-3
@@ -175,6 +175,7 @@ impl Agent {
|
||||
pub(super) async fn process_user_input(
|
||||
&self,
|
||||
message: &IncomingMessage,
|
||||
tenant: crate::tenant::TenantCtx,
|
||||
session: Arc<Mutex<Session>>,
|
||||
thread_id: Uuid,
|
||||
content: &str,
|
||||
@@ -351,7 +352,7 @@ impl Agent {
|
||||
|
||||
if let Some(intent) = self.router.route_command(&temp_message) {
|
||||
// Explicit command like /status, /job, /list - handle directly
|
||||
return self.handle_job_or_command(intent, message).await;
|
||||
return self.handle_job_or_command(intent, message, &tenant).await;
|
||||
}
|
||||
|
||||
// Natural language goes through the agentic loop
|
||||
@@ -462,7 +463,7 @@ impl Agent {
|
||||
|
||||
// Run the agentic tool execution loop
|
||||
let result = self
|
||||
.run_agentic_loop(message, session.clone(), thread_id, turn_messages)
|
||||
.run_agentic_loop(message, tenant, session.clone(), thread_id, turn_messages)
|
||||
.await;
|
||||
|
||||
// Re-acquire lock and check if interrupted
|
||||
@@ -1444,7 +1445,13 @@ impl Agent {
|
||||
|
||||
// Continue the agentic loop (a tool was already executed this turn)
|
||||
let result = self
|
||||
.run_agentic_loop(message, session.clone(), thread_id, context_messages)
|
||||
.run_agentic_loop(
|
||||
message,
|
||||
self.tenant_ctx(&message.user_id).await,
|
||||
session.clone(),
|
||||
thread_id,
|
||||
context_messages,
|
||||
)
|
||||
.await;
|
||||
|
||||
// Handle the result
|
||||
|
||||
@@ -36,6 +36,10 @@ pub struct AgentConfig {
|
||||
/// Whether the deployment is multi-tenant (multiple users sharing one
|
||||
/// instance). Auto-detected from GATEWAY_USER_TOKENS presence.
|
||||
pub multi_tenant: bool,
|
||||
/// Maximum concurrent LLM calls per user. None = use default (4).
|
||||
pub max_llm_concurrent_per_user: Option<usize>,
|
||||
/// Maximum concurrent jobs per user. None = use default (3).
|
||||
pub max_jobs_concurrent_per_user: Option<usize>,
|
||||
}
|
||||
|
||||
impl AgentConfig {
|
||||
@@ -60,6 +64,8 @@ impl AgentConfig {
|
||||
default_timezone: "UTC".to_string(),
|
||||
max_tokens_per_job: 0,
|
||||
multi_tenant: false,
|
||||
max_llm_concurrent_per_user: None,
|
||||
max_jobs_concurrent_per_user: None,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -123,6 +129,8 @@ impl AgentConfig {
|
||||
// Auto-detected from GATEWAY_USER_TOKENS presence. Not a separate
|
||||
// knob — multi-tenant mode is always implied by configuring user tokens.
|
||||
multi_tenant: optional_env("GATEWAY_USER_TOKENS")?.is_some(),
|
||||
max_llm_concurrent_per_user: parse_option_env("TENANT_MAX_LLM_CONCURRENT")?,
|
||||
max_jobs_concurrent_per_user: parse_option_env("TENANT_MAX_JOBS_CONCURRENT")?,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -69,6 +69,7 @@ pub mod service;
|
||||
pub mod settings;
|
||||
pub mod setup;
|
||||
pub mod skills;
|
||||
pub mod tenant;
|
||||
pub mod timezone;
|
||||
pub mod tools;
|
||||
pub mod tracing_fmt;
|
||||
|
||||
@@ -913,6 +913,10 @@ async fn async_main() -> anyhow::Result<()> {
|
||||
},
|
||||
builder: components.builder,
|
||||
llm_backend: config.llm.backend.clone(),
|
||||
tenant_rates: Arc::new(ironclaw::tenant::TenantRateRegistry::new(
|
||||
config.agent.max_llm_concurrent_per_user.unwrap_or(4),
|
||||
config.agent.max_jobs_concurrent_per_user.unwrap_or(3),
|
||||
)),
|
||||
};
|
||||
|
||||
let channels_for_warnings = Arc::clone(&channels);
|
||||
|
||||
+906
@@ -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));
|
||||
}
|
||||
}
|
||||
@@ -565,6 +565,7 @@ impl TestHarnessBuilder {
|
||||
sandbox_readiness: crate::agent::routine_engine::SandboxReadiness::DisabledByConfig,
|
||||
builder: None,
|
||||
llm_backend: "nearai".to_string(),
|
||||
tenant_rates: std::sync::Arc::new(crate::tenant::TenantRateRegistry::new(4, 3)),
|
||||
};
|
||||
|
||||
TestHarness {
|
||||
|
||||
+3
-3
@@ -20,7 +20,6 @@ use crate::agent::scheduler::WorkerMessage;
|
||||
use crate::agent::task::TaskOutput;
|
||||
use crate::channels::web::types::SseEvent;
|
||||
use crate::context::{ContextManager, JobState};
|
||||
use crate::db::Database;
|
||||
use crate::error::Error;
|
||||
use crate::hooks::HookRegistry;
|
||||
use crate::llm::{
|
||||
@@ -28,6 +27,7 @@ use crate::llm::{
|
||||
ToolSelection,
|
||||
};
|
||||
use crate::safety::SafetyLayer;
|
||||
use crate::tenant::AdminScope;
|
||||
use crate::tools::execute::process_tool_result;
|
||||
use crate::tools::rate_limiter::RateLimitResult;
|
||||
use crate::tools::{
|
||||
@@ -44,7 +44,7 @@ pub struct WorkerDeps {
|
||||
pub llm: Arc<dyn LlmProvider>,
|
||||
pub safety: Arc<SafetyLayer>,
|
||||
pub tools: Arc<ToolRegistry>,
|
||||
pub store: Option<Arc<dyn Database>>,
|
||||
pub store: Option<AdminScope>,
|
||||
pub hooks: Arc<HookRegistry>,
|
||||
pub timeout: Duration,
|
||||
pub use_planning: bool,
|
||||
@@ -93,7 +93,7 @@ impl Worker {
|
||||
&self.deps.tools
|
||||
}
|
||||
|
||||
fn store(&self) -> Option<&Arc<dyn Database>> {
|
||||
fn store(&self) -> Option<&AdminScope> {
|
||||
self.deps.store.as_ref()
|
||||
}
|
||||
|
||||
|
||||
@@ -337,14 +337,14 @@ mod tests {
|
||||
SchedulerDeps {
|
||||
tools: registry.clone(),
|
||||
extension_manager: extension_manager.clone(),
|
||||
store: Some(db.clone()),
|
||||
store: Some(ironclaw::tenant::AdminScope::new(db.clone())),
|
||||
hooks: Arc::new(HookRegistry::new()),
|
||||
},
|
||||
));
|
||||
|
||||
Arc::new(RoutineEngine::new(
|
||||
RoutineConfig::default(),
|
||||
db,
|
||||
ironclaw::tenant::AdminScope::new(db),
|
||||
llm,
|
||||
ws,
|
||||
notify_tx,
|
||||
@@ -448,7 +448,7 @@ mod tests {
|
||||
|
||||
let engine = Arc::new(RoutineEngine::new(
|
||||
RoutineConfig::default(),
|
||||
db.clone(),
|
||||
ironclaw::tenant::AdminScope::new(db.clone()),
|
||||
llm,
|
||||
ws,
|
||||
notify_tx,
|
||||
@@ -527,7 +527,7 @@ mod tests {
|
||||
|
||||
let engine = Arc::new(RoutineEngine::new(
|
||||
RoutineConfig::default(),
|
||||
db.clone(),
|
||||
ironclaw::tenant::AdminScope::new(db.clone()),
|
||||
llm,
|
||||
ws,
|
||||
notify_tx,
|
||||
@@ -614,7 +614,7 @@ mod tests {
|
||||
|
||||
let engine = Arc::new(RoutineEngine::new(
|
||||
RoutineConfig::default(),
|
||||
db.clone(),
|
||||
ironclaw::tenant::AdminScope::new(db.clone()),
|
||||
llm,
|
||||
ws,
|
||||
notify_tx,
|
||||
@@ -723,7 +723,7 @@ mod tests {
|
||||
|
||||
let engine = Arc::new(RoutineEngine::new(
|
||||
RoutineConfig::default(),
|
||||
db.clone(),
|
||||
ironclaw::tenant::AdminScope::new(db.clone()),
|
||||
llm,
|
||||
ws,
|
||||
notify_tx,
|
||||
@@ -866,7 +866,7 @@ mod tests {
|
||||
|
||||
let engine = Arc::new(RoutineEngine::new(
|
||||
RoutineConfig::default(),
|
||||
db.clone(),
|
||||
ironclaw::tenant::AdminScope::new(db.clone()),
|
||||
llm,
|
||||
ws,
|
||||
notify_tx,
|
||||
@@ -1049,7 +1049,7 @@ mod tests {
|
||||
|
||||
let engine = Arc::new(RoutineEngine::new(
|
||||
RoutineConfig::default(),
|
||||
Arc::clone(&db),
|
||||
ironclaw::tenant::AdminScope::new(Arc::clone(&db)),
|
||||
llm,
|
||||
ws,
|
||||
notify_tx,
|
||||
@@ -1171,7 +1171,7 @@ mod tests {
|
||||
|
||||
let engine = Arc::new(RoutineEngine::new(
|
||||
RoutineConfig::default(),
|
||||
db.clone(),
|
||||
ironclaw::tenant::AdminScope::new(db.clone()),
|
||||
llm,
|
||||
ws,
|
||||
notify_tx,
|
||||
@@ -1279,7 +1279,7 @@ mod tests {
|
||||
|
||||
let engine = Arc::new(RoutineEngine::new(
|
||||
config,
|
||||
db.clone(),
|
||||
ironclaw::tenant::AdminScope::new(db.clone()),
|
||||
llm,
|
||||
ws,
|
||||
notify_tx,
|
||||
|
||||
@@ -201,6 +201,7 @@ mod tests {
|
||||
sandbox_readiness: ironclaw::agent::SandboxReadiness::DisabledByConfig,
|
||||
builder: None,
|
||||
llm_backend: "nearai".to_string(),
|
||||
tenant_rates: std::sync::Arc::new(ironclaw::tenant::TenantRateRegistry::new(4, 3)),
|
||||
};
|
||||
|
||||
let gateway = Arc::new(TestChannel::new());
|
||||
|
||||
@@ -265,6 +265,7 @@ impl GatewayWorkflowHarness {
|
||||
sandbox_readiness: ironclaw::agent::SandboxReadiness::DisabledByConfig,
|
||||
builder: None,
|
||||
llm_backend: "nearai".to_string(),
|
||||
tenant_rates: std::sync::Arc::new(ironclaw::tenant::TenantRateRegistry::new(4, 3)),
|
||||
},
|
||||
channels,
|
||||
None,
|
||||
|
||||
@@ -642,7 +642,7 @@ impl TestRigBuilder {
|
||||
let (notify_tx, _notify_rx) = tokio::sync::mpsc::channel(16);
|
||||
let engine = Arc::new(RoutineEngine::new(
|
||||
routine_config,
|
||||
Arc::clone(db_arc),
|
||||
ironclaw::tenant::AdminScope::new(Arc::clone(db_arc)),
|
||||
components.llm.clone(),
|
||||
Arc::clone(ws),
|
||||
notify_tx,
|
||||
@@ -762,6 +762,7 @@ impl TestRigBuilder {
|
||||
sandbox_readiness: ironclaw::agent::SandboxReadiness::Available, // tests don't use real Docker
|
||||
builder: None,
|
||||
llm_backend: "nearai".to_string(),
|
||||
tenant_rates: std::sync::Arc::new(ironclaw::tenant::TenantRateRegistry::new(4, 3)),
|
||||
};
|
||||
|
||||
// 7. Create TestChannel and ChannelManager.
|
||||
|
||||
Reference in New Issue
Block a user