mirror of
https://github.com/outbackdingo/optimclaw.git
synced 2026-08-25 14:53:34 +00:00
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:
co-authored by
Claude Opus 4.6
parent
86d1143064
commit
4c043bf057
+136
-37
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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};
|
||||||
|
|||||||
@@ -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>,
|
||||||
|
|||||||
@@ -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);
|
||||||
|
|||||||
@@ -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
@@ -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
|
||||||
|
|||||||
@@ -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,
|
||||||
},
|
},
|
||||||
));
|
));
|
||||||
|
|
||||||
|
|||||||
@@ -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 { .. }
|
||||||
|
|||||||
@@ -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
@@ -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")?,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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(),
|
||||||
|
)?,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
@@ -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
@@ -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
@@ -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()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -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;
|
||||||
|
|||||||
@@ -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
@@ -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)
|
||||||
|
|||||||
@@ -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
@@ -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));
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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
@@ -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()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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.
|
||||||
|
|||||||
Reference in New Issue
Block a user