Merge branch 'staging' into fix/approval-thread-safety

This commit is contained in:
Zaki Manian
2026-03-24 10:11:38 -07:00
committed by GitHub
84 changed files with 6588 additions and 1756 deletions
+1 -1
View File
@@ -21,7 +21,7 @@
},
{
"name": "feishu_app_secret",
"prompt": "Enter your Feishu/Lark App Secret",
"prompt": "Enter your Feishu/Lark App Secret (from your app settings at open.feishu.cn)",
"optional": false
},
{
+10 -7
View File
@@ -157,8 +157,8 @@ pub struct AgentDeps {
pub hooks: Arc<HookRegistry>,
/// Cost enforcement guardrails (daily budget, hourly rate limits).
pub cost_guard: Arc<crate::agent::cost_guard::CostGuard>,
/// SSE broadcast sender for live job event streaming to the web gateway.
pub sse_tx: Option<tokio::sync::broadcast::Sender<crate::channels::web::types::SseEvent>>,
/// SSE manager for live job event streaming to the web gateway.
pub sse_tx: Option<Arc<crate::channels::web::sse::SseManager>>,
/// HTTP interceptor for trace recording/replay.
pub http_interceptor: Option<Arc<dyn crate::llm::recording::HttpInterceptor>>,
/// Audio transcription middleware for voice messages.
@@ -169,6 +169,9 @@ pub struct AgentDeps {
pub sandbox_readiness: crate::agent::routine_engine::SandboxReadiness,
/// Software builder for self-repair tool rebuilding.
pub builder: Option<Arc<dyn crate::tools::SoftwareBuilder>>,
/// Resolved LLM backend identifier (e.g., "nearai", "openai", "groq").
/// Used by `/model` persistence to determine which env var to update.
pub llm_backend: String,
}
/// The main agent that coordinates all components.
@@ -235,8 +238,8 @@ impl Agent {
hooks: deps.hooks.clone(),
},
);
if let Some(ref tx) = deps.sse_tx {
scheduler.set_sse_sender(tx.clone());
if let Some(ref sse) = deps.sse_tx {
scheduler.set_sse_sender(Arc::clone(sse));
}
if let Some(ref interceptor) = deps.http_interceptor {
scheduler.set_http_interceptor(Arc::clone(interceptor));
@@ -1136,9 +1139,9 @@ impl Agent {
&& let Submission::UserInput { ref content } = submission
&& let Some(engine) = self.routine_engine().await
{
let fired = engine
.check_event_triggers(&message.user_id, &message.channel, content)
.await;
// Use post-hook content so that BeforeInbound hooks that rewrite
// input are respected by event trigger matching.
let fired = engine.check_event_triggers(message, content).await;
if fired > 0 {
tracing::debug!(
channel = %message.channel,
+49 -3
View File
@@ -841,12 +841,50 @@ impl Agent {
.await
{
tracing::warn!("Failed to persist model to DB: {}", e);
} else {
tracing::debug!("Persisted selected_model to DB: {}", model);
}
} else {
tracing::warn!("No database store available — model choice will not persist to DB");
}
// 2. Update TOML config file if it exists (sync I/O in spawn_blocking).
// 2. Update .env and TOML config file (sync I/O in spawn_blocking).
let model_owned = model.to_string();
let backend = self.deps.llm_backend.clone();
if let Err(e) = tokio::task::spawn_blocking(move || {
// 2a. Update the backend-specific model env var in ~/.ironclaw/.env.
//
// Env vars have the HIGHEST priority in LlmConfig::resolve_model()
// (env var > TOML > DB > default). If the .env file has e.g.
// NEARAI_MODEL=old-model, it shadows everything else. We must
// update this var or the /model change is invisible on restart.
let registry = crate::llm::ProviderRegistry::load();
let model_env = registry.model_env_var(&backend);
let env_var_prefix = format!("{}=", model_env);
// Only update the .env file if the var is actually set there
// (avoid injecting new vars the user never configured).
let env_path = crate::bootstrap::ironclaw_env_path();
let env_has_var = std::fs::read_to_string(&env_path)
.ok()
.is_some_and(|content| {
content.lines().any(|line| {
let trimmed = line.trim_start();
!trimmed.starts_with('#') && trimmed.starts_with(&env_var_prefix)
})
});
if env_has_var {
if let Err(e) = crate::bootstrap::upsert_bootstrap_var(model_env, &model_owned) {
tracing::warn!("Failed to update {} in .env: {}", model_env, e);
} else {
tracing::debug!("Updated {} in .env to {}", model_env, model_owned);
}
}
// 2b. Update (or create) the TOML config file.
//
// The TOML overlay has higher priority than DB settings on
// startup, so it MUST stay in sync with the DB.
let toml_path = crate::settings::Settings::default_toml_path();
match crate::settings::Settings::load_toml(&toml_path) {
Ok(Some(mut settings)) => {
@@ -856,7 +894,15 @@ impl Agent {
}
}
Ok(None) => {
// No config file on disk; nothing to update.
// No config file yet — create one so the model choice
// survives restarts even when the DB is unavailable.
let settings = crate::settings::Settings {
selected_model: Some(model_owned),
..Default::default()
};
if let Err(e) = settings.save_toml(&toml_path) {
tracing::warn!("Failed to create config.toml for model persistence: {}", e);
}
}
Err(e) => {
tracing::warn!("Failed to load config.toml for model persistence: {}", e);
@@ -865,7 +911,7 @@ impl Agent {
})
.await
{
tracing::warn!("Model TOML persistence task failed: {}", e);
tracing::warn!("Model persistence task failed: {}", e);
}
}
}
+25 -4
View File
@@ -1098,15 +1098,23 @@ pub(crate) fn extract_suggestions(text: &str) -> (String, Vec<String>) {
Regex::new(r"(?s)<suggestions>\s*(.*?)\s*</suggestions>").expect("valid regex") // safety: constant pattern
});
// Find the position of the last closing code fence to avoid matching inside code blocks
let last_code_fence = text.rfind("```").unwrap_or(0);
// Build a sorted list of code fence positions to determine open/close pairing.
// A position is "inside" a fenced block when it falls between an odd-numbered
// fence (opening) and the next even-numbered fence (closing).
let fence_positions: Vec<usize> = text.match_indices("```").map(|(pos, _)| pos).collect();
// Find all matches, take the last one that's after the last code fence
let is_inside_fence = |pos: usize| -> bool {
// Count how many fences appear before `pos`. If odd, we're inside a fence.
let count = fence_positions.iter().take_while(|&&fp| fp <= pos).count();
count % 2 == 1
};
// Find all matches, take the last one that's outside any code fence
let mut best_match: Option<regex::Match<'_>> = None;
let mut best_capture: Option<String> = None;
for caps in RE.captures_iter(text) {
if let (Some(full), Some(inner)) = (caps.get(0), caps.get(1))
&& full.start() >= last_code_fence
&& !is_inside_fence(full.start())
{
best_match = Some(full);
best_capture = Some(inner.as_str().to_string());
@@ -1225,6 +1233,7 @@ mod tests {
document_extraction: None,
sandbox_readiness: crate::agent::routine_engine::SandboxReadiness::DisabledByConfig,
builder: None,
llm_backend: "nearai".to_string(),
};
Agent::new(
@@ -2092,6 +2101,7 @@ mod tests {
document_extraction: None,
sandbox_readiness: crate::agent::routine_engine::SandboxReadiness::DisabledByConfig,
builder: None,
llm_backend: "nearai".to_string(),
};
Agent::new(
@@ -2212,6 +2222,7 @@ mod tests {
document_extraction: None,
sandbox_readiness: crate::agent::routine_engine::SandboxReadiness::DisabledByConfig,
builder: None,
llm_backend: "nearai".to_string(),
};
Agent::new(
@@ -2345,6 +2356,16 @@ mod tests {
assert!(suggestions.is_empty()); // safety: test
}
#[test]
fn test_extract_suggestions_inside_unclosed_code_fence() {
// Regression: odd number of fences (unclosed fence) must still be
// treated as "inside a code block".
let input = "```\ncode\n<suggestions>[\"bar\"]</suggestions>";
let (text, suggestions) = super::extract_suggestions(input);
assert_eq!(text, input); // safety: test
assert!(suggestions.is_empty()); // safety: test
}
#[test]
fn test_extract_suggestions_after_code_fence() {
let input = "```\ncode\n```\nAnswer.\n<suggestions>[\"foo\"]</suggestions>";
+22 -12
View File
@@ -44,7 +44,7 @@ pub struct JobMonitorRoute {
/// the main agent's context window).
pub fn spawn_job_monitor(
job_id: Uuid,
event_rx: broadcast::Receiver<(Uuid, SseEvent)>,
event_rx: broadcast::Receiver<(Uuid, String, SseEvent)>,
inject_tx: mpsc::Sender<IncomingMessage>,
route: JobMonitorRoute,
) -> JoinHandle<()> {
@@ -56,7 +56,7 @@ pub fn spawn_job_monitor(
/// jobs don't stay `InProgress` forever in the `ContextManager`.
pub fn spawn_job_monitor_with_context(
job_id: Uuid,
mut event_rx: broadcast::Receiver<(Uuid, SseEvent)>,
mut event_rx: broadcast::Receiver<(Uuid, String, SseEvent)>,
inject_tx: mpsc::Sender<IncomingMessage>,
route: JobMonitorRoute,
context_manager: Option<Arc<ContextManager>>,
@@ -68,7 +68,7 @@ pub fn spawn_job_monitor_with_context(
loop {
match event_rx.recv().await {
Ok((ev_job_id, event)) => {
Ok((ev_job_id, _user_id, event)) => {
if ev_job_id != job_id {
continue;
}
@@ -162,7 +162,7 @@ pub fn spawn_job_monitor_with_context(
/// inject messages into) but we still need to free the `max_jobs` slot.
pub fn spawn_completion_watcher(
job_id: Uuid,
mut event_rx: broadcast::Receiver<(Uuid, SseEvent)>,
mut event_rx: broadcast::Receiver<(Uuid, String, SseEvent)>,
context_manager: Arc<ContextManager>,
) -> JoinHandle<()> {
let short_id = job_id.to_string()[..8].to_string();
@@ -170,7 +170,9 @@ pub fn spawn_completion_watcher(
tokio::spawn(async move {
loop {
match event_rx.recv().await {
Ok((ev_job_id, SseEvent::JobResult { status, .. })) if ev_job_id == job_id => {
Ok((ev_job_id, _user_id, SseEvent::JobResult { status, .. }))
if ev_job_id == job_id =>
{
let target = if status == "completed" {
JobState::Completed
} else {
@@ -227,7 +229,7 @@ mod tests {
#[tokio::test]
async fn test_monitor_forwards_assistant_messages() {
let (event_tx, _) = broadcast::channel::<(Uuid, SseEvent)>(16);
let (event_tx, _) = broadcast::channel::<(Uuid, String, SseEvent)>(16);
let (inject_tx, mut inject_rx) = mpsc::channel::<IncomingMessage>(16);
let job_id = Uuid::new_v4();
@@ -237,6 +239,7 @@ mod tests {
event_tx
.send((
job_id,
"test-user".to_string(),
SseEvent::JobMessage {
job_id: job_id.to_string(),
role: "assistant".to_string(),
@@ -259,7 +262,7 @@ mod tests {
#[tokio::test]
async fn test_monitor_ignores_other_jobs() {
let (event_tx, _) = broadcast::channel::<(Uuid, SseEvent)>(16);
let (event_tx, _) = broadcast::channel::<(Uuid, String, SseEvent)>(16);
let (inject_tx, mut inject_rx) = mpsc::channel::<IncomingMessage>(16);
let job_id = Uuid::new_v4();
@@ -270,6 +273,7 @@ mod tests {
event_tx
.send((
other_job_id,
"test-user".to_string(),
SseEvent::JobMessage {
job_id: other_job_id.to_string(),
role: "assistant".to_string(),
@@ -289,7 +293,7 @@ mod tests {
#[tokio::test]
async fn test_monitor_exits_on_job_result() {
let (event_tx, _) = broadcast::channel::<(Uuid, SseEvent)>(16);
let (event_tx, _) = broadcast::channel::<(Uuid, String, SseEvent)>(16);
let (inject_tx, mut inject_rx) = mpsc::channel::<IncomingMessage>(16);
let job_id = Uuid::new_v4();
@@ -299,6 +303,7 @@ mod tests {
event_tx
.send((
job_id,
"test-user".to_string(),
SseEvent::JobResult {
job_id: job_id.to_string(),
status: "completed".to_string(),
@@ -324,7 +329,7 @@ mod tests {
#[tokio::test]
async fn test_monitor_skips_tool_events() {
let (event_tx, _) = broadcast::channel::<(Uuid, SseEvent)>(16);
let (event_tx, _) = broadcast::channel::<(Uuid, String, SseEvent)>(16);
let (inject_tx, mut inject_rx) = mpsc::channel::<IncomingMessage>(16);
let job_id = Uuid::new_v4();
@@ -334,6 +339,7 @@ mod tests {
event_tx
.send((
job_id,
"test-user".to_string(),
SseEvent::JobToolUse {
job_id: job_id.to_string(),
tool_name: "shell".to_string(),
@@ -346,6 +352,7 @@ mod tests {
event_tx
.send((
job_id,
"test-user".to_string(),
SseEvent::JobMessage {
job_id: job_id.to_string(),
role: "user".to_string(),
@@ -402,7 +409,7 @@ mod tests {
.await
.unwrap();
let (event_tx, _) = broadcast::channel::<(Uuid, SseEvent)>(16);
let (event_tx, _) = broadcast::channel::<(Uuid, String, SseEvent)>(16);
let (inject_tx, mut inject_rx) = mpsc::channel::<IncomingMessage>(16);
let handle = spawn_job_monitor_with_context(
@@ -417,6 +424,7 @@ mod tests {
event_tx
.send((
job_id,
"test-user".to_string(),
SseEvent::JobResult {
job_id: job_id.to_string(),
status: "completed".to_string(),
@@ -450,7 +458,7 @@ mod tests {
.await
.unwrap();
let (event_tx, _) = broadcast::channel::<(Uuid, SseEvent)>(16);
let (event_tx, _) = broadcast::channel::<(Uuid, String, SseEvent)>(16);
let (inject_tx, mut inject_rx) = mpsc::channel::<IncomingMessage>(16);
let handle = spawn_job_monitor_with_context(
@@ -465,6 +473,7 @@ mod tests {
event_tx
.send((
job_id,
"test-user".to_string(),
SseEvent::JobResult {
job_id: job_id.to_string(),
status: "failed".to_string(),
@@ -498,12 +507,13 @@ mod tests {
.await
.unwrap();
let (event_tx, _) = broadcast::channel::<(Uuid, SseEvent)>(16);
let (event_tx, _) = broadcast::channel::<(Uuid, String, SseEvent)>(16);
let handle = spawn_completion_watcher(job_id, event_tx.subscribe(), Arc::clone(&cm));
event_tx
.send((
job_id,
"test-user".to_string(),
SseEvent::JobResult {
job_id: job_id.to_string(),
status: "completed".to_string(),
+204 -21
View File
@@ -24,7 +24,7 @@ use crate::agent::Scheduler;
use crate::agent::routine::{
NotifyConfig, Routine, RoutineAction, RoutineRun, RunStatus, Trigger, next_cron_fire,
};
use crate::channels::OutgoingResponse;
use crate::channels::{IncomingMessage, OutgoingResponse};
use crate::config::RoutineConfig;
use crate::context::{JobContext, JobState};
use crate::db::Database;
@@ -56,6 +56,40 @@ pub enum SandboxReadiness {
DockerUnavailable,
}
/// Check whether an event-triggered routine's user/channel filters match an
/// incoming message.
///
/// Returns `true` if:
/// - The routine has an `Event` trigger (non-Event routines always return `false`)
/// - The routine's `user_id` matches the message's user scope
/// - The routine's channel filter (if any) matches the message channel
/// case-insensitively
///
/// This is a pure function extracted from `check_event_triggers` so the
/// filter logic can be unit-tested without async infrastructure.
pub(crate) fn routine_matches_message(routine: &Routine, message: &IncomingMessage) -> bool {
// Only Event-triggered routines can match incoming messages.
if !matches!(routine.trigger, Trigger::Event { .. }) {
return false;
}
// User ownership filter — only fire routines scoped to this user.
if routine.user_id != message.user_id {
return false;
}
// Channel filter (case-insensitive, matching emit_system_event behavior)
if let Trigger::Event {
channel: Some(ch), ..
} = &routine.trigger
&& !ch.eq_ignore_ascii_case(&message.channel)
{
return false;
}
true
}
/// The routine execution engine.
pub struct RoutineEngine {
config: RoutineConfig,
@@ -167,10 +201,7 @@ impl RoutineEngine {
}
/// Check incoming message against event triggers. Returns number of routines fired.
///
/// Accepts only the three fields needed for matching (user scope, channel,
/// message content) so callers never need to clone a full `IncomingMessage`.
pub async fn check_event_triggers(&self, user_id: &str, channel: &str, content: &str) -> usize {
pub async fn check_event_triggers(&self, message: &IncomingMessage, content: &str) -> usize {
let cache = self.event_cache.read().await;
// Early return if there are no message matchers at all.
@@ -208,16 +239,24 @@ impl RoutineEngine {
EventMatcher::System { .. } => continue,
};
if routine.user_id != user_id {
continue;
}
// Channel filter
if let Trigger::Event {
channel: Some(ch), ..
} = &routine.trigger
&& ch != channel
{
// User ownership + channel filter (extracted for testability).
if !routine_matches_message(routine, message) {
// User mismatch is expected for multi-user setups — keep at
// trace to avoid one log per routine per inbound message.
if routine.user_id != message.user_id {
tracing::trace!(
routine = %routine.name,
routine_user = %routine.user_id,
message_user = %message.user_id,
"Skipped: user scope mismatch"
);
} else {
tracing::debug!(
routine = %routine.name,
channel = %message.channel,
"Skipped: channel mismatch"
);
}
continue;
}
@@ -228,14 +267,14 @@ impl RoutineEngine {
// Cooldown check
if !self.check_cooldown(routine) {
tracing::trace!(routine = %routine.name, "Skipped: cooldown active");
tracing::debug!(routine = %routine.name, "Skipped: cooldown active");
continue;
}
// Concurrent run check (using batch-loaded counts)
let running_count = concurrent_counts.get(&routine.id).copied().unwrap_or(0);
if running_count >= routine.guardrails.max_concurrent as i64 {
tracing::trace!(routine = %routine.name, "Skipped: max concurrent reached");
tracing::debug!(routine = %routine.name, "Skipped: max concurrent reached");
continue;
}
@@ -1305,6 +1344,19 @@ async fn execute_lightweight(
}
}
/// Sanitize a user-controlled string before interpolation into an LLM prompt.
/// Strips newlines (which could break prompt structure) and truncates to a
/// reasonable length to limit abuse surface.
fn sanitize_prompt_field(value: &str) -> String {
const MAX_LEN: usize = 128;
value
.chars()
.filter(|&c| c != '\n' && c != '\r')
.take(MAX_LEN)
.map(|c| if c == '`' { '\'' } else { c })
.collect()
}
fn build_lightweight_prompt(
prompt: &str,
context_parts: &[String],
@@ -1323,14 +1375,16 @@ fn build_lightweight_prompt(
);
if let Some(channel) = notify.channel.as_deref() {
let sanitized = sanitize_prompt_field(channel);
full_prompt.push_str(&format!(
"The configured delivery channel for this routine is `{channel}`.\n"
"The configured delivery channel for this routine is `{sanitized}`.\n"
));
}
if let Some(user) = notify.user.as_deref() {
let sanitized = sanitize_prompt_field(user);
full_prompt.push_str(&format!(
"The configured delivery target for this routine is `{user}`.\n"
"The configured delivery target for this routine is `{sanitized}`.\n"
));
}
@@ -1440,6 +1494,7 @@ fn handle_text_response(
/// This is a simplified version of the full dispatcher loop:
/// - Max 3-5 iterations (configurable)
/// - Sequential tool execution (not parallel)
/// - Uses the owner's live autonomous tool scope when lightweight tools are enabled
/// - Auto-approval of non-Always tools
/// - No hooks or approval dialogs
async fn execute_lightweight_with_tools(
@@ -1765,6 +1820,13 @@ pub fn spawn_cron_ticker(
engine.check_cron_triggers().await;
let mut ticker = tokio::time::interval(interval);
ticker.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
// Periodic event cache refresh so web/CLI mutations are picked up
// without requiring tool-path code to call refresh_event_cache().
// Uses wall-clock elapsed time so the refresh cadence is stable
// regardless of the cron tick interval configuration.
let refresh_interval = Duration::from_secs(60);
let mut last_refresh = tokio::time::Instant::now();
loop {
ticker.tick().await;
@@ -1772,7 +1834,11 @@ pub fn spawn_cron_ticker(
// never races with FullJobWatcher instances from this process.
engine.sync_dispatched_runs().await;
engine.check_cron_triggers().await;
engine.sync_dispatched_runs().await;
if last_refresh.elapsed() >= refresh_interval {
engine.refresh_event_cache().await;
last_refresh = tokio::time::Instant::now();
}
}
})
}
@@ -1838,7 +1904,13 @@ fn strip_html_tags(s: &str) -> String {
#[cfg(test)]
mod tests {
use crate::agent::routine::{NotifyConfig, RunStatus};
use chrono::Utc;
use uuid::Uuid;
use crate::agent::routine::{
NotifyConfig, Routine, RoutineAction, RoutineGuardrails, RunStatus, Trigger,
};
use crate::channels::IncomingMessage;
use crate::config::RoutineConfig;
#[test]
@@ -2036,6 +2108,117 @@ mod tests {
}
}
/// Helper to build a test routine with the given user_id and trigger.
fn make_routine(user_id: &str, trigger: Trigger) -> Routine {
Routine {
id: Uuid::new_v4(),
name: "test".to_string(),
description: String::new(),
user_id: user_id.to_string(),
enabled: true,
trigger,
action: RoutineAction::Lightweight {
prompt: String::new(),
context_paths: vec![],
max_tokens: 1000,
use_tools: false,
max_tool_rounds: 0,
},
guardrails: RoutineGuardrails::default(),
notify: Default::default(),
last_run_at: None,
next_fire_at: None,
run_count: 0,
consecutive_failures: 0,
state: serde_json::Value::Null,
created_at: Utc::now(),
updated_at: Utc::now(),
}
}
/// Helper to build a test IncomingMessage.
fn make_message(user_id: &str, channel: &str, content: &str) -> IncomingMessage {
IncomingMessage {
id: Uuid::new_v4(),
channel: channel.to_string(),
user_id: user_id.to_string(),
owner_id: user_id.to_string(),
sender_id: user_id.to_string(),
user_name: None,
content: content.to_string(),
thread_id: None,
conversation_scope_id: None,
received_at: Utc::now(),
metadata: serde_json::Value::Null,
timezone: None,
attachments: vec![],
is_internal: false,
}
}
/// Regression test for issue #1051: event triggers used case-sensitive
/// channel comparison, so "Telegram" != "telegram" caused silent mismatch.
/// Tests the actual `routine_matches_message` function used in `check_event_triggers`.
#[test]
fn test_channel_filter_is_case_insensitive() {
let routine = make_routine(
"user1",
Trigger::Event {
pattern: ".*".to_string(),
channel: Some("Telegram".to_string()),
},
);
let msg = make_message("user1", "telegram", "hello");
// Case-insensitive channel match must succeed
assert!(super::routine_matches_message(&routine, &msg));
// Exact case must also work
let msg_exact = make_message("user1", "Telegram", "hello");
assert!(super::routine_matches_message(&routine, &msg_exact));
// Different channel must not match
let msg_wrong = make_message("user1", "discord", "hello");
assert!(!super::routine_matches_message(&routine, &msg_wrong));
}
/// Regression test for issue #1051: event triggers did not filter by
/// user_id, so routines from user A could fire on messages from user B.
/// Tests the actual `routine_matches_message` function used in `check_event_triggers`.
#[test]
fn test_event_trigger_requires_user_match() {
let routine = make_routine(
"alice",
Trigger::Event {
pattern: ".*".to_string(),
channel: None,
},
);
// Different user must not match
let msg_bob = make_message("bob", "telegram", "hello");
assert!(!super::routine_matches_message(&routine, &msg_bob));
// Same user must match
let msg_alice = make_message("alice", "telegram", "hello");
assert!(super::routine_matches_message(&routine, &msg_alice));
}
/// When no channel filter is set, any channel should match (given user matches).
#[test]
fn test_no_channel_filter_matches_any_channel() {
let routine = make_routine(
"user1",
Trigger::Event {
pattern: ".*".to_string(),
channel: None,
},
);
let msg = make_message("user1", "whatever_channel", "hello");
assert!(super::routine_matches_message(&routine, &msg));
}
#[test]
fn test_routine_tool_denylist_blocks_self_management_tools() {
let denylisted = vec![
+5 -6
View File
@@ -9,7 +9,6 @@ use tokio::task::JoinHandle;
use uuid::Uuid;
use crate::agent::task::{Task, TaskContext, TaskOutput};
use crate::channels::web::types::SseEvent;
use crate::config::AgentConfig;
use crate::context::{ContextManager, JobContext, JobState};
use crate::db::Database;
@@ -67,8 +66,8 @@ pub struct Scheduler {
extension_manager: Option<Arc<ExtensionManager>>,
store: Option<Arc<dyn Database>>,
hooks: Arc<HookRegistry>,
/// SSE broadcast sender for live job event streaming.
sse_tx: Option<tokio::sync::broadcast::Sender<SseEvent>>,
/// SSE manager for live job event streaming.
sse_tx: Option<Arc<crate::channels::web::sse::SseManager>>,
/// HTTP interceptor for trace recording/replay (propagated to workers).
http_interceptor: Option<Arc<dyn crate::llm::recording::HttpInterceptor>>,
/// Running jobs (main LLM-driven jobs).
@@ -102,9 +101,9 @@ impl Scheduler {
}
}
/// Set the SSE broadcast sender for live job event streaming.
pub fn set_sse_sender(&mut self, tx: tokio::sync::broadcast::Sender<SseEvent>) {
self.sse_tx = Some(tx);
/// Set the SSE manager for live job event streaming.
pub fn set_sse_sender(&mut self, sse: Arc<crate::channels::web::sse::SseManager>) {
self.sse_tx = Some(sse);
}
/// Set the HTTP interceptor for trace recording/replay.
+1 -1
View File
@@ -1702,7 +1702,7 @@ impl Agent {
};
match ext_mgr
.configure_token(&pending.extension_name, token)
.configure_token(&pending.extension_name, token, &message.user_id)
.await
{
Ok(result) if result.activated => {
+30 -2
View File
@@ -327,7 +327,7 @@ impl AppBuilder {
.with_search_config(&self.config.search);
if let Some(ref emb) = embeddings {
ws = ws.with_embeddings_cached(emb.clone(), emb_cache_config);
ws = ws.with_embeddings_cached(emb.clone(), emb_cache_config.clone());
}
// Wire workspace-level settings (read scopes, memory layers)
@@ -341,7 +341,35 @@ impl AppBuilder {
}
ws = ws.with_memory_layers(self.config.workspace.memory_layers.clone());
let ws = Arc::new(ws);
tools.register_memory_tools(Arc::clone(&ws));
// Detect multi-tenant mode: when GATEWAY_USER_TOKENS is configured,
// each authenticated user needs their own workspace scope. Use
// WorkspacePool (which implements WorkspaceResolver) to create
// per-user workspaces on demand instead of sharing the startup
// workspace across all users.
let is_multi_tenant = self
.config
.channels
.gateway
.as_ref()
.is_some_and(|gw| gw.user_tokens.is_some());
if is_multi_tenant {
let pool = Arc::new(crate::channels::web::server::WorkspacePool::new(
Arc::clone(db),
embeddings.clone(),
emb_cache_config,
self.config.search.clone(),
self.config.workspace.clone(),
));
tools.register_memory_tools_with_resolver(pool);
tracing::info!(
"Memory tools configured with per-user workspace resolver (multi-tenant mode)"
);
} else {
tools.register_memory_tools(Arc::clone(&ws));
}
Some(ws)
} else {
None
+9 -2
View File
@@ -333,6 +333,9 @@ async fn webhook_handler(
let channel_name = channel.channel_name();
// Track whether any authentication was performed and passed.
let mut did_authenticate = false;
// Check if secret is required
if state.router.requires_secret(channel_name).await {
// Get the secret header name for this channel (from capabilities or default)
@@ -382,6 +385,7 @@ async fn webhook_handler(
);
}
tracing::debug!(channel = %channel_name, "Webhook secret validated");
did_authenticate = true;
}
None => {
tracing::warn!(
@@ -433,6 +437,7 @@ async fn webhook_handler(
);
}
tracing::debug!(channel = %channel_name, "Ed25519 signature verified");
did_authenticate = true;
}
_ => {
tracing::warn!(
@@ -484,6 +489,7 @@ async fn webhook_handler(
);
}
tracing::debug!(channel = %channel_name, "HMAC-SHA256 signature verified");
did_authenticate = true;
}
_ => {
tracing::warn!(
@@ -510,8 +516,9 @@ async fn webhook_handler(
})
.collect();
// Call the WASM channel
let secret_validated = state.router.requires_secret(channel_name).await;
// Call the WASM channel. `did_authenticate` was set above by whichever
// auth guard (secret / Ed25519 / HMAC) successfully validated the request.
let secret_validated = did_authenticate;
tracing::info!(
channel = %channel_name,
+383 -22
View File
@@ -1,17 +1,133 @@
//! Bearer token authentication middleware for the web gateway.
//!
//! Supports multi-user mode: each token maps to a `UserIdentity` that carries
//! the user_id. The identity is inserted into request extensions so downstream
//! handlers can extract it via `AuthenticatedUser`.
use std::collections::HashMap;
use axum::{
extract::{Request, State},
http::{HeaderMap, Method, StatusCode},
extract::{FromRequestParts, Request, State},
http::{HeaderMap, Method, StatusCode, request::Parts},
middleware::Next,
response::{IntoResponse, Response},
};
use sha2::{Digest, Sha256};
use subtle::ConstantTimeEq;
/// Shared auth state injected via axum middleware state.
/// Identity resolved from a bearer token.
#[derive(Debug, Clone)]
pub struct UserIdentity {
pub user_id: String,
/// Additional user scopes this identity can read from.
pub workspace_read_scopes: Vec<String>,
}
/// Hash a token with SHA-256 for constant-size, timing-safe storage.
fn hash_token(token: &str) -> [u8; 32] {
let mut hasher = Sha256::new();
hasher.update(token.as_bytes());
hasher.finalize().into()
}
/// Multi-user auth state: maps token hashes to user identities.
///
/// Tokens are SHA-256 hashed on construction so they are never stored in
/// plaintext. Authentication compares fixed-size (32-byte) digests using
/// constant-time comparison, eliminating both length-oracle timing leaks
/// and accidental token exposure in memory dumps.
///
/// In single-user mode (the default), contains exactly one entry.
#[derive(Clone)]
pub struct AuthState {
pub token: String,
pub struct MultiAuthState {
/// Maps SHA-256(token) → identity. Tokens are never stored in cleartext.
hashed_tokens: Vec<([u8; 32], UserIdentity)>,
/// Original first token kept only for single-user startup printing.
/// Not used for authentication.
display_token: Option<String>,
}
impl MultiAuthState {
/// Create a single-user auth state (backwards compatible).
pub fn single(token: String, user_id: String) -> Self {
let hash = hash_token(&token);
Self {
hashed_tokens: vec![(
hash,
UserIdentity {
user_id,
workspace_read_scopes: Vec::new(),
},
)],
display_token: Some(token),
}
}
/// Create a multi-user auth state from a map of tokens to identities.
pub fn multi(tokens: HashMap<String, UserIdentity>) -> Self {
let hashed_tokens: Vec<([u8; 32], UserIdentity)> = tokens
.into_iter()
.map(|(tok, identity)| (hash_token(&tok), identity))
.collect();
Self {
hashed_tokens,
display_token: None,
}
}
/// Authenticate a token, returning the associated identity if valid.
///
/// Uses SHA-256 hashing + constant-time comparison (`subtle::ConstantTimeEq`)
/// to prevent timing side-channels. Both the candidate and stored tokens are
/// hashed to 32-byte digests, eliminating length-oracle leaks. Iterates all
/// entries regardless of match to avoid early-exit timing differences.
/// O(n) in the number of configured users — negligible for typical
/// deployments (< 10 users).
pub fn authenticate(&self, candidate: &str) -> Option<&UserIdentity> {
let candidate_hash = hash_token(candidate);
let mut matched: Option<&UserIdentity> = None;
for (stored_hash, identity) in &self.hashed_tokens {
if bool::from(candidate_hash.ct_eq(stored_hash)) {
matched = Some(identity);
}
}
matched
}
/// Get the first token for backwards-compatible printing at startup.
///
/// Only available in single-user mode; returns `None` in multi-user mode
/// to avoid exposing tokens.
pub fn first_token(&self) -> Option<&str> {
self.display_token.as_deref()
}
/// Get the first user identity (for single-user fallback).
pub fn first_identity(&self) -> Option<&UserIdentity> {
self.hashed_tokens.first().map(|(_, id)| id)
}
}
/// Axum extractor that provides the authenticated user identity.
///
/// Only available on routes behind `auth_middleware`. Extracts the
/// `UserIdentity` that the middleware inserted into request extensions.
pub struct AuthenticatedUser(pub UserIdentity);
impl<S> FromRequestParts<S> for AuthenticatedUser
where
S: Send + Sync,
{
type Rejection = (StatusCode, &'static str);
async fn from_request_parts(parts: &mut Parts, _state: &S) -> Result<Self, Self::Rejection> {
parts
.extensions
.get::<UserIdentity>()
.cloned()
.map(AuthenticatedUser)
.ok_or((StatusCode::UNAUTHORIZED, "Not authenticated"))
}
}
/// Whether query-string token auth is allowed for this request.
@@ -51,29 +167,34 @@ fn query_token(request: &Request) -> Option<String> {
/// Auth middleware that validates bearer token from header or query param.
///
/// SSE connections can't set headers from `EventSource`, so we also accept
/// `?token=xxx` as a query parameter, but only on SSE endpoints.
/// `?token=xxx` as a query parameter, but only on SSE/WS endpoints.
///
/// On successful authentication, inserts the matching `UserIdentity` into
/// request extensions for downstream extraction via `AuthenticatedUser`.
pub async fn auth_middleware(
State(auth): State<AuthState>,
State(auth): State<MultiAuthState>,
headers: HeaderMap,
request: Request,
mut request: Request,
next: Next,
) -> Response {
// Try Authorization header first (constant-time comparison).
// Try Authorization header first.
// RFC 6750 Section 2.1: auth-scheme comparison is case-insensitive.
if let Some(auth_header) = headers.get("authorization")
&& let Ok(value) = auth_header.to_str()
&& value.len() > 7
&& value[..7].eq_ignore_ascii_case("Bearer ")
&& bool::from(value.as_bytes()[7..].ct_eq(auth.token.as_bytes()))
&& let Some(identity) = auth.authenticate(&value[7..])
{
request.extensions_mut().insert(identity.clone());
return next.run(request).await;
}
// Fall back to query parameter, but only for SSE endpoints (constant-time comparison).
// Fall back to query parameter, but only for SSE/WS endpoints.
if allows_query_token_auth(&request)
&& let Some(token) = query_token(&request)
&& bool::from(token.as_bytes().ct_eq(auth.token.as_bytes()))
&& let Some(identity) = auth.authenticate(&token)
{
request.extensions_mut().insert(identity.clone());
return next.run(request).await;
}
@@ -83,15 +204,61 @@ pub async fn auth_middleware(
#[cfg(test)]
mod tests {
use super::*;
use crate::testing::credentials::{TEST_AUTH_SECRET_TOKEN, TEST_BEARER_TOKEN};
use crate::testing::credentials::TEST_AUTH_SECRET_TOKEN;
#[test]
fn test_auth_state_clone() {
let state = AuthState {
token: TEST_BEARER_TOKEN.to_string(),
};
let cloned = state.clone();
assert_eq!(cloned.token, TEST_BEARER_TOKEN);
fn test_multi_auth_state_single() {
let state = MultiAuthState::single("tok-123".to_string(), "alice".to_string());
let identity = state.authenticate("tok-123");
assert!(identity.is_some());
assert_eq!(identity.unwrap().user_id, "alice");
}
#[test]
fn test_multi_auth_state_reject_wrong_token() {
let state = MultiAuthState::single("tok-123".to_string(), "alice".to_string());
assert!(state.authenticate("wrong-token").is_none());
}
#[test]
fn test_multi_auth_state_multi_users() {
let mut tokens = HashMap::new();
tokens.insert(
"tok-alice".to_string(),
UserIdentity {
user_id: "alice".to_string(),
workspace_read_scopes: Vec::new(),
},
);
tokens.insert(
"tok-bob".to_string(),
UserIdentity {
user_id: "bob".to_string(),
workspace_read_scopes: Vec::new(),
},
);
let state = MultiAuthState::multi(tokens);
let alice = state.authenticate("tok-alice").unwrap();
assert_eq!(alice.user_id, "alice");
let bob = state.authenticate("tok-bob").unwrap();
assert_eq!(bob.user_id, "bob");
assert!(state.authenticate("tok-charlie").is_none());
}
#[test]
fn test_multi_auth_state_first_token() {
let state = MultiAuthState::single("my-token".to_string(), "user1".to_string());
assert_eq!(state.first_token(), Some("my-token"));
}
#[test]
fn test_multi_auth_state_first_identity() {
let state = MultiAuthState::single("my-token".to_string(), "user1".to_string());
let identity = state.first_identity().unwrap();
assert_eq!(identity.user_id, "user1");
}
use axum::Router;
@@ -107,9 +274,7 @@ mod tests {
/// Router with streaming endpoints (query auth allowed) and regular
/// endpoints (query auth rejected).
fn test_app(token: &str) -> Router {
let state = AuthState {
token: token.to_string(),
};
let state = MultiAuthState::single(token.to_string(), "test-user".to_string());
Router::new()
.route("/api/chat/events", get(dummy_handler))
.route("/api/logs/events", get(dummy_handler))
@@ -306,4 +471,200 @@ mod tests {
let resp = app.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
}
// --- Multi-tenant auth integration tests ---
/// Handler that extracts `AuthenticatedUser` and returns the resolved user_id.
async fn identity_handler(AuthenticatedUser(identity): AuthenticatedUser) -> String {
identity.user_id
}
/// Handler that extracts `AuthenticatedUser` and returns workspace_read_scopes as JSON.
async fn scopes_handler(AuthenticatedUser(identity): AuthenticatedUser) -> String {
serde_json::to_string(&identity.workspace_read_scopes).unwrap()
}
/// Build a multi-user router where each token maps to a distinct identity.
fn multi_user_app(tokens: HashMap<String, UserIdentity>) -> Router {
let state = MultiAuthState::multi(tokens);
Router::new()
.route("/api/chat/events", get(identity_handler))
.route("/api/chat/send", post(identity_handler))
.route("/api/scopes", get(scopes_handler))
.layer(middleware::from_fn_with_state(state, auth_middleware))
}
fn two_user_tokens() -> HashMap<String, UserIdentity> {
let mut tokens = HashMap::new();
tokens.insert(
"tok-alice".to_string(),
UserIdentity {
user_id: "alice".to_string(),
workspace_read_scopes: vec!["shared".to_string()],
},
);
tokens.insert(
"tok-bob".to_string(),
UserIdentity {
user_id: "bob".to_string(),
workspace_read_scopes: vec!["shared".to_string(), "alice".to_string()],
},
);
tokens
}
#[tokio::test]
async fn test_multi_user_alice_token_resolves_to_alice() {
let app = multi_user_app(two_user_tokens());
let req = Request::builder()
.uri("/api/chat/events")
.header("Authorization", "Bearer tok-alice")
.body(Body::empty())
.unwrap();
let resp = app.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
let body = axum::body::to_bytes(resp.into_body(), 1024).await.unwrap();
assert_eq!(body, "alice");
}
#[tokio::test]
async fn test_multi_user_bob_token_resolves_to_bob() {
let app = multi_user_app(two_user_tokens());
let req = Request::builder()
.uri("/api/chat/events")
.header("Authorization", "Bearer tok-bob")
.body(Body::empty())
.unwrap();
let resp = app.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
let body = axum::body::to_bytes(resp.into_body(), 1024).await.unwrap();
assert_eq!(body, "bob");
}
#[tokio::test]
async fn test_multi_user_sequential_tokens_resolve_independently() {
// Send both alice and bob tokens sequentially and verify each gets
// the correct identity — guards against token map corruption.
let tokens = two_user_tokens();
let app1 = multi_user_app(tokens.clone());
let req = Request::builder()
.uri("/api/chat/events")
.header("Authorization", "Bearer tok-alice")
.body(Body::empty())
.unwrap();
let resp = app1.oneshot(req).await.unwrap();
let body = axum::body::to_bytes(resp.into_body(), 1024).await.unwrap();
assert_eq!(body, "alice");
let app2 = multi_user_app(tokens);
let req = Request::builder()
.uri("/api/chat/events")
.header("Authorization", "Bearer tok-bob")
.body(Body::empty())
.unwrap();
let resp = app2.oneshot(req).await.unwrap();
let body = axum::body::to_bytes(resp.into_body(), 1024).await.unwrap();
assert_eq!(body, "bob");
}
#[tokio::test]
async fn test_multi_user_unknown_token_rejected() {
let app = multi_user_app(two_user_tokens());
let req = Request::builder()
.uri("/api/chat/events")
.header("Authorization", "Bearer tok-charlie")
.body(Body::empty())
.unwrap();
let resp = app.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
}
#[tokio::test]
async fn test_multi_user_workspace_read_scopes_propagated() {
let app = multi_user_app(two_user_tokens());
// Alice has ["shared"]
let req = Request::builder()
.uri("/api/scopes")
.header("Authorization", "Bearer tok-alice")
.body(Body::empty())
.unwrap();
let resp = app.oneshot(req).await.unwrap();
let body = axum::body::to_bytes(resp.into_body(), 1024).await.unwrap();
let scopes: Vec<String> = serde_json::from_slice(&body).unwrap();
assert_eq!(scopes, vec!["shared"]);
}
#[tokio::test]
async fn test_multi_user_bob_has_two_scopes() {
let app = multi_user_app(two_user_tokens());
// Bob has ["shared", "alice"]
let req = Request::builder()
.uri("/api/scopes")
.header("Authorization", "Bearer tok-bob")
.body(Body::empty())
.unwrap();
let resp = app.oneshot(req).await.unwrap();
let body = axum::body::to_bytes(resp.into_body(), 1024).await.unwrap();
let scopes: Vec<String> = serde_json::from_slice(&body).unwrap();
assert_eq!(scopes, vec!["shared", "alice"]);
}
#[tokio::test]
async fn test_multi_user_query_param_resolves_correct_identity() {
let app = multi_user_app(two_user_tokens());
let req = Request::builder()
.uri("/api/chat/events?token=tok-bob")
.body(Body::empty())
.unwrap();
let resp = app.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
let body = axum::body::to_bytes(resp.into_body(), 1024).await.unwrap();
assert_eq!(body, "bob");
}
#[tokio::test]
async fn test_multi_user_post_with_bearer_resolves_identity() {
let app = multi_user_app(two_user_tokens());
let req = Request::builder()
.method(Method::POST)
.uri("/api/chat/send")
.header("Authorization", "Bearer tok-alice")
.body(Body::empty())
.unwrap();
let resp = app.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
let body = axum::body::to_bytes(resp.into_body(), 1024).await.unwrap();
assert_eq!(body, "alice");
}
#[tokio::test]
async fn test_multi_user_empty_scopes_for_single_user() {
// Single-user mode creates identity with empty workspace_read_scopes.
let state = MultiAuthState::single("tok-only".to_string(), "solo".to_string());
let app = Router::new()
.route("/api/scopes", get(scopes_handler))
.layer(middleware::from_fn_with_state(state, auth_middleware));
let req = Request::builder()
.uri("/api/scopes")
.header("Authorization", "Bearer tok-only")
.body(Body::empty())
.unwrap();
let resp = app.oneshot(req).await.unwrap();
let body = axum::body::to_bytes(resp.into_body(), 1024).await.unwrap();
let scopes: Vec<String> = serde_json::from_slice(&body).unwrap();
assert!(scopes.is_empty());
}
#[tokio::test]
async fn test_prefix_and_extension_tokens_rejected() {
// Verifies that prefix/suffix variants of valid tokens are rejected.
// Note: the constant-time property is enforced structurally by use of
// subtle::ConstantTimeEq and cannot be verified via outcome testing.
let state = MultiAuthState::single("long-secret-token".to_string(), "user".to_string());
assert!(state.authenticate("long-secret").is_none());
assert!(state.authenticate("long-secret-token-extra").is_none());
}
}
+62 -35
View File
@@ -12,22 +12,24 @@ use serde::Deserialize;
use uuid::Uuid;
use crate::channels::IncomingMessage;
use crate::channels::web::auth::AuthenticatedUser;
use crate::channels::web::server::GatewayState;
use crate::channels::web::types::*;
use crate::channels::web::util::{build_turns_from_db_messages, truncate_preview};
pub async fn chat_send_handler(
State(state): State<Arc<GatewayState>>,
AuthenticatedUser(identity): AuthenticatedUser,
Json(req): Json<SendMessageRequest>,
) -> Result<(StatusCode, Json<SendMessageResponse>), (StatusCode, String)> {
if !state.chat_rate_limiter.check() {
if !state.chat_rate_limiter.check(&identity.user_id) {
return Err((
StatusCode::TOO_MANY_REQUESTS,
"Rate limit exceeded. Try again shortly.".to_string(),
));
}
let mut msg = IncomingMessage::new("gateway", &state.user_id, &req.content);
let mut msg = IncomingMessage::new("gateway", &identity.user_id, &req.content);
if let Some(ref thread_id) = req.thread_id {
msg = msg.with_thread(thread_id);
@@ -74,6 +76,7 @@ pub async fn chat_send_handler(
pub async fn chat_approval_handler(
State(state): State<Arc<GatewayState>>,
AuthenticatedUser(identity): AuthenticatedUser,
Json(req): Json<ApprovalRequest>,
) -> Result<(StatusCode, Json<SendMessageResponse>), (StatusCode, String)> {
let (approved, always) = match req.action.as_str() {
@@ -109,7 +112,7 @@ pub async fn chat_approval_handler(
)
})?;
let mut msg = IncomingMessage::new("gateway", &state.user_id, content);
let mut msg = IncomingMessage::new("gateway", &identity.user_id, content);
if let Some(ref thread_id) = req.thread_id {
msg = msg.with_thread(thread_id);
@@ -150,6 +153,7 @@ pub async fn chat_approval_handler(
/// The token never touches the LLM, chat history, or SSE stream.
pub async fn chat_auth_token_handler(
State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Json(req): Json<AuthTokenRequest>,
) -> Result<Json<ActionResponse>, (StatusCode, String)> {
let ext_mgr = state.extension_manager.as_ref().ok_or((
@@ -158,7 +162,7 @@ pub async fn chat_auth_token_handler(
))?;
match ext_mgr
.configure_token(&req.extension_name, &req.token)
.configure_token(&req.extension_name, &req.token, &user.user_id)
.await
{
Ok(result) => {
@@ -169,20 +173,26 @@ pub async fn chat_auth_token_handler(
resp.instructions = result.verification.as_ref().map(|v| v.instructions.clone());
if result.verification.is_some() {
state.sse.broadcast(SseEvent::AuthRequired {
extension_name: req.extension_name.clone(),
instructions: Some(result.message),
auth_url: None,
setup_url: None,
});
state.sse.broadcast_for_user(
&user.user_id,
SseEvent::AuthRequired {
extension_name: req.extension_name.clone(),
instructions: Some(result.message),
auth_url: None,
setup_url: None,
},
);
} else {
clear_auth_mode(&state).await;
clear_auth_mode(&state, &user.user_id).await;
state.sse.broadcast(SseEvent::AuthCompleted {
extension_name: req.extension_name.clone(),
success: true,
message: result.message,
});
state.sse.broadcast_for_user(
&user.user_id,
SseEvent::AuthCompleted {
extension_name: req.extension_name.clone(),
success: true,
message: result.message,
},
);
}
Ok(Json(resp))
@@ -190,12 +200,15 @@ pub async fn chat_auth_token_handler(
Err(e) => {
let msg = e.to_string();
if matches!(e, crate::extensions::ExtensionError::ValidationFailed(_)) {
state.sse.broadcast(SseEvent::AuthRequired {
extension_name: req.extension_name.clone(),
instructions: Some(msg.clone()),
auth_url: None,
setup_url: None,
});
state.sse.broadcast_for_user(
&user.user_id,
SseEvent::AuthRequired {
extension_name: req.extension_name.clone(),
instructions: Some(msg.clone()),
auth_url: None,
setup_url: None,
},
);
}
Ok(Json(ActionResponse::fail(msg)))
}
@@ -205,16 +218,17 @@ pub async fn chat_auth_token_handler(
/// Cancel an in-progress auth flow.
pub async fn chat_auth_cancel_handler(
State(state): State<Arc<GatewayState>>,
AuthenticatedUser(identity): AuthenticatedUser,
Json(_req): Json<AuthCancelRequest>,
) -> Result<Json<ActionResponse>, (StatusCode, String)> {
clear_auth_mode(&state).await;
clear_auth_mode(&state, &identity.user_id).await;
Ok(Json(ActionResponse::ok("Auth cancelled")))
}
/// Clear pending auth mode on the active thread.
pub async fn clear_auth_mode(state: &GatewayState) {
pub async fn clear_auth_mode(state: &GatewayState, user_id: &str) {
if let Some(ref sm) = state.session_manager {
let session = sm.get_or_create_session(&state.user_id).await;
let session = sm.get_or_create_session(user_id).await;
let mut sess = session.lock().await;
if let Some(thread_id) = sess.active_thread
&& let Some(thread) = sess.threads.get_mut(&thread_id)
@@ -226,8 +240,9 @@ pub async fn clear_auth_mode(state: &GatewayState) {
pub async fn chat_events_handler(
State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
) -> Result<impl IntoResponse, (StatusCode, String)> {
state.sse.subscribe().ok_or((
state.sse.subscribe(Some(user.user_id)).ok_or((
StatusCode::SERVICE_UNAVAILABLE,
"Too many connections".to_string(),
))
@@ -237,6 +252,7 @@ pub async fn chat_ws_handler(
headers: axum::http::HeaderMap,
ws: WebSocketUpgrade,
State(state): State<Arc<GatewayState>>,
AuthenticatedUser(identity): AuthenticatedUser,
) -> Result<impl IntoResponse, (StatusCode, String)> {
// Validate Origin header to prevent cross-site WebSocket hijacking.
let origin = headers
@@ -262,7 +278,9 @@ pub async fn chat_ws_handler(
"WebSocket origin not allowed".to_string(),
));
}
Ok(ws.on_upgrade(move |socket| crate::channels::web::ws::handle_ws_connection(socket, state)))
Ok(ws.on_upgrade(move |socket| {
crate::channels::web::ws::handle_ws_connection(socket, state, identity)
}))
}
#[derive(Deserialize)]
@@ -274,6 +292,7 @@ pub struct HistoryQuery {
pub async fn chat_history_handler(
State(state): State<Arc<GatewayState>>,
AuthenticatedUser(identity): AuthenticatedUser,
Query(query): Query<HistoryQuery>,
) -> Result<Json<HistoryResponse>, (StatusCode, String)> {
let session_manager = state.session_manager.as_ref().ok_or((
@@ -281,7 +300,9 @@ pub async fn chat_history_handler(
"Session manager not available".to_string(),
))?;
let session = session_manager.get_or_create_session(&state.user_id).await;
let session = session_manager
.get_or_create_session(&identity.user_id)
.await;
let limit = query.limit.unwrap_or(50);
let before_cursor = query
@@ -314,7 +335,7 @@ pub async fn chat_history_handler(
&& let Some(ref store) = state.store
{
let owned = store
.conversation_belongs_to_user(thread_id, &state.user_id)
.conversation_belongs_to_user(thread_id, &identity.user_id)
.await
.unwrap_or(false);
if !owned {
@@ -434,24 +455,27 @@ pub async fn chat_history_handler(
pub async fn chat_threads_handler(
State(state): State<Arc<GatewayState>>,
AuthenticatedUser(identity): AuthenticatedUser,
) -> Result<Json<ThreadListResponse>, (StatusCode, String)> {
let session_manager = state.session_manager.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE,
"Session manager not available".to_string(),
))?;
let session = session_manager.get_or_create_session(&state.user_id).await;
let session = session_manager
.get_or_create_session(&identity.user_id)
.await;
// Try DB first for persistent thread list
if let Some(ref store) = state.store {
// Auto-create assistant thread if it doesn't exist
let assistant_id = store
.get_or_create_assistant_conversation(&state.user_id, "gateway")
.get_or_create_assistant_conversation(&identity.user_id, "gateway")
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
if let Ok(summaries) = store
.list_conversations_all_channels(&state.user_id, 50)
.list_conversations_all_channels(&identity.user_id, 50)
.await
{
let mut assistant_thread = None;
@@ -534,13 +558,16 @@ pub async fn chat_threads_handler(
pub async fn chat_new_thread_handler(
State(state): State<Arc<GatewayState>>,
AuthenticatedUser(identity): AuthenticatedUser,
) -> Result<Json<ThreadInfo>, (StatusCode, String)> {
let session_manager = state.session_manager.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE,
"Session manager not available".to_string(),
))?;
let session = session_manager.get_or_create_session(&state.user_id).await;
let session = session_manager
.get_or_create_session(&identity.user_id)
.await;
let (thread_id, info) = {
let mut sess = session.lock().await;
let thread = sess.create_thread();
@@ -562,12 +589,12 @@ pub async fn chat_new_thread_handler(
// so that the subsequent loadThreads() call from the frontend sees it.
if let Some(ref store) = state.store {
match store
.ensure_conversation(thread_id, "gateway", &state.user_id, None)
.ensure_conversation(thread_id, "gateway", &identity.user_id, None)
.await
{
Ok(true) => {}
Ok(false) => tracing::warn!(
user = %state.user_id,
user = %identity.user_id,
thread_id = %thread_id,
"Skipped persisting new thread due to ownership/channel conflict"
),
+8 -3
View File
@@ -8,11 +8,13 @@ use axum::{
http::StatusCode,
};
use crate::channels::web::auth::AuthenticatedUser;
use crate::channels::web::server::GatewayState;
use crate::channels::web::types::*;
pub async fn extensions_list_handler(
State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
) -> Result<Json<ExtensionListResponse>, (StatusCode, String)> {
let ext_mgr = state.extension_manager.as_ref().ok_or((
StatusCode::NOT_IMPLEMENTED,
@@ -20,7 +22,7 @@ pub async fn extensions_list_handler(
))?;
let installed = ext_mgr
.list(None, false)
.list(None, false, &user.user_id)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
@@ -80,6 +82,7 @@ pub async fn extensions_list_handler(
pub async fn extensions_tools_handler(
State(state): State<Arc<GatewayState>>,
AuthenticatedUser(_user): AuthenticatedUser,
) -> Result<Json<ToolListResponse>, (StatusCode, String)> {
let registry = state.tool_registry.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE,
@@ -100,6 +103,7 @@ pub async fn extensions_tools_handler(
pub async fn extensions_install_handler(
State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Json(req): Json<InstallExtensionRequest>,
) -> Result<Json<ActionResponse>, (StatusCode, String)> {
let ext_mgr = state.extension_manager.as_ref().ok_or((
@@ -116,7 +120,7 @@ pub async fn extensions_install_handler(
});
match ext_mgr
.install(&req.name, req.url.as_deref(), kind_hint)
.install(&req.name, req.url.as_deref(), kind_hint, &user.user_id)
.await
{
Ok(result) => Ok(Json(ActionResponse::ok(result.message))),
@@ -126,6 +130,7 @@ pub async fn extensions_install_handler(
pub async fn extensions_remove_handler(
State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Path(name): Path<String>,
) -> Result<Json<ActionResponse>, (StatusCode, String)> {
let ext_mgr = state.extension_manager.as_ref().ok_or((
@@ -133,7 +138,7 @@ pub async fn extensions_remove_handler(
"Extension manager not available (secrets store required)".to_string(),
))?;
match ext_mgr.remove(&name).await {
match ext_mgr.remove(&name, &user.user_id).await {
Ok(message) => Ok(Json(ActionResponse::ok(message))),
Err(e) => Ok(Json(ActionResponse::fail(e.to_string()))),
}
+400 -277
View File
@@ -11,11 +11,13 @@ use axum::{
use serde::Deserialize;
use uuid::Uuid;
use crate::channels::web::auth::AuthenticatedUser;
use crate::channels::web::server::GatewayState;
use crate::channels::web::types::*;
pub async fn jobs_list_handler(
State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
) -> Result<Json<JobListResponse>, (StatusCode, String)> {
let store = state.store.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE,
@@ -25,8 +27,8 @@ pub async fn jobs_list_handler(
let mut jobs: Vec<JobInfo> = Vec::new();
let mut seen_ids: HashSet<Uuid> = HashSet::new();
// Fetch sandbox jobs from database.
match store.list_sandbox_jobs().await {
// Fetch sandbox jobs scoped to this user.
match store.list_sandbox_jobs_for_user(&user.user_id).await {
Ok(sandbox_jobs) => {
for j in &sandbox_jobs {
let ui_state = match j.status.as_str() {
@@ -50,8 +52,8 @@ pub async fn jobs_list_handler(
}
}
// Fetch agent (non-sandbox) jobs from database, deduplicating by ID.
match store.list_agent_jobs().await {
// Fetch agent (non-sandbox) jobs scoped to this user, deduplicating by ID.
match store.list_agent_jobs_for_user(&user.user_id).await {
Ok(agent_jobs) => {
for j in &agent_jobs {
if seen_ids.contains(&j.id) {
@@ -80,6 +82,7 @@ pub async fn jobs_list_handler(
pub async fn jobs_summary_handler(
State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
) -> Result<Json<JobSummaryResponse>, (StatusCode, String)> {
let store = state.store.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE,
@@ -93,8 +96,8 @@ pub async fn jobs_summary_handler(
let mut failed = 0;
let mut stuck = 0;
// Sandbox job counts.
match store.sandbox_job_summary().await {
// Sandbox job counts scoped to this user.
match store.sandbox_job_summary_for_user(&user.user_id).await {
Ok(s) => {
total += s.total;
pending += s.creating;
@@ -107,8 +110,8 @@ pub async fn jobs_summary_handler(
}
}
// Agent job counts.
match store.agent_job_summary().await {
// Agent job counts scoped to this user.
match store.agent_job_summary_for_user(&user.user_id).await {
Ok(s) => {
total += s.total;
pending += s.pending;
@@ -134,6 +137,7 @@ pub async fn jobs_summary_handler(
pub async fn jobs_detail_handler(
State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Path(id): Path<String>,
) -> Result<Json<JobDetailResponse>, (StatusCode, String)> {
let store = state.store.as_ref().ok_or((
@@ -145,169 +149,213 @@ pub async fn jobs_detail_handler(
.map_err(|_| (StatusCode::BAD_REQUEST, "Invalid job ID".to_string()))?;
// Try sandbox job from DB first.
if let Ok(Some(job)) = store.get_sandbox_job(job_id).await {
let browse_id = std::path::Path::new(&job.project_dir)
.file_name()
.map(|n| n.to_string_lossy().to_string())
.unwrap_or_else(|| job.id.to_string());
match store.get_sandbox_job(job_id).await {
Ok(Some(job)) => {
if job.user_id != user.user_id {
return Err((StatusCode::NOT_FOUND, "Job not found".to_string()));
}
let browse_id = std::path::Path::new(&job.project_dir)
.file_name()
.map(|n| n.to_string_lossy().to_string())
.unwrap_or_else(|| job.id.to_string());
let ui_state = match job.status.as_str() {
"creating" => "pending",
"running" => "in_progress",
s => s,
};
let ui_state = match job.status.as_str() {
"creating" => "pending",
"running" => "in_progress",
s => s,
};
let elapsed_secs = job.started_at.map(|start| {
let end = job.completed_at.unwrap_or_else(chrono::Utc::now);
(end - start).num_seconds().max(0) as u64
});
// Synthesize transitions from timestamps.
let mut transitions = Vec::new();
if let Some(started) = job.started_at {
transitions.push(TransitionInfo {
from: "creating".to_string(),
to: "running".to_string(),
timestamp: started.to_rfc3339(),
reason: None,
let elapsed_secs = job.started_at.map(|start| {
let end = job.completed_at.unwrap_or_else(chrono::Utc::now);
(end - start).num_seconds().max(0) as u64
});
}
if let Some(completed) = job.completed_at {
transitions.push(TransitionInfo {
from: "running".to_string(),
to: job.status.clone(),
timestamp: completed.to_rfc3339(),
reason: job.failure_reason.clone(),
});
}
let mode = store.get_sandbox_job_mode(job.id).await.ok().flatten();
let is_claude_code = mode.as_deref() == Some("claude_code");
// Synthesize transitions from timestamps.
let mut transitions = Vec::new();
if let Some(started) = job.started_at {
transitions.push(TransitionInfo {
from: "creating".to_string(),
to: "running".to_string(),
timestamp: started.to_rfc3339(),
reason: None,
});
}
if let Some(completed) = job.completed_at {
transitions.push(TransitionInfo {
from: "running".to_string(),
to: job.status.clone(),
timestamp: completed.to_rfc3339(),
reason: job.failure_reason.clone(),
});
}
return Ok(Json(JobDetailResponse {
id: job.id,
title: job.task.clone(),
description: String::new(),
state: ui_state.to_string(),
user_id: job.user_id.clone(),
created_at: job.created_at.to_rfc3339(),
started_at: job.started_at.map(|dt| dt.to_rfc3339()),
completed_at: job.completed_at.map(|dt| dt.to_rfc3339()),
elapsed_secs,
project_dir: Some(job.project_dir.clone()),
browse_url: Some(format!("/projects/{}/", browse_id)),
job_mode: mode.filter(|m| m != "worker"),
transitions,
can_restart: state.job_manager.is_some(),
can_prompt: is_claude_code && state.prompt_queue.is_some(),
job_kind: Some("sandbox".to_string()),
}));
let mode = store.get_sandbox_job_mode(job.id).await.ok().flatten();
let is_claude_code = mode.as_deref() == Some("claude_code");
return Ok(Json(JobDetailResponse {
id: job.id,
title: job.task.clone(),
description: String::new(),
state: ui_state.to_string(),
user_id: job.user_id.clone(),
created_at: job.created_at.to_rfc3339(),
started_at: job.started_at.map(|dt| dt.to_rfc3339()),
completed_at: job.completed_at.map(|dt| dt.to_rfc3339()),
elapsed_secs,
project_dir: Some(job.project_dir.clone()),
browse_url: Some(format!("/projects/{}/", browse_id)),
job_mode: mode.filter(|m| m != "worker"),
transitions,
can_restart: state.job_manager.is_some(),
can_prompt: is_claude_code && state.prompt_queue.is_some(),
job_kind: Some("sandbox".to_string()),
}));
}
Ok(None) => {}
Err(e) => {
return Err((
StatusCode::INTERNAL_SERVER_ERROR,
format!("Database error: {}", e),
));
}
}
// Fall back to agent job from DB.
if let Ok(Some(ctx)) = store.get_job(job_id).await {
let elapsed_secs = ctx.started_at.map(|start| {
let end = ctx.completed_at.unwrap_or_else(chrono::Utc::now);
(end - start).num_seconds().max(0) as u64
});
match store.get_job(job_id).await {
Ok(Some(ctx)) => {
if ctx.user_id != user.user_id {
return Err((StatusCode::NOT_FOUND, "Job not found".to_string()));
}
let elapsed_secs = ctx.started_at.map(|start| {
let end = ctx.completed_at.unwrap_or_else(chrono::Utc::now);
(end - start).num_seconds().max(0) as u64
});
// Only show prompt bar for jobs that have a running worker (Pending/InProgress).
// Stuck jobs have no active worker loop, so messages would be silently dropped.
let is_promptable = matches!(
ctx.state,
crate::context::JobState::Pending | crate::context::JobState::InProgress
);
return Ok(Json(JobDetailResponse {
id: ctx.job_id,
title: ctx.title.clone(),
description: ctx.description.clone(),
state: ctx.state.to_string(),
user_id: ctx.user_id.clone(),
created_at: ctx.created_at.to_rfc3339(),
started_at: ctx.started_at.map(|dt| dt.to_rfc3339()),
completed_at: ctx.completed_at.map(|dt| dt.to_rfc3339()),
elapsed_secs,
project_dir: None,
browse_url: None,
job_mode: None,
transitions: Vec::new(),
can_restart: state.scheduler.is_some(),
can_prompt: is_promptable && state.scheduler.is_some(),
job_kind: Some("agent".to_string()),
}));
// Only show prompt bar for jobs that have a running worker (Pending/InProgress).
// Stuck jobs have no active worker loop, so messages would be silently dropped.
let is_promptable = matches!(
ctx.state,
crate::context::JobState::Pending | crate::context::JobState::InProgress
);
Ok(Json(JobDetailResponse {
id: ctx.job_id,
title: ctx.title.clone(),
description: ctx.description.clone(),
state: ctx.state.to_string(),
user_id: ctx.user_id.clone(),
created_at: ctx.created_at.to_rfc3339(),
started_at: ctx.started_at.map(|dt| dt.to_rfc3339()),
completed_at: ctx.completed_at.map(|dt| dt.to_rfc3339()),
elapsed_secs,
project_dir: None,
browse_url: None,
job_mode: None,
transitions: Vec::new(),
can_restart: state.scheduler.is_some(),
can_prompt: is_promptable && state.scheduler.is_some(),
job_kind: Some("agent".to_string()),
}))
}
Ok(None) => Err((StatusCode::NOT_FOUND, "Job not found".to_string())),
Err(e) => Err((
StatusCode::INTERNAL_SERVER_ERROR,
format!("Database error: {}", e),
)),
}
Err((StatusCode::NOT_FOUND, "Job not found".to_string()))
}
pub async fn jobs_cancel_handler(
State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Path(id): Path<String>,
) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
let job_id = Uuid::parse_str(&id)
.map_err(|_| (StatusCode::BAD_REQUEST, "Invalid job ID".to_string()))?;
// Try sandbox job cancellation.
if let Some(ref store) = state.store
&& let Ok(Some(job)) = store.get_sandbox_job(job_id).await
{
if job.status == "running" || job.status == "creating" {
// Stop the container if we have a job manager.
if let Some(ref jm) = state.job_manager
&& let Err(e) = jm.stop_job(job_id).await
{
tracing::warn!(job_id = %job_id, error = %e, "Failed to stop container during cancellation");
if let Some(ref store) = state.store {
match store.get_sandbox_job(job_id).await {
Ok(Some(job)) => {
if job.user_id != user.user_id {
return Err((StatusCode::NOT_FOUND, "Job not found".to_string()));
}
if job.status == "running" || job.status == "creating" {
if let Some(ref jm) = state.job_manager
&& let Err(e) = jm.stop_job(job_id).await
{
tracing::warn!(job_id = %job_id, error = %e, "Failed to stop container during cancellation");
}
store
.update_sandbox_job_status(
job_id,
"failed",
Some(false),
Some("Cancelled by user"),
None,
Some(chrono::Utc::now()),
)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
}
return Ok(Json(serde_json::json!({
"status": "cancelled",
"job_id": job_id,
})));
}
Ok(None) => {}
Err(e) => {
return Err((
StatusCode::INTERNAL_SERVER_ERROR,
format!("Database error: {}", e),
));
}
store
.update_sandbox_job_status(
job_id,
"failed",
Some(false),
Some("Cancelled by user"),
None,
Some(chrono::Utc::now()),
)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
}
return Ok(Json(serde_json::json!({
"status": "cancelled",
"job_id": job_id,
})));
}
// Fall back to agent job cancellation: stop the worker via the scheduler
// (which updates the in-memory ContextManager AND aborts the task handle),
// then persist the status to the DB as a fallback.
if let Some(ref store) = state.store
&& let Ok(Some(job)) = store.get_job(job_id).await
{
if job.state.is_active() {
// Try to stop via scheduler (aborts the worker task + updates
// in-memory ContextManager). This is best-effort — the job may
// not be in the scheduler map if it already finished.
if let Some(ref slot) = state.scheduler
&& let Some(ref scheduler) = *slot.read().await
{
let _ = scheduler.stop(job_id).await;
}
if let Some(ref store) = state.store {
match store.get_job(job_id).await {
Ok(Some(job)) => {
if job.user_id != user.user_id {
return Err((StatusCode::NOT_FOUND, "Job not found".to_string()));
}
if job.state.is_active() {
// Try to stop via scheduler (aborts the worker task + updates
// in-memory ContextManager). This is best-effort — the job may
// not be in the scheduler map if it already finished.
if let Some(ref slot) = state.scheduler
&& let Some(ref scheduler) = *slot.read().await
{
let _ = scheduler.stop(job_id).await;
}
// Always persist cancellation to the DB so the state is
// consistent even if the scheduler wasn't available or the
// job wasn't in its in-memory map.
store
.update_job_status(
job_id,
crate::context::JobState::Cancelled,
Some("Cancelled by user"),
)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
// Always persist cancellation to the DB so the state is
// consistent even if the scheduler wasn't available or the
// job wasn't in its in-memory map.
store
.update_job_status(
job_id,
crate::context::JobState::Cancelled,
Some("Cancelled by user"),
)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
}
return Ok(Json(serde_json::json!({
"status": "cancelled",
"job_id": job_id,
})));
}
Ok(None) => {}
Err(e) => {
return Err((
StatusCode::INTERNAL_SERVER_ERROR,
format!("Database error: {}", e),
));
}
}
return Ok(Json(serde_json::json!({
"status": "cancelled",
"job_id": job_id,
})));
}
Err((StatusCode::NOT_FOUND, "Job not found".to_string()))
@@ -315,6 +363,7 @@ pub async fn jobs_cancel_handler(
pub async fn jobs_restart_handler(
State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Path(id): Path<String>,
) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
let store = state.store.as_ref().ok_or((
@@ -326,146 +375,166 @@ pub async fn jobs_restart_handler(
.map_err(|_| (StatusCode::BAD_REQUEST, "Invalid job ID".to_string()))?;
// Try sandbox job restart first.
if let Ok(Some(old_job)) = store.get_sandbox_job(old_job_id).await {
if old_job.status != "interrupted" && old_job.status != "failed" {
match store.get_sandbox_job(old_job_id).await {
Ok(Some(old_job)) => {
if old_job.user_id != user.user_id {
return Err((StatusCode::NOT_FOUND, "Job not found".to_string()));
}
if old_job.status != "interrupted" && old_job.status != "failed" {
return Err((
StatusCode::CONFLICT,
format!("Cannot restart job in state '{}'", old_job.status),
));
}
let jm = state.job_manager.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE,
"Sandbox not enabled".to_string(),
))?;
// Enrich the task with failure context.
let task = if let Some(ref reason) = old_job.failure_reason {
format!(
"Previous attempt failed: {}. Retry: {}",
reason, old_job.task
)
} else {
old_job.task.clone()
};
let new_job_id = Uuid::new_v4();
let now = chrono::Utc::now();
let record = crate::history::SandboxJobRecord {
id: new_job_id,
task: task.clone(),
status: "creating".to_string(),
user_id: old_job.user_id.clone(),
project_dir: old_job.project_dir.clone(),
success: None,
failure_reason: None,
created_at: now,
started_at: None,
completed_at: None,
credential_grants_json: old_job.credential_grants_json.clone(),
};
store
.save_sandbox_job(&record)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
let mode = match store.get_sandbox_job_mode(old_job_id).await {
Ok(Some(m)) if m == "claude_code" => {
crate::orchestrator::job_manager::JobMode::ClaudeCode
}
_ => crate::orchestrator::job_manager::JobMode::Worker,
};
let credential_grants: Vec<crate::orchestrator::auth::CredentialGrant> =
serde_json::from_str(&old_job.credential_grants_json).unwrap_or_else(|e| {
tracing::warn!(
job_id = %old_job.id,
"Failed to deserialize credential grants from stored job: {}. \
Restarted job will have no credentials.",
e
);
vec![]
});
let project_dir = std::path::PathBuf::from(&old_job.project_dir);
let _token = jm
.create_job(
new_job_id,
&task,
Some(project_dir),
mode,
credential_grants,
)
.await
.map_err(|e| {
(
StatusCode::INTERNAL_SERVER_ERROR,
format!("Failed to create container: {}", e),
)
})?;
store
.update_sandbox_job_status(new_job_id, "running", None, None, Some(now), None)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
return Ok(Json(serde_json::json!({
"status": "restarted",
"old_job_id": old_job_id,
"new_job_id": new_job_id,
})));
}
Ok(None) => {}
Err(e) => {
return Err((
StatusCode::CONFLICT,
format!("Cannot restart job in state '{}'", old_job.status),
StatusCode::INTERNAL_SERVER_ERROR,
format!("Database error: {}", e),
));
}
let jm = state.job_manager.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE,
"Sandbox not enabled".to_string(),
))?;
// Enrich the task with failure context.
let task = if let Some(ref reason) = old_job.failure_reason {
format!(
"Previous attempt failed: {}. Retry: {}",
reason, old_job.task
)
} else {
old_job.task.clone()
};
let new_job_id = Uuid::new_v4();
let now = chrono::Utc::now();
let record = crate::history::SandboxJobRecord {
id: new_job_id,
task: task.clone(),
status: "creating".to_string(),
user_id: old_job.user_id.clone(),
project_dir: old_job.project_dir.clone(),
success: None,
failure_reason: None,
created_at: now,
started_at: None,
completed_at: None,
credential_grants_json: old_job.credential_grants_json.clone(),
};
store
.save_sandbox_job(&record)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
let mode = match store.get_sandbox_job_mode(old_job_id).await {
Ok(Some(m)) if m == "claude_code" => {
crate::orchestrator::job_manager::JobMode::ClaudeCode
}
_ => crate::orchestrator::job_manager::JobMode::Worker,
};
let credential_grants: Vec<crate::orchestrator::auth::CredentialGrant> =
serde_json::from_str(&old_job.credential_grants_json).unwrap_or_else(|e| {
tracing::warn!(
job_id = %old_job.id,
"Failed to deserialize credential grants from stored job: {}. \
Restarted job will have no credentials.",
e
);
vec![]
});
let project_dir = std::path::PathBuf::from(&old_job.project_dir);
let _token = jm
.create_job(
new_job_id,
&task,
Some(project_dir),
mode,
credential_grants,
)
.await
.map_err(|e| {
(
StatusCode::INTERNAL_SERVER_ERROR,
format!("Failed to create container: {}", e),
)
})?;
store
.update_sandbox_job_status(new_job_id, "running", None, None, Some(now), None)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
return Ok(Json(serde_json::json!({
"status": "restarted",
"old_job_id": old_job_id,
"new_job_id": new_job_id,
})));
}
// Try agent job restart: dispatch a new job via the scheduler.
if let Ok(Some(old_job)) = store.get_job(old_job_id).await {
if old_job.state.is_active() {
return Err((
StatusCode::CONFLICT,
format!("Cannot restart job in state '{}'", old_job.state),
));
match store.get_job(old_job_id).await {
Ok(Some(old_job)) => {
if old_job.user_id != user.user_id {
return Err((StatusCode::NOT_FOUND, "Job not found".to_string()));
}
if old_job.state.is_active() {
return Err((
StatusCode::CONFLICT,
format!("Cannot restart job in state '{}'", old_job.state),
));
}
let slot = state.scheduler.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE,
"Scheduler not available".to_string(),
))?;
let scheduler_guard = slot.read().await;
let scheduler = scheduler_guard.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE,
"Agent not started yet".to_string(),
))?;
// Look up failure reason (O(1) point lookup).
let failure_reason = store
.get_agent_job_failure_reason(old_job_id)
.await
.ok()
.flatten()
.unwrap_or_default();
let title = if !failure_reason.is_empty() {
format!(
"Previous attempt failed: {}. Retry: {}",
failure_reason, old_job.title
)
} else {
old_job.title.clone()
};
let new_job_id = scheduler
.dispatch_job(&old_job.user_id, &title, &old_job.description, None)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
Ok(Json(serde_json::json!({
"status": "restarted",
"old_job_id": old_job_id,
"new_job_id": new_job_id,
})))
}
let slot = state.scheduler.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE,
"Scheduler not available".to_string(),
))?;
let scheduler_guard = slot.read().await;
let scheduler = scheduler_guard.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE,
"Agent not started yet".to_string(),
))?;
// Look up failure reason (O(1) point lookup).
let failure_reason = store
.get_agent_job_failure_reason(old_job_id)
.await
.ok()
.flatten()
.unwrap_or_default();
let title = if !failure_reason.is_empty() {
format!(
"Previous attempt failed: {}. Retry: {}",
failure_reason, old_job.title
)
} else {
old_job.title.clone()
};
let new_job_id = scheduler
.dispatch_job(&old_job.user_id, &title, &old_job.description, None)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
return Ok(Json(serde_json::json!({
"status": "restarted",
"old_job_id": old_job_id,
"new_job_id": new_job_id,
})));
Ok(None) => Err((StatusCode::NOT_FOUND, "Job not found".to_string())),
Err(e) => Err((
StatusCode::INTERNAL_SERVER_ERROR,
format!("Database error: {}", e),
)),
}
Err((StatusCode::NOT_FOUND, "Job not found".to_string()))
}
/// Submit a follow-up prompt to a running job.
@@ -476,6 +545,7 @@ pub async fn jobs_restart_handler(
/// - Worker-mode sandbox jobs → not supported (no mechanism to inject)
pub async fn jobs_prompt_handler(
State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Path(id): Path<String>,
Json(body): Json<serde_json::Value>,
) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
@@ -494,10 +564,15 @@ pub async fn jobs_prompt_handler(
let done = body.get("done").and_then(|v| v.as_bool()).unwrap_or(false);
// Try sandbox job path: check if we have a sandbox record for this ID.
// Try sandbox job path first: verify ownership, then route to Claude Code or reject.
if let Some(ref s) = state.store
&& let Ok(Some(_)) = s.get_sandbox_job(job_id).await
&& let Ok(Some(sandbox_job)) = s.get_sandbox_job(job_id).await
{
// Verify ownership.
if sandbox_job.user_id != user.user_id {
return Err((StatusCode::NOT_FOUND, "Job not found".to_string()));
}
// It's a sandbox job. Check if Claude Code mode.
let mode = s.get_sandbox_job_mode(job_id).await.ok().flatten();
if mode.as_deref() == Some("claude_code") {
@@ -522,7 +597,26 @@ pub async fn jobs_prompt_handler(
}
}
// Try agent job path: send via scheduler.
// Try agent job path: verify ownership, then send via scheduler.
if let Some(ref store) = state.store {
match store.get_job(job_id).await {
Ok(Some(agent_job)) => {
if agent_job.user_id != user.user_id {
return Err((StatusCode::NOT_FOUND, "Job not found".to_string()));
}
}
Ok(None) => {
return Err((StatusCode::NOT_FOUND, "Job not found".to_string()));
}
Err(e) => {
return Err((
StatusCode::INTERNAL_SERVER_ERROR,
format!("Database error: {}", e),
));
}
}
}
let slot = state.scheduler.as_ref().ok_or((
StatusCode::NOT_IMPLEMENTED,
"Agent job prompts require the scheduler to be configured".to_string(),
@@ -550,6 +644,7 @@ pub async fn jobs_prompt_handler(
/// Load persisted job events for a job (for history replay on page open).
pub async fn jobs_events_handler(
State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Path(id): Path<String>,
) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
let store = state.store.as_ref().ok_or((
@@ -561,6 +656,24 @@ pub async fn jobs_events_handler(
.parse()
.map_err(|_| (StatusCode::BAD_REQUEST, "Invalid job ID".to_string()))?;
// Verify ownership before returning events.
match store.get_sandbox_job(job_id).await {
Ok(Some(job)) => {
if job.user_id != user.user_id {
return Err((StatusCode::NOT_FOUND, "Job not found".to_string()));
}
}
Ok(None) => {
return Err((StatusCode::NOT_FOUND, "Job not found".to_string()));
}
Err(e) => {
return Err((
StatusCode::INTERNAL_SERVER_ERROR,
format!("Database error: {}", e),
));
}
}
let events = store
.list_job_events(job_id, None)
.await
@@ -593,6 +706,7 @@ pub struct FilePathQuery {
pub async fn job_files_list_handler(
State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Path(id): Path<String>,
Query(query): Query<FilePathQuery>,
) -> Result<Json<ProjectFilesResponse>, (StatusCode, String)> {
@@ -610,6 +724,10 @@ pub async fn job_files_list_handler(
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?
.ok_or((StatusCode::NOT_FOUND, "Job not found".to_string()))?;
if job.user_id != user.user_id {
return Err((StatusCode::NOT_FOUND, "Job not found".to_string()));
}
let base = std::path::PathBuf::from(&job.project_dir);
let rel_path = query.path.as_deref().unwrap_or("");
let target = base.join(rel_path);
@@ -656,6 +774,7 @@ pub async fn job_files_list_handler(
pub async fn job_files_read_handler(
State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Path(id): Path<String>,
Query(query): Query<FilePathQuery>,
) -> Result<Json<ProjectFileReadResponse>, (StatusCode, String)> {
@@ -673,6 +792,10 @@ pub async fn job_files_read_handler(
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?
.ok_or((StatusCode::NOT_FOUND, "Job not found".to_string()))?;
if job.user_id != user.user_id {
return Err((StatusCode::NOT_FOUND, "Job not found".to_string()));
}
let path = query.path.as_deref().ok_or((
StatusCode::BAD_REQUEST,
"path parameter required".to_string(),
+92 -21
View File
@@ -9,8 +9,27 @@ use axum::{
};
use serde::Deserialize;
use crate::channels::web::auth::{AuthenticatedUser, UserIdentity};
use crate::channels::web::server::GatewayState;
use crate::channels::web::types::*;
use crate::workspace::Workspace;
/// Resolve the workspace for the authenticated user.
///
/// Prefers `workspace_pool` (multi-user mode) when available, falling back
/// to the single-user `state.workspace`.
pub(crate) async fn resolve_workspace(
state: &GatewayState,
user: &UserIdentity,
) -> Result<Arc<Workspace>, (StatusCode, String)> {
if let Some(ref pool) = state.workspace_pool {
return Ok(pool.get_or_create(user).await);
}
state.workspace.as_ref().cloned().ok_or((
StatusCode::SERVICE_UNAVAILABLE,
"Workspace not available".to_string(),
))
}
#[derive(Deserialize)]
pub struct TreeQuery {
@@ -20,12 +39,10 @@ pub struct TreeQuery {
pub async fn memory_tree_handler(
State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Query(_query): Query<TreeQuery>,
) -> Result<Json<MemoryTreeResponse>, (StatusCode, String)> {
let workspace = state.workspace.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE,
"Workspace not available".to_string(),
))?;
let workspace = resolve_workspace(&state, &user).await?;
// Build tree from list_all (flat list of all paths)
let all_paths = workspace
@@ -68,12 +85,10 @@ pub struct ListQuery {
pub async fn memory_list_handler(
State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Query(query): Query<ListQuery>,
) -> Result<Json<MemoryListResponse>, (StatusCode, String)> {
let workspace = state.workspace.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE,
"Workspace not available".to_string(),
))?;
let workspace = resolve_workspace(&state, &user).await?;
let path = query.path.as_deref().unwrap_or("");
let entries = workspace
@@ -104,12 +119,10 @@ pub struct ReadQuery {
pub async fn memory_read_handler(
State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Query(query): Query<ReadQuery>,
) -> Result<Json<MemoryReadResponse>, (StatusCode, String)> {
let workspace = state.workspace.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE,
"Workspace not available".to_string(),
))?;
let workspace = resolve_workspace(&state, &user).await?;
let doc = workspace
.read(&query.path)
@@ -123,17 +136,75 @@ pub async fn memory_read_handler(
}))
}
// memory_write_handler lives in server.rs (layer-aware version with append,
// privacy redirect, and proper error status codes).
pub async fn memory_write_handler(
State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Json(req): Json<MemoryWriteRequest>,
) -> Result<Json<MemoryWriteResponse>, (StatusCode, String)> {
let workspace = resolve_workspace(&state, &user).await?;
// Route through layer-aware methods when a layer is specified.
//
// Note: unlike MemoryWriteTool, this endpoint does NOT block writes to
// identity files (IDENTITY.md, SOUL.md, etc.). The HTTP API is an
// authenticated admin interface; the supervisor uses it to seed identity
// files at startup. Identity-file protection is enforced at the tool
// layer (LLM-facing) where the write originates from an untrusted agent.
if let Some(ref layer_name) = req.layer {
let result = if req.append {
workspace
.append_to_layer(layer_name, &req.path, &req.content, req.force)
.await
} else {
workspace
.write_to_layer(layer_name, &req.path, &req.content, req.force)
.await
}
.map_err(|e| {
use crate::error::WorkspaceError;
let status = match &e {
WorkspaceError::LayerNotFound { .. } => StatusCode::BAD_REQUEST,
WorkspaceError::LayerReadOnly { .. } => StatusCode::FORBIDDEN,
WorkspaceError::PrivacyRedirectFailed => StatusCode::UNPROCESSABLE_ENTITY,
_ => StatusCode::INTERNAL_SERVER_ERROR,
};
(status, e.to_string())
})?;
return Ok(Json(MemoryWriteResponse {
path: req.path,
status: "written",
redirected: Some(result.redirected),
actual_layer: Some(result.actual_layer),
}));
}
// Non-layer path: honor the append field
if req.append {
workspace
.append(&req.path, &req.content)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
} else {
workspace
.write(&req.path, &req.content)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
}
Ok(Json(MemoryWriteResponse {
path: req.path,
status: "written",
redirected: None,
actual_layer: None,
}))
}
pub async fn memory_search_handler(
State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Json(req): Json<MemorySearchRequest>,
) -> Result<Json<MemorySearchResponse>, (StatusCode, String)> {
let workspace = state.workspace.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE,
"Workspace not available".to_string(),
))?;
let workspace = resolve_workspace(&state, &user).await?;
let limit = req.limit.unwrap_or(10);
let results = workspace
@@ -142,10 +213,10 @@ pub async fn memory_search_handler(
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
let hits: Vec<SearchHit> = results
.into_iter()
.iter()
.map(|r| SearchHit {
path: r.document_path,
content: r.content,
path: r.document_id.to_string(),
content: r.content.clone(),
score: r.score as f64,
})
.collect();
+3 -12
View File
@@ -1,13 +1,10 @@
//! Handler modules for the web gateway API.
//!
//! Each module groups related endpoint handlers by domain.
//!
//! # Migration status
//!
//! `skills` is the canonical implementation used by `server.rs`.
//! The remaining modules are in-progress migrations from inline server.rs
//! handlers; their functions are not yet wired up, hence the `dead_code` allow.
pub mod jobs;
pub mod memory;
pub mod routines;
pub mod skills;
// Modules not yet wired into server.rs router -- suppress dead_code until
@@ -17,12 +14,6 @@ pub mod chat;
#[allow(dead_code)]
pub mod extensions;
#[allow(dead_code)]
pub mod jobs;
#[allow(dead_code)]
pub mod memory;
#[allow(dead_code)]
pub mod routines;
#[allow(dead_code)]
pub mod settings;
#[allow(dead_code)]
pub mod static_files;
+44 -5
View File
@@ -11,12 +11,14 @@ use serde::Deserialize;
use uuid::Uuid;
use crate::agent::routine::{Trigger, next_cron_fire};
use crate::channels::web::auth::AuthenticatedUser;
use crate::channels::web::server::GatewayState;
use crate::channels::web::types::*;
use crate::error::RoutineError;
pub async fn routines_list_handler(
State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
) -> Result<Json<RoutineListResponse>, (StatusCode, String)> {
let store = state.store.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE,
@@ -24,7 +26,7 @@ pub async fn routines_list_handler(
))?;
let routines = store
.list_all_routines()
.list_routines(&user.user_id)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
@@ -35,6 +37,7 @@ pub async fn routines_list_handler(
pub async fn routines_summary_handler(
State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
) -> Result<Json<RoutineSummaryResponse>, (StatusCode, String)> {
let store = state.store.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE,
@@ -42,7 +45,7 @@ pub async fn routines_summary_handler(
))?;
let routines = store
.list_all_routines()
.list_routines(&user.user_id)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
@@ -78,6 +81,7 @@ pub async fn routines_summary_handler(
pub async fn routines_detail_handler(
State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Path(id): Path<String>,
) -> Result<Json<RoutineDetailResponse>, (StatusCode, String)> {
let store = state.store.as_ref().ok_or((
@@ -94,6 +98,10 @@ pub async fn routines_detail_handler(
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?
.ok_or((StatusCode::NOT_FOUND, "Routine not found".to_string()))?;
if routine.user_id != user.user_id {
return Err((StatusCode::NOT_FOUND, "Routine not found".to_string()));
}
let runs = store
.list_routine_runs(routine_id, 20)
.await
@@ -106,7 +114,7 @@ pub async fn routines_detail_handler(
trigger_type: run.trigger_type.clone(),
started_at: run.started_at.to_rfc3339(),
completed_at: run.completed_at.map(|dt| dt.to_rfc3339()),
status: format!("{:?}", run.status),
status: run.status.to_string(),
result_summary: run.result_summary.clone(),
tokens_used: run.tokens_used,
job_id: run.job_id,
@@ -137,6 +145,7 @@ pub async fn routines_detail_handler(
pub async fn routines_trigger_handler(
State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Path(id): Path<String>,
) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
// Clone the Arc out of the lock to avoid holding the RwLock across .await.
@@ -152,7 +161,7 @@ pub async fn routines_trigger_handler(
.map_err(|_| (StatusCode::BAD_REQUEST, "Invalid routine ID".to_string()))?;
let run_id = engine
.fire_manual(routine_id, Some(&state.user_id))
.fire_manual(routine_id, Some(&user.user_id))
.await
.map_err(|e| (routine_error_status(&e), e.to_string()))?;
@@ -170,6 +179,7 @@ pub struct ToggleRequest {
pub async fn routines_toggle_handler(
State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Path(id): Path<String>,
body: Option<Json<ToggleRequest>>,
) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
@@ -187,6 +197,10 @@ pub async fn routines_toggle_handler(
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?
.ok_or((StatusCode::NOT_FOUND, "Routine not found".to_string()))?;
if routine.user_id != user.user_id {
return Err((StatusCode::NOT_FOUND, "Routine not found".to_string()));
}
let was_enabled = routine.enabled;
// If a specific value was provided, use it; otherwise toggle.
routine.enabled = match body {
@@ -230,6 +244,7 @@ pub async fn routines_toggle_handler(
pub async fn routines_delete_handler(
State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Path(id): Path<String>,
) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
let store = state.store.as_ref().ok_or((
@@ -240,6 +255,17 @@ pub async fn routines_delete_handler(
let routine_id = Uuid::parse_str(&id)
.map_err(|_| (StatusCode::BAD_REQUEST, "Invalid routine ID".to_string()))?;
// Verify ownership before deleting.
let routine = store
.get_routine(routine_id)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?
.ok_or((StatusCode::NOT_FOUND, "Routine not found".to_string()))?;
if routine.user_id != user.user_id {
return Err((StatusCode::NOT_FOUND, "Routine not found".to_string()));
}
let deleted = store
.delete_routine(routine_id)
.await
@@ -261,8 +287,10 @@ pub async fn routines_delete_handler(
}
}
#[allow(dead_code)] // Used by server.rs inline version; kept in sync here for future migration.
pub async fn routines_runs_handler(
State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Path(id): Path<String>,
) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
let store = state.store.as_ref().ok_or((
@@ -273,6 +301,17 @@ pub async fn routines_runs_handler(
let routine_id = Uuid::parse_str(&id)
.map_err(|_| (StatusCode::BAD_REQUEST, "Invalid routine ID".to_string()))?;
// Verify ownership before listing runs.
let routine = store
.get_routine(routine_id)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?
.ok_or((StatusCode::NOT_FOUND, "Routine not found".to_string()))?;
if routine.user_id != user.user_id {
return Err((StatusCode::NOT_FOUND, "Routine not found".to_string()));
}
let runs = store
.list_routine_runs(routine_id, 50)
.await
@@ -285,7 +324,7 @@ pub async fn routines_runs_handler(
trigger_type: run.trigger_type.clone(),
started_at: run.started_at.to_rfc3339(),
completed_at: run.completed_at.map(|dt| dt.to_rfc3339()),
status: format!("{:?}", run.status),
status: run.status.to_string(),
result_summary: run.result_summary.clone(),
tokens_used: run.tokens_used,
job_id: run.job_id,
+13 -6
View File
@@ -8,17 +8,19 @@ use axum::{
http::StatusCode,
};
use crate::channels::web::auth::AuthenticatedUser;
use crate::channels::web::server::GatewayState;
use crate::channels::web::types::*;
pub async fn settings_list_handler(
State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
) -> Result<Json<SettingsListResponse>, StatusCode> {
let store = state
.store
.as_ref()
.ok_or(StatusCode::SERVICE_UNAVAILABLE)?;
let rows = store.list_settings(&state.user_id).await.map_err(|e| {
let rows = store.list_settings(&user.user_id).await.map_err(|e| {
tracing::error!("Failed to list settings: {}", e);
StatusCode::INTERNAL_SERVER_ERROR
})?;
@@ -37,6 +39,7 @@ pub async fn settings_list_handler(
pub async fn settings_get_handler(
State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Path(key): Path<String>,
) -> Result<Json<SettingResponse>, StatusCode> {
let store = state
@@ -44,7 +47,7 @@ pub async fn settings_get_handler(
.as_ref()
.ok_or(StatusCode::SERVICE_UNAVAILABLE)?;
let row = store
.get_setting_full(&state.user_id, &key)
.get_setting_full(&user.user_id, &key)
.await
.map_err(|e| {
tracing::error!("Failed to get setting '{}': {}", key, e);
@@ -61,6 +64,7 @@ pub async fn settings_get_handler(
pub async fn settings_set_handler(
State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Path(key): Path<String>,
Json(body): Json<SettingWriteRequest>,
) -> Result<StatusCode, StatusCode> {
@@ -69,7 +73,7 @@ pub async fn settings_set_handler(
.as_ref()
.ok_or(StatusCode::SERVICE_UNAVAILABLE)?;
store
.set_setting(&state.user_id, &key, &body.value)
.set_setting(&user.user_id, &key, &body.value)
.await
.map_err(|e| {
tracing::error!("Failed to set setting '{}': {}", key, e);
@@ -81,6 +85,7 @@ pub async fn settings_set_handler(
pub async fn settings_delete_handler(
State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Path(key): Path<String>,
) -> Result<StatusCode, StatusCode> {
let store = state
@@ -88,7 +93,7 @@ pub async fn settings_delete_handler(
.as_ref()
.ok_or(StatusCode::SERVICE_UNAVAILABLE)?;
store
.delete_setting(&state.user_id, &key)
.delete_setting(&user.user_id, &key)
.await
.map_err(|e| {
tracing::error!("Failed to delete setting '{}': {}", key, e);
@@ -100,12 +105,13 @@ pub async fn settings_delete_handler(
pub async fn settings_export_handler(
State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
) -> Result<Json<SettingsExportResponse>, StatusCode> {
let store = state
.store
.as_ref()
.ok_or(StatusCode::SERVICE_UNAVAILABLE)?;
let settings = store.get_all_settings(&state.user_id).await.map_err(|e| {
let settings = store.get_all_settings(&user.user_id).await.map_err(|e| {
tracing::error!("Failed to export settings: {}", e);
StatusCode::INTERNAL_SERVER_ERROR
})?;
@@ -115,6 +121,7 @@ pub async fn settings_export_handler(
pub async fn settings_import_handler(
State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Json(body): Json<SettingsImportRequest>,
) -> Result<StatusCode, StatusCode> {
let store = state
@@ -122,7 +129,7 @@ pub async fn settings_import_handler(
.as_ref()
.ok_or(StatusCode::SERVICE_UNAVAILABLE)?;
store
.set_all_settings(&state.user_id, &body.settings)
.set_all_settings(&user.user_id, &body.settings)
.await
.map_err(|e| {
tracing::error!("Failed to import settings: {}", e);
+9
View File
@@ -8,11 +8,13 @@ use axum::{
http::StatusCode,
};
use crate::channels::web::auth::AuthenticatedUser;
use crate::channels::web::server::GatewayState;
use crate::channels::web::types::*;
pub async fn skills_list_handler(
State(state): State<Arc<GatewayState>>,
AuthenticatedUser(_user): AuthenticatedUser,
) -> Result<Json<SkillListResponse>, (StatusCode, String)> {
let registry = state.skill_registry.as_ref().ok_or((
StatusCode::NOT_IMPLEMENTED,
@@ -45,6 +47,7 @@ pub async fn skills_list_handler(
pub async fn skills_search_handler(
State(state): State<Arc<GatewayState>>,
AuthenticatedUser(_user): AuthenticatedUser,
Json(req): Json<SkillSearchRequest>,
) -> Result<Json<SkillSearchResponse>, (StatusCode, String)> {
let registry = state.skill_registry.as_ref().ok_or((
@@ -119,6 +122,7 @@ pub async fn skills_search_handler(
pub async fn skills_install_handler(
State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
headers: axum::http::HeaderMap,
Json(req): Json<SkillInstallRequest>,
) -> Result<Json<ActionResponse>, (StatusCode, String)> {
@@ -135,6 +139,8 @@ pub async fn skills_install_handler(
));
}
tracing::info!(user_id = %user.user_id, skill = %req.name, "skill install requested");
let registry = state.skill_registry.as_ref().ok_or((
StatusCode::NOT_IMPLEMENTED,
"Skills system not enabled".to_string(),
@@ -219,6 +225,7 @@ pub async fn skills_install_handler(
pub async fn skills_remove_handler(
State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
headers: axum::http::HeaderMap,
Path(name): Path<String>,
) -> Result<Json<ActionResponse>, (StatusCode, String)> {
@@ -234,6 +241,8 @@ pub async fn skills_remove_handler(
));
}
tracing::info!(user_id = %user.user_id, skill = %name, "skill remove requested");
let registry = state.skill_registry.as_ref().ok_or((
StatusCode::NOT_IMPLEMENTED,
"Skills system not enabled".to_string(),
@@ -7,6 +7,7 @@ use axum::{
};
use crate::bootstrap::ironclaw_base_dir;
use crate::channels::web::auth::AuthenticatedUser;
use crate::channels::web::types::*;
// --- Static file handlers ---
@@ -113,6 +114,7 @@ use crate::channels::web::server::GatewayState;
pub async fn logs_events_handler(
State(state): State<Arc<GatewayState>>,
AuthenticatedUser(_user): AuthenticatedUser,
) -> Result<
Sse<impl futures::Stream<Item = Result<Event, Infallible>> + Send + 'static>,
(StatusCode, String),
@@ -152,6 +154,7 @@ pub async fn logs_events_handler(
pub async fn gateway_status_handler(
State(state): State<Arc<GatewayState>>,
AuthenticatedUser(_user): AuthenticatedUser,
) -> Json<GatewayStatusResponse> {
let sse_connections = state.sse.connection_count();
let ws_connections = state
+90 -22
View File
@@ -31,6 +31,9 @@ pub mod ws;
/// [`TestGatewayBuilder`](test_helpers::TestGatewayBuilder).
pub mod test_helpers;
#[cfg(test)]
mod tests;
use std::net::SocketAddr;
use std::sync::Arc;
@@ -52,6 +55,7 @@ use crate::workspace::Workspace;
use self::log_layer::{LogBroadcaster, LogLevelHandle};
use self::auth::MultiAuthState;
use self::server::GatewayState;
use self::sse::SseManager;
use self::types::SseEvent;
@@ -60,14 +64,15 @@ use self::types::SseEvent;
pub struct GatewayChannel {
config: GatewayConfig,
state: Arc<GatewayState>,
/// The actual auth token in use (generated or from config).
auth_token: String,
/// Multi-user auth state (replaces bare auth_token).
auth: MultiAuthState,
}
impl GatewayChannel {
/// Create a new gateway channel.
///
/// If no auth token is configured, generates a random one and prints it.
/// Builds a single-user `MultiAuthState` from the config.
pub fn new(config: GatewayConfig) -> Self {
let auth_token = config.auth_token.clone().unwrap_or_else(|| {
use rand::RngCore;
@@ -77,10 +82,13 @@ impl GatewayChannel {
bytes.iter().map(|b| format!("{b:02x}")).collect()
});
let auth = MultiAuthState::single(auth_token, config.user_id.clone());
let state = Arc::new(GatewayState {
msg_tx: tokio::sync::RwLock::new(None),
sse: SseManager::new(),
sse: Arc::new(SseManager::new()),
workspace: None,
workspace_pool: None,
session_manager: None,
log_broadcaster: None,
log_level_handle: None,
@@ -90,13 +98,13 @@ impl GatewayChannel {
job_manager: None,
prompt_queue: None,
scheduler: None,
user_id: config.user_id.clone(),
default_user_id: config.user_id.clone(),
shutdown_tx: tokio::sync::RwLock::new(None),
ws_tracker: Some(Arc::new(ws::WsConnectionTracker::new())),
llm_provider: None,
skill_registry: None,
skill_catalog: None,
chat_rate_limiter: server::RateLimiter::new(30, 60),
chat_rate_limiter: server::PerUserRateLimiter::new(30, 60),
oauth_rate_limiter: server::RateLimiter::new(10, 60),
webhook_rate_limiter: server::RateLimiter::new(10, 60),
registry_entries: Vec::new(),
@@ -109,7 +117,46 @@ impl GatewayChannel {
Self {
config,
state,
auth_token,
auth,
}
}
/// Create a gateway channel with a pre-built multi-user auth state.
pub fn new_multi_auth(config: GatewayConfig, auth: MultiAuthState) -> Self {
let state = Arc::new(GatewayState {
msg_tx: tokio::sync::RwLock::new(None),
sse: Arc::new(SseManager::new()),
workspace: None,
workspace_pool: None,
session_manager: None,
log_broadcaster: None,
log_level_handle: None,
extension_manager: None,
tool_registry: None,
store: None,
job_manager: None,
prompt_queue: None,
scheduler: None,
default_user_id: config.user_id.clone(),
shutdown_tx: tokio::sync::RwLock::new(None),
ws_tracker: Some(Arc::new(ws::WsConnectionTracker::new())),
llm_provider: None,
skill_registry: None,
skill_catalog: None,
chat_rate_limiter: server::PerUserRateLimiter::new(30, 60),
oauth_rate_limiter: server::RateLimiter::new(10, 60),
registry_entries: Vec::new(),
cost_guard: None,
routine_engine: Arc::new(tokio::sync::RwLock::new(None)),
startup_time: std::time::Instant::now(),
webhook_rate_limiter: server::RateLimiter::new(10, 60),
active_config: server::ActiveConfigSnapshot::default(),
});
Self {
config,
state,
auth,
}
}
@@ -118,8 +165,9 @@ impl GatewayChannel {
let mut new_state = GatewayState {
msg_tx: tokio::sync::RwLock::new(None),
// Preserve the existing broadcast channel so sender handles remain valid.
sse: SseManager::from_sender(self.state.sse.sender()),
sse: Arc::new(SseManager::from_sender(self.state.sse.sender())),
workspace: self.state.workspace.clone(),
workspace_pool: self.state.workspace_pool.clone(),
session_manager: self.state.session_manager.clone(),
log_broadcaster: self.state.log_broadcaster.clone(),
log_level_handle: self.state.log_level_handle.clone(),
@@ -129,13 +177,13 @@ impl GatewayChannel {
job_manager: self.state.job_manager.clone(),
prompt_queue: self.state.prompt_queue.clone(),
scheduler: self.state.scheduler.clone(),
user_id: self.state.user_id.clone(),
default_user_id: self.state.default_user_id.clone(),
shutdown_tx: tokio::sync::RwLock::new(None),
ws_tracker: self.state.ws_tracker.clone(),
llm_provider: self.state.llm_provider.clone(),
skill_registry: self.state.skill_registry.clone(),
skill_catalog: self.state.skill_catalog.clone(),
chat_rate_limiter: server::RateLimiter::new(30, 60),
chat_rate_limiter: server::PerUserRateLimiter::new(30, 60),
oauth_rate_limiter: server::RateLimiter::new(10, 60),
webhook_rate_limiter: server::RateLimiter::new(10, 60),
registry_entries: self.state.registry_entries.clone(),
@@ -260,9 +308,15 @@ impl GatewayChannel {
self
}
/// Get the auth token (for printing to console on startup).
/// Inject the per-user workspace pool for multi-user mode.
pub fn with_workspace_pool(mut self, pool: Arc<server::WorkspacePool>) -> Self {
self.rebuild_state(|s| s.workspace_pool = Some(pool));
self
}
/// Get the first auth token (for printing to console on startup).
pub fn auth_token(&self) -> &str {
&self.auth_token
self.auth.first_token().unwrap_or("")
}
/// Get a reference to the shared gateway state (for the agent to push SSE events).
@@ -291,7 +345,7 @@ impl Channel for GatewayChannel {
),
})?;
server::start_server(addr, self.state.clone(), self.auth_token.clone()).await?;
server::start_server(addr, self.state.clone(), self.auth.clone()).await?;
Ok(Box::pin(ReceiverStream::new(rx)))
}
@@ -311,10 +365,13 @@ impl Channel for GatewayChannel {
}
};
self.state.sse.broadcast(SseEvent::Response {
content: response.content,
thread_id,
});
self.state.sse.broadcast_for_user(
&msg.user_id,
SseEvent::Response {
content: response.content,
thread_id,
},
);
Ok(())
}
@@ -427,13 +484,21 @@ impl Channel for GatewayChannel {
},
};
self.state.sse.broadcast(event);
// Scope events to the user when user_id is available in metadata.
// When user_id is missing (heartbeat, routines), events go to all
// subscribers. In multi-tenant mode this leaks status across users.
if let Some(uid) = metadata.get("user_id").and_then(|v| v.as_str()) {
self.state.sse.broadcast_for_user(uid, event);
} else {
tracing::debug!("Status event missing user_id in metadata; broadcasting globally");
self.state.sse.broadcast(event);
}
Ok(())
}
async fn broadcast(
&self,
_user_id: &str,
user_id: &str,
response: OutgoingResponse,
) -> Result<(), ChannelError> {
let thread_id = match response.thread_id {
@@ -445,10 +510,13 @@ impl Channel for GatewayChannel {
return Ok(());
}
};
self.state.sse.broadcast(SseEvent::Response {
content: response.content,
thread_id,
});
self.state.sse.broadcast_for_user(
user_id,
SseEvent::Response {
content: response.content,
thread_id,
},
);
Ok(())
}
+2 -1
View File
@@ -463,9 +463,10 @@ fn build_tool_request(
pub async fn chat_completions_handler(
State(state): State<Arc<GatewayState>>,
super::auth::AuthenticatedUser(user): super::auth::AuthenticatedUser,
Json(req): Json<OpenAiChatRequest>,
) -> Result<impl IntoResponse, (StatusCode, Json<OpenAiErrorResponse>)> {
if !state.chat_rate_limiter.check() {
if !state.chat_rate_limiter.check(&user.user_id) {
return Err(openai_error(
StatusCode::TOO_MANY_REQUESTS,
"Rate limit exceeded. Please try again later.",
+520 -373
View File
File diff suppressed because it is too large Load Diff
+129 -29
View File
@@ -17,9 +17,25 @@ use crate::channels::web::types::SseEvent;
/// Prevents resource exhaustion from connection flooding.
const MAX_CONNECTIONS: u64 = 100;
/// Envelope for broadcast events: carries an optional user scope.
///
/// `user_id = None` means the event is global (e.g. Heartbeat) and delivered
/// to all subscribers. `user_id = Some(id)` means the event is only delivered
/// to subscribers that match that user_id.
#[derive(Debug, Clone)]
pub(crate) struct ScopedEvent {
pub(crate) user_id: Option<String>,
pub(crate) event: SseEvent,
}
/// Manages SSE broadcast to all connected browser tabs.
///
/// In multi-user mode, events are scoped by user_id so that each subscriber
/// only receives events intended for their user (plus global events like
/// Heartbeat). In single-user mode, all events are delivered to all subscribers
/// (backwards compatible).
pub struct SseManager {
tx: broadcast::Sender<SseEvent>,
tx: broadcast::Sender<ScopedEvent>,
connection_count: Arc<AtomicU64>,
max_connections: u64,
}
@@ -45,7 +61,7 @@ impl SseManager {
/// only be called before the server starts accepting connections (i.e.,
/// during startup wiring). Calling it after connections are established
/// will break connection tracking and allow exceeding `MAX_CONNECTIONS`.
pub fn from_sender(tx: broadcast::Sender<SseEvent>) -> Self {
pub(crate) fn from_sender(tx: broadcast::Sender<ScopedEvent>) -> Self {
Self {
tx,
connection_count: Arc::new(AtomicU64::new(0)),
@@ -53,15 +69,28 @@ impl SseManager {
}
}
/// Broadcast an event to all connected clients.
pub fn broadcast(&self, event: SseEvent) {
// Ignore send errors (no receivers is fine)
let _ = self.tx.send(event);
/// Get a clone of the broadcast sender for use by other components.
pub(crate) fn sender(&self) -> broadcast::Sender<ScopedEvent> {
self.tx.clone()
}
/// Get a clone of the broadcast sender for use by other components.
pub fn sender(&self) -> broadcast::Sender<SseEvent> {
self.tx.clone()
/// Broadcast an event to all connected clients (global/unscoped).
pub fn broadcast(&self, event: SseEvent) {
let _ = self.tx.send(ScopedEvent {
user_id: None,
event,
});
}
/// Broadcast an event scoped to a specific user.
///
/// Only subscribers for this user_id (or unscoped subscribers) will
/// receive the event.
pub fn broadcast_for_user(&self, user_id: &str, event: SseEvent) {
let _ = self.tx.send(ScopedEvent {
user_id: Some(user_id.to_string()),
event,
});
}
/// Get current number of active connections.
@@ -71,11 +100,15 @@ impl SseManager {
/// Create a raw broadcast subscription for non-SSE consumers (e.g. WebSocket).
///
/// Returns a stream of `SseEvent` values and increments/decrements the
/// connection counter on creation/drop, just like `subscribe()` does for SSE.
/// When `user_id` is `Some`, only events scoped to that user (or global
/// events) are delivered. When `None`, all events are delivered (single-user
/// backwards compatibility).
///
/// Returns `None` if the maximum connection limit has been reached.
pub fn subscribe_raw(&self) -> Option<impl Stream<Item = SseEvent> + Send + 'static + use<>> {
pub fn subscribe_raw(
&self,
user_id: Option<String>,
) -> Option<impl Stream<Item = SseEvent> + Send + 'static + use<>> {
// Atomically increment only if below the limit. This prevents
// concurrent callers from overshooting max_connections.
let counter = Arc::clone(&self.connection_count);
@@ -91,7 +124,19 @@ impl SseManager {
.ok()?;
let rx = self.tx.subscribe();
let stream = BroadcastStream::new(rx).filter_map(|result| result.ok());
let stream = BroadcastStream::new(rx).filter_map(move |result| match result {
Ok(scoped) => {
// Global events (user_id=None) always pass through.
// Scoped events only pass if the subscriber matches (or subscriber is unscoped).
match (&user_id, &scoped.user_id) {
(_, None) => Some(scoped.event), // global -> all
(None, _) => Some(scoped.event), // unscoped subscriber -> all
(Some(sub), Some(ev)) if sub == ev => Some(scoped.event), // match
_ => None, // different user -> skip
}
}
Err(_) => None,
});
Some(CountedStream {
inner: stream,
@@ -101,9 +146,13 @@ impl SseManager {
/// Create a new SSE stream for a client connection.
///
/// When `user_id` is `Some`, only events for that user (or global events)
/// are delivered. When `None`, all events are delivered.
///
/// Returns `None` if the maximum connection limit has been reached.
pub fn subscribe(
&self,
user_id: Option<String>,
) -> Option<Sse<impl Stream<Item = Result<Event, Infallible>> + Send + 'static + use<>>> {
// Atomically increment only if below the limit.
let counter = Arc::clone(&self.connection_count);
@@ -120,9 +169,23 @@ impl SseManager {
let rx = self.tx.subscribe();
let stream = BroadcastStream::new(rx)
.filter_map(|result| result.ok())
.map(|event| {
let data = serde_json::to_string(&event).unwrap_or_default();
.filter_map(move |result| match result {
Ok(scoped) => match (&user_id, &scoped.user_id) {
(_, None) => Some(scoped.event),
(None, _) => Some(scoped.event),
(Some(sub), Some(ev)) if sub == ev => Some(scoped.event),
_ => None,
},
Err(_) => None,
})
.filter_map(|event| {
let data = match serde_json::to_string(&event) {
Ok(s) => s,
Err(e) => {
tracing::warn!("Failed to serialize SSE event: {}", e);
return None;
}
};
let event_type = match &event {
SseEvent::Response { .. } => "response",
SseEvent::Thinking { .. } => "thinking",
@@ -147,7 +210,7 @@ impl SseManager {
SseEvent::TurnCost { .. } => "turn_cost",
SseEvent::ExtensionStatus { .. } => "extension_status",
};
Ok(Event::default().event(event_type).data(data))
Some(Ok(Event::default().event(event_type).data(data)))
});
// Wrap in a stream that decrements on drop
@@ -215,16 +278,14 @@ mod tests {
#[tokio::test]
async fn test_broadcast_to_receiver() {
let manager = SseManager::new();
let mut rx = BroadcastStream::new(manager.tx.subscribe());
let mut stream = Box::pin(manager.subscribe_raw(None).expect("should subscribe"));
manager.broadcast(SseEvent::Status {
message: "test".to_string(),
thread_id: None,
});
let event = rx.next().await;
assert!(event.is_some());
let event = event.unwrap().unwrap();
let event = stream.next().await.unwrap();
match event {
SseEvent::Status { message, .. } => assert_eq!(message, "test"),
_ => panic!("unexpected event type"),
@@ -234,7 +295,7 @@ mod tests {
#[tokio::test]
async fn test_subscribe_raw_receives_events() {
let manager = SseManager::new();
let mut stream = Box::pin(manager.subscribe_raw().expect("should subscribe"));
let mut stream = Box::pin(manager.subscribe_raw(None).expect("should subscribe"));
assert_eq!(manager.connection_count(), 1);
@@ -254,7 +315,7 @@ mod tests {
async fn test_subscribe_raw_decrements_on_drop() {
let manager = SseManager::new();
{
let _stream = Box::pin(manager.subscribe_raw().expect("should subscribe"));
let _stream = Box::pin(manager.subscribe_raw(None).expect("should subscribe"));
assert_eq!(manager.connection_count(), 1);
}
// Stream dropped, counter should decrement
@@ -264,8 +325,8 @@ mod tests {
#[tokio::test]
async fn test_subscribe_raw_multiple_subscribers() {
let manager = SseManager::new();
let mut s1 = Box::pin(manager.subscribe_raw().expect("should subscribe"));
let mut s2 = Box::pin(manager.subscribe_raw().expect("should subscribe"));
let mut s1 = Box::pin(manager.subscribe_raw(None).expect("should subscribe"));
let mut s2 = Box::pin(manager.subscribe_raw(None).expect("should subscribe"));
assert_eq!(manager.connection_count(), 2);
manager.broadcast(SseEvent::Heartbeat);
@@ -286,12 +347,51 @@ mod tests {
let mut manager = SseManager::new();
manager.max_connections = 2; // Low limit for testing
let _s1 = Box::pin(manager.subscribe_raw().expect("first should succeed"));
let _s2 = Box::pin(manager.subscribe_raw().expect("second should succeed"));
let _s1 = Box::pin(manager.subscribe_raw(None).expect("first should succeed"));
let _s2 = Box::pin(manager.subscribe_raw(None).expect("second should succeed"));
assert_eq!(manager.connection_count(), 2);
// Third should be rejected
assert!(manager.subscribe_raw().is_none());
assert!(manager.subscribe().is_none());
assert!(manager.subscribe_raw(None).is_none());
assert!(manager.subscribe(None).is_none());
}
#[tokio::test]
async fn test_scoped_events_filtered_by_user() {
let manager = SseManager::new();
let mut alice = Box::pin(
manager
.subscribe_raw(Some("alice".to_string()))
.expect("subscribe"),
);
let mut bob = Box::pin(
manager
.subscribe_raw(Some("bob".to_string()))
.expect("subscribe"),
);
// Send event scoped to alice
manager.broadcast_for_user(
"alice",
SseEvent::Status {
message: "alice only".to_string(),
thread_id: None,
},
);
// Send global event
manager.broadcast(SseEvent::Heartbeat);
// Alice gets her scoped event
let e = alice.next().await.unwrap();
assert!(matches!(e, SseEvent::Status { .. }));
// Alice also gets the global heartbeat
let e = alice.next().await.unwrap();
assert!(matches!(e, SseEvent::Heartbeat));
// Bob only gets the global heartbeat (alice's event was filtered)
let e = bob.next().await.unwrap(); // safety: test-only
assert!(matches!(e, SseEvent::Heartbeat)); // safety: test assertion
}
}
+3 -3
View File
@@ -4265,9 +4265,9 @@ function renderRoutineDetail(routine) {
+ '<th>Trigger</th><th>Started</th><th>Completed</th><th>Status</th><th>Summary</th><th>Tokens</th>'
+ '</tr></thead><tbody>';
for (const run of routine.recent_runs) {
const runStatusClass = run.status === 'Ok' ? 'completed'
: run.status === 'Failed' ? 'failed'
: run.status === 'Attention' ? 'stuck'
const runStatusClass = run.status === 'ok' ? 'completed'
: run.status === 'failed' ? 'failed'
: run.status === 'attention' ? 'stuck'
: 'in_progress';
html += '<tr>'
+ '<td>' + escapeHtml(run.trigger_type) + '</td>'
+23 -6
View File
@@ -10,7 +10,8 @@ use std::sync::Arc;
use tokio::sync::mpsc;
use crate::channels::IncomingMessage;
use crate::channels::web::server::{GatewayState, RateLimiter, start_server};
use crate::channels::web::auth::MultiAuthState;
use crate::channels::web::server::{GatewayState, PerUserRateLimiter, RateLimiter, start_server};
use crate::channels::web::sse::SseManager;
use crate::channels::web::ws::WsConnectionTracker;
@@ -64,8 +65,9 @@ impl TestGatewayBuilder {
pub fn build(self) -> Arc<GatewayState> {
Arc::new(GatewayState {
msg_tx: tokio::sync::RwLock::new(self.msg_tx),
sse: SseManager::new(),
sse: Arc::new(SseManager::new()),
workspace: None,
workspace_pool: None,
session_manager: None,
log_broadcaster: None,
log_level_handle: None,
@@ -74,14 +76,14 @@ impl TestGatewayBuilder {
store: None,
job_manager: None,
prompt_queue: None,
user_id: self.user_id,
default_user_id: self.user_id,
shutdown_tx: tokio::sync::RwLock::new(None),
ws_tracker: Some(Arc::new(WsConnectionTracker::new())),
llm_provider: self.llm_provider,
skill_registry: None,
skill_catalog: None,
scheduler: None,
chat_rate_limiter: RateLimiter::new(30, 60),
chat_rate_limiter: PerUserRateLimiter::new(30, 60),
oauth_rate_limiter: RateLimiter::new(10, 60),
webhook_rate_limiter: RateLimiter::new(10, 60),
registry_entries: Vec::new(),
@@ -98,11 +100,26 @@ impl TestGatewayBuilder {
self,
auth_token: &str,
) -> Result<(SocketAddr, Arc<GatewayState>), crate::error::ChannelError> {
let auth = MultiAuthState::single(auth_token.to_string(), "test-user".to_string());
let state = self.build();
let addr: SocketAddr = "127.0.0.1:0"
.parse()
.expect("hard-coded address must parse");
let bound = start_server(addr, state.clone(), auth_token.to_string()).await?;
.expect("hard-coded address must parse"); // safety: constant literal
let bound = start_server(addr, state.clone(), auth).await?;
Ok((bound, state))
}
/// Build the state and start a gateway server with multi-user auth.
/// Returns the bound address and the shared state.
pub async fn start_multi(
self,
auth: MultiAuthState,
) -> Result<(SocketAddr, Arc<GatewayState>), crate::error::ChannelError> {
let state = self.build();
let addr: SocketAddr = "127.0.0.1:0"
.parse()
.expect("hard-coded address must parse"); // safety: constant literal
let bound = start_server(addr, state.clone(), auth).await?;
Ok((bound, state))
}
}
+3
View File
@@ -0,0 +1,3 @@
//! Integration tests for the web gateway module.
mod multi_tenant;
+796
View File
@@ -0,0 +1,796 @@
//! Multi-tenant isolation tests for the web gateway.
//!
//! Tests cover workspace pool scoping, job handler isolation, and auth
//! enforcement on protected endpoints. Uses `LibSqlBackend::new_local()`
//! with a temporary directory for a real (but ephemeral) database.
use std::collections::HashMap;
use std::sync::Arc;
use std::time::Duration;
use axum::Router;
use axum::body::Body;
use axum::http::{Method, Request, StatusCode};
use axum::middleware;
use axum::routing::{delete, get, post};
use tower::ServiceExt;
use uuid::Uuid;
use crate::channels::web::auth::{
AuthenticatedUser, MultiAuthState, UserIdentity, auth_middleware,
};
use crate::channels::web::server::{
ActiveConfigSnapshot, GatewayState, PerUserRateLimiter, PromptQueue, RateLimiter, WorkspacePool,
};
use crate::channels::web::sse::SseManager;
// ── Helpers ────────────────────────────────────────────────────────────
/// Create a two-user `MultiAuthState` for alice and bob.
fn two_user_auth() -> MultiAuthState {
let mut tokens = HashMap::new();
tokens.insert(
"tok-alice".to_string(),
UserIdentity {
user_id: "alice".to_string(),
workspace_read_scopes: vec!["shared".to_string()],
},
);
tokens.insert(
"tok-bob".to_string(),
UserIdentity {
user_id: "bob".to_string(),
workspace_read_scopes: vec!["shared".to_string(), "alice".to_string()],
},
);
MultiAuthState::multi(tokens)
}
/// Build a `GatewayState` with configurable store and prompt queue.
fn build_state(
store: Option<Arc<dyn crate::db::Database>>,
prompt_queue: Option<PromptQueue>,
) -> Arc<GatewayState> {
Arc::new(GatewayState {
msg_tx: tokio::sync::RwLock::new(None),
sse: Arc::new(SseManager::new()),
workspace: None,
workspace_pool: None,
session_manager: None,
log_broadcaster: None,
log_level_handle: None,
extension_manager: None,
tool_registry: None,
store,
job_manager: None,
prompt_queue,
default_user_id: "test".to_string(),
shutdown_tx: tokio::sync::RwLock::new(None),
ws_tracker: None,
llm_provider: None,
skill_registry: None,
skill_catalog: None,
scheduler: None,
chat_rate_limiter: PerUserRateLimiter::new(30, 60),
oauth_rate_limiter: RateLimiter::new(10, 60),
webhook_rate_limiter: RateLimiter::new(10, 60),
registry_entries: Vec::new(),
cost_guard: None,
routine_engine: Arc::new(tokio::sync::RwLock::new(None)),
startup_time: std::time::Instant::now(),
active_config: ActiveConfigSnapshot::default(),
})
}
/// Create a libSQL-backed test database in a temporary directory.
///
/// Returns the database and a `TempDir` guard — the database file is
/// deleted when the guard is dropped.
#[cfg(feature = "libsql")]
async fn test_db() -> (Arc<dyn crate::db::Database>, tempfile::TempDir) {
use crate::db::Database;
let dir = tempfile::tempdir().expect("failed to create temp dir"); // safety: test-only
let path = dir.path().join("test.db");
let backend = crate::db::libsql::LibSqlBackend::new_local(&path)
.await
.expect("failed to create test LibSqlBackend"); // safety: test-only
backend
.run_migrations()
.await
.expect("failed to run migrations"); // safety: test-only
(Arc::new(backend) as Arc<dyn crate::db::Database>, dir)
}
/// Build a minimal Routine for testing.
fn make_routine(user_id: &str, name: &str) -> crate::agent::routine::Routine {
let now = chrono::Utc::now();
crate::agent::routine::Routine {
id: Uuid::new_v4(),
name: name.to_string(),
description: format!("Test routine: {name}"),
user_id: user_id.to_string(),
enabled: true,
trigger: crate::agent::routine::Trigger::Cron {
schedule: "0 9 * * *".to_string(),
timezone: None,
},
action: crate::agent::routine::RoutineAction::Lightweight {
prompt: "hello".to_string(),
context_paths: vec![],
max_tokens: 1024,
use_tools: false,
max_tool_rounds: 3,
},
guardrails: crate::agent::routine::RoutineGuardrails {
cooldown: Duration::from_secs(60),
max_concurrent: 1,
dedup_window: None,
},
notify: crate::agent::routine::NotifyConfig {
channel: None,
user: None,
on_success: false,
on_failure: true,
on_attention: true,
},
last_run_at: None,
next_fire_at: None,
run_count: 0,
consecutive_failures: 0,
state: serde_json::json!({}),
created_at: now,
updated_at: now,
}
}
/// Build a minimal SandboxJobRecord for testing.
fn make_sandbox_job(user_id: &str, task: &str) -> crate::history::SandboxJobRecord {
let now = chrono::Utc::now();
crate::history::SandboxJobRecord {
id: Uuid::new_v4(),
task: task.to_string(),
status: "completed".to_string(),
user_id: user_id.to_string(),
project_dir: format!("/tmp/test-{}", Uuid::new_v4()),
success: Some(true),
failure_reason: None,
created_at: now,
started_at: Some(now),
completed_at: Some(now),
credential_grants_json: "[]".to_string(),
}
}
// ═══════════════════════════════════════════════════════════════════════
// WorkspacePool Tests
// ═══════════════════════════════════════════════════════════════════════
#[cfg(feature = "libsql")]
mod workspace_pool {
use super::*;
use crate::config::{WorkspaceConfig, WorkspaceSearchConfig};
use crate::workspace::EmbeddingCacheConfig;
use crate::workspace::layer::MemoryLayer;
#[tokio::test]
async fn test_workspace_pool_applies_search_config() {
let (db, _dir) = test_db().await;
let search_config = WorkspaceSearchConfig {
rrf_k: 42,
..Default::default()
};
let pool = WorkspacePool::new(
db,
None,
EmbeddingCacheConfig::default(),
search_config,
WorkspaceConfig::default(),
);
let identity = UserIdentity {
user_id: "alice".to_string(),
workspace_read_scopes: vec![],
};
let ws = pool.get_or_create(&identity).await;
assert_eq!(ws.user_id(), "alice");
}
#[tokio::test]
async fn test_workspace_pool_applies_memory_layers() {
let (db, _dir) = test_db().await;
let layers = vec![MemoryLayer {
name: "shared-layer".to_string(),
scope: "shared".to_string(),
writable: false,
sensitivity: Default::default(),
}];
let ws_config = WorkspaceConfig {
memory_layers: layers,
read_scopes: vec![],
};
let pool = WorkspacePool::new(
db,
None,
EmbeddingCacheConfig::default(),
WorkspaceSearchConfig::default(),
ws_config,
);
let identity = UserIdentity {
user_id: "alice".to_string(),
workspace_read_scopes: vec![],
};
let ws = pool.get_or_create(&identity).await;
// Memory layer scope "shared" should appear in read_user_ids.
assert!(
ws.read_user_ids().contains(&"shared".to_string()),
"expected 'shared' in read_user_ids, got {:?}",
ws.read_user_ids()
);
}
#[tokio::test]
async fn test_workspace_pool_applies_identity_read_scopes() {
let (db, _dir) = test_db().await;
let pool = WorkspacePool::new(
db,
None,
EmbeddingCacheConfig::default(),
WorkspaceSearchConfig::default(),
WorkspaceConfig::default(),
);
let identity = UserIdentity {
user_id: "bob".to_string(),
workspace_read_scopes: vec!["alice".to_string(), "shared".to_string()],
};
let ws = pool.get_or_create(&identity).await;
assert_eq!(ws.user_id(), "bob");
assert!(
ws.read_user_ids().contains(&"alice".to_string()),
"expected 'alice' in read_user_ids from identity scopes"
);
assert!(
ws.read_user_ids().contains(&"shared".to_string()),
"expected 'shared' in read_user_ids from identity scopes"
);
}
#[tokio::test]
async fn test_workspace_pool_caches_per_user() {
let (db, _dir) = test_db().await;
let pool = WorkspacePool::new(
db,
None,
EmbeddingCacheConfig::default(),
WorkspaceSearchConfig::default(),
WorkspaceConfig::default(),
);
let alice_id = UserIdentity {
user_id: "alice".to_string(),
workspace_read_scopes: vec![],
};
let bob_id = UserIdentity {
user_id: "bob".to_string(),
workspace_read_scopes: vec![],
};
let alice_ws1 = pool.get_or_create(&alice_id).await;
let alice_ws2 = pool.get_or_create(&alice_id).await;
let bob_ws = pool.get_or_create(&bob_id).await;
// Same user gets the same Arc.
assert!(Arc::ptr_eq(&alice_ws1, &alice_ws2));
// Different users get different instances.
assert!(!Arc::ptr_eq(&alice_ws1, &bob_ws));
assert_eq!(alice_ws1.user_id(), "alice");
assert_eq!(bob_ws.user_id(), "bob");
}
#[tokio::test]
async fn test_workspace_pool_combines_global_and_identity_scopes() {
let (db, _dir) = test_db().await;
let ws_config = WorkspaceConfig {
memory_layers: vec![],
read_scopes: vec!["global-shared".to_string()],
};
let pool = WorkspacePool::new(
db,
None,
EmbeddingCacheConfig::default(),
WorkspaceSearchConfig::default(),
ws_config,
);
let identity = UserIdentity {
user_id: "alice".to_string(),
workspace_read_scopes: vec!["token-scope".to_string()],
};
let ws = pool.get_or_create(&identity).await;
let scopes = ws.read_user_ids();
// Primary scope
assert!(scopes.contains(&"alice".to_string()));
// Global config scope
assert!(
scopes.contains(&"global-shared".to_string()),
"expected global scope 'global-shared', got {:?}",
scopes
);
// Token identity scope
assert!(
scopes.contains(&"token-scope".to_string()),
"expected token scope 'token-scope', got {:?}",
scopes
);
}
}
// ═══════════════════════════════════════════════════════════════════════
// Jobs Handler Isolation Tests
// ═══════════════════════════════════════════════════════════════════════
#[cfg(feature = "libsql")]
mod jobs_isolation {
use super::*;
use crate::channels::web::handlers::jobs::{
jobs_cancel_handler, jobs_prompt_handler, jobs_restart_handler, jobs_summary_handler,
};
// SandboxStore methods are accessed through the Database supertrait.
/// Build a router with job endpoints behind multi-user auth.
fn jobs_router(state: Arc<GatewayState>, auth: MultiAuthState) -> Router {
Router::new()
.route("/api/jobs/summary", get(jobs_summary_handler))
.route("/api/jobs/{id}/cancel", post(jobs_cancel_handler))
.route("/api/jobs/{id}/restart", post(jobs_restart_handler))
.route("/api/jobs/{id}/prompt", post(jobs_prompt_handler))
.layer(middleware::from_fn_with_state(auth, auth_middleware))
.with_state(state)
}
#[tokio::test]
async fn test_jobs_summary_scoped_to_user() {
let (db, _dir) = test_db().await;
// Insert sandbox jobs for alice and bob.
let alice_job = make_sandbox_job("alice", "alice task");
let bob_job = make_sandbox_job("bob", "bob task");
db.save_sandbox_job(&alice_job).await.unwrap();
db.save_sandbox_job(&bob_job).await.unwrap();
let state = build_state(Some(db), None);
let auth = two_user_auth();
let app = jobs_router(state, auth);
// Alice should see 1 job.
let req = Request::builder()
.uri("/api/jobs/summary")
.header("Authorization", "Bearer tok-alice")
.body(Body::empty())
.unwrap();
let resp = app.clone().oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
let body: serde_json::Value =
serde_json::from_slice(&axum::body::to_bytes(resp.into_body(), 4096).await.unwrap())
.unwrap();
assert_eq!(body["total"], 1, "alice should see only her own jobs");
// Bob should see 1 job.
let req = Request::builder()
.uri("/api/jobs/summary")
.header("Authorization", "Bearer tok-bob")
.body(Body::empty())
.unwrap();
let resp = app.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
let body: serde_json::Value =
serde_json::from_slice(&axum::body::to_bytes(resp.into_body(), 4096).await.unwrap())
.unwrap();
assert_eq!(body["total"], 1, "bob should see only his own jobs");
}
#[tokio::test]
async fn test_jobs_restart_rejects_other_user() {
let (db, _dir) = test_db().await;
// Insert a failed sandbox job owned by alice.
let mut alice_job = make_sandbox_job("alice", "alice task");
alice_job.status = "failed".to_string();
alice_job.success = Some(false);
db.save_sandbox_job(&alice_job).await.unwrap();
let state = build_state(Some(db), None);
let auth = two_user_auth();
let app = jobs_router(state, auth);
// Bob tries to restart alice's job.
let req = Request::builder()
.method(Method::POST)
.uri(format!("/api/jobs/{}/restart", alice_job.id))
.header("Authorization", "Bearer tok-bob")
.body(Body::empty())
.unwrap();
let resp = app.oneshot(req).await.unwrap();
assert_eq!(
resp.status(),
StatusCode::NOT_FOUND,
"bob should not be able to restart alice's job"
);
}
#[tokio::test]
async fn test_jobs_prompt_works_for_agent_jobs() {
let (db, _dir) = test_db().await;
// Insert a running sandbox job owned by alice in claude_code mode.
let mut alice_job = make_sandbox_job("alice", "prompt test");
alice_job.status = "running".to_string();
alice_job.success = None;
alice_job.completed_at = None;
db.save_sandbox_job(&alice_job).await.unwrap();
db.update_sandbox_job_mode(alice_job.id, "claude_code")
.await
.unwrap();
let prompt_queue: PromptQueue =
Arc::new(tokio::sync::Mutex::new(std::collections::HashMap::new()));
let state = build_state(Some(db), Some(prompt_queue.clone()));
let auth = two_user_auth();
let app = jobs_router(state, auth);
// Alice prompts her own job.
let req = Request::builder()
.method(Method::POST)
.uri(format!("/api/jobs/{}/prompt", alice_job.id))
.header("Authorization", "Bearer tok-alice")
.header("Content-Type", "application/json")
.body(Body::from(
serde_json::to_string(&serde_json::json!({"content": "hello"})).unwrap(),
))
.unwrap();
let resp = app.oneshot(req).await.unwrap();
assert_eq!(
resp.status(),
StatusCode::OK,
"alice should be able to prompt her own job"
);
// Verify prompt was enqueued.
let queue = prompt_queue.lock().await;
assert!(
queue.contains_key(&alice_job.id),
"prompt queue should contain alice's job"
);
}
#[tokio::test]
async fn test_jobs_prompt_rejects_other_user() {
let (db, _dir) = test_db().await;
let mut alice_job = make_sandbox_job("alice", "alice task");
alice_job.status = "running".to_string();
alice_job.success = None;
alice_job.completed_at = None;
db.save_sandbox_job(&alice_job).await.unwrap();
db.update_sandbox_job_mode(alice_job.id, "claude_code")
.await
.unwrap();
let prompt_queue: PromptQueue =
Arc::new(tokio::sync::Mutex::new(std::collections::HashMap::new()));
let state = build_state(Some(db), Some(prompt_queue));
let auth = two_user_auth();
let app = jobs_router(state, auth);
// Bob tries to prompt alice's job.
let req = Request::builder()
.method(Method::POST)
.uri(format!("/api/jobs/{}/prompt", alice_job.id))
.header("Authorization", "Bearer tok-bob")
.header("Content-Type", "application/json")
.body(Body::from(
serde_json::to_string(&serde_json::json!({"content": "sneaky"})).unwrap(),
))
.unwrap();
let resp = app.oneshot(req).await.unwrap();
assert_eq!(
resp.status(),
StatusCode::NOT_FOUND,
"bob should not be able to prompt alice's job"
);
}
#[tokio::test]
async fn test_jobs_cancel_rejects_other_user() {
let (db, _dir) = test_db().await;
let mut alice_job = make_sandbox_job("alice", "alice running");
alice_job.status = "running".to_string();
alice_job.success = None;
alice_job.completed_at = None;
db.save_sandbox_job(&alice_job).await.unwrap();
let state = build_state(Some(db), None);
let auth = two_user_auth();
let app = jobs_router(state, auth);
// Bob tries to cancel alice's job.
let req = Request::builder()
.method(Method::POST)
.uri(format!("/api/jobs/{}/cancel", alice_job.id))
.header("Authorization", "Bearer tok-bob")
.body(Body::empty())
.unwrap();
let resp = app.oneshot(req).await.unwrap();
assert_eq!(
resp.status(),
StatusCode::NOT_FOUND,
"bob should not be able to cancel alice's job"
);
}
}
// ═══════════════════════════════════════════════════════════════════════
// Routines Isolation Tests
// ═══════════════════════════════════════════════════════════════════════
#[cfg(feature = "libsql")]
mod routines_isolation {
use super::*;
use crate::channels::web::handlers::routines::{
routines_delete_handler, routines_detail_handler, routines_list_handler,
routines_summary_handler, routines_toggle_handler,
};
// RoutineStore methods are accessed through the Database supertrait.
fn routines_router(state: Arc<GatewayState>, auth: MultiAuthState) -> Router {
Router::new()
.route("/api/routines", get(routines_list_handler))
.route("/api/routines/summary", get(routines_summary_handler))
.route("/api/routines/{id}", get(routines_detail_handler))
.route("/api/routines/{id}/toggle", post(routines_toggle_handler))
.route("/api/routines/{id}", delete(routines_delete_handler))
.layer(middleware::from_fn_with_state(auth, auth_middleware))
.with_state(state)
}
#[tokio::test]
async fn test_routines_isolation() {
let (db, _dir) = test_db().await;
// Create routines for alice and bob.
let alice_routine = make_routine("alice", "alice-daily");
let bob_routine = make_routine("bob", "bob-daily");
db.create_routine(&alice_routine).await.unwrap();
db.create_routine(&bob_routine).await.unwrap();
let state = build_state(Some(db), None);
let auth = two_user_auth();
let app = routines_router(state, auth);
// Alice sees only her routine in the list.
let req = Request::builder()
.uri("/api/routines")
.header("Authorization", "Bearer tok-alice")
.body(Body::empty())
.unwrap();
let resp = app.clone().oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
let body: serde_json::Value =
serde_json::from_slice(&axum::body::to_bytes(resp.into_body(), 8192).await.unwrap())
.unwrap();
let routines = body["routines"].as_array().unwrap();
assert_eq!(routines.len(), 1, "alice should see only her routines");
assert_eq!(routines[0]["name"], "alice-daily");
// Bob sees only his routine.
let req = Request::builder()
.uri("/api/routines")
.header("Authorization", "Bearer tok-bob")
.body(Body::empty())
.unwrap();
let resp = app.clone().oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
let body: serde_json::Value =
serde_json::from_slice(&axum::body::to_bytes(resp.into_body(), 8192).await.unwrap())
.unwrap();
let routines = body["routines"].as_array().unwrap();
assert_eq!(routines.len(), 1, "bob should see only his routines");
assert_eq!(routines[0]["name"], "bob-daily");
// Bob cannot view alice's routine detail.
let req = Request::builder()
.uri(format!("/api/routines/{}", alice_routine.id))
.header("Authorization", "Bearer tok-bob")
.body(Body::empty())
.unwrap();
let resp = app.clone().oneshot(req).await.unwrap();
assert_eq!(
resp.status(),
StatusCode::NOT_FOUND,
"bob should not see alice's routine detail"
);
// Bob cannot toggle alice's routine.
let req = Request::builder()
.method(Method::POST)
.uri(format!("/api/routines/{}/toggle", alice_routine.id))
.header("Authorization", "Bearer tok-bob")
.body(Body::empty())
.unwrap();
let resp = app.clone().oneshot(req).await.unwrap();
assert_eq!(
resp.status(),
StatusCode::NOT_FOUND,
"bob should not toggle alice's routine"
);
// Bob cannot delete alice's routine.
let req = Request::builder()
.method(Method::DELETE)
.uri(format!("/api/routines/{}", alice_routine.id))
.header("Authorization", "Bearer tok-bob")
.body(Body::empty())
.unwrap();
let resp = app.oneshot(req).await.unwrap();
assert_eq!(
resp.status(),
StatusCode::NOT_FOUND,
"bob should not delete alice's routine"
);
}
}
// ═══════════════════════════════════════════════════════════════════════
// Handler Auth Enforcement Tests
// ═══════════════════════════════════════════════════════════════════════
mod auth_enforcement {
use super::*;
/// Dummy handler that extracts `AuthenticatedUser` — if the auth middleware
/// rejects the request, this handler is never reached.
async fn authed_handler(AuthenticatedUser(_user): AuthenticatedUser) -> &'static str {
"ok"
}
/// Build a router with the real auth middleware and dummy handlers at all
/// the paths we want to verify require authentication.
fn auth_test_router(auth: MultiAuthState) -> Router {
let state = build_state(None, None);
Router::new()
// Routines
.route("/api/routines", get(authed_handler))
.route("/api/routines/summary", get(authed_handler))
.route("/api/routines/{id}", get(authed_handler))
.route("/api/routines/{id}/toggle", post(authed_handler))
.route("/api/routines/{id}", delete(authed_handler))
// Skills
.route("/api/skills", get(authed_handler))
.route("/api/skills/search", post(authed_handler))
.route("/api/skills/install", post(authed_handler))
.route("/api/skills/{name}", delete(authed_handler))
// Logs
.route("/api/logs/events", get(authed_handler))
.route("/api/logs/level", get(authed_handler).put(authed_handler))
// Gateway status
.route("/api/gateway/status", get(authed_handler))
.layer(middleware::from_fn_with_state(auth, auth_middleware))
.with_state(state)
}
/// Send a request without auth and assert it returns UNAUTHORIZED.
async fn assert_requires_auth(app: &Router, method: Method, uri: &str) {
let req = Request::builder()
.method(method.clone())
.uri(uri)
.body(Body::empty())
.unwrap();
let resp = app.clone().oneshot(req).await.unwrap();
assert_eq!(
resp.status(),
StatusCode::UNAUTHORIZED,
"{} {} should require auth",
method,
uri
);
}
/// Send a request with a valid token and assert it succeeds.
async fn assert_passes_with_token(app: &Router, method: Method, uri: &str, token: &str) {
let req = Request::builder()
.method(method.clone())
.uri(uri)
.header("Authorization", format!("Bearer {token}"))
.body(Body::empty())
.unwrap();
let resp = app.clone().oneshot(req).await.unwrap();
assert_eq!(
resp.status(),
StatusCode::OK,
"{} {} should pass with valid token",
method,
uri
);
}
#[tokio::test]
async fn test_routines_handlers_require_auth() {
let auth = MultiAuthState::single("secret-tok".to_string(), "user".to_string());
let app = auth_test_router(auth);
let id = Uuid::new_v4();
assert_requires_auth(&app, Method::GET, "/api/routines").await;
assert_requires_auth(&app, Method::GET, "/api/routines/summary").await;
assert_requires_auth(&app, Method::GET, &format!("/api/routines/{id}")).await;
assert_requires_auth(&app, Method::POST, &format!("/api/routines/{id}/toggle")).await;
assert_requires_auth(&app, Method::DELETE, &format!("/api/routines/{id}")).await;
}
#[tokio::test]
async fn test_skills_handlers_require_auth() {
let auth = MultiAuthState::single("secret-tok".to_string(), "user".to_string());
let app = auth_test_router(auth);
assert_requires_auth(&app, Method::GET, "/api/skills").await;
assert_requires_auth(&app, Method::POST, "/api/skills/search").await;
assert_requires_auth(&app, Method::POST, "/api/skills/install").await;
assert_requires_auth(&app, Method::DELETE, "/api/skills/test-skill").await;
}
#[tokio::test]
async fn test_logs_handlers_require_auth() {
let auth = MultiAuthState::single("secret-tok".to_string(), "user".to_string());
let app = auth_test_router(auth);
assert_requires_auth(&app, Method::GET, "/api/logs/events").await;
assert_requires_auth(&app, Method::GET, "/api/logs/level").await;
assert_requires_auth(&app, Method::PUT, "/api/logs/level").await;
}
#[tokio::test]
async fn test_gateway_status_requires_auth() {
let auth = MultiAuthState::single("secret-tok".to_string(), "user".to_string());
let app = auth_test_router(auth);
assert_requires_auth(&app, Method::GET, "/api/gateway/status").await;
}
#[tokio::test]
async fn test_valid_token_passes_all_endpoints() {
let auth = MultiAuthState::single("secret-tok".to_string(), "user".to_string());
let app = auth_test_router(auth);
let id = Uuid::new_v4();
assert_passes_with_token(&app, Method::GET, "/api/routines", "secret-tok").await;
assert_passes_with_token(&app, Method::GET, "/api/skills", "secret-tok").await;
assert_passes_with_token(&app, Method::GET, "/api/logs/events", "secret-tok").await;
assert_passes_with_token(&app, Method::GET, "/api/gateway/status", "secret-tok").await;
assert_passes_with_token(
&app,
Method::GET,
&format!("/api/routines/{id}"),
"secret-tok",
)
.await;
}
#[tokio::test]
async fn test_wrong_token_rejected_on_all_endpoints() {
let auth = MultiAuthState::single("secret-tok".to_string(), "user".to_string());
let app = auth_test_router(auth);
// Wrong token should be rejected.
let req = Request::builder()
.uri("/api/routines")
.header("Authorization", "Bearer wrong-tok")
.body(Body::empty())
.unwrap();
let resp = app.clone().oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
let req = Request::builder()
.uri("/api/gateway/status")
.header("Authorization", "Bearer wrong-tok")
.body(Body::empty())
.unwrap();
let resp = app.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
}
}
+24 -13
View File
@@ -62,7 +62,11 @@ impl Default for WsConnectionTracker {
///
/// When either task ends (client disconnect or broadcast closed), both are
/// cleaned up.
pub async fn handle_ws_connection(socket: WebSocket, state: Arc<GatewayState>) {
pub async fn handle_ws_connection(
socket: WebSocket,
state: Arc<GatewayState>,
user: crate::channels::web::auth::UserIdentity,
) {
let (mut ws_sink, mut ws_stream) = socket.split();
// Track connection
@@ -71,9 +75,9 @@ pub async fn handle_ws_connection(socket: WebSocket, state: Arc<GatewayState>) {
}
let tracker_for_drop = state.ws_tracker.clone();
// Subscribe to broadcast events (same source as SSE).
// Subscribe to broadcast events (same source as SSE), scoped to this user.
// Reject if we've hit the connection limit.
let Some(raw_stream) = state.sse.subscribe_raw() else {
let Some(raw_stream) = state.sse.subscribe_raw(Some(user.user_id.clone())) else {
tracing::warn!("WebSocket rejected: too many connections");
// Decrement the WS tracker we already incremented above.
if let Some(ref tracker) = tracker_for_drop {
@@ -117,7 +121,7 @@ pub async fn handle_ws_connection(socket: WebSocket, state: Arc<GatewayState>) {
});
// Receiver task: read client frames and route to agent
let user_id = state.user_id.clone();
let user_id = user.user_id;
while let Some(Ok(frame)) = ws_stream.next().await {
match frame {
Message::Text(text) => {
@@ -263,10 +267,14 @@ async fn handle_client_message(
token,
} => {
if let Some(ref ext_mgr) = state.extension_manager {
match ext_mgr.configure_token(&extension_name, &token).await {
match ext_mgr
.configure_token(&extension_name, &token, user_id)
.await
{
Ok(result) => {
if result.verification.is_some() {
state.sse.broadcast(
state.sse.broadcast_for_user(
user_id,
crate::channels::web::types::SseEvent::AuthRequired {
extension_name: extension_name.clone(),
instructions: Some(result.message),
@@ -275,8 +283,9 @@ async fn handle_client_message(
},
);
} else {
crate::channels::web::server::clear_auth_mode(state).await;
state.sse.broadcast(
crate::channels::web::server::clear_auth_mode(state, user_id).await;
state.sse.broadcast_for_user(
user_id,
crate::channels::web::types::SseEvent::AuthCompleted {
extension_name,
success: true,
@@ -288,7 +297,8 @@ async fn handle_client_message(
Err(e) => {
let msg = format!("Auth failed: {}", e);
if matches!(e, crate::extensions::ExtensionError::ValidationFailed(_)) {
state.sse.broadcast(
state.sse.broadcast_for_user(
user_id,
crate::channels::web::types::SseEvent::AuthRequired {
extension_name: extension_name.clone(),
instructions: Some(msg.clone()),
@@ -311,7 +321,7 @@ async fn handle_client_message(
}
}
WsClientMessage::AuthCancel { .. } => {
crate::channels::web::server::clear_auth_mode(state).await;
crate::channels::web::server::clear_auth_mode(state, user_id).await;
}
WsClientMessage::Ping => {
let _ = direct_tx.send(WsServerMessage::Pong).await;
@@ -498,8 +508,9 @@ mod tests {
GatewayState {
msg_tx: tokio::sync::RwLock::new(msg_tx),
sse: SseManager::new(),
sse: Arc::new(SseManager::new()),
workspace: None,
workspace_pool: None,
session_manager: None,
log_broadcaster: None,
log_level_handle: None,
@@ -509,13 +520,13 @@ mod tests {
job_manager: None,
prompt_queue: None,
scheduler: None,
user_id: "test".to_string(),
default_user_id: "test".to_string(),
shutdown_tx: tokio::sync::RwLock::new(None),
ws_tracker: Some(Arc::new(WsConnectionTracker::new())),
llm_provider: None,
skill_registry: None,
skill_catalog: None,
chat_rate_limiter: crate::channels::web::server::RateLimiter::new(30, 60),
chat_rate_limiter: crate::channels::web::server::PerUserRateLimiter::new(30, 60),
oauth_rate_limiter: crate::channels::web::server::RateLimiter::new(10, 60),
webhook_rate_limiter: crate::channels::web::server::RateLimiter::new(10, 60),
registry_entries: Vec::new(),
+2 -2
View File
@@ -447,8 +447,8 @@ pub struct PendingOAuthFlow {
pub user_id: String,
/// Secrets store reference for token persistence.
pub secrets: Arc<dyn SecretsStore + Send + Sync>,
/// SSE broadcast sender for notifying the web UI.
pub sse_sender: Option<tokio::sync::broadcast::Sender<crate::channels::web::types::SseEvent>>,
/// SSE broadcast manager for notifying the web UI.
pub sse_manager: Option<Arc<crate::channels::web::sse::SseManager>>,
/// Gateway auth token for authenticating with the platform token exchange proxy.
pub gateway_token: Option<String>,
/// Additional form params for the token exchange request.
+66 -10
View File
@@ -10,7 +10,7 @@ use clap::Subcommand;
use uuid::Uuid;
use crate::agent::routine::{
NotifyConfig, Routine, RoutineAction, RoutineGuardrails, Trigger, next_cron_fire,
NotifyConfig, Routine, RoutineAction, RoutineGuardrails, RunStatus, Trigger, next_cron_fire,
};
use crate::db::Database;
@@ -251,15 +251,26 @@ async fn list(
);
println!("{}", "-".repeat(130));
// Fetch last-run status for all routines in a single batch query
let routine_ids: Vec<Uuid> = filtered.iter().map(|r| r.id).collect();
let last_run_results = db
.batch_get_last_run_status(&routine_ids)
.await
.unwrap_or_default();
for r in &filtered {
let status = if r.enabled {
if r.consecutive_failures > 0 {
format!("err({})", r.consecutive_failures)
} else {
"active".to_string()
}
} else {
let last_run_status = last_run_results.get(&r.id).copied();
let status = if !r.enabled {
"disabled".to_string()
} else if last_run_status == Some(RunStatus::Running) {
"running".to_string()
} else if r.consecutive_failures > 0 {
format!("err({})", r.consecutive_failures)
} else if last_run_status == Some(RunStatus::Attention) {
"attention".to_string()
} else {
"active".to_string()
};
let next_fire = r
@@ -340,8 +351,8 @@ async fn create(
prompt: prompt.to_string(),
context_paths: Vec::new(),
max_tokens: 4096,
use_tools: false,
max_tool_rounds: 0,
use_tools: true,
max_tool_rounds: 3,
},
guardrails: RoutineGuardrails {
cooldown: std::time::Duration::from_secs(cooldown_secs),
@@ -685,6 +696,7 @@ fn truncate(s: &str, max_chars: usize) -> String {
#[cfg(test)]
mod tests {
use super::*;
use crate::agent::routine::RoutineAction;
#[test]
fn format_relative_future() {
@@ -743,4 +755,48 @@ mod tests {
assert!(notify.on_failure); // safety: test-only assertion
assert!(!notify.on_success); // safety: test-only assertion
}
#[cfg(feature = "libsql")]
#[tokio::test]
async fn cli_create_defaults_lightweight_routines_to_tools_enabled() {
let harness = crate::testing::TestHarnessBuilder::new().build().await;
let db = harness.db.clone();
run_routines_command(
RoutinesCommand::Create {
name: "cli-digest".to_string(),
schedule: "0 0 9 * * *".to_string(),
prompt: "Prepare the morning digest.".to_string(),
description: "CLI created routine".to_string(),
timezone: Some("UTC".to_string()),
cooldown: 300,
notify_channel: None,
},
db.clone(),
"user1",
)
.await
.expect("create routine");
let routine = db
.get_routine_by_name("user1", "cli-digest")
.await
.expect("get routine by name")
.expect("cli-digest should exist");
match routine.action {
RoutineAction::Lightweight {
use_tools,
max_tool_rounds,
..
} => {
assert!(
use_tools,
"CLI-created lightweight routines should default to tools"
);
assert_eq!(max_tool_rounds, 3);
}
other => panic!("expected lightweight action, got {other:?}"),
}
}
}
+142
View File
@@ -2,6 +2,7 @@ use std::collections::HashMap;
use std::path::PathBuf;
use secrecy::SecretString;
use serde::Deserialize;
use crate::bootstrap::ironclaw_base_dir;
use crate::config::helpers::{optional_env, parse_bool_env, parse_optional_env};
@@ -45,6 +46,26 @@ pub struct GatewayConfig {
/// Bearer token for authentication. Random hex generated at startup if unset.
pub auth_token: Option<String>,
pub user_id: String,
/// Additional user scopes for workspace reads.
///
/// When set, the workspace will be able to read (search, read, list) from
/// these additional user scopes while writes remain isolated to `user_id`.
/// Parsed from `WORKSPACE_READ_SCOPES` (comma-separated).
pub workspace_read_scopes: Vec<String>,
/// Memory layer definitions (JSON in env var, or from external config).
pub memory_layers: Vec<crate::workspace::layer::MemoryLayer>,
/// Multi-user token map. When set, each token maps to a user identity.
/// Parsed from `GATEWAY_USER_TOKENS` (JSON string). When absent, falls back
/// to single-user mode via `auth_token` + `user_id`.
pub user_tokens: Option<HashMap<String, UserTokenConfig>>,
}
/// Per-user token configuration for multi-user mode.
#[derive(Debug, Clone, Deserialize)]
pub struct UserTokenConfig {
pub user_id: String,
#[serde(default)]
pub workspace_read_scopes: Vec<String>,
}
/// Signal channel configuration (signal-cli daemon HTTP/JSON-RPC).
@@ -115,6 +136,118 @@ impl ChannelsConfig {
.or_else(|| cs.gateway_user_id.clone())
.unwrap_or_else(|| owner_id.to_string());
let memory_layers: Vec<crate::workspace::layer::MemoryLayer> =
match optional_env("MEMORY_LAYERS")? {
Some(json_str) => {
serde_json::from_str(&json_str).map_err(|e| ConfigError::InvalidValue {
key: "MEMORY_LAYERS".to_string(),
message: format!("must be valid JSON array of layer objects: {e}"),
})?
}
None => crate::workspace::layer::MemoryLayer::default_for_user(&user_id),
};
// Validate layer names and scopes
for layer in &memory_layers {
if layer.name.trim().is_empty() {
return Err(ConfigError::InvalidValue {
key: "MEMORY_LAYERS".to_string(),
message: "layer name must not be empty".to_string(),
});
}
if layer.name.len() > 64 {
return Err(ConfigError::InvalidValue {
key: "MEMORY_LAYERS".to_string(),
message: format!("layer name '{}' exceeds 64 characters", layer.name),
});
}
if !layer
.name
.chars()
.all(|c| c.is_ascii_alphanumeric() || c == '_' || c == '-')
{
return Err(ConfigError::InvalidValue {
key: "MEMORY_LAYERS".to_string(),
message: format!(
"layer name '{}' contains invalid characters \
(allowed: a-z, A-Z, 0-9, _, -)",
layer.name
),
});
}
if layer.scope.trim().is_empty() {
return Err(ConfigError::InvalidValue {
key: "MEMORY_LAYERS".to_string(),
message: format!("layer '{}' has an empty scope", layer.name),
});
}
}
// Check for duplicate layer names
{
let mut seen = std::collections::HashSet::new();
for layer in &memory_layers {
if !seen.insert(&layer.name) {
return Err(ConfigError::InvalidValue {
key: "MEMORY_LAYERS".to_string(),
message: format!("duplicate layer name '{}'", layer.name),
});
}
}
}
let user_tokens: Option<HashMap<String, UserTokenConfig>> =
match optional_env("GATEWAY_USER_TOKENS")? {
Some(json_str) => {
let tokens: HashMap<String, UserTokenConfig> = serde_json::from_str(
&json_str,
)
.map_err(|e| ConfigError::InvalidValue {
key: "GATEWAY_USER_TOKENS".to_string(),
message: format!(
"must be valid JSON object mapping tokens to user configs: {e}"
),
})?;
if tokens.is_empty() {
return Err(ConfigError::InvalidValue {
key: "GATEWAY_USER_TOKENS".to_string(),
message:
"token map is empty — remove the variable to use single-user mode"
.to_string(),
});
}
for (tok, cfg) in &tokens {
if cfg.user_id.trim().is_empty() {
return Err(ConfigError::InvalidValue {
key: "GATEWAY_USER_TOKENS".to_string(),
message: format!(
"token '{}...' has an empty user_id",
&tok[..tok.len().min(8)]
),
});
}
}
Some(tokens)
}
None => None,
};
let workspace_read_scopes: Vec<String> = optional_env("WORKSPACE_READ_SCOPES")?
.map(|s| {
s.split(',')
.map(|s| s.trim().to_string())
.filter(|s| !s.is_empty())
.collect()
})
.unwrap_or_default();
for scope in &workspace_read_scopes {
if scope.len() > 128 {
return Err(ConfigError::InvalidValue {
key: "WORKSPACE_READ_SCOPES".to_string(),
message: format!("scope '{}...' exceeds 128 characters", &scope[..32]),
});
}
}
Some(GatewayConfig {
host: optional_env("GATEWAY_HOST")?
.or_else(|| cs.gateway_host.clone())
@@ -126,6 +259,9 @@ impl ChannelsConfig {
auth_token: optional_env("GATEWAY_AUTH_TOKEN")?
.or_else(|| cs.gateway_auth_token.clone()),
user_id,
workspace_read_scopes,
memory_layers,
user_tokens,
})
} else {
None
@@ -281,6 +417,9 @@ mod tests {
port: 3000,
auth_token: Some("tok-abc".to_string()),
user_id: "default".to_string(),
workspace_read_scopes: vec![],
memory_layers: vec![],
user_tokens: None,
};
assert_eq!(cfg.host, "127.0.0.1");
assert_eq!(cfg.port, 3000);
@@ -295,6 +434,9 @@ mod tests {
port: 3001,
auth_token: None,
user_id: "anon".to_string(),
workspace_read_scopes: vec![],
memory_layers: vec![],
user_tokens: None,
};
assert!(cfg.auth_token.is_none());
}
+12 -4
View File
@@ -406,7 +406,7 @@ impl LlmConfig {
// Resolve extra headers
let extra_headers = if let Some(env_var) = extra_headers_env {
optional_env(env_var)?
.map(|val| parse_extra_headers(&val))
.map(|val| parse_extra_headers_with_key(&val, env_var))
.transpose()?
.unwrap_or_default()
} else {
@@ -475,7 +475,10 @@ impl LlmConfig {
///
/// Format: `Key1:Value1,Key2:Value2` (colon-separated, not `=`, because
/// header values often contain `=`).
fn parse_extra_headers(val: &str) -> Result<Vec<(String, String)>, ConfigError> {
fn parse_extra_headers_with_key(
val: &str,
env_var_name: &str,
) -> Result<Vec<(String, String)>, ConfigError> {
if val.trim().is_empty() {
return Ok(Vec::new());
}
@@ -488,14 +491,14 @@ fn parse_extra_headers(val: &str) -> Result<Vec<(String, String)>, ConfigError>
}
let Some((key, value)) = pair.split_once(':') else {
return Err(ConfigError::InvalidValue {
key: "LLM_EXTRA_HEADERS".to_string(),
key: env_var_name.to_string(),
message: format!("malformed header entry '{}', expected Key:Value", pair),
});
};
let key = key.trim();
if key.is_empty() {
return Err(ConfigError::InvalidValue {
key: "LLM_EXTRA_HEADERS".to_string(),
key: env_var_name.to_string(),
message: format!("empty header name in entry '{}'", pair),
});
}
@@ -536,6 +539,11 @@ mod tests {
use crate::settings::Settings;
use crate::testing::credentials::*;
/// Convenience wrapper for tests — uses "TEST_HEADERS" as the env var name.
fn parse_extra_headers(val: &str) -> Result<Vec<(String, String)>, ConfigError> {
parse_extra_headers_with_key(val, "TEST_HEADERS")
}
/// Clear all openai-compatible-related env vars.
fn clear_openai_compatible_env() {
// SAFETY: Only called under ENV_MUTEX in tests.
+69
View File
@@ -230,6 +230,49 @@ impl JobStore for LibSqlBackend {
Ok(jobs)
}
async fn list_agent_jobs_for_user(
&self,
user_id: &str,
) -> Result<Vec<AgentJobRecord>, DatabaseError> {
let conn = self.connect().await?;
let mut rows = conn
.query(
r#"
SELECT id, title, status, user_id, failure_reason,
created_at, started_at, completed_at
FROM agent_jobs WHERE source = 'direct' AND user_id = ?1
ORDER BY created_at DESC
"#,
params![user_id],
)
.await
.map_err(|e| DatabaseError::Query(e.to_string()))?;
let mut jobs = Vec::new();
while let Some(row) = rows
.next()
.await
.map_err(|e| DatabaseError::Query(e.to_string()))?
{
let id_str = get_text(&row, 0);
let Ok(id) = id_str.parse() else {
tracing::warn!("Skipping agent job with invalid UUID: {}", id_str);
continue;
};
jobs.push(AgentJobRecord {
id,
title: get_text(&row, 1),
status: get_text(&row, 2),
user_id: get_text(&row, 3),
failure_reason: get_opt_text(&row, 4),
created_at: get_ts(&row, 5),
started_at: get_opt_ts(&row, 6),
completed_at: get_opt_ts(&row, 7),
});
}
Ok(jobs)
}
async fn get_agent_job_failure_reason(
&self,
id: Uuid,
@@ -277,6 +320,32 @@ impl JobStore for LibSqlBackend {
Ok(summary)
}
async fn agent_job_summary_for_user(
&self,
user_id: &str,
) -> Result<AgentJobSummary, DatabaseError> {
let conn = self.connect().await?;
let mut rows = conn
.query(
"SELECT status, COUNT(*) as cnt FROM agent_jobs WHERE source = 'direct' AND user_id = ?1 GROUP BY status",
params![user_id],
)
.await
.map_err(|e| DatabaseError::Query(e.to_string()))?;
let mut summary = AgentJobSummary::default();
while let Some(row) = rows
.next()
.await
.map_err(|e| DatabaseError::Query(e.to_string()))?
{
let status = get_text(&row, 0);
let count = get_i64(&row, 1) as usize;
summary.add_count(&status, count);
}
Ok(summary)
}
async fn save_action(&self, job_id: Uuid, action: &ActionRecord) -> Result<(), DatabaseError> {
let conn = self.connect().await?;
let duration_ms = action.duration.as_millis() as i64;
+50
View File
@@ -462,6 +462,56 @@ impl RoutineStore for LibSqlBackend {
Ok(counts)
}
async fn batch_get_last_run_status(
&self,
routine_ids: &[Uuid],
) -> Result<HashMap<Uuid, RunStatus>, DatabaseError> {
if routine_ids.is_empty() {
return Ok(HashMap::new());
}
let conn = self.connect().await?;
// SQLite doesn't support ANY($1), so we query all latest runs and filter in memory.
// Uses a subquery to pick only the most recent run per routine.
let mut rows = conn
.query(
"SELECT routine_id, status FROM routine_runs r1
WHERE started_at = (
SELECT MAX(started_at) FROM routine_runs r2
WHERE r2.routine_id = r1.routine_id
)
GROUP BY routine_id",
params![],
)
.await
.map_err(|e| {
DatabaseError::Query(format!("Failed to batch get last run status: {}", e))
})?;
let routine_id_set: HashSet<Uuid> = routine_ids.iter().copied().collect();
let mut statuses = HashMap::new();
while let Some(row) = rows
.next()
.await
.map_err(|e| DatabaseError::Query(e.to_string()))?
{
let id_str: String = get_text(&row, 0);
let id = Uuid::parse_str(&id_str)
.map_err(|e| DatabaseError::Query(format!("Invalid routine UUID: {}", e)))?;
if routine_id_set.contains(&id) {
let status_str: String = get_text(&row, 1);
if let std::result::Result::Ok(status) = status_str.parse::<RunStatus>() {
statuses.insert(id, status);
}
}
}
Ok(statuses)
}
async fn link_routine_run_to_job(
&self,
run_id: Uuid,
+17
View File
@@ -409,7 +409,15 @@ pub trait JobStore: Send + Sync {
async fn mark_job_stuck(&self, id: Uuid) -> Result<(), DatabaseError>;
async fn get_stuck_jobs(&self) -> Result<Vec<Uuid>, DatabaseError>;
async fn list_agent_jobs(&self) -> Result<Vec<AgentJobRecord>, DatabaseError>;
async fn list_agent_jobs_for_user(
&self,
user_id: &str,
) -> Result<Vec<AgentJobRecord>, DatabaseError>;
async fn agent_job_summary(&self) -> Result<AgentJobSummary, DatabaseError>;
async fn agent_job_summary_for_user(
&self,
user_id: &str,
) -> Result<AgentJobSummary, DatabaseError>;
/// Get the failure reason for a single agent job (O(1) lookup).
async fn get_agent_job_failure_reason(&self, id: Uuid)
-> Result<Option<String>, DatabaseError>;
@@ -520,6 +528,15 @@ pub trait RoutineStore: Send + Sync {
&self,
routine_ids: &[Uuid],
) -> Result<HashMap<Uuid, i64>, DatabaseError>;
/// Fetch the last run status for multiple routines in a single query.
/// Returns a map from routine_id to its most recent RunStatus.
/// Routines with no runs are omitted from the result.
async fn batch_get_last_run_status(
&self,
routine_ids: &[Uuid],
) -> Result<HashMap<Uuid, RunStatus>, DatabaseError>;
async fn link_routine_run_to_job(
&self,
run_id: Uuid,
+22
View File
@@ -249,10 +249,24 @@ impl JobStore for PgBackend {
self.store.list_agent_jobs().await
}
async fn list_agent_jobs_for_user(
&self,
user_id: &str,
) -> Result<Vec<AgentJobRecord>, DatabaseError> {
self.store.list_agent_jobs_for_user(user_id).await
}
async fn agent_job_summary(&self) -> Result<AgentJobSummary, DatabaseError> {
self.store.agent_job_summary().await
}
async fn agent_job_summary_for_user(
&self,
user_id: &str,
) -> Result<AgentJobSummary, DatabaseError> {
self.store.agent_job_summary_for_user(user_id).await
}
async fn get_agent_job_failure_reason(
&self,
id: Uuid,
@@ -496,6 +510,14 @@ impl RoutineStore for PgBackend {
.await
}
async fn batch_get_last_run_status(
&self,
routine_ids: &[Uuid],
) -> Result<std::collections::HashMap<Uuid, crate::agent::routine::RunStatus>, DatabaseError>
{
self.store.batch_get_last_run_status(routine_ids).await
}
async fn link_routine_run_to_job(
&self,
run_id: Uuid,
+326 -233
View File
File diff suppressed because it is too large Load Diff
+87
View File
@@ -842,6 +842,38 @@ impl Store {
.collect())
}
pub async fn list_agent_jobs_for_user(
&self,
user_id: &str,
) -> Result<Vec<AgentJobRecord>, DatabaseError> {
let conn = self.conn().await?;
let rows = conn
.query(
r#"
SELECT id, title, status, user_id, failure_reason,
created_at, started_at, completed_at
FROM agent_jobs WHERE source = 'direct' AND user_id = $1
ORDER BY created_at DESC
"#,
&[&user_id],
)
.await?;
Ok(rows
.iter()
.map(|r| AgentJobRecord {
id: r.get("id"),
title: r.get("title"),
status: r.get("status"),
user_id: r.get::<_, Option<String>>("user_id").unwrap_or_default(),
created_at: r.get("created_at"),
started_at: r.get("started_at"),
completed_at: r.get("completed_at"),
failure_reason: r.get("failure_reason"),
})
.collect())
}
/// Get the failure reason for a single agent job.
pub async fn get_agent_job_failure_reason(
&self,
@@ -875,6 +907,27 @@ impl Store {
}
Ok(summary)
}
pub async fn agent_job_summary_for_user(
&self,
user_id: &str,
) -> Result<AgentJobSummary, DatabaseError> {
let conn = self.conn().await?;
let rows = conn
.query(
"SELECT status, COUNT(*) as cnt FROM agent_jobs WHERE source = 'direct' AND user_id = $1 GROUP BY status",
&[&user_id],
)
.await?;
let mut summary = AgentJobSummary::default();
for row in &rows {
let status: String = row.get("status");
let count: i64 = row.get("cnt");
summary.add_count(&status, count as usize);
}
Ok(summary)
}
}
// ==================== Job Events ====================
@@ -1350,6 +1403,40 @@ impl Store {
Ok(counts)
}
/// Batch-load the most recent run status for multiple routines in a single query.
/// Uses a window function to pick only the latest run per routine.
#[cfg(feature = "postgres")]
pub async fn batch_get_last_run_status(
&self,
routine_ids: &[Uuid],
) -> Result<HashMap<Uuid, RunStatus>, DatabaseError> {
if routine_ids.is_empty() {
return Ok(HashMap::new());
}
let conn = self.conn().await?;
let rows = conn
.query(
"SELECT DISTINCT ON (routine_id) routine_id, status
FROM routine_runs
WHERE routine_id = ANY($1)
ORDER BY routine_id, started_at DESC",
&[&routine_ids],
)
.await?;
let mut statuses = HashMap::new();
for row in rows {
let id: Uuid = row.get("routine_id");
let status_str: String = row.get("status");
if let std::result::Result::Ok(status) = status_str.parse::<RunStatus>() {
statuses.insert(id, status);
}
}
Ok(statuses)
}
/// Link a routine run to a dispatched job.
pub async fn link_routine_run_to_job(
&self,
+19 -52
View File
@@ -107,14 +107,21 @@ impl GithubCopilotProvider {
body: &impl Serialize,
) -> Result<R, LlmError> {
let url = self.api_url();
// Map token exchange failures to RequestFailed (retryable) rather than
// AuthFailed (non-retryable), since transient network errors during
// exchange should be retried by RetryProvider.
// Distinguish permanent auth errors (non-retryable) from transient
// network failures (retryable) so RetryProvider handles them correctly.
let token = self.token_manager.get_token().await.map_err(|e| {
tracing::warn!(error = %e, "Copilot: token exchange failed");
LlmError::RequestFailed {
provider: "github_copilot".to_string(),
reason: format!("Token exchange failed: {e}"),
match &e {
crate::llm::github_copilot_auth::GithubCopilotAuthError::AccessDenied
| crate::llm::github_copilot_auth::GithubCopilotAuthError::Expired => {
LlmError::AuthFailed {
provider: "github_copilot".to_string(),
}
}
_ => LlmError::RequestFailed {
provider: "github_copilot".to_string(),
reason: format!("Token exchange failed: {e}"),
},
}
})?;
@@ -157,54 +164,14 @@ impl GithubCopilotProvider {
);
if status.as_u16() == 401 {
// Invalidate the cached session token and retry once with a
// fresh exchange — stale tokens are the most common 401 cause.
tracing::warn!("Copilot: 401 Unauthorized — invalidating session token, retrying");
// Invalidate the cached session token so the next attempt
// (driven by RetryProvider) gets a fresh one. We don't retry
// inline to avoid nested retries with the outer RetryProvider.
tracing::warn!("Copilot: 401 Unauthorized — invalidating session token for retry");
self.token_manager.invalidate().await;
let fresh = self.token_manager.get_token().await.map_err(|e| {
tracing::warn!(error = %e, "Copilot: re-exchange after 401 failed");
LlmError::RequestFailed {
provider: "github_copilot".to_string(),
reason: format!("Token re-exchange after 401 failed: {e}"),
}
})?;
let mut retry_req = self
.client
.post(&url)
.bearer_auth(fresh.expose_secret())
.header("Content-Type", "application/json");
for (key, value) in &self.extra_headers {
retry_req = retry_req.header(key.as_str(), value.as_str());
}
let retry =
retry_req
.json(body)
.send()
.await
.map_err(|e| LlmError::RequestFailed {
provider: "github_copilot".to_string(),
reason: format!("Retry after 401 failed: {e}"),
})?;
if retry.status().is_success() {
let text = retry.text().await.map_err(|e| LlmError::RequestFailed {
provider: "github_copilot".to_string(),
reason: format!("Failed to read retry response body: {e}"),
})?;
return serde_json::from_str(&text).map_err(|e| {
let truncated = crate::agent::truncate_for_preview(&text, 512);
LlmError::InvalidResponse {
provider: "github_copilot".to_string(),
reason: format!("JSON parse error: {e}. Raw: {truncated}"),
}
});
}
let retry_status = retry.status();
tracing::warn!(
status = %retry_status,
"Copilot: 401 retry also failed"
);
return Err(LlmError::AuthFailed {
return Err(LlmError::RequestFailed {
provider: "github_copilot".to_string(),
reason: "HTTP 401 Unauthorized".to_string(),
});
}
if status.as_u16() == 429 {
+61 -15
View File
@@ -589,15 +589,46 @@ async fn async_main() -> anyhow::Result<()> {
// ── Gateway channel ────────────────────────────────────────────────
let mut gateway_url: Option<String> = None;
let mut sse_sender: Option<
tokio::sync::broadcast::Sender<ironclaw::channels::web::types::SseEvent>,
> = None;
let mut sse_manager: Option<std::sync::Arc<ironclaw::channels::web::sse::SseManager>> = None;
if let Some(ref gw_config) = config.channels.gateway {
let mut gw =
GatewayChannel::new(gw_config.clone()).with_llm_provider(Arc::clone(&components.llm));
// Build multi-user auth state if user_tokens is configured, else single-user.
let mut gw = if let Some(ref user_tokens) = gw_config.user_tokens {
use ironclaw::channels::web::auth::{MultiAuthState, UserIdentity};
let tokens = user_tokens
.iter()
.map(|(token, cfg)| {
(
token.clone(),
UserIdentity {
user_id: cfg.user_id.clone(),
workspace_read_scopes: cfg.workspace_read_scopes.clone(),
},
)
})
.collect();
let auth = MultiAuthState::multi(tokens);
GatewayChannel::new_multi_auth(gw_config.clone(), auth)
} else {
GatewayChannel::new(gw_config.clone())
};
gw = gw.with_llm_provider(Arc::clone(&components.llm));
if let Some(ref ws) = components.workspace {
gw = gw.with_workspace(Arc::clone(ws));
}
// Create per-user workspace pool for multi-user mode.
if let Some(ref db) = components.db {
let emb_cache_config = ironclaw::workspace::EmbeddingCacheConfig {
max_entries: config.embeddings.cache_size,
};
let pool = Arc::new(ironclaw::channels::web::server::WorkspacePool::new(
Arc::clone(db),
components.embeddings.clone(),
emb_cache_config,
config.search.clone(),
config.workspace.clone(),
));
gw = gw.with_workspace_pool(pool);
}
gw = gw.with_session_manager(Arc::clone(&session_manager));
gw = gw.with_log_broadcaster(Arc::clone(&log_broadcaster));
gw = gw.with_log_level_handle(Arc::clone(&log_level_handle));
@@ -648,8 +679,12 @@ async fn async_main() -> anyhow::Result<()> {
let mut rx = tx.subscribe();
let gw_state = Arc::clone(gw.state());
tokio::spawn(async move {
while let Ok((_job_id, event)) = rx.recv().await {
gw_state.sse.broadcast(event);
while let Ok((_job_id, user_id, event)) = rx.recv().await {
if user_id.is_empty() {
gw_state.sse.broadcast(event);
} else {
gw_state.sse.broadcast_for_user(&user_id, event);
}
}
});
}
@@ -691,7 +726,7 @@ async fn async_main() -> anyhow::Result<()> {
// Capture SSE sender and routine engine slot before moving gw into channels.
// IMPORTANT: This must come after all `with_*` calls since `rebuild_state`
// creates a new SseManager, which would orphan this sender.
sse_sender = Some(gw.state().sse.sender());
sse_manager = Some(Arc::clone(&gw.state().sse));
channel_names.push("gateway".to_string());
channels.add(Box::new(gw)).await;
}
@@ -754,6 +789,14 @@ async fn async_main() -> anyhow::Result<()> {
.register_message_tools(Arc::clone(&channels), components.extension_manager.clone())
.await;
// Default user ID for extension operations (single-user mode).
let ext_user_id = config
.channels
.gateway
.as_ref()
.map(|g| g.user_id.clone())
.unwrap_or_else(|| "default".to_string());
// Wire up channel runtime for hot-activation of WASM channels.
if let Some(ref ext_mgr) = components.extension_manager
&& let Some((rt, ps, router)) = wasm_channel_runtime_state.take()
@@ -774,12 +817,14 @@ async fn async_main() -> anyhow::Result<()> {
// Auto-activate WASM channels that were active in a previous session.
// Relay channels are handled separately below via restore_relay_channels().
let persisted = ext_mgr.load_persisted_active_channels().await;
let persisted = ext_mgr.load_persisted_active_channels(&ext_user_id).await;
for name in &persisted {
if active_at_startup.contains(name) || ext_mgr.is_relay_channel(name).await {
if active_at_startup.contains(name)
|| ext_mgr.is_relay_channel(name, &ext_user_id).await
{
continue;
}
match ext_mgr.activate(name).await {
match ext_mgr.activate(name, &ext_user_id).await {
Ok(result) => {
tracing::debug!(
channel = %name,
@@ -804,14 +849,14 @@ async fn async_main() -> anyhow::Result<()> {
ext_mgr
.set_relay_channel_manager(Arc::clone(&channels))
.await;
ext_mgr.restore_relay_channels().await;
ext_mgr.restore_relay_channels(&ext_user_id).await;
}
// Wire SSE sender into extension manager for broadcasting status events.
if let Some(ref ext_mgr) = components.extension_manager
&& let Some(ref sender) = sse_sender
&& let Some(ref sse) = sse_manager
{
ext_mgr.set_sse_sender(sender.clone()).await;
ext_mgr.set_sse_sender(Arc::clone(sse)).await;
}
// Snapshot memory for trace recording before the agent starts
@@ -849,7 +894,7 @@ async fn async_main() -> anyhow::Result<()> {
skills_config: config.skills.clone(),
hooks: components.hooks,
cost_guard: components.cost_guard,
sse_tx: sse_sender,
sse_tx: sse_manager,
http_interceptor,
transcription: config.transcription.create_provider().map(|p| {
Arc::new(ironclaw::llm::transcription::TranscriptionMiddleware::new(
@@ -867,6 +912,7 @@ async fn async_main() -> anyhow::Result<()> {
ironclaw::agent::routine_engine::SandboxReadiness::DockerUnavailable
},
builder: components.builder,
llm_backend: config.llm.backend.clone(),
};
let channels_for_warnings = Arc::clone(&channels);
+53 -6
View File
@@ -40,7 +40,8 @@ pub struct OrchestratorState {
pub job_manager: Arc<ContainerJobManager>,
pub token_store: TokenStore,
/// Broadcast channel for job events (consumed by the web gateway SSE).
pub job_event_tx: Option<broadcast::Sender<(Uuid, SseEvent)>>,
/// Tuple: (job_id, user_id, event).
pub job_event_tx: Option<broadcast::Sender<(Uuid, String, SseEvent)>>,
/// Buffered follow-up prompts for sandbox jobs, keyed by job_id.
pub prompt_queue: Arc<Mutex<HashMap<Uuid, VecDeque<PendingPrompt>>>>,
/// Database handle for persisting job events.
@@ -49,6 +50,9 @@ pub struct OrchestratorState {
pub secrets_store: Option<Arc<dyn SecretsStore + Send + Sync>>,
/// User ID for secret lookups (single-tenant, typically "default").
pub user_id: String,
/// In-memory cache of job_id → user_id for SSE scoping. Populated when
/// sandbox jobs are created, avoiding a DB round-trip on every job event.
pub job_owner_cache: Arc<std::sync::RwLock<HashMap<Uuid, String>>>,
}
/// The orchestrator's internal API server.
@@ -351,9 +355,45 @@ async fn job_event_handler(
},
};
// Broadcast via the channel (if configured)
// Broadcast via the channel (if configured).
// Look up the job owner from the in-memory cache (populated at job creation).
if let Some(ref tx) = state.job_event_tx {
let _ = tx.send((job_id, sse_event));
let cached_uid = state
.job_owner_cache
.read()
.unwrap_or_else(|e| e.into_inner())
.get(&job_id)
.cloned();
let user_id = match cached_uid {
Some(uid) => uid,
None => {
// Cache miss: fall back to DB lookup and populate cache.
let uid = match state.store.as_ref() {
Some(store) => store
.get_sandbox_job(job_id)
.await
.ok()
.flatten()
.map(|j| j.user_id),
None => None,
};
if let Some(ref uid) = uid {
state
.job_owner_cache
.write()
.unwrap_or_else(|e| e.into_inner())
.insert(job_id, uid.clone());
}
uid.unwrap_or_default()
}
};
if user_id.is_empty() {
let _ = tx.send((job_id, String::new(), sse_event));
} else {
let _ = tx.send((job_id, user_id, sse_event));
}
}
Ok(StatusCode::OK)
@@ -480,6 +520,7 @@ mod tests {
store: None,
secrets_store: None,
user_id: "default".to_string(),
job_owner_cache: Arc::new(std::sync::RwLock::new(HashMap::new())),
}
}
@@ -709,6 +750,7 @@ mod tests {
store: None,
secrets_store: Some(secrets_store),
user_id: "default".to_string(),
job_owner_cache: Arc::new(std::sync::RwLock::new(HashMap::new())),
};
let router = OrchestratorApi::router(state);
@@ -744,6 +786,7 @@ mod tests {
store: None,
secrets_store: None,
user_id: "default".to_string(),
job_owner_cache: Arc::new(std::sync::RwLock::new(HashMap::new())),
};
let job_id = Uuid::new_v4();
@@ -769,8 +812,10 @@ mod tests {
let resp = router.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
let (recv_id, event) = rx.recv().await.unwrap();
let (recv_id, recv_uid, event) = rx.recv().await.unwrap();
assert_eq!(recv_id, job_id);
// No store configured, so user_id falls back to empty string.
assert_eq!(recv_uid, "");
match event {
SseEvent::JobMessage {
job_id: jid,
@@ -799,6 +844,7 @@ mod tests {
store: None,
secrets_store: None,
user_id: "default".to_string(),
job_owner_cache: Arc::new(std::sync::RwLock::new(HashMap::new())),
};
let job_id = Uuid::new_v4();
@@ -824,7 +870,7 @@ mod tests {
let resp = router.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
let (_recv_id, event) = rx.recv().await.unwrap();
let (_recv_id, _recv_uid, event) = rx.recv().await.unwrap();
match event {
SseEvent::JobToolUse { tool_name, .. } => {
assert_eq!(tool_name, "shell");
@@ -847,6 +893,7 @@ mod tests {
store: None,
secrets_store: None,
user_id: "default".to_string(),
job_owner_cache: Arc::new(std::sync::RwLock::new(HashMap::new())),
};
let job_id = Uuid::new_v4();
@@ -869,7 +916,7 @@ mod tests {
let resp = router.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
let (_recv_id, event) = rx.recv().await.unwrap();
let (_recv_id, _recv_uid, event) = rx.recv().await.unwrap();
// Unknown event types fall through to JobStatus
assert!(matches!(event, SseEvent::JobStatus { .. }));
}
+2 -1
View File
@@ -63,7 +63,7 @@ fn resolve_orchestrator_port() -> u16 {
/// Result of orchestrator setup, containing all handles needed by the agent.
pub struct OrchestratorSetup {
pub container_job_manager: Option<Arc<ContainerJobManager>>,
pub job_event_tx: Option<broadcast::Sender<(Uuid, SseEvent)>>,
pub job_event_tx: Option<broadcast::Sender<(Uuid, String, SseEvent)>>,
pub prompt_queue: Arc<Mutex<HashMap<Uuid, VecDeque<api::PendingPrompt>>>>,
pub docker_status: crate::sandbox::DockerStatus,
}
@@ -134,6 +134,7 @@ pub async fn setup_orchestrator(
store: db.cloned(),
secrets_store: secrets_store.cloned(),
user_id: "default".to_string(),
job_owner_cache: Arc::new(std::sync::RwLock::new(std::collections::HashMap::new())),
};
tokio::spawn(async move {
+108
View File
@@ -1297,6 +1297,92 @@ mod tests {
assert_eq!(loaded.heartbeat.interval_secs, 900);
}
/// Regression: /model writes a single key ("selected_model") to the DB via
/// set_setting(). On restart, get_all_settings() returns ALL keys including
/// wizard-written defaults. The single-key update must survive the full
/// from_db_map() round trip.
#[test]
fn db_single_key_model_update_survives_roundtrip() {
// Step 1: Wizard writes full settings to DB (including selected_model
// from initial setup).
let wizard_settings = Settings {
llm_backend: Some("nearai".to_string()),
selected_model: Some("old-wizard-model".to_string()),
..Default::default()
};
let mut db: std::collections::HashMap<String, serde_json::Value> =
wizard_settings.to_db_map();
// Step 2: User runs /model new-model — persist_selected_model writes
// a single key, overwriting the wizard value.
db.insert(
"selected_model".to_string(),
serde_json::Value::String("new-model".to_string()),
);
// Step 3: On restart, from_db_map() rebuilds Settings from the full
// DB map.
let restored = Settings::from_db_map(&db);
assert_eq!(
restored.selected_model,
Some("new-model".to_string()),
"/model change must survive DB round trip"
);
}
/// Regression: TOML overlay must not clobber a DB-persisted selected_model
/// when the TOML file matches the DB. This is the normal case after /model
/// successfully writes to both DB and TOML.
#[test]
fn toml_overlay_preserves_matching_model() {
// DB settings with new model from /model command.
let mut db_settings = Settings {
llm_backend: Some("nearai".to_string()),
selected_model: Some("new-model".to_string()),
..Default::default()
};
// TOML also updated by /model command to the same value.
let toml_settings = Settings {
selected_model: Some("new-model".to_string()),
..Default::default()
};
db_settings.merge_from(&toml_settings);
assert_eq!(
db_settings.selected_model,
Some("new-model".to_string()),
"TOML overlay must not clobber matching model"
);
}
/// Regression: when /model updates DB but TOML write fails, a stale TOML
/// file would overwrite the DB value. This test documents the priority:
/// TOML > DB (by design). persist_selected_model MUST update the TOML.
#[test]
fn stale_toml_overwrites_db_model() {
// DB has the new model from /model.
let mut db_settings = Settings {
selected_model: Some("new-model".to_string()),
..Default::default()
};
// TOML still has the old model (write failed or was not attempted).
let stale_toml = Settings {
selected_model: Some("old-model".to_string()),
..Default::default()
};
db_settings.merge_from(&stale_toml);
// This documents the current priority: TOML wins over DB.
// The fix in persist_selected_model ensures TOML is always updated.
assert_eq!(
db_settings.selected_model,
Some("old-model".to_string()),
"TOML overlay has higher priority than DB (by design)"
);
}
/// Regression test: /model command must persist selected_model to TOML config.
/// Prior to the fix, `set_model()` only changed the in-memory provider and the
/// choice was lost on restart.
@@ -1322,6 +1408,28 @@ mod tests {
assert_eq!(reloaded.selected_model, Some("new-model".to_string()));
}
/// Regression: /model must create config.toml when it doesn't exist, so the
/// model survives restarts. Previously the Ok(None) case was a no-op.
#[test]
fn toml_created_when_missing_for_model_persist() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("config.toml");
// No config.toml yet (fresh install, no wizard).
assert!(Settings::load_toml(&path).unwrap().is_none());
// Simulate what persist_selected_model now does for the Ok(None) case.
let settings = Settings {
selected_model: Some("new-model".to_string()),
..Default::default()
};
settings.save_toml(&path).unwrap();
// Verify the model survived.
let loaded = Settings::load_toml(&path).unwrap().unwrap();
assert_eq!(loaded.selected_model, Some("new-model".to_string()));
}
#[test]
fn toml_missing_file_returns_none() {
let result = Settings::load_toml(std::path::Path::new("/tmp/nonexistent_config.toml"));
+1
View File
@@ -563,6 +563,7 @@ impl TestHarnessBuilder {
document_extraction: None,
sandbox_readiness: crate::agent::routine_engine::SandboxReadiness::DisabledByConfig,
builder: None,
llm_backend: "nearai".to_string(),
};
TestHarness {
+17 -17
View File
@@ -130,7 +130,7 @@ impl Tool for ToolInstallTool {
async fn execute(
&self,
params: serde_json::Value,
_ctx: &JobContext,
ctx: &JobContext,
) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now();
@@ -150,7 +150,7 @@ impl Tool for ToolInstallTool {
let result = self
.manager
.install(name, url, kind_hint)
.install(name, url, kind_hint, &ctx.user_id)
.await
.map_err(|e| ToolError::ExecutionFailed(e.to_string()))?;
@@ -205,7 +205,7 @@ impl Tool for ToolAuthTool {
async fn execute(
&self,
params: serde_json::Value,
_ctx: &JobContext,
ctx: &JobContext,
) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now();
@@ -213,13 +213,13 @@ impl Tool for ToolAuthTool {
let result = self
.manager
.auth(name)
.auth(name, &ctx.user_id)
.await
.map_err(|e| ToolError::ExecutionFailed(e.to_string()))?;
// Auto-activate after successful auth so tools are available immediately
if result.is_authenticated() {
match self.manager.activate(name).await {
match self.manager.activate(name, &ctx.user_id).await {
Ok(activate_result) => {
let output = serde_json::json!({
"status": "authenticated_and_activated",
@@ -304,13 +304,13 @@ impl Tool for ToolActivateTool {
async fn execute(
&self,
params: serde_json::Value,
_ctx: &JobContext,
ctx: &JobContext,
) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now();
let name = require_str(&params, "name")?;
match self.manager.activate(name).await {
match self.manager.activate(name, &ctx.user_id).await {
Ok(result) => {
let output = serde_json::to_value(&result)
.unwrap_or_else(|_| serde_json::json!({"error": "serialization failed"}));
@@ -329,12 +329,12 @@ impl Tool for ToolActivateTool {
// Activation failed due to missing auth; initiate auth flow
// so the agent loop can show the auth card.
match self.manager.auth(name).await {
match self.manager.auth(name, &ctx.user_id).await {
Ok(auth_result) if auth_result.is_authenticated() => {
// Auth succeeded (e.g. env var was set); retry activation.
let result = self
.manager
.activate(name)
.activate(name, &ctx.user_id)
.await
.map_err(|e| ToolError::ExecutionFailed(e.to_string()))?;
let output = serde_json::to_value(&result).unwrap_or_else(
@@ -404,7 +404,7 @@ impl Tool for ToolListTool {
async fn execute(
&self,
params: serde_json::Value,
_ctx: &JobContext,
ctx: &JobContext,
) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now();
@@ -425,7 +425,7 @@ impl Tool for ToolListTool {
let extensions = self
.manager
.list(kind_filter, include_available)
.list(kind_filter, include_available, &ctx.user_id)
.await
.map_err(|e| ToolError::ExecutionFailed(e.to_string()))?;
@@ -477,7 +477,7 @@ impl Tool for ToolRemoveTool {
async fn execute(
&self,
params: serde_json::Value,
_ctx: &JobContext,
ctx: &JobContext,
) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now();
@@ -485,7 +485,7 @@ impl Tool for ToolRemoveTool {
let message = self
.manager
.remove(name)
.remove(name, &ctx.user_id)
.await
.map_err(|e| ToolError::ExecutionFailed(e.to_string()))?;
@@ -541,7 +541,7 @@ impl Tool for ToolUpgradeTool {
async fn execute(
&self,
params: serde_json::Value,
_ctx: &JobContext,
ctx: &JobContext,
) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now();
@@ -549,7 +549,7 @@ impl Tool for ToolUpgradeTool {
let result = self
.manager
.upgrade(name)
.upgrade(name, &ctx.user_id)
.await
.map_err(|e| ToolError::ExecutionFailed(e.to_string()))?;
@@ -603,7 +603,7 @@ impl Tool for ExtensionInfoTool {
async fn execute(
&self,
params: serde_json::Value,
_ctx: &JobContext,
ctx: &JobContext,
) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now();
@@ -611,7 +611,7 @@ impl Tool for ExtensionInfoTool {
let info = self
.manager
.extension_info(name)
.extension_info(name, &ctx.user_id)
.await
.map_err(|e| ToolError::ExecutionFailed(e.to_string()))?;
+2 -2
View File
@@ -85,7 +85,7 @@ pub struct CreateJobTool {
job_manager: Option<Arc<ContainerJobManager>>,
store: Option<Arc<dyn Database>>,
/// Broadcast sender for job events (used to subscribe a monitor).
event_tx: Option<tokio::sync::broadcast::Sender<(Uuid, SseEvent)>>,
event_tx: Option<tokio::sync::broadcast::Sender<(Uuid, String, SseEvent)>>,
/// Injection channel for pushing messages into the agent loop.
inject_tx: Option<tokio::sync::mpsc::Sender<IncomingMessage>>,
/// Encrypted secrets store for validating credential grants.
@@ -120,7 +120,7 @@ impl CreateJobTool {
/// monitor that forwards Claude Code output to the main agent loop.
pub fn with_monitor_deps(
mut self,
event_tx: tokio::sync::broadcast::Sender<(Uuid, SseEvent)>,
event_tx: tokio::sync::broadcast::Sender<(Uuid, String, SseEvent)>,
inject_tx: tokio::sync::mpsc::Sender<IncomingMessage>,
) -> Self {
self.event_tx = Some(event_tx);
+283 -45
View File
@@ -21,6 +21,35 @@ use crate::context::JobContext;
use crate::tools::tool::{Tool, ToolError, ToolOutput, require_str};
use crate::workspace::{Workspace, paths};
// ── WorkspaceResolver ──────────────────────────────────────────────
/// Resolves a workspace for a given user ID.
///
/// In single-user mode, always returns the same workspace.
/// In multi-tenant mode, creates per-user workspaces on demand.
#[async_trait]
pub trait WorkspaceResolver: Send + Sync {
async fn resolve(&self, user_id: &str) -> Arc<Workspace>;
}
/// Returns a fixed workspace regardless of user ID (single-user mode).
pub struct FixedWorkspaceResolver {
workspace: Arc<Workspace>,
}
impl FixedWorkspaceResolver {
pub fn new(workspace: Arc<Workspace>) -> Self {
Self { workspace }
}
}
#[async_trait]
impl WorkspaceResolver for FixedWorkspaceResolver {
async fn resolve(&self, _user_id: &str) -> Arc<Workspace> {
Arc::clone(&self.workspace)
}
}
/// Detect paths that are clearly local filesystem references, not workspace-memory docs.
///
/// Examples:
@@ -62,13 +91,20 @@ fn map_write_err(e: crate::error::WorkspaceError) -> ToolError {
/// The agent should call this tool before answering questions about
/// prior work, decisions, preferences, or any historical context.
pub struct MemorySearchTool {
workspace: Arc<Workspace>,
resolver: Arc<dyn WorkspaceResolver>,
}
impl MemorySearchTool {
/// Create a new memory search tool.
pub fn new(workspace: Arc<Workspace>) -> Self {
Self { workspace }
/// Create a new memory search tool with a workspace resolver.
pub fn new(resolver: Arc<dyn WorkspaceResolver>) -> Self {
Self { resolver }
}
/// Create from a fixed workspace (backward compatibility).
pub fn from_workspace(workspace: Arc<Workspace>) -> Self {
Self {
resolver: Arc::new(FixedWorkspaceResolver::new(workspace)),
}
}
}
@@ -107,7 +143,7 @@ impl Tool for MemorySearchTool {
async fn execute(
&self,
params: serde_json::Value,
_ctx: &JobContext,
ctx: &JobContext,
) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now();
@@ -119,8 +155,8 @@ impl Tool for MemorySearchTool {
.unwrap_or(5)
.min(20) as usize;
let results = self
.workspace
let workspace = self.resolver.resolve(&ctx.user_id).await;
let results = workspace
.search(query, limit)
.await
.map_err(|e| ToolError::ExecutionFailed(format!("Search failed: {}", e)))?;
@@ -151,13 +187,20 @@ impl Tool for MemorySearchTool {
/// Use this to persist important information that should be remembered
/// across sessions: decisions, preferences, facts, lessons learned.
pub struct MemoryWriteTool {
workspace: Arc<Workspace>,
resolver: Arc<dyn WorkspaceResolver>,
}
impl MemoryWriteTool {
/// Create a new memory write tool.
pub fn new(workspace: Arc<Workspace>) -> Self {
Self { workspace }
/// Create a new memory write tool with a workspace resolver.
pub fn new(resolver: Arc<dyn WorkspaceResolver>) -> Self {
Self { resolver }
}
/// Create from a fixed workspace (backward compatibility).
pub fn from_workspace(workspace: Arc<Workspace>) -> Self {
Self {
resolver: Arc::new(FixedWorkspaceResolver::new(workspace)),
}
}
}
@@ -231,19 +274,21 @@ impl Tool for MemoryWriteTool {
)));
}
let workspace = self.resolver.resolve(&ctx.user_id).await;
// Bootstrap target: clear BOOTSTRAP.md to mark first-run ritual complete.
// Handled early because it accepts empty content (unlike other targets).
if target == "bootstrap" {
// Write empty content to effectively disable the bootstrap injection.
// system_prompt_for_context() skips empty files.
self.workspace
workspace
.write(paths::BOOTSTRAP, "")
.await
.map_err(map_write_err)?;
// Also set the in-memory flag so BOOTSTRAP.md injection stops
// immediately without waiting for a restart.
self.workspace.mark_bootstrap_completed();
workspace.mark_bootstrap_completed();
let output = serde_json::json!({
"status": "cleared",
@@ -289,12 +334,12 @@ impl Tool for MemoryWriteTool {
// Otherwise, use default workspace methods (which include injection scanning).
let layer_result = if let Some(layer_name) = layer {
let result = if append {
self.workspace
workspace
.append_to_layer(layer_name, &resolved_path, content, force)
.await
.map_err(map_write_err)?
} else {
self.workspace
workspace
.write_to_layer(layer_name, &resolved_path, content, force)
.await
.map_err(map_write_err)?
@@ -307,31 +352,33 @@ impl Tool for MemoryWriteTool {
match target {
"memory" => {
if append {
self.workspace
workspace
.append_memory(content)
.await
.map_err(map_write_err)?;
} else {
self.workspace
workspace
.write(paths::MEMORY, content)
.await
.map_err(map_write_err)?;
}
}
"daily_log" => {
self.workspace
let tz = crate::timezone::parse_timezone(&ctx.user_timezone)
.unwrap_or(chrono_tz::Tz::UTC);
workspace
.append_daily_log_tz(content, tz)
.await
.map_err(map_write_err)?;
}
_ => {
if append {
self.workspace
workspace
.append(&resolved_path, content)
.await
.map_err(map_write_err)?;
} else {
self.workspace
workspace
.write(&resolved_path, content)
.await
.map_err(map_write_err)?;
@@ -361,12 +408,12 @@ impl Tool for MemoryWriteTool {
};
let mut synced_docs: Vec<&str> = Vec::new();
if normalized_path == paths::PROFILE {
match self.workspace.sync_profile_documents().await {
match workspace.sync_profile_documents().await {
Ok(true) => {
tracing::info!("profile write: synced USER.md + assistant-directives.md");
synced_docs.extend_from_slice(&[paths::USER, paths::ASSISTANT_DIRECTIVES]);
self.workspace.mark_bootstrap_completed();
workspace.mark_bootstrap_completed();
let toml_path = crate::settings::Settings::default_toml_path();
if let Ok(Some(mut settings)) = crate::settings::Settings::load_toml(&toml_path)
&& !settings.profile_onboarding_completed
@@ -416,13 +463,20 @@ impl Tool for MemoryWriteTool {
///
/// Use this to read the full content of any file in the workspace.
pub struct MemoryReadTool {
workspace: Arc<Workspace>,
resolver: Arc<dyn WorkspaceResolver>,
}
impl MemoryReadTool {
/// Create a new memory read tool.
pub fn new(workspace: Arc<Workspace>) -> Self {
Self { workspace }
/// Create a new memory read tool with a workspace resolver.
pub fn new(resolver: Arc<dyn WorkspaceResolver>) -> Self {
Self { resolver }
}
/// Create from a fixed workspace (backward compatibility).
pub fn from_workspace(workspace: Arc<Workspace>) -> Self {
Self {
resolver: Arc::new(FixedWorkspaceResolver::new(workspace)),
}
}
}
@@ -456,7 +510,7 @@ impl Tool for MemoryReadTool {
async fn execute(
&self,
params: serde_json::Value,
_ctx: &JobContext,
ctx: &JobContext,
) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now();
@@ -470,8 +524,8 @@ impl Tool for MemoryReadTool {
)));
}
let doc = self
.workspace
let workspace = self.resolver.resolve(&ctx.user_id).await;
let doc = workspace
.read(path)
.await
.map_err(|e| ToolError::ExecutionFailed(format!("Read failed: {}", e)))?;
@@ -495,20 +549,27 @@ impl Tool for MemoryReadTool {
///
/// Returns a hierarchical view of files and directories with configurable depth.
pub struct MemoryTreeTool {
workspace: Arc<Workspace>,
resolver: Arc<dyn WorkspaceResolver>,
}
impl MemoryTreeTool {
/// Create a new memory tree tool.
pub fn new(workspace: Arc<Workspace>) -> Self {
Self { workspace }
/// Create a new memory tree tool with a workspace resolver.
pub fn new(resolver: Arc<dyn WorkspaceResolver>) -> Self {
Self { resolver }
}
/// Create from a fixed workspace (backward compatibility).
pub fn from_workspace(workspace: Arc<Workspace>) -> Self {
Self {
resolver: Arc::new(FixedWorkspaceResolver::new(workspace)),
}
}
/// Recursively build tree structure.
///
/// Returns a compact format where directories end with `/` and may have children.
async fn build_tree(
&self,
workspace: &Arc<Workspace>,
path: &str,
current_depth: usize,
max_depth: usize,
@@ -517,8 +578,7 @@ impl MemoryTreeTool {
return Ok(Vec::new());
}
let entries = self
.workspace
let entries = workspace
.list(path)
.await
.map_err(|e| ToolError::ExecutionFailed(format!("Tree failed: {}", e)))?;
@@ -533,8 +593,13 @@ impl MemoryTreeTool {
};
if entry.is_directory && current_depth < max_depth {
let children =
Box::pin(self.build_tree(&entry.path, current_depth + 1, max_depth)).await?;
let children = Box::pin(Self::build_tree(
workspace,
&entry.path,
current_depth + 1,
max_depth,
))
.await?;
if children.is_empty() {
result.push(serde_json::Value::String(display_path));
} else {
@@ -584,7 +649,7 @@ impl Tool for MemoryTreeTool {
async fn execute(
&self,
params: serde_json::Value,
_ctx: &JobContext,
ctx: &JobContext,
) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now();
@@ -596,7 +661,8 @@ impl Tool for MemoryTreeTool {
.unwrap_or(1)
.clamp(1, 10) as usize;
let tree = self.build_tree(path, 1, depth).await?;
let workspace = self.resolver.resolve(&ctx.user_id).await;
let tree = Self::build_tree(&workspace, path, 1, depth).await?;
// Compact output: just the tree array
Ok(ToolOutput::success(
@@ -650,7 +716,7 @@ mod tests {
#[test]
fn test_memory_search_schema() {
let workspace = make_test_workspace();
let tool = MemorySearchTool::new(workspace);
let tool = MemorySearchTool::from_workspace(workspace);
assert_eq!(tool.name(), "memory_search");
assert!(!tool.requires_sanitization());
@@ -668,7 +734,7 @@ mod tests {
#[test]
fn test_memory_write_schema() {
let workspace = make_test_workspace();
let tool = MemoryWriteTool::new(workspace);
let tool = MemoryWriteTool::from_workspace(workspace);
assert_eq!(tool.name(), "memory_write");
@@ -681,7 +747,7 @@ mod tests {
#[test]
fn test_memory_read_schema() {
let workspace = make_test_workspace();
let tool = MemoryReadTool::new(workspace);
let tool = MemoryReadTool::from_workspace(workspace);
assert_eq!(tool.name(), "memory_read");
@@ -698,7 +764,7 @@ mod tests {
#[test]
fn test_memory_tree_schema() {
let workspace = make_test_workspace();
let tool = MemoryTreeTool::new(workspace);
let tool = MemoryTreeTool::from_workspace(workspace);
assert_eq!(tool.name(), "memory_tree");
@@ -711,7 +777,7 @@ mod tests {
#[tokio::test]
async fn test_memory_write_rejects_injection_to_identity_file() {
let workspace = make_test_workspace();
let tool = MemoryWriteTool::new(workspace);
let tool = MemoryWriteTool::from_workspace(workspace);
let ctx = JobContext::default();
let params = serde_json::json!({
@@ -733,4 +799,176 @@ mod tests {
}
}
}
// Regression tests for per-user workspace scoping (multi-tenant mode).
// See: https://github.com/nearai/ironclaw/pull/1118
// Bug: memory tools used a single startup workspace regardless of which
// user was chatting. Fix: resolve workspace per-request via JobContext.user_id.
#[cfg(feature = "postgres")]
mod resolver_tests {
use super::*;
fn make_test_workspace_for_user(user_id: &str) -> Arc<Workspace> {
Arc::new(Workspace::new(
user_id,
deadpool_postgres::Pool::builder(deadpool_postgres::Manager::new(
tokio_postgres::Config::new(),
tokio_postgres::NoTls,
))
.build()
.unwrap(),
))
}
#[tokio::test]
async fn test_fixed_workspace_resolver_ignores_user_id() {
let ws = make_test_workspace_for_user("alice");
let resolver = FixedWorkspaceResolver::new(Arc::clone(&ws));
let ws_alice = resolver.resolve("alice").await;
let ws_bob = resolver.resolve("bob").await;
// Both should return the exact same Arc (pointer equality)
assert!(Arc::ptr_eq(&ws_alice, &ws_bob));
assert_eq!(ws_alice.user_id(), "alice");
}
/// Tracking resolver that records which user_ids were requested.
struct TrackingWorkspaceResolver {
inner: FixedWorkspaceResolver,
resolved_users: std::sync::Mutex<Vec<String>>,
}
impl TrackingWorkspaceResolver {
fn new(workspace: Arc<Workspace>) -> Self {
Self {
inner: FixedWorkspaceResolver::new(workspace),
resolved_users: std::sync::Mutex::new(Vec::new()),
}
}
fn resolved_users(&self) -> Vec<String> {
self.resolved_users.lock().unwrap().clone()
}
}
#[async_trait]
impl WorkspaceResolver for TrackingWorkspaceResolver {
async fn resolve(&self, user_id: &str) -> Arc<Workspace> {
self.resolved_users
.lock()
.unwrap()
.push(user_id.to_string());
self.inner.resolve(user_id).await
}
}
#[tokio::test]
async fn test_memory_search_uses_job_context_user_id() {
let ws = make_test_workspace_for_user("default");
let tracker = Arc::new(TrackingWorkspaceResolver::new(ws));
let tool = MemorySearchTool::new(tracker.clone() as Arc<dyn WorkspaceResolver>);
// Execute with user_id "alice"
let ctx_alice = JobContext::with_user("alice", "test", "test");
let params = serde_json::json!({"query": "test"});
// The search will fail (no real DB) but we only care about resolver call
let _ = tool.execute(params, &ctx_alice).await;
// Execute with user_id "bob"
let ctx_bob = JobContext::with_user("bob", "test", "test");
let params = serde_json::json!({"query": "test"});
let _ = tool.execute(params, &ctx_bob).await;
let resolved = tracker.resolved_users();
assert_eq!(resolved, vec!["alice", "bob"]);
}
#[tokio::test]
async fn test_memory_write_uses_job_context_user_id() {
let ws = make_test_workspace_for_user("default");
let tracker = Arc::new(TrackingWorkspaceResolver::new(ws));
let tool = MemoryWriteTool::new(tracker.clone() as Arc<dyn WorkspaceResolver>);
// Execute with user_id "alice"
let ctx_alice = JobContext::with_user("alice", "test", "test");
let params = serde_json::json!({
"content": "remember this",
"target": "daily_log",
});
let _ = tool.execute(params, &ctx_alice).await;
// Execute with user_id "bob"
let ctx_bob = JobContext::with_user("bob", "test", "test");
let params = serde_json::json!({
"content": "remember that",
"target": "daily_log",
});
let _ = tool.execute(params, &ctx_bob).await;
let resolved = tracker.resolved_users();
assert_eq!(resolved, vec!["alice", "bob"]);
}
}
#[cfg(feature = "libsql")]
mod per_user_resolver_tests {
use super::*;
async fn make_test_db() -> Arc<dyn crate::db::Database> {
use crate::db::libsql::LibSqlBackend;
let temp_dir = tempfile::tempdir().expect("tempdir");
let db_path = temp_dir.path().join("resolver_test.db");
let backend = LibSqlBackend::new_local(&db_path)
.await
.expect("LibSqlBackend");
<LibSqlBackend as crate::db::Database>::run_migrations(&backend)
.await
.expect("migrations");
// Leak the tempdir so it outlives the test (cleaned up on process exit).
std::mem::forget(temp_dir);
Arc::new(backend)
}
#[tokio::test]
async fn test_workspace_pool_resolver_returns_different_workspaces() {
let db = make_test_db().await;
let pool = crate::channels::web::server::WorkspacePool::new(
db,
None,
crate::workspace::EmbeddingCacheConfig::default(),
crate::config::WorkspaceSearchConfig::default(),
crate::config::WorkspaceConfig::default(),
);
let ws_alice = pool.resolve("alice").await;
let ws_bob = pool.resolve("bob").await;
// Different user IDs should get different workspaces
assert_eq!(ws_alice.user_id(), "alice");
assert_eq!(ws_bob.user_id(), "bob");
assert!(!Arc::ptr_eq(&ws_alice, &ws_bob));
}
#[tokio::test]
async fn test_workspace_pool_resolver_caches_workspace() {
let db = make_test_db().await;
let pool = crate::channels::web::server::WorkspacePool::new(
db,
None,
crate::workspace::EmbeddingCacheConfig::default(),
crate::config::WorkspaceSearchConfig::default(),
crate::config::WorkspaceConfig::default(),
);
let ws1 = pool.resolve("alice").await;
let ws2 = pool.resolve("alice").await;
// Same user_id should return the same cached Arc (pointer equality)
assert!(Arc::ptr_eq(&ws1, &ws2));
}
}
}
+1 -1
View File
@@ -6,7 +6,7 @@ mod file;
mod http;
mod job;
mod json;
mod memory;
pub mod memory;
mod message;
pub mod path_utils;
mod restart;
+83 -10
View File
@@ -140,7 +140,8 @@ fn execution_properties() -> Value {
},
"use_tools": {
"type": "boolean",
"description": "Only applies to lightweight mode. When true, safe non-approval tools are available."
"default": true,
"description": "Only applies to lightweight mode. New lightweight routines default this to true; when enabled, the routine can use the owner's live autonomous tool scope."
},
"max_tool_rounds": {
"type": "integer",
@@ -290,7 +291,7 @@ fn routine_request_discovery_schema() -> Value {
fn lightweight_execution_variant() -> Value {
serde_json::json!({
"type": "object",
"description": "Default lightweight execution. Applies when execution is omitted or execution.mode='lightweight'.",
"description": "Default lightweight execution. Applies when execution is omitted or execution.mode='lightweight'. New lightweight routines default to tools enabled unless execution.use_tools=false is set.",
"properties": {
"mode": {
"type": "string",
@@ -304,7 +305,8 @@ fn lightweight_execution_variant() -> Value {
},
"use_tools": {
"type": "boolean",
"description": "When true, safe non-approval tools are available."
"default": true,
"description": "Defaults to true for new lightweight routines. When enabled, the routine can use the owner's live autonomous tool scope."
},
"max_tool_rounds": {
"type": "integer",
@@ -335,7 +337,7 @@ fn full_job_execution_variant() -> Value {
fn execution_discovery_schema() -> Value {
serde_json::json!({
"type": "object",
"description": "Optional execution settings. Omit this block for the default lightweight mode.",
"description": "Optional execution settings. Omit this block for the default lightweight mode with tools enabled.",
"properties": execution_properties(),
"oneOf": [
lightweight_execution_variant(),
@@ -408,7 +410,8 @@ fn routine_create_tool_summary() -> ToolDiscoverySummary {
"execution.mode='full_job' uses the owner's live autonomous tool scope and ignores use_tools, max_tool_rounds, and context_paths.".into(),
],
notes: vec![
"Omitting execution defaults to lightweight mode.".into(),
"Omitting execution defaults to lightweight mode with tools enabled.".into(),
"Set execution.use_tools=false to keep a new lightweight routine text-only.".into(),
"Omitting delivery.user falls back to the owner's last-seen notification target.".into(),
"advanced.cooldown_secs defaults to 300.".into(),
"Legacy flat aliases are still accepted for compatibility, but grouped fields are preferred.".into(),
@@ -605,7 +608,8 @@ fn routine_create_schema(include_compatibility_aliases: bool) -> Value {
}
pub(crate) fn routine_create_parameters_schema() -> Value {
routine_create_schema(false)
static CACHE: OnceLock<Value> = OnceLock::new();
CACHE.get_or_init(|| routine_create_schema(false)).clone()
}
fn routine_create_discovery_schema() -> Value {
@@ -852,11 +856,15 @@ fn parse_execution_mode(value: Option<String>) -> Result<NormalizedExecutionMode
}
}
fn parse_routine_execution(params: &Value) -> Result<NormalizedExecutionRequest, ToolError> {
fn parse_routine_execution(
params: &Value,
default_use_tools: bool,
) -> Result<NormalizedExecutionRequest, ToolError> {
let mode = parse_execution_mode(string_field(params, "execution", "mode", &["action_type"]))?;
let context_paths =
string_array_field(params, "execution", "context_paths", &["context_paths"]);
let use_tools = bool_field(params, "execution", "use_tools", &["use_tools"]).unwrap_or(false);
let use_tools =
bool_field(params, "execution", "use_tools", &["use_tools"]).unwrap_or(default_use_tools);
let max_tool_rounds = u64_field(params, "execution", "max_tool_rounds", &["max_tool_rounds"])
.unwrap_or(3)
.clamp(1, crate::agent::routine::MAX_TOOL_ROUNDS_LIMIT as u64)
@@ -888,7 +896,7 @@ fn parse_routine_create_request(
.unwrap_or("")
.to_string();
let trigger = parse_routine_trigger(params)?;
let execution = parse_routine_execution(params)?;
let execution = parse_routine_execution(params, true)?;
let delivery = parse_routine_delivery(params);
let cooldown_secs =
u64_field(params, "advanced", "cooldown_secs", &["cooldown_secs"]).unwrap_or(300);
@@ -1007,7 +1015,8 @@ fn event_emit_schema(include_source_alias: bool) -> Value {
}
pub(crate) fn event_emit_parameters_schema() -> Value {
event_emit_schema(false)
static CACHE: OnceLock<Value> = OnceLock::new();
CACHE.get_or_init(|| event_emit_schema(false)).clone()
}
fn event_emit_discovery_schema() -> Value {
@@ -1863,6 +1872,56 @@ mod tests {
);
}
#[test]
fn parses_lightweight_create_with_tools_enabled_by_default() {
let params = serde_json::json!({
"name": "manual-check",
"prompt": "Inspect the repo for issues.",
"request": {
"kind": "manual"
}
});
let parsed = parse_routine_create_request(&params).expect("parse default lightweight");
assert!(
matches!(parsed.execution.mode, NormalizedExecutionMode::Lightweight),
"expected lightweight execution mode",
);
assert!(
parsed.execution.use_tools,
"new lightweight routines should default use_tools=true",
);
assert_eq!(parsed.execution.max_tool_rounds, 3);
}
#[test]
fn parses_lightweight_create_with_explicit_tools_disabled() {
let params = serde_json::json!({
"name": "manual-check",
"prompt": "Inspect the repo for issues.",
"request": {
"kind": "manual"
},
"execution": {
"use_tools": false
}
});
let parsed =
parse_routine_create_request(&params).expect("parse lightweight with tools disabled");
assert!(
matches!(parsed.execution.mode, NormalizedExecutionMode::Lightweight),
"expected lightweight execution mode",
);
assert!(
!parsed.execution.use_tools,
"explicit use_tools=false should be preserved",
);
assert_eq!(parsed.execution.max_tool_rounds, 3);
}
#[test]
fn parses_context_paths_with_trim_drop_empty_and_stable_dedupe() {
let params = serde_json::json!({
@@ -2201,6 +2260,20 @@ mod tests {
.any(|rule| rule.contains("request.kind='cron'")),
"summary should explain cron requirement",
);
assert!(
summary
.notes
.iter()
.any(|note| note.contains("lightweight mode with tools enabled")),
"summary should mention the new lightweight default",
);
assert!(
summary
.notes
.iter()
.any(|note| note.contains("execution.use_tools=false")),
"summary should mention the text-only opt-out",
);
assert!(
summary
.notes
+32 -6
View File
@@ -334,15 +334,37 @@ impl ToolRegistry {
tracing::debug!("Registered 5 development tools");
}
/// Register memory tools with a workspace.
/// Register memory tools with a workspace resolver.
///
/// Memory tools require a workspace resolver for persistence. Call this after
/// `register_builtin_tools()` if you have a workspace available.
pub fn register_memory_tools_with_resolver(
&self,
resolver: Arc<dyn crate::tools::builtin::memory::WorkspaceResolver>,
) {
self.register_sync(Arc::new(MemorySearchTool::new(Arc::clone(&resolver))));
self.register_sync(Arc::new(MemoryWriteTool::new(Arc::clone(&resolver))));
self.register_sync(Arc::new(MemoryReadTool::new(Arc::clone(&resolver))));
self.register_sync(Arc::new(MemoryTreeTool::new(resolver)));
tracing::debug!("Registered 4 memory tools");
}
/// Register memory tools with a fixed workspace (backward compatibility).
///
/// Memory tools require a workspace for persistence. Call this after
/// `register_builtin_tools()` if you have a workspace available.
pub fn register_memory_tools(&self, workspace: Arc<Workspace>) {
self.register_sync(Arc::new(MemorySearchTool::new(Arc::clone(&workspace))));
self.register_sync(Arc::new(MemoryWriteTool::new(Arc::clone(&workspace))));
self.register_sync(Arc::new(MemoryReadTool::new(Arc::clone(&workspace))));
self.register_sync(Arc::new(MemoryTreeTool::new(workspace)));
self.register_sync(Arc::new(MemorySearchTool::from_workspace(Arc::clone(
&workspace,
))));
self.register_sync(Arc::new(MemoryWriteTool::from_workspace(Arc::clone(
&workspace,
))));
self.register_sync(Arc::new(MemoryReadTool::from_workspace(Arc::clone(
&workspace,
))));
self.register_sync(Arc::new(MemoryTreeTool::from_workspace(workspace)));
tracing::debug!("Registered 4 memory tools");
}
@@ -361,7 +383,11 @@ impl ToolRegistry {
job_manager: Option<Arc<ContainerJobManager>>,
store: Option<Arc<dyn Database>>,
job_event_tx: Option<
tokio::sync::broadcast::Sender<(uuid::Uuid, crate::channels::web::types::SseEvent)>,
tokio::sync::broadcast::Sender<(
uuid::Uuid,
String,
crate::channels::web::types::SseEvent,
)>,
>,
inject_tx: Option<tokio::sync::mpsc::Sender<crate::channels::IncomingMessage>>,
prompt_queue: Option<PromptQueue>,
+20 -102
View File
@@ -47,12 +47,6 @@ pub struct CapabilitiesFile {
#[serde(default)]
pub description: Option<String>,
/// JSON Schema for the tool's input parameters.
/// Used as the `Tool::parameters_schema()` return value.
/// If omitted, a permissive fallback is used (with a warning).
#[serde(default)]
pub parameters: Option<serde_json::Value>,
/// Extension version (semver).
#[serde(default)]
pub version: Option<String>,
@@ -103,9 +97,6 @@ pub struct CapabilitiesFile {
/// Maximum length for the description field to prevent memory abuse.
const MAX_DESCRIPTION_CHARS: usize = 4096;
/// Maximum serialized size of the parameters schema JSON.
const MAX_PARAMETERS_SCHEMA_BYTES: usize = 64 * 1024;
impl CapabilitiesFile {
/// Parse from JSON string.
pub fn from_json(json: &str) -> Result<Self, serde_json::Error> {
@@ -135,18 +126,6 @@ impl CapabilitiesFile {
);
self.description = Some(truncated.to_string());
}
// Drop oversized parameters schema (issue #977)
if let Some(ref params) = self.parameters {
let size = params.to_string().len();
if size > MAX_PARAMETERS_SCHEMA_BYTES {
tracing::warn!(
"Capabilities parameters schema dropped ({} bytes exceeds {} limit)",
size,
MAX_PARAMETERS_SCHEMA_BYTES,
);
self.parameters = None;
}
}
}
/// Merge nested `capabilities` wrapper into top-level fields.
@@ -171,7 +150,6 @@ impl CapabilitiesFile {
if let Some(inner) = self.capabilities.take() {
let inner = inner.resolve_nested_inner(depth + 1);
self.description = self.description.or(inner.description);
self.parameters = self.parameters.or(inner.parameters);
self.http = self.http.or(inner.http);
self.secrets = self.secrets.or(inner.secrets);
self.tool_invoke = self.tool_invoke.or(inner.tool_invoke);
@@ -1424,26 +1402,12 @@ mod tests {
);
}
// ── Tool description and parameters schema ──────────────────────────
// ── Tool description ────────────────────────────────────────────────
#[test]
fn test_parse_description_and_parameters() {
fn test_parse_description() {
let json = r#"{
"description": "Search the web using Brave Search API",
"parameters": {
"type": "object",
"properties": {
"query": {
"type": "string",
"description": "Search query"
},
"count": {
"type": "integer",
"description": "Number of results"
}
},
"required": ["query"]
}
"description": "Search the web using Brave Search API"
}"#;
let caps = CapabilitiesFile::from_json(json).unwrap();
@@ -1451,28 +1415,10 @@ mod tests {
caps.description.as_deref(),
Some("Search the web using Brave Search API")
);
let params = caps.parameters.unwrap();
assert_eq!(params["type"], "object");
assert!(params["properties"]["query"].is_object());
assert_eq!(params["required"][0], "query");
}
#[test]
fn test_parse_description_only() {
let json = r#"{
"description": "A tool without explicit parameters schema"
}"#;
let caps = CapabilitiesFile::from_json(json).unwrap();
assert_eq!(
caps.description.as_deref(),
Some("A tool without explicit parameters schema")
);
assert!(caps.parameters.is_none());
}
#[test]
fn test_parse_without_description_or_parameters() {
fn test_parse_without_description() {
let json = r#"{
"http": {
"allowlist": [{ "host": "api.example.com" }]
@@ -1484,24 +1430,28 @@ mod tests {
caps.description.is_none(),
"description should be None when not provided"
);
assert!(
caps.parameters.is_none(),
"parameters should be None when not provided"
);
}
#[test]
fn test_parameters_field_silently_ignored() {
// Backward compat: old capabilities files with "parameters" still parse.
let json = r#"{
"description": "A tool",
"parameters": {
"type": "object",
"properties": { "action": { "type": "string" } }
}
}"#;
let caps = CapabilitiesFile::from_json(json).unwrap();
assert_eq!(caps.description.as_deref(), Some("A tool"));
}
#[test]
fn test_resolve_nested_description_promoted() {
let json = r#"{
"capabilities": {
"description": "Inner tool description",
"parameters": {
"type": "object",
"properties": {
"input": { "type": "string" }
},
"required": ["input"]
}
"description": "Inner tool description"
}
}"#;
@@ -1511,10 +1461,6 @@ mod tests {
Some("Inner tool description"),
"description should be promoted from inner capabilities"
);
assert!(
caps.parameters.is_some(),
"parameters should be promoted from inner capabilities"
);
}
#[test]
@@ -1564,32 +1510,4 @@ mod tests {
desc.len()
);
}
/// Regression test for issue #977: oversized parameters schema is dropped.
#[test]
fn test_oversized_parameters_schema_dropped() {
// Build a parameters schema larger than MAX_PARAMETERS_SCHEMA_BYTES
let mut properties = serde_json::Map::new();
for i in 0..2000 {
properties.insert(
format!("field_{i}"),
serde_json::json!({
"type": "string",
"description": "x".repeat(50)
}),
);
}
let schema = serde_json::json!({
"type": "object",
"properties": properties,
});
let json = serde_json::json!({
"parameters": schema,
});
let caps = CapabilitiesFile::from_json(&json.to_string()).unwrap();
assert!(
caps.parameters.is_none(),
"oversized parameters schema should be dropped"
);
}
}
+36 -58
View File
@@ -123,73 +123,51 @@ impl WasmToolLoader {
}
let wasm_bytes = fs::read(wasm_path).await?;
// Read capabilities (optional) and extract OAuth refresh config,
// tool description, and parameter schema.
let (capabilities, oauth_refresh, description, schema) =
if let Some(cap_path) = capabilities_path {
if cap_path.exists() {
let cap_bytes = fs::read(cap_path).await?;
let cap_file = CapabilitiesFile::from_bytes(&cap_bytes)
.map_err(|e| WasmLoadError::InvalidCapabilities(e.to_string()))?;
cap_file.validate(name);
// Read capabilities (optional) and extract OAuth refresh config
// and tool description. Parameter schema is auto-derived from the
// WASM module's schema() export (see WasmToolSchemas::compact_schema).
let (capabilities, oauth_refresh, description) = if let Some(cap_path) = capabilities_path {
if cap_path.exists() {
let cap_bytes = fs::read(cap_path).await?;
let cap_file = CapabilitiesFile::from_bytes(&cap_bytes)
.map_err(|e| WasmLoadError::InvalidCapabilities(e.to_string()))?;
cap_file.validate(name);
// Check WIT version compatibility
check_wit_version_compat(
name,
cap_file.wit_version.as_deref(),
crate::tools::wasm::WIT_TOOL_VERSION,
)?;
// Check WIT version compatibility
check_wit_version_compat(
name,
cap_file.wit_version.as_deref(),
crate::tools::wasm::WIT_TOOL_VERSION,
)?;
let caps = cap_file.to_capabilities();
let oauth = resolve_oauth_refresh_config(&cap_file);
let desc = cap_file.description.clone();
// Validate parameters schema before accepting it.
let params = cap_file.parameters.clone().and_then(|p| {
let errors = crate::tools::validate_tool_schema(&p, name);
if errors.is_empty() {
Some(p)
} else {
tracing::warn!(
tool = name,
?errors,
"Invalid parameters schema in capabilities.json, \
using permissive fallback"
);
None
}
});
if desc.is_none() {
tracing::warn!(
tool = name,
path = %cap_path.display(),
"Capabilities file missing \"description\" field; \
tool will use generic fallback description"
);
}
if params.is_none() && cap_file.parameters.is_none() {
tracing::warn!(
tool = name,
path = %cap_path.display(),
"Capabilities file missing \"parameters\" field; \
tool will accept any JSON object (permissive fallback)"
);
}
(caps, oauth, desc, params)
} else {
let caps = cap_file.to_capabilities();
let oauth = resolve_oauth_refresh_config(&cap_file);
let desc = cap_file.description.clone();
if desc.is_none() {
tracing::warn!(
tool = name,
path = %cap_path.display(),
"Capabilities file not found, using default (no permissions)"
"Capabilities file missing \"description\" field; \
tool will use generic fallback description"
);
(Capabilities::default(), None, None, None)
}
(caps, oauth, desc)
} else {
tracing::warn!(
tool = name,
"No capabilities file for WASM tool; \
tool will use generic fallback description and accept any JSON object"
path = %cap_path.display(),
"Capabilities file not found, using default (no permissions)"
);
(Capabilities::default(), None, None, None)
};
(Capabilities::default(), None, None)
}
} else {
tracing::warn!(
tool = name,
"No capabilities file for WASM tool; \
tool will use generic fallback description"
);
(Capabilities::default(), None, None)
};
// Register the tool
self.registry
@@ -200,7 +178,7 @@ impl WasmToolLoader {
capabilities,
limits: None,
description: description.as_deref(),
schema,
schema: None,
secrets_store: self.secrets_store.clone(),
oauth_refresh,
})
+243 -42
View File
@@ -656,12 +656,125 @@ impl WasmToolSchemas {
}
fn new(discovery: serde_json::Value) -> Self {
let advertised = Self::compact_schema(&discovery);
Self {
advertised: Self::permissive_schema(),
advertised,
discovery,
}
}
/// Derive a compact advertised schema from the full discovery schema.
///
/// Collects properties from top-level `properties` and from
/// `oneOf`/`anyOf`/`allOf` variants. Keeps only properties that are in
/// the top-level `required` array or carry an `enum`/`const` constraint.
/// For properties defined via `const` across multiple variants (e.g.
/// `"action": {"const": "get_repo"}` in each `oneOf` branch), the `const`
/// values are merged into a single `enum` array.
///
/// Variant-level `required` fields (e.g. `owner`, `repo` required within
/// each `oneOf` variant but not top-level) are intentionally omitted from
/// the compact schema — the LLM can discover them via
/// `tool_info(detail: "schema")`.
///
/// At most `MAX_COMPACT_PROPERTIES` properties are collected to bound
/// allocations from adversarial schemas.
fn compact_schema(discovery: &serde_json::Value) -> serde_json::Value {
const MAX_COMPACT_PROPERTIES: usize = 100;
let required: std::collections::HashSet<String> = discovery
.get("required")
.and_then(|r| r.as_array())
.map(|arr| {
arr.iter()
.filter_map(|v| v.as_str().map(String::from))
.collect()
})
.unwrap_or_default();
// Collect properties from top-level and oneOf/anyOf/allOf variants.
// For properties with `const` across variants, merge into an `enum`.
let mut all_properties = serde_json::Map::new();
// Track const values per property to merge into enum.
let mut const_values: std::collections::HashMap<String, Vec<serde_json::Value>> =
std::collections::HashMap::new();
if let Some(props) = discovery.get("properties").and_then(|p| p.as_object()) {
for (k, v) in props {
if all_properties.len() >= MAX_COMPACT_PROPERTIES {
break;
}
all_properties.insert(k.clone(), v.clone());
}
}
for key in ["oneOf", "anyOf", "allOf"] {
if let Some(variants) = discovery.get(key).and_then(|v| v.as_array()) {
for variant in variants {
if let Some(props) = variant.get("properties").and_then(|p| p.as_object()) {
for (k, v) in props {
if all_properties.len() >= MAX_COMPACT_PROPERTIES
&& !all_properties.contains_key(k)
{
continue;
}
// Track const values for merging into enum.
if let Some(c) = v.get("const") {
const_values.entry(k.clone()).or_default().push(c.clone());
}
all_properties.entry(k.clone()).or_insert_with(|| v.clone());
}
}
}
}
}
// Merge collected const values into enum arrays.
for (name, values) in &const_values {
if values.len() > 1
&& let Some(prop) = all_properties.get_mut(name)
{
let mut merged = prop.clone();
if let Some(obj) = merged.as_object_mut() {
obj.remove("const");
obj.insert("enum".to_string(), serde_json::Value::Array(values.clone()));
}
*prop = merged;
}
}
if all_properties.is_empty() {
return Self::permissive_schema();
}
let kept: serde_json::Map<String, serde_json::Value> = all_properties
.into_iter()
.filter(|(name, prop)| {
required.contains(name) || prop.get("enum").is_some() || prop.get("const").is_some()
})
.collect();
if kept.is_empty() {
return Self::permissive_schema();
}
let kept_required: Vec<serde_json::Value> = required
.iter()
.filter(|name| kept.contains_key(name.as_str()))
.map(|name| serde_json::Value::String(name.clone()))
.collect();
let mut result = serde_json::json!({
"type": "object",
"properties": kept,
"additionalProperties": true,
});
if !kept_required.is_empty() {
result["required"] = serde_json::Value::Array(kept_required);
}
result
}
fn with_override(&self, schema: serde_json::Value) -> Self {
Self {
advertised: schema.clone(),
@@ -1655,7 +1768,7 @@ mod tests {
}
#[tokio::test]
async fn test_advertised_schema_stays_permissive_until_sidecar_override() {
async fn test_advertised_schema_auto_compacted_from_discovery() {
let discovery_schema = serde_json::json!({
"type": "object",
"properties": {
@@ -1675,42 +1788,7 @@ mod tests {
wrapper.schemas = super::WasmToolSchemas::new(discovery_schema.clone());
wrapper.description = "Search documents".to_string();
// Advertised schema stays permissive; discovery holds the typed schema
assert_eq!(
wrapper.parameters_schema(),
serde_json::json!({
"type": "object",
"properties": {},
"additionalProperties": true
})
);
assert_eq!(wrapper.discovery_schema(), discovery_schema);
// Raw description is clean — no tool_info hint baked in
assert!(!wrapper.description().contains("tool_info"));
// But schema() composes the hint at display time when advertised is permissive
let schema = wrapper.schema();
assert!(
schema.description.contains("tool_info"),
"schema().description should contain tool_info hint: {}",
schema.description
);
assert!(
schema.description.contains("include_schema: true"),
"hint should mention include_schema: true: {}",
schema.description
);
// After sidecar override, both schemas match and hint disappears
let wrapper = wrapper.with_schema(serde_json::json!({
"type": "object",
"properties": {
"query": { "type": "string" }
},
"required": ["query"]
}));
// Advertised schema is auto-compacted: keeps required props, drops optional
assert_eq!(
wrapper.parameters_schema(),
serde_json::json!({
@@ -1718,20 +1796,143 @@ mod tests {
"properties": {
"query": { "type": "string" }
},
"required": ["query"]
"required": ["query"],
"additionalProperties": true
})
);
assert_eq!(wrapper.discovery_schema(), wrapper.parameters_schema());
// Discovery retains the full schema
assert_eq!(wrapper.discovery_schema(), discovery_schema);
// With typed schema, schema() should NOT include tool_info hint
// Compacted schema has typed properties, so no tool_info hint needed
let schema = wrapper.schema();
assert!(
!schema.description.contains("tool_info"),
"schema().description should not contain tool_info hint when typed: {}",
"schema().description should not contain tool_info hint when auto-compacted: {}",
schema.description
);
}
#[test]
fn test_compact_schema_keeps_required_and_enum_properties() {
let schema = serde_json::json!({
"type": "object",
"properties": {
"action": {
"type": "string",
"enum": ["list", "get", "create"],
"description": "The operation"
},
"query": { "type": "string" },
"limit": { "type": "integer" },
"format": {
"type": "string",
"enum": ["json", "csv"]
}
},
"required": ["action"]
});
let compacted = super::WasmToolSchemas::compact_schema(&schema);
let props = compacted["properties"].as_object().unwrap();
// action: required + enum → kept
assert!(props.contains_key("action"));
// format: has enum → kept
assert!(props.contains_key("format"));
// query: not required, no enum → dropped
assert!(!props.contains_key("query"));
// limit: not required, no enum → dropped
assert!(!props.contains_key("limit"));
// additionalProperties lets the LLM still pass dropped props
assert_eq!(compacted["additionalProperties"], true);
assert_eq!(compacted["required"], serde_json::json!(["action"]));
}
#[test]
fn test_compact_schema_falls_back_to_permissive_when_empty() {
// No required, no enum → permissive fallback
let schema = serde_json::json!({
"type": "object",
"properties": {
"query": { "type": "string" },
"limit": { "type": "integer" }
}
});
let compacted = super::WasmToolSchemas::compact_schema(&schema);
assert!(compacted["properties"].as_object().unwrap().is_empty());
}
#[test]
fn test_compact_schema_handles_no_properties() {
let schema = serde_json::json!({ "type": "object" });
let compacted = super::WasmToolSchemas::compact_schema(&schema);
assert!(compacted["properties"].as_object().unwrap().is_empty());
}
#[test]
fn test_compact_schema_handles_oneof_variants() {
// GitHub-style schema: oneOf with no top-level properties, const per variant
let schema = serde_json::json!({
"type": "object",
"required": ["action"],
"oneOf": [
{
"properties": {
"action": { "const": "get_repo" },
"owner": { "type": "string" },
"repo": { "type": "string" }
},
"required": ["action", "owner", "repo"]
},
{
"properties": {
"action": { "const": "list_issues" },
"owner": { "type": "string" },
"repo": { "type": "string" },
"state": { "type": "string", "enum": ["open", "closed", "all"] }
},
"required": ["action", "owner", "repo"]
}
]
});
let compacted = super::WasmToolSchemas::compact_schema(&schema);
let props = compacted["properties"].as_object().unwrap();
// action: required + const values merged into enum → kept
let action = &props["action"];
assert!(
action.get("enum").is_some(),
"action const values should be merged into enum: {action}"
);
let action_enum = action["enum"].as_array().unwrap();
assert!(
action_enum.contains(&serde_json::json!("get_repo")),
"enum should contain get_repo"
);
assert!(
action_enum.contains(&serde_json::json!("list_issues")),
"enum should contain list_issues"
);
assert!(
action.get("const").is_none(),
"const should be removed after merging into enum"
);
// state: has enum → kept
assert!(
props.contains_key("state"),
"state should be kept (has enum)"
);
// owner/repo: not in top-level required, no enum → intentionally dropped
// (variant-level required is omitted; discoverable via tool_info)
assert!(!props.contains_key("owner"), "owner should be dropped");
assert!(!props.contains_key("repo"), "repo should be dropped");
assert_eq!(compacted["additionalProperties"], true);
assert_eq!(compacted["required"], serde_json::json!(["action"]));
}
#[test]
fn test_capabilities_default() {
let caps = Capabilities::default();
+20 -8
View File
@@ -111,10 +111,23 @@ impl Tunnel for CloudflareTunnel {
}
}
// Drain stderr in the background to prevent SIGPIPE/buffer stalls.
tokio::spawn(async move { while let Ok(Some(_)) = reader.next_line().await {} });
if let Ok(mut guard) = self.url.write() {
*guard = Some(public_url.clone());
}
// Drain stdout silently.
// We took ownership of cloudflared's stderr pipe above to parse the URL.
// cloudflared continues writing logs for its entire lifetime. If we drop
// the reader, the pipe closes and cloudflared gets SIGPIPE on its next
// write. We can't just store the reader without reading — the OS pipe
// buffer fills up and cloudflared blocks. So we drain it in a background
// task. The task exits naturally when cloudflared is killed (EOF).
let drain_handle = tokio::spawn(async move {
while let Ok(Some(line)) = reader.next_line().await {
tracing::trace!("cloudflared: {line}");
}
});
// Drain stdout silently to prevent SIGPIPE/buffer stalls.
if let Some(stdout) = stdout {
tokio::spawn(async move {
let mut out_reader = tokio::io::BufReader::new(stdout).lines();
@@ -122,12 +135,11 @@ impl Tunnel for CloudflareTunnel {
});
}
if let Ok(mut guard) = self.url.write() {
*guard = Some(public_url.clone());
}
let mut guard = self.proc.lock().await;
*guard = Some(TunnelProcess { child });
*guard = Some(TunnelProcess {
child,
_pipe_drain: Some(drain_handle),
});
Ok(public_url)
}
+18 -5
View File
@@ -73,6 +73,7 @@ impl Tunnel for CustomTunnel {
let stderr = child.stderr.take();
let mut public_url = format!("http://{local_host}:{local_port}");
let mut drain_handle: Option<tokio::task::JoinHandle<()>> = None;
if self.url_pattern.is_some()
&& let Some(stdout) = stdout
@@ -103,17 +104,26 @@ impl Tunnel for CustomTunnel {
Err(_) => {}
}
}
// Drain remaining stdout to prevent SIGPIPE/buffer stalls.
tokio::spawn(async move { while let Ok(Some(_)) = reader.next_line().await {} });
// We took ownership of the process's stdout pipe above to parse the
// URL. The process may continue writing to stdout for its lifetime.
// If we drop the reader, the pipe closes and the process gets SIGPIPE.
// We can't just store the reader without reading — the OS pipe buffer
// fills up and the process blocks. So we drain it in a background task.
// The task exits naturally when the process is killed (EOF).
drain_handle = Some(tokio::spawn(async move {
while let Ok(Some(line)) = reader.next_line().await {
tracing::trace!("custom-tunnel: {line}");
}
}));
} else if let Some(stdout) = stdout {
// No url_pattern: still drain stdout to prevent pipe stalls.
// No url_pattern: still drain stdout to prevent SIGPIPE/buffer stalls.
tokio::spawn(async move {
let mut reader = tokio::io::BufReader::new(stdout).lines();
while let Ok(Some(_)) = reader.next_line().await {}
});
}
// Drain stderr silently.
// Drain stderr to prevent SIGPIPE/buffer stalls.
if let Some(stderr) = stderr {
tokio::spawn(async move {
let mut reader = tokio::io::BufReader::new(stderr).lines();
@@ -126,7 +136,10 @@ impl Tunnel for CustomTunnel {
}
let mut guard = self.proc.lock().await;
*guard = Some(TunnelProcess { child });
*guard = Some(TunnelProcess {
child,
_pipe_drain: drain_handle,
});
Ok(public_url)
}
+120 -16
View File
@@ -66,6 +66,10 @@ pub trait Tunnel: Send + Sync {
/// Wraps a spawned tunnel child process.
pub(crate) struct TunnelProcess {
pub child: tokio::process::Child,
/// Background task that drains the process's output pipe (stdout or stderr).
/// Must stay alive or the process dies (SIGPIPE from closed pipe) or hangs
/// (OS pipe buffer fills up, blocking the process's writes).
pub _pipe_drain: Option<tokio::task::JoinHandle<()>>,
}
pub(crate) type SharedProcess = Arc<Mutex<Option<TunnelProcess>>>;
@@ -182,6 +186,22 @@ pub fn create_tunnel(config: &TunnelProviderConfig) -> Result<Option<Box<dyn Tun
// ── Managed tunnel startup ───────────────────────────────────────
/// Determine which local address the tunnel should forward traffic to.
///
/// Prefers the webhook server (`HTTP_PORT`) since that's where webhook routes
/// (Telegram, etc.) are served. Falls back to the gateway port if configured,
/// otherwise defaults to 0.0.0.0:8080 (the same fallback the webhook server
/// uses in main.rs when no HTTP config is present).
fn resolve_tunnel_target(channels: &crate::config::ChannelsConfig) -> (&str, u16) {
if let Some(ref http) = channels.http {
return (http.host.as_str(), http.port);
}
if let Some(ref gw) = channels.gateway {
return (gw.host.as_str(), gw.port);
}
("0.0.0.0", 8080)
}
/// Start a managed tunnel if configured and no static URL is already set.
///
/// Returns the (potentially mutated) config with `tunnel.public_url` set,
@@ -201,28 +221,17 @@ pub async fn start_managed_tunnel(
return (config, None);
};
let gateway_port = config
.channels
.gateway
.as_ref()
.map(|g| g.port)
.unwrap_or(3000);
let gateway_host = config
.channels
.gateway
.as_ref()
.map(|g| g.host.as_str())
.unwrap_or("127.0.0.1");
let (tunnel_host, tunnel_port) = resolve_tunnel_target(&config.channels);
match create_tunnel(provider_config) {
Ok(Some(tunnel)) => {
tracing::debug!(
"Starting {} tunnel on {}:{}...",
tunnel.name(),
gateway_host,
gateway_port
tunnel_host,
tunnel_port
);
match tunnel.start(gateway_host, gateway_port).await {
match tunnel.start(tunnel_host, tunnel_port).await {
Ok(url) => {
tracing::debug!("Tunnel started: {}", url);
config.tunnel.public_url = Some(url);
@@ -383,10 +392,105 @@ mod tests {
{
let mut guard = proc.lock().await;
*guard = Some(TunnelProcess { child });
*guard = Some(TunnelProcess {
child,
_pipe_drain: None,
});
}
kill_shared(&proc).await.unwrap();
assert!(proc.lock().await.is_none());
}
// ── Port selection regression tests ──────────────────────────────
fn base_channels() -> crate::config::ChannelsConfig {
crate::config::ChannelsConfig {
cli: crate::config::CliConfig { enabled: false },
http: None,
gateway: None,
signal: None,
wasm_channels_dir: std::env::temp_dir().join("ironclaw-test-channels"),
wasm_channels_enabled: false,
wasm_channel_owner_ids: std::collections::HashMap::new(),
}
}
fn channels_with_http(host: &str, port: u16) -> crate::config::ChannelsConfig {
let mut c = base_channels();
c.http = Some(crate::config::HttpConfig {
host: host.to_string(),
port,
webhook_secret: None,
user_id: "test".to_string(),
});
c.gateway = Some(crate::config::GatewayConfig {
host: "127.0.0.1".to_string(),
port: 3000,
auth_token: None,
user_id: "test".to_string(),
});
c
}
fn channels_gateway_only(host: &str, port: u16) -> crate::config::ChannelsConfig {
let mut c = base_channels();
c.gateway = Some(crate::config::GatewayConfig {
host: host.to_string(),
port,
auth_token: None,
user_id: "test".to_string(),
});
c
}
fn channels_neither() -> crate::config::ChannelsConfig {
base_channels()
}
#[test]
fn tunnel_target_prefers_http_port() {
let channels = channels_with_http("0.0.0.0", 8080);
let (host, port) = resolve_tunnel_target(&channels);
assert_eq!(host, "0.0.0.0"); // safety: test-only
assert_eq!(port, 8080); // safety: test-only
}
#[test]
fn tunnel_target_falls_back_to_gateway() {
let channels = channels_gateway_only("10.0.0.1", 4000);
let (host, port) = resolve_tunnel_target(&channels);
assert_eq!(host, "10.0.0.1"); // safety: test-only
assert_eq!(port, 4000); // safety: test-only
}
#[test]
fn tunnel_target_defaults_to_webhook_fallback() {
let channels = channels_neither();
let (host, port) = resolve_tunnel_target(&channels);
// Matches the webhook server's hardcoded fallback in main.rs
assert_eq!(host, "0.0.0.0"); // safety: test-only
assert_eq!(port, 8080); // safety: test-only
}
#[test]
fn tunnel_target_http_takes_priority_over_gateway() {
let channels = channels_with_http("192.168.1.1", 9090);
let (host, port) = resolve_tunnel_target(&channels);
// Should use HTTP config, not gateway's 127.0.0.1:3000
assert_eq!(host, "192.168.1.1"); // safety: test-only
assert_eq!(port, 9090); // safety: test-only
}
#[test]
fn tunnel_target_no_http_no_gateway_matches_webhook_fallback() {
// When HTTP_PORT is not set and gateway is not configured (e.g. WASM
// channels exist but no explicit HTTP config), the webhook server in
// main.rs binds to 0.0.0.0:8080 as a hardcoded fallback. The tunnel
// must target the same address so webhook traffic reaches the right
// server.
let channels = channels_neither();
let (host, port) = resolve_tunnel_target(&channels);
assert_eq!((host, port), ("0.0.0.0", 8080)); // safety: test-only
}
}
+21 -10
View File
@@ -110,12 +110,24 @@ impl Tunnel for NgrokTunnel {
}
}
// Drain stdout silently — ngrok only emits low-level connection events
// to stdout; the pipe must be consumed to prevent SIGPIPE/buffer stalls.
tokio::spawn(async move { while let Ok(Some(_)) = reader.next_line().await {} });
if let Ok(mut guard) = self.url.write() {
*guard = Some(public_url.clone());
}
// Drain stderr silently — with --log stdout all meaningful output goes
// to stdout; stderr only needs to be consumed to prevent pipe stalls.
// We took ownership of ngrok's stdout pipe above to parse the URL.
// ngrok continues writing logs to stdout for its entire lifetime.
// If we drop the reader, the pipe closes and ngrok gets SIGPIPE on
// its next write → process dies. We can't just store the reader
// without reading — the OS pipe buffer (~64KB) fills up and ngrok
// blocks. So we drain it in a background task. The task exits
// naturally when ngrok is killed (EOF on the pipe).
let drain_handle = tokio::spawn(async move {
while let Ok(Some(line)) = reader.next_line().await {
tracing::trace!("ngrok: {line}");
}
});
// Drain stderr silently to prevent SIGPIPE/buffer stalls.
if let Some(stderr) = stderr {
tokio::spawn(async move {
let mut err_reader = tokio::io::BufReader::new(stderr).lines();
@@ -123,12 +135,11 @@ impl Tunnel for NgrokTunnel {
});
}
if let Ok(mut guard) = self.url.write() {
*guard = Some(public_url.clone());
}
let mut guard = self.proc.lock().await;
*guard = Some(TunnelProcess { child });
*guard = Some(TunnelProcess {
child,
_pipe_drain: Some(drain_handle),
});
Ok(public_url)
}
+4 -4
View File
@@ -48,8 +48,8 @@ pub struct WorkerDeps {
pub hooks: Arc<HookRegistry>,
pub timeout: Duration,
pub use_planning: bool,
/// SSE broadcast sender for live job event streaming to the web gateway.
pub sse_tx: Option<tokio::sync::broadcast::Sender<SseEvent>>,
/// SSE manager for live job event streaming to the web gateway.
pub sse_tx: Option<Arc<crate::channels::web::sse::SseManager>>,
/// Approval context for tool execution. When `None`, all non-`Never` tools are
/// blocked (legacy behavior). When `Some`, the context determines which tools
/// are pre-approved for autonomous execution.
@@ -138,7 +138,7 @@ impl Worker {
}
// Broadcast SSE for live web UI updates
if let Some(ref tx) = self.deps.sse_tx {
if let Some(ref sse) = self.deps.sse_tx {
let job_id_str = job_id.to_string();
let event = match event_type {
"message" => Some(SseEvent::JobMessage {
@@ -203,7 +203,7 @@ impl Worker {
_ => None,
};
if let Some(event) = event {
let _ = tx.send(event);
sse.broadcast(event);
}
}
}
+31 -104
View File
@@ -3,14 +3,13 @@
//! Avoids redundant HTTP calls for identical texts by caching embeddings
//! in memory keyed by `SHA-256(model_name + "\0" + text)`.
//!
//! Follows the same cache pattern as `llm::response_cache::CachedProvider`:
//! `HashMap` + `last_accessed` tracking + manual LRU eviction.
//! Uses `lru::LruCache` for O(1) insertion, lookup, and eviction.
use std::collections::HashMap;
use std::num::NonZeroUsize;
use std::sync::{Arc, Mutex};
use std::time::Instant;
use async_trait::async_trait;
use lru::LruCache;
use sha2::{Digest, Sha256};
use crate::workspace::embeddings::{EmbeddingError, EmbeddingProvider};
@@ -22,8 +21,7 @@ pub struct EmbeddingCacheConfig {
///
/// Approximate raw embedding payload: `max_entries × dimension × 4 bytes`.
/// At 10,000 entries × 1536 floats ≈ 58 MB (payload only; actual memory
/// is higher due to HashMap buckets, `[u8; 32]` hash keys, `Vec`/`Instant`
/// per-entry overhead).
/// is higher due to per-entry overhead in the linked-list LRU).
pub max_entries: usize,
}
@@ -35,11 +33,6 @@ impl Default for EmbeddingCacheConfig {
}
}
struct CacheEntry {
embedding: Vec<f32>,
last_accessed: Instant,
}
/// Embedding provider wrapper that caches results in memory.
///
/// Thread-safe via `std::sync::Mutex`. The lock is **never held**
@@ -47,8 +40,7 @@ struct CacheEntry {
/// so a synchronous mutex is cheaper than `tokio::sync::Mutex`.
pub struct CachedEmbeddingProvider {
inner: Arc<dyn EmbeddingProvider>,
cache: Mutex<HashMap<[u8; 32], CacheEntry>>,
config: EmbeddingCacheConfig,
cache: Mutex<LruCache<[u8; 32], Vec<f32>>>,
}
impl CachedEmbeddingProvider {
@@ -56,19 +48,18 @@ impl CachedEmbeddingProvider {
///
/// `config.max_entries` is clamped to at least 1.
pub fn new(inner: Arc<dyn EmbeddingProvider>, config: EmbeddingCacheConfig) -> Self {
let config = EmbeddingCacheConfig {
max_entries: config.max_entries.max(1),
};
if config.max_entries > 100_000 {
let max_entries = config.max_entries.max(1);
if max_entries > 100_000 {
tracing::warn!(
max_entries = config.max_entries,
max_entries,
"Embedding cache size exceeds 100,000 entries; memory usage may be significant"
);
}
// safety: max_entries >= 1 due to .max(1) above
let cap = NonZeroUsize::new(max_entries).expect("clamped to >= 1"); // safety: always >= 1
Self {
inner,
cache: Mutex::new(HashMap::with_capacity(config.max_entries.min(1024))),
config,
cache: Mutex::new(LruCache::new(cap)),
}
}
@@ -100,49 +91,6 @@ impl CachedEmbeddingProvider {
hasher.update(text.as_bytes());
hasher.finalize().into()
}
/// Evict the least-recently-used entry if at capacity (single-entry path).
// TODO: O(n) scan per eviction. If max_entries grows large, switch to
// an ordered data structure (e.g. `IndexMap` with swap_remove, or a
// linked-list LRU like the `lru` crate).
fn evict_lru(cache: &mut HashMap<[u8; 32], CacheEntry>, max_entries: usize) {
while cache.len() >= max_entries {
let oldest_key = cache
.iter()
.min_by_key(|(_, entry)| entry.last_accessed)
.map(|(k, _)| *k);
if let Some(k) = oldest_key {
cache.remove(&k);
} else {
break;
}
}
}
/// Evict the `k` oldest entries in O(n) average time via partial selection.
///
/// Used by `embed_batch` to avoid the O(n×m) cost of calling
/// `evict_lru` per insert.
fn evict_k_oldest(cache: &mut HashMap<[u8; 32], CacheEntry>, k: usize) {
if k == 0 || cache.is_empty() {
return;
}
if k >= cache.len() {
cache.clear();
return;
}
// Partial selection: find the k oldest in O(n) average via
// select_nth_unstable_by_key, then remove the first k entries.
let mut entries: Vec<([u8; 32], Instant)> = cache
.iter()
.map(|(key, entry)| (*key, entry.last_accessed))
.collect();
entries.select_nth_unstable_by_key(k - 1, |(_, t)| *t);
for (key, _) in entries.into_iter().take(k) {
cache.remove(&key);
}
}
}
#[async_trait]
@@ -162,39 +110,32 @@ impl EmbeddingProvider for CachedEmbeddingProvider {
async fn embed(&self, text: &str) -> Result<Vec<f32>, EmbeddingError> {
let key = self.cache_key(text);
// Check cache (short critical section)
// Check cache (short critical section). LruCache::get promotes the
// entry to most-recently-used automatically.
{
let mut guard = self.cache.lock().unwrap_or_else(|e| e.into_inner());
if let Some(entry) = guard.get_mut(&key) {
entry.last_accessed = Instant::now();
if let Some(embedding) = guard.get(&key) {
tracing::trace!("embedding cache hit");
return Ok(entry.embedding.clone());
return Ok(embedding.clone());
}
}
// Lock released before HTTP call.
// NOTE: Thundering herd — multiple concurrent callers with the same
// uncached key will each call the inner provider. This is acceptable:
// embeddings are idempotent and the last writer wins in the HashMap.
// embeddings are idempotent and the last writer wins in the LruCache.
let embedding = self.inner.embed(text).await?;
// Store result. Re-check under lock: another concurrent caller may
// have inserted this key while the lock was released for the HTTP call.
// Store result under lock. Re-check first: another concurrent caller
// may have already cached this key while the lock was released.
{
let mut guard = self.cache.lock().unwrap_or_else(|e| e.into_inner());
if let Some(entry) = guard.get_mut(&key) {
// Thundering herd — another caller already cached it.
// Just touch timestamp; skip the clone.
entry.last_accessed = Instant::now();
if guard.get(&key).is_some() {
// Thundering herd — another caller beat us. LruCache::get
// already promoted it to most-recently-used; skip the clone.
tracing::trace!("embedding cache: concurrent insert, skipping clone");
} else {
Self::evict_lru(&mut guard, self.config.max_entries);
guard.insert(
key,
CacheEntry {
embedding: embedding.clone(),
last_accessed: Instant::now(),
},
);
guard.push(key, embedding.clone());
}
}
@@ -214,11 +155,9 @@ impl EmbeddingProvider for CachedEmbeddingProvider {
{
let mut guard = self.cache.lock().unwrap_or_else(|e| e.into_inner());
let now = Instant::now();
for (i, key) in keys.iter().enumerate() {
if let Some(entry) = guard.get_mut(key) {
entry.last_accessed = now;
results[i] = Some(entry.embedding.clone());
if let Some(embedding) = guard.get(key) {
results[i] = Some(embedding.clone());
} else {
miss_indices.push(i);
}
@@ -228,7 +167,6 @@ impl EmbeddingProvider for CachedEmbeddingProvider {
if miss_indices.is_empty() {
tracing::trace!(count = texts.len(), "embedding batch: all cache hits");
// All slots populated from cache hits
return results
.into_iter()
.enumerate()
@@ -260,29 +198,18 @@ impl EmbeddingProvider for CachedEmbeddingProvider {
"embedding batch: partial cache"
);
// Cache FIRST (clone only the cacheable subset), then move originals
// into results. This avoids cloning capacity-skipped embeddings entirely.
// Cache only the last `cap` new embeddings — caching more than the
// cache capacity wastes clone work on entries that are immediately evicted.
{
let mut guard = self.cache.lock().unwrap_or_else(|e| e.into_inner());
let cacheable = miss_indices.len().min(self.config.max_entries);
let skip = miss_indices.len() - cacheable;
let need_to_evict = (guard.len() + cacheable).saturating_sub(self.config.max_entries);
if need_to_evict > 0 {
Self::evict_k_oldest(&mut guard, need_to_evict);
}
let now = Instant::now();
let cap = guard.cap().get();
let skip = miss_indices.len().saturating_sub(cap);
for (&orig_idx, emb) in miss_indices[skip..].iter().zip(&new_embeddings[skip..]) {
guard.insert(
keys[orig_idx],
CacheEntry {
embedding: emb.clone(),
last_accessed: now,
},
);
guard.push(keys[orig_idx], emb.clone());
}
}
// Move originals into results (zero-copy for all, including cached ones).
// Move originals into results (zero-copy).
for (orig_idx, emb) in miss_indices.iter().copied().zip(new_embeddings) {
results[orig_idx] = Some(emb);
}
+191
View File
@@ -0,0 +1,191 @@
//! Tests for batch_get_last_run_status (#1469 N+1 fix).
//!
//! Verifies:
//! 1. Empty input returns empty map
//! 2. Returns the most recent run status per routine
//! 3. Routines with no runs are omitted from result
//! 4. Multiple routines with different statuses are correctly returned
#[cfg(feature = "libsql")]
mod tests {
use std::sync::Arc;
use chrono::{Duration, Utc};
use uuid::Uuid;
use ironclaw::agent::routine::{
Routine, RoutineAction, RoutineGuardrails, RoutineRun, RunStatus, Trigger,
};
use ironclaw::db::Database;
async fn create_test_db() -> (Arc<dyn Database>, tempfile::TempDir) {
use ironclaw::db::libsql::LibSqlBackend;
let temp_dir = tempfile::tempdir().expect("tempdir");
let db_path = temp_dir.path().join("test.db");
let backend = LibSqlBackend::new_local(&db_path)
.await
.expect("LibSqlBackend");
backend.run_migrations().await.expect("migrations");
let db: Arc<dyn Database> = Arc::new(backend);
(db, temp_dir)
}
fn make_routine(id: Uuid) -> Routine {
Routine {
id,
name: format!("test-routine-{}", id),
description: "Test routine".to_string(),
user_id: "default".to_string(),
enabled: true,
trigger: Trigger::Manual,
action: RoutineAction::FullJob {
title: "Test job".to_string(),
description: "Test description".to_string(),
max_iterations: 5,
},
guardrails: RoutineGuardrails {
cooldown: std::time::Duration::from_secs(0),
max_concurrent: 1,
dedup_window: None,
},
notify: Default::default(),
last_run_at: None,
next_fire_at: None,
run_count: 0,
consecutive_failures: 0,
state: serde_json::json!({}),
created_at: Utc::now(),
updated_at: Utc::now(),
}
}
fn make_run(
routine_id: Uuid,
status: RunStatus,
started_at: chrono::DateTime<chrono::Utc>,
) -> RoutineRun {
RoutineRun {
id: Uuid::new_v4(),
routine_id,
trigger_type: "manual".to_string(),
trigger_detail: None,
started_at,
completed_at: if status == RunStatus::Running {
None
} else {
Some(Utc::now())
},
status,
result_summary: None,
tokens_used: None,
job_id: None,
created_at: Utc::now(),
}
}
#[tokio::test]
async fn test_batch_get_last_run_status_empty_input() {
let (db, _tmp) = create_test_db().await;
let result = db
.batch_get_last_run_status(&[])
.await
.expect("batch query");
assert!(result.is_empty());
}
#[tokio::test]
async fn test_batch_get_last_run_status_returns_latest() {
let (db, _tmp) = create_test_db().await;
let routine_id = Uuid::new_v4();
db.create_routine(&make_routine(routine_id))
.await
.expect("create routine");
// Create an older run with Ok status
let older_run = make_run(routine_id, RunStatus::Ok, Utc::now() - Duration::hours(2));
db.create_routine_run(&older_run)
.await
.expect("create older run");
db.complete_routine_run(older_run.id, RunStatus::Ok, None, None)
.await
.expect("complete older run");
// Create a newer run with Attention status
let newer_run = make_run(
routine_id,
RunStatus::Attention,
Utc::now() - Duration::hours(1),
);
db.create_routine_run(&newer_run)
.await
.expect("create newer run");
db.complete_routine_run(newer_run.id, RunStatus::Attention, None, None)
.await
.expect("complete newer run");
let result = db
.batch_get_last_run_status(&[routine_id])
.await
.expect("batch query");
assert_eq!(result.get(&routine_id), Some(&RunStatus::Attention));
}
#[tokio::test]
async fn test_batch_get_last_run_status_omits_routines_without_runs() {
let (db, _tmp) = create_test_db().await;
let with_runs = Uuid::new_v4();
let without_runs = Uuid::new_v4();
db.create_routine(&make_routine(with_runs))
.await
.expect("create routine");
db.create_routine(&make_routine(without_runs))
.await
.expect("create routine");
let run = make_run(with_runs, RunStatus::Ok, Utc::now());
db.create_routine_run(&run).await.expect("create run");
db.complete_routine_run(run.id, RunStatus::Ok, None, None)
.await
.expect("complete run");
let result = db
.batch_get_last_run_status(&[with_runs, without_runs])
.await
.expect("batch query");
assert_eq!(result.get(&with_runs), Some(&RunStatus::Ok));
assert_eq!(result.get(&without_runs), None);
}
#[tokio::test]
async fn test_batch_get_last_run_status_multiple_routines() {
let (db, _tmp) = create_test_db().await;
let r1 = Uuid::new_v4();
let r2 = Uuid::new_v4();
db.create_routine(&make_routine(r1))
.await
.expect("create r1");
db.create_routine(&make_routine(r2))
.await
.expect("create r2");
let run1 = make_run(r1, RunStatus::Running, Utc::now());
db.create_routine_run(&run1).await.expect("create run1");
let run2 = make_run(r2, RunStatus::Failed, Utc::now());
db.create_routine_run(&run2).await.expect("create run2");
db.complete_routine_run(run2.id, RunStatus::Failed, None, None)
.await
.expect("complete run2");
let result = db
.batch_get_last_run_status(&[r1, r2])
.await
.expect("batch query");
assert_eq!(result.get(&r1), Some(&RunStatus::Running));
assert_eq!(result.get(&r2), Some(&RunStatus::Failed));
}
}
+1 -1
View File
@@ -661,7 +661,7 @@ mod advanced {
.await
.expect("failed to inject test token");
let activate_result = ext_mgr.activate("mock-notion").await;
let activate_result = ext_mgr.activate("mock-notion", "default").await;
assert!(
activate_result.is_ok(),
"activation failed: {:?}",
+47 -6
View File
@@ -205,11 +205,11 @@ mod tests {
}
// -----------------------------------------------------------------------
// Test 5: routine_manual_create
// Test 5: routine_manual_create_defaults_to_tools_enabled
// -----------------------------------------------------------------------
#[tokio::test]
async fn routine_manual_create() {
async fn routine_manual_create_defaults_to_tools_enabled() {
let trace = LlmTrace::from_file(concat!(
env!("CARGO_MANIFEST_DIR"),
"/tests/fixtures/llm_traces/tools/routine_manual_create.json"
@@ -237,8 +237,8 @@ mod tests {
assert!(matches!(routine.trigger, Trigger::Manual));
assert!(
matches!(&routine.action, RoutineAction::Lightweight { use_tools, .. } if !*use_tools),
"manual routine should default to lightweight without tools: {:?}",
matches!(&routine.action, RoutineAction::Lightweight { use_tools, .. } if *use_tools),
"manual routine should default to lightweight with tools enabled: {:?}",
routine.action
);
@@ -246,7 +246,48 @@ mod tests {
}
// -----------------------------------------------------------------------
// Test 6: routine_history
// Test 6: routine_manual_create_explicit_no_tools
// -----------------------------------------------------------------------
#[tokio::test]
async fn routine_manual_create_explicit_no_tools() {
let trace = LlmTrace::from_file(concat!(
env!("CARGO_MANIFEST_DIR"),
"/tests/fixtures/llm_traces/tools/routine_manual_create_no_tools.json"
))
.expect("failed to load routine_manual_create_no_tools.json");
let rig = TestRigBuilder::new()
.with_trace(trace.clone())
.with_auto_approve_tools(true)
.build()
.await;
rig.send_message("Create a manual routine for quiet text-only bug triage")
.await;
let responses = rig.wait_for_responses(1, Duration::from_secs(15)).await;
rig.verify_trace_expects(&trace, &responses);
let routine = rig
.database()
.get_routine_by_name("test-user", "manual-triage-no-tools")
.await
.expect("get_routine_by_name")
.expect("manual-triage-no-tools should exist");
assert!(matches!(routine.trigger, Trigger::Manual));
assert!(
matches!(&routine.action, RoutineAction::Lightweight { use_tools, .. } if !*use_tools),
"manual routine should preserve explicit use_tools=false: {:?}",
routine.action
);
rig.shutdown();
}
// -----------------------------------------------------------------------
// Test 7: routine_history
// -----------------------------------------------------------------------
#[tokio::test]
@@ -283,7 +324,7 @@ mod tests {
}
// -----------------------------------------------------------------------
// Test 7: routine_system_event_emit
// Test 8: routine_system_event_emit
// -----------------------------------------------------------------------
#[tokio::test]
+12 -32
View File
@@ -561,11 +561,7 @@ mod tests {
"deploy to production now",
);
let fired = engine
.check_event_triggers(
&matching_msg.user_id,
&matching_msg.channel,
&matching_msg.content,
)
.check_event_triggers(&matching_msg, &matching_msg.content)
.await;
assert!(
fired >= 1,
@@ -584,11 +580,7 @@ mod tests {
"check the staging environment",
);
let fired_neg = engine
.check_event_triggers(
&non_matching_msg.user_id,
&non_matching_msg.channel,
&non_matching_msg.content,
)
.check_event_triggers(&non_matching_msg, &non_matching_msg.content)
.await;
assert_eq!(fired_neg, 0, "Expected 0 routines fired on non-match");
}
@@ -652,7 +644,7 @@ mod tests {
"deploy to production now",
);
let guest_fired = engine
.check_event_triggers(&guest_msg.user_id, &guest_msg.channel, &guest_msg.content)
.check_event_triggers(&guest_msg, &guest_msg.content)
.await;
assert_eq!(
guest_fired, 0,
@@ -677,7 +669,7 @@ mod tests {
"deploy to production now",
);
let owner_fired = engine
.check_event_triggers(&owner_msg.user_id, &owner_msg.channel, &owner_msg.content)
.check_event_triggers(&owner_msg, &owner_msg.content)
.await;
assert!(
owner_fired >= 1,
@@ -906,9 +898,7 @@ mod tests {
"default",
"test-cooldown trigger",
);
let fired1 = engine
.check_event_triggers(&msg.user_id, &msg.channel, &msg.content)
.await;
let fired1 = engine.check_event_triggers(&msg, &msg.content).await;
assert!(fired1 >= 1, "First fire should work");
// Give spawn time, then update last_run_at to simulate recent execution.
@@ -923,9 +913,7 @@ mod tests {
engine.refresh_event_cache().await;
// Second fire should be blocked by cooldown.
let fired2 = engine
.check_event_triggers(&msg.user_id, &msg.channel, &msg.content)
.await;
let fired2 = engine.check_event_triggers(&msg, &msg.content).await;
assert_eq!(fired2, 0, "Second fire should be blocked by cooldown");
}
@@ -1095,9 +1083,7 @@ mod tests {
engine.refresh_event_cache().await;
let msg = IncomingMessage::new("test", "default", "DISABLE_ME");
let fired_before = engine
.check_event_triggers(&msg.user_id, &msg.channel, &msg.content)
.await;
let fired_before = engine.check_event_triggers(&msg, &msg.content).await;
assert!(fired_before >= 1, "Expected routine to fire before disable");
// Simulate what routines_toggle_handler now does: update DB, then refresh.
@@ -1106,9 +1092,7 @@ mod tests {
db.update_routine(&routine).await.expect("update_routine");
engine.refresh_event_cache().await;
let fired_after = engine
.check_event_triggers(&msg.user_id, &msg.channel, &msg.content)
.await;
let fired_after = engine.check_event_triggers(&msg, &msg.content).await;
assert_eq!(
fired_after, 0,
"Disabled routine must not fire after cache refresh"
@@ -1134,10 +1118,7 @@ mod tests {
let msg = IncomingMessage::new("test", "default", "DELETE_ME");
assert!(
engine
.check_event_triggers(&msg.user_id, &msg.channel, &msg.content)
.await
>= 1,
engine.check_event_triggers(&msg, &msg.content).await >= 1,
"Expected routine to fire before delete"
);
@@ -1146,9 +1127,7 @@ mod tests {
engine.refresh_event_cache().await;
assert_eq!(
engine
.check_event_triggers(&msg.user_id, &msg.channel, &msg.content)
.await,
engine.check_event_triggers(&msg, &msg.content).await,
0,
"Deleted routine must not fire after cache refresh"
);
@@ -1462,8 +1441,9 @@ mod tests {
db.create_routine(&routine).await.expect("create_routine");
engine.refresh_event_cache().await;
let trigger_msg = IncomingMessage::new("test", "default", "owner-gate");
let fired = engine
.check_event_triggers("default", "test", "owner-gate")
.check_event_triggers(&trigger_msg, &trigger_msg.content)
.await;
assert_eq!(fired, 1, "expected one matching event routine");
+1
View File
@@ -200,6 +200,7 @@ mod tests {
document_extraction: None,
sandbox_readiness: ironclaw::agent::SandboxReadiness::DisabledByConfig,
builder: None,
llm_backend: "nearai".to_string(),
};
let gateway = Arc::new(TestChannel::new());
@@ -0,0 +1,39 @@
{
"model_name": "test-routine-manual-create-no-tools",
"expects": {
"tools_used": ["routine_create"],
"all_tools_succeeded": true,
"min_responses": 1
},
"steps": [
{
"response": {
"type": "tool_calls",
"tool_calls": [
{
"id": "call_rc_manual_2",
"name": "routine_create",
"arguments": {
"name": "manual-triage-no-tools",
"trigger_type": "manual",
"prompt": "Summarize the latest bug reports when this routine is fired.",
"execution": {
"use_tools": false
}
}
}
],
"input_tokens": 90,
"output_tokens": 24
}
},
{
"response": {
"type": "text",
"content": "Created the manual-triage-no-tools routine. It will only run when explicitly fired and stay text-only.",
"input_tokens": 140,
"output_tokens": 18
}
}
]
}
+1 -1
View File
@@ -216,7 +216,7 @@ async fn extension_manager_with_process_manager_constructs() {
);
// Verify the manager is functional — list returns Ok.
let result = manager.list(None, false).await;
let result = manager.list(None, false, "test").await;
assert!(result.is_ok(), "list should succeed on empty manager");
assert!(result.unwrap().is_empty());
}
File diff suppressed because it is too large Load Diff
+240
View File
@@ -0,0 +1,240 @@
//! Tests proving that multi-tenant system prompts are broken.
//!
//! Bug: In multi-tenant mode, the agent loop uses `self.workspace()` which
//! returns a single shared workspace (user_id="default"). Identity files
//! (IDENTITY.md, SOUL.md, USER.md) seeded under per-user IDs ("alice",
//! "bob") are invisible to this workspace, so the system prompt is
//! empty/wrong.
//!
//! These tests:
//! 1. Seed identity files for two users (alice, bob) in the database
//! 2. Send messages as each user
//! 3. Verify the system prompt in captured LLM requests contains the
//! correct user's identity
//! 4. Verify user A's identity doesn't leak into user B's prompt
//!
//! All tests are expected to FAIL until the bug is fixed.
#[cfg(feature = "libsql")]
mod support;
#[cfg(feature = "libsql")]
mod tests {
use std::sync::Arc;
use std::time::Duration;
use ironclaw::channels::IncomingMessage;
use ironclaw::llm::Role;
use ironclaw::workspace::Workspace;
use crate::support::test_rig::TestRigBuilder;
use crate::support::trace_llm::{LlmTrace, TraceResponse, TraceStep};
const TIMEOUT: Duration = Duration::from_secs(15);
const ALICE_USER_ID: &str = "alice";
const BOB_USER_ID: &str = "bob";
const ALICE_IDENTITY: &str = "You are Alice's personal assistant. \
Alice is a software engineer who lives in Seattle.";
const BOB_IDENTITY: &str = "You are Bob's personal assistant. \
Bob is a marine biologist who lives in Miami.";
/// Create a simple trace that returns a canned text response.
/// We need one step per message we plan to send.
fn simple_trace(num_steps: usize) -> LlmTrace {
let steps: Vec<TraceStep> = (0..num_steps)
.map(|i| TraceStep {
request_hint: None,
response: TraceResponse::Text {
content: format!("Response {}", i),
input_tokens: 100,
output_tokens: 10,
},
expected_tool_results: Vec::new(),
})
.collect();
// Create separate turns for each step so the trace replays correctly.
let turns: Vec<crate::support::trace_llm::TraceTurn> = steps
.into_iter()
.enumerate()
.map(|(i, step)| crate::support::trace_llm::TraceTurn {
user_input: format!("message {}", i),
steps: vec![step],
expects: Default::default(),
})
.collect();
LlmTrace::new("test-model", turns)
}
/// Seed identity files for a user by creating a workspace scoped to that
/// user and writing IDENTITY.md.
async fn seed_identity(db: &Arc<dyn ironclaw::db::Database>, user_id: &str, content: &str) {
let ws = Workspace::new_with_db(user_id, db.clone());
ws.write("IDENTITY.md", content)
.await
.unwrap_or_else(|e| panic!("Failed to seed IDENTITY.md for {user_id}: {e}"));
}
/// Extract the system prompt from captured LLM requests.
///
/// The system prompt is the first message with role=System in the first
/// LLM request for a given turn.
fn extract_system_prompt(requests: &[Vec<ironclaw::llm::ChatMessage>]) -> Option<String> {
requests.last().and_then(|msgs| {
msgs.iter()
.find(|m| matches!(m.role, Role::System))
.map(|m| m.content.clone())
})
}
// -----------------------------------------------------------------------
// Test 1: Alice's identity should appear in system prompt when messaging
// as Alice.
// -----------------------------------------------------------------------
#[tokio::test]
async fn alice_system_prompt_contains_alice_identity() {
let trace = simple_trace(1);
let rig = TestRigBuilder::new().with_trace(trace).build().await;
// Seed alice's identity into the database
let db = rig.database();
seed_identity(db, ALICE_USER_ID, ALICE_IDENTITY).await;
// Send a message AS alice (using her user_id)
let msg = IncomingMessage::new("test", ALICE_USER_ID, "Hello, who am I?");
rig.send_incoming(msg).await;
let _responses = rig.wait_for_responses(1, TIMEOUT).await;
// The system prompt sent to the LLM should contain Alice's identity
let requests = rig.captured_llm_requests();
let system_prompt =
extract_system_prompt(&requests).expect("Expected a system prompt in the LLM request");
assert!(
system_prompt.contains("Alice is a software engineer"),
"System prompt should contain Alice's identity when messaging as Alice.\n\
Actual system prompt:\n{system_prompt}"
);
rig.shutdown();
}
// -----------------------------------------------------------------------
// Test 2: Bob's identity should appear in system prompt when messaging
// as Bob.
// -----------------------------------------------------------------------
#[tokio::test]
async fn bob_system_prompt_contains_bob_identity() {
let trace = simple_trace(1);
let rig = TestRigBuilder::new().with_trace(trace).build().await;
// Seed bob's identity into the database
let db = rig.database();
seed_identity(db, BOB_USER_ID, BOB_IDENTITY).await;
// Send a message AS bob
let msg = IncomingMessage::new("test", BOB_USER_ID, "Hello, who am I?");
rig.send_incoming(msg).await;
let _responses = rig.wait_for_responses(1, TIMEOUT).await;
// The system prompt should contain Bob's identity
let requests = rig.captured_llm_requests();
let system_prompt =
extract_system_prompt(&requests).expect("Expected a system prompt in the LLM request");
assert!(
system_prompt.contains("Bob is a marine biologist"),
"System prompt should contain Bob's identity when messaging as Bob.\n\
Actual system prompt:\n{system_prompt}"
);
rig.shutdown();
}
// -----------------------------------------------------------------------
// Test 3: Alice's identity must NOT appear in Bob's system prompt.
// -----------------------------------------------------------------------
#[tokio::test]
async fn alice_identity_does_not_leak_into_bob_prompt() {
let trace = simple_trace(1);
let rig = TestRigBuilder::new().with_trace(trace).build().await;
// Seed BOTH users' identities
let db = rig.database();
seed_identity(db, ALICE_USER_ID, ALICE_IDENTITY).await;
seed_identity(db, BOB_USER_ID, BOB_IDENTITY).await;
// Send a message AS bob
let msg = IncomingMessage::new("test", BOB_USER_ID, "Tell me about myself");
rig.send_incoming(msg).await;
let _responses = rig.wait_for_responses(1, TIMEOUT).await;
// Bob's prompt must NOT contain Alice's identity
let requests = rig.captured_llm_requests();
let system_prompt = extract_system_prompt(&requests);
if let Some(ref prompt) = system_prompt {
assert!(
!prompt.contains("Alice is a software engineer"),
"Alice's identity LEAKED into Bob's system prompt!\n\
System prompt:\n{prompt}"
);
}
// Also verify Bob's identity IS present (compound check)
let prompt = system_prompt.expect("Expected a system prompt in the LLM request");
assert!(
prompt.contains("Bob is a marine biologist"),
"Bob's own identity should be in his system prompt.\n\
Actual system prompt:\n{prompt}"
);
rig.shutdown();
}
// -----------------------------------------------------------------------
// Test 4: Bob's identity must NOT appear in Alice's system prompt.
// -----------------------------------------------------------------------
#[tokio::test]
async fn bob_identity_does_not_leak_into_alice_prompt() {
let trace = simple_trace(1);
let rig = TestRigBuilder::new().with_trace(trace).build().await;
// Seed BOTH users' identities
let db = rig.database();
seed_identity(db, ALICE_USER_ID, ALICE_IDENTITY).await;
seed_identity(db, BOB_USER_ID, BOB_IDENTITY).await;
// Send a message AS alice
let msg = IncomingMessage::new("test", ALICE_USER_ID, "Tell me about myself");
rig.send_incoming(msg).await;
let _responses = rig.wait_for_responses(1, TIMEOUT).await;
// Alice's prompt must NOT contain Bob's identity
let requests = rig.captured_llm_requests();
let system_prompt = extract_system_prompt(&requests);
if let Some(ref prompt) = system_prompt {
assert!(
!prompt.contains("Bob is a marine biologist"),
"Bob's identity LEAKED into Alice's system prompt!\n\
System prompt:\n{prompt}"
);
}
// Also verify Alice's identity IS present
let prompt = system_prompt.expect("Expected a system prompt in the LLM request");
assert!(
prompt.contains("Alice is a software engineer"),
"Alice's own identity should be in her system prompt.\n\
Actual system prompt:\n{prompt}"
);
rig.shutdown();
}
}
+22 -13
View File
@@ -191,8 +191,9 @@ async fn start_test_server_with_provider(
) -> (SocketAddr, Arc<GatewayState>) {
let state = Arc::new(GatewayState {
msg_tx: tokio::sync::RwLock::new(None),
sse: SseManager::new(),
sse: Arc::new(SseManager::new()),
workspace: None,
workspace_pool: None,
session_manager: None,
log_broadcaster: None,
log_level_handle: None,
@@ -202,13 +203,13 @@ async fn start_test_server_with_provider(
job_manager: None,
prompt_queue: None,
scheduler: None,
user_id: "test-user".to_string(),
default_user_id: "test-user".to_string(),
shutdown_tx: tokio::sync::RwLock::new(None),
ws_tracker: Some(Arc::new(WsConnectionTracker::new())),
llm_provider: Some(llm_provider),
skill_registry: None,
skill_catalog: None,
chat_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(30, 60),
chat_rate_limiter: ironclaw::channels::web::server::PerUserRateLimiter::new(30, 60),
oauth_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(10, 60),
webhook_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(10, 60),
registry_entries: Vec::new(),
@@ -218,8 +219,12 @@ async fn start_test_server_with_provider(
active_config: ironclaw::channels::web::server::ActiveConfigSnapshot::default(),
});
let auth = ironclaw::channels::web::auth::MultiAuthState::single(
AUTH_TOKEN.to_string(),
"test-user".to_string(),
);
let addr: SocketAddr = "127.0.0.1:0".parse().unwrap();
let bound_addr = start_server(addr, state.clone(), AUTH_TOKEN.to_string())
let bound_addr = start_server(addr, state.clone(), auth)
.await
.expect("Failed to start test server");
@@ -684,8 +689,9 @@ async fn test_no_llm_provider_returns_503() {
// Create state WITHOUT llm_provider
let state = Arc::new(GatewayState {
msg_tx: tokio::sync::RwLock::new(None),
sse: SseManager::new(),
sse: Arc::new(SseManager::new()),
workspace: None,
workspace_pool: None,
session_manager: None,
log_broadcaster: None,
log_level_handle: None,
@@ -695,13 +701,13 @@ async fn test_no_llm_provider_returns_503() {
job_manager: None,
prompt_queue: None,
scheduler: None,
user_id: "test-user".to_string(),
default_user_id: "test-user".to_string(),
shutdown_tx: tokio::sync::RwLock::new(None),
ws_tracker: Some(Arc::new(WsConnectionTracker::new())),
llm_provider: None, // No LLM!
skill_registry: None,
skill_catalog: None,
chat_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(30, 60),
chat_rate_limiter: ironclaw::channels::web::server::PerUserRateLimiter::new(30, 60),
oauth_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(10, 60),
webhook_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(10, 60),
registry_entries: Vec::new(),
@@ -711,10 +717,12 @@ async fn test_no_llm_provider_returns_503() {
active_config: ironclaw::channels::web::server::ActiveConfigSnapshot::default(),
});
let auth = ironclaw::channels::web::auth::MultiAuthState::single(
AUTH_TOKEN.to_string(),
"test-user".to_string(),
);
let addr: SocketAddr = "127.0.0.1:0".parse().unwrap();
let bound_addr = start_server(addr, state, AUTH_TOKEN.to_string())
.await
.unwrap();
let bound_addr = start_server(addr, state, auth).await.unwrap();
let url = format!("http://{}/v1/chat/completions", bound_addr);
let resp = client()
@@ -741,9 +749,10 @@ async fn test_chat_completions_body_too_large() {
let state = ironclaw::channels::web::test_helpers::TestGatewayBuilder::new()
.llm_provider(llm_provider)
.build();
let auth_state = ironclaw::channels::web::auth::AuthState {
token: AUTH_TOKEN.to_string(),
};
let auth_state = ironclaw::channels::web::auth::MultiAuthState::single(
AUTH_TOKEN.to_string(),
"test-user".to_string(),
);
let app = Router::new()
.route(
+12 -6
View File
@@ -13,8 +13,11 @@ use ironclaw::agent::routine_engine::RoutineEngine;
use ironclaw::agent::{Agent, AgentDeps, SessionManager as AgentSessionManager};
use ironclaw::app::{AppBuilder, AppBuilderFlags};
use ironclaw::channels::IncomingMessage;
use ironclaw::channels::web::auth::MultiAuthState;
use ironclaw::channels::web::log_layer::LogBroadcaster;
use ironclaw::channels::web::server::{GatewayState, RateLimiter, start_server};
use ironclaw::channels::web::server::{
GatewayState, PerUserRateLimiter, RateLimiter, start_server,
};
use ironclaw::channels::web::sse::SseManager;
use ironclaw::channels::web::ws::WsConnectionTracker;
use ironclaw::config::{Config, RegistryProviderConfig, RoutineConfig};
@@ -211,8 +214,9 @@ impl GatewayWorkflowHarness {
let gateway_state = Arc::new(GatewayState {
msg_tx: tokio::sync::RwLock::new(Some(gw_tx)),
sse: SseManager::new(),
sse: Arc::new(SseManager::new()),
workspace: components.workspace.clone(),
workspace_pool: None,
session_manager: Some(Arc::clone(&agent_session_manager)),
log_broadcaster: None,
log_level_handle: None,
@@ -222,13 +226,13 @@ impl GatewayWorkflowHarness {
job_manager: None,
prompt_queue: None,
scheduler: Some(scheduler_slot.clone()),
user_id: user_id.clone(),
default_user_id: user_id.clone(),
shutdown_tx: tokio::sync::RwLock::new(None),
ws_tracker: Some(Arc::new(WsConnectionTracker::new())),
llm_provider: Some(Arc::clone(&components.llm)),
skill_registry: components.skill_registry.clone(),
skill_catalog: components.skill_catalog.clone(),
chat_rate_limiter: RateLimiter::new(120, 60),
chat_rate_limiter: PerUserRateLimiter::new(120, 60),
oauth_rate_limiter: RateLimiter::new(10, 60),
webhook_rate_limiter: RateLimiter::new(10, 60),
registry_entries: Vec::new(),
@@ -254,12 +258,13 @@ impl GatewayWorkflowHarness {
skills_config: components.config.skills.clone(),
hooks: components.hooks,
cost_guard: components.cost_guard,
sse_tx: Some(gateway_state.sse.sender()),
sse_tx: None,
http_interceptor: None,
transcription: None,
document_extraction: None,
sandbox_readiness: ironclaw::agent::SandboxReadiness::DisabledByConfig,
builder: None,
llm_backend: "nearai".to_string(),
},
channels,
None,
@@ -288,10 +293,11 @@ impl GatewayWorkflowHarness {
}
let auth_token = "gateway-test-token".to_string();
let auth = MultiAuthState::single(auth_token.clone(), user_id.clone());
let addr = start_server(
"127.0.0.1:0".parse().expect("valid localhost addr"),
Arc::clone(&gateway_state),
auth_token.clone(),
auth,
)
.await
.expect("failed to start gateway server");
+5 -11
View File
@@ -701,7 +701,7 @@ impl TestRigBuilder {
let wasm_bytes = tokio::fs::read(&spec.wasm_path)
.await
.unwrap_or_else(|e| panic!("read {}: {e}", spec.wasm_path.display()));
let (capabilities, description, schema) =
let (capabilities, description) =
if let Some(cap_path) = &spec.capabilities_path {
if cap_path.exists() {
let cap_bytes = tokio::fs::read(cap_path)
@@ -709,16 +709,12 @@ impl TestRigBuilder {
.unwrap_or_else(|e| panic!("read {}: {e}", cap_path.display()));
let cap_file = CapabilitiesFile::from_bytes(&cap_bytes)
.expect("parse capabilities.json");
(
cap_file.to_capabilities(),
cap_file.description.clone(),
cap_file.parameters.clone(),
)
(cap_file.to_capabilities(), cap_file.description.clone())
} else {
(Capabilities::default(), None, None)
(Capabilities::default(), None)
}
} else {
(Capabilities::default(), None, None)
(Capabilities::default(), None)
};
let prepared = runtime
@@ -730,9 +726,6 @@ impl TestRigBuilder {
if let Some(desc) = description {
wrapper = wrapper.with_description(desc);
}
if let Some(s) = schema {
wrapper = wrapper.with_schema(s);
}
if let Some(interceptor) = &http_interceptor {
wrapper = wrapper.with_http_interceptor(Arc::clone(interceptor));
}
@@ -768,6 +761,7 @@ impl TestRigBuilder {
document_extraction: None,
sandbox_readiness: ironclaw::agent::SandboxReadiness::Available, // tests don't use real Docker
builder: None,
llm_backend: "nearai".to_string(),
};
// 7. Create TestChannel and ChannelManager.
+9 -4
View File
@@ -39,8 +39,9 @@ async fn start_test_server() -> (
let state = Arc::new(GatewayState {
msg_tx: tokio::sync::RwLock::new(Some(agent_tx)),
sse: SseManager::new(),
sse: Arc::new(SseManager::new()),
workspace: None,
workspace_pool: None,
session_manager: None,
log_broadcaster: None,
log_level_handle: None,
@@ -50,13 +51,13 @@ async fn start_test_server() -> (
job_manager: None,
prompt_queue: None,
scheduler: None,
user_id: "test-user".to_string(),
default_user_id: "test-user".to_string(),
shutdown_tx: tokio::sync::RwLock::new(None),
ws_tracker: Some(Arc::new(WsConnectionTracker::new())),
llm_provider: None,
skill_registry: None,
skill_catalog: None,
chat_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(30, 60),
chat_rate_limiter: ironclaw::channels::web::server::PerUserRateLimiter::new(30, 60),
oauth_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(10, 60),
webhook_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(10, 60),
registry_entries: Vec::new(),
@@ -66,8 +67,12 @@ async fn start_test_server() -> (
active_config: ironclaw::channels::web::server::ActiveConfigSnapshot::default(),
});
let auth = ironclaw::channels::web::auth::MultiAuthState::single(
AUTH_TOKEN.to_string(),
"test-user".to_string(),
);
let addr: SocketAddr = "127.0.0.1:0".parse().unwrap();
let bound_addr = start_server(addr, state.clone(), AUTH_TOKEN.to_string())
let bound_addr = start_server(addr, state.clone(), auth)
.await
.expect("Failed to start test server");
@@ -1,6 +1,7 @@
{
"version": "0.2.1",
"wit_version": "0.3.0",
"description": "Manage GitHub repositories, issues, pull requests, reviews, and workflows. Supports listing, creating, commenting, merging PRs, and triggering GitHub Actions.",
"capabilities": {
"webhook": {
"hmac_secret_name": "github_webhook_secret",
+2 -1
View File
@@ -1,6 +1,7 @@
{
"version": "0.2.0",
"wit_version": "0.3.0",
"description": "Read, search, send, draft, and reply to emails via Gmail. Supports Gmail search query syntax (is:unread, from:, subject:, after:, etc.).",
"http": {
"allowlist": [
{
@@ -53,7 +54,7 @@
},
{
"name": "google_oauth_client_secret",
"prompt": "Google OAuth Client Secret"
"prompt": "Google OAuth Client Secret (from console.cloud.google.com/apis/credentials)"
}
]
}
@@ -1,6 +1,7 @@
{
"version": "0.2.0",
"wit_version": "0.3.0",
"description": "View, create, update, and delete Google Calendar events. Supports timed events, all-day events, attendees, locations, and free text search.",
"http": {
"allowlist": [
{
@@ -52,7 +53,7 @@
},
{
"name": "google_oauth_client_secret",
"prompt": "Google OAuth Client Secret"
"prompt": "Google OAuth Client Secret (from console.cloud.google.com/apis/credentials)"
}
]
}
@@ -1,6 +1,7 @@
{
"version": "0.2.0",
"wit_version": "0.3.0",
"description": "Create, read, edit, and format Google Docs documents. Supports text insert/delete/replace, formatting (bold, italic, font, color, size), paragraph styling, tables, and lists.",
"http": {
"allowlist": [
{
@@ -52,7 +53,7 @@
},
{
"name": "google_oauth_client_secret",
"prompt": "Google OAuth Client Secret"
"prompt": "Google OAuth Client Secret (from console.cloud.google.com/apis/credentials)"
}
]
}
@@ -1,6 +1,7 @@
{
"version": "0.2.0",
"wit_version": "0.3.0",
"description": "Search, access, upload, share, and organize files and folders in Google Drive. Supports personal drives and shared (organizational) drives.",
"http": {
"allowlist": [
{
@@ -57,7 +58,7 @@
},
{
"name": "google_oauth_client_secret",
"prompt": "Google OAuth Client Secret"
"prompt": "Google OAuth Client Secret (from console.cloud.google.com/apis/credentials)"
}
]
}
@@ -1,6 +1,7 @@
{
"version": "0.2.0",
"wit_version": "0.3.0",
"description": "Create, read, write, and format Google Sheets spreadsheets. Supports cell operations using A1 notation, sheet (tab) management, and cell formatting.",
"http": {
"allowlist": [
{
@@ -52,7 +53,7 @@
},
{
"name": "google_oauth_client_secret",
"prompt": "Google OAuth Client Secret"
"prompt": "Google OAuth Client Secret (from console.cloud.google.com/apis/credentials)"
}
]
}
@@ -1,6 +1,7 @@
{
"version": "0.2.0",
"wit_version": "0.3.0",
"description": "Create, read, edit, and format Google Slides presentations. Supports slide management, text operations, shapes, images, text formatting, and paragraph alignment.",
"http": {
"allowlist": [
{
@@ -52,7 +53,7 @@
},
{
"name": "google_oauth_client_secret",
"prompt": "Google OAuth Client Secret"
"prompt": "Google OAuth Client Secret (from console.cloud.google.com/apis/credentials)"
}
]
}
@@ -1,6 +1,7 @@
{
"version": "0.1.0",
"wit_version": "0.3.0",
"description": "Fetch pre-extracted web content from Brave Search for grounding LLM answers. Returns actual page content (text chunks, tables, code) relevant to the query, ready for RAG or fact-checking.",
"capabilities": {
"http": {
"allowlist": [
+2 -1
View File
@@ -1,6 +1,7 @@
{
"version": "0.2.0",
"wit_version": "0.3.0",
"description": "Send messages, list channels, read history, add reactions, and get user information in Slack.",
"http": {
"allowlist": [
{
@@ -57,7 +58,7 @@
},
{
"name": "slack_oauth_client_secret",
"prompt": "Slack OAuth Client Secret"
"prompt": "Slack OAuth Client Secret (from api.slack.com/apps > Basic Information)"
}
]
}
@@ -1,6 +1,7 @@
{
"version": "0.2.0",
"wit_version": "0.3.0",
"description": "Read and send messages from a Telegram user account. Supports contacts, chat history, message search, sending, forwarding, and deletion via encrypted MTProto.",
"http": {
"allowlist": [
{
@@ -35,7 +36,7 @@
},
{
"name": "telegram_api_hash",
"prompt": "Telegram API Hash"
"prompt": "Telegram API Hash (from my.telegram.org/apps — alphanumeric string)"
}
]
}
@@ -2,40 +2,6 @@
"version": "0.2.0",
"wit_version": "0.3.0",
"description": "Search the web using Brave Search. Returns titles, URLs, descriptions, and publication dates for matching web pages. Supports filtering by country, language, and freshness. Authentication is handled via the 'brave_api_key' secret injected by the host.",
"parameters": {
"type": "object",
"properties": {
"query": {
"type": "string",
"description": "The search query to look up on the web"
},
"count": {
"type": "integer",
"description": "Number of results to return (1-20, default 5)",
"minimum": 1,
"maximum": 20,
"default": 5
},
"country": {
"type": "string",
"description": "2-letter uppercase country code to bias results (e.g. 'US', 'DE', 'JP')"
},
"search_lang": {
"type": "string",
"description": "2-letter lowercase language code for search results (e.g. 'en', 'de', 'fr')"
},
"ui_lang": {
"type": "string",
"description": "Locale in language-region format (e.g. 'en-US', 'de-DE')"
},
"freshness": {
"type": "string",
"description": "Filter by discovery time: 'pd' (past day), 'pw' (past week), 'pm' (past month), 'py' (past year), or date range 'YYYY-MM-DDtoYYYY-MM-DD'"
}
},
"required": ["query"],
"additionalProperties": false
},
"capabilities": {
"http": {
"allowlist": [