Compare commits

...
Author SHA1 Message Date
[email protected] 78b3b327c1 chore: trigger CI 2026-03-23 16:08:49 -07:00
[email protected]andClaude Opus 4.6 32c5a86dd3 fix: address review findings — token hashing, broadcast scoping, error handling
Security fixes:
- Hash tokens with SHA-256 at construction time so authentication
  compares fixed-size 32-byte digests, eliminating length-oracle
  timing leaks
- Scope auth SSE broadcasts per-user in chat_auth_token_handler —
  AuthRequired/AuthCompleted events were leaking across tenants
- Propagate DB errors in restart handlers instead of silently
  swallowing via `if let Ok(Some(...))` pattern

Code quality:
- Log SSE serialization failures instead of silently producing empty
  strings via unwrap_or_default()
- Remove dead `pub type AuthState = MultiAuthState` alias
- Replace `.unwrap()` with `Arc::clone(db)` in app.rs multi-tenant
  workspace setup (db is guaranteed Some in context, but unwrap
  violates project convention)
- Fix telegram setup test to inject UserIdentity into request
  extensions (handler now requires AuthenticatedUser)
- Add safety comments on test-only expect/unwrap calls for CI
- Apply cargo fmt to fix pre-existing formatting

Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]>
2026-03-23 16:06:21 -07:00
Andrew PreeceandClaude Opus 4.6 6f4050dafa fix: second-pass multi-tenant audit — scope SSE broadcasts, DB queries, dead handlers
Second audit pass applying learned patterns across the codebase:

- OAuth callback SSE broadcasts now use broadcast_for_user (lines 773, 912)
- jobs_list_handler uses list_agent_jobs_for_user instead of fetching
  all users' jobs and filtering in Rust
- list_agent_jobs_for_user added to Database trait + postgres + libsql
- Dead handler files (extensions.rs, static_files.rs) hardened with
  AuthenticatedUser to prevent auth regression if migrated

Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]>
2026-03-23 18:01:30 +00:00
Andrew PreeceandClaude Opus 4.6 9ac03c3e62 fix: comprehensive multi-tenant isolation audit
Address all review findings from @serrrfirat plus 7 additional gaps
found via full security audit:

Reviewer findings (5):
- WorkspacePool now applies search config, memory layers, embedding
  cache, identity read scopes, and global config scopes (was bare)
- jobs_summary_handler uses per-user queries instead of global counters
- jobs_prompt_handler restructured to not 404 agent jobs + ownership check
- jobs_restart_handler agent branch now verifies user ownership
- agent_job_summary_for_user added to Database trait + both backends

Audit findings (7):
- Delete dead handlers/memory.rs (stale copies with no auth)
- Add AuthenticatedUser to logs_events, logs_level_get, logs_level_set
- Add AuthenticatedUser to extensions_tools_handler, gateway_status_handler
- Add auth + ownership checks to all 6 routines handlers
- Add auth to all 4 skills handlers with audit logging on mutations
- Scope extension setup SSE broadcast to user (broadcast_for_user)
- Fix pre-existing test compilation errors in extensions/manager.rs

17 new multi-tenant isolation tests covering:
- WorkspacePool config propagation and scope merging
- Jobs handler per-user isolation (summary, restart, prompt, cancel)
- Routines handler auth enforcement and cross-user rejection
- Auth middleware enforcement on logs, skills, status endpoints

Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]>
2026-03-23 18:01:30 +00:00
Andrew PreeceandClaude Opus 4.6 76431159c8 fix: scope memory tools per-user in multi-tenant mode
Memory tools (search, write, read, tree) held a single workspace
created at startup with GATEWAY_USER_ID. In multi-tenant mode, all
users' tool calls searched the default user's scope.

Add WorkspaceResolver trait that resolves workspaces per-request using
JobContext.user_id. In single-user mode, returns the startup workspace.
In multi-tenant mode (GATEWAY_USER_TOKENS configured), creates and
caches per-user workspaces on demand.

Includes regression tests for workspace resolution and user isolation.

Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]>
2026-03-23 18:01:30 +00:00
Andrew Preece b394c0cfdd feat: multi-tenant auth with per-user scoping
Multi-user authentication and authorization for IronClaw gateway:
- Token-based auth mapping tokens to user IDs via GATEWAY_USER_TOKENS
- Per-user SSE broadcast scoping
- Per-user rate limiting with poisoned lock recovery
- Handler auth and ownership checks for jobs, settings, routines
- Extension secrets scoped per-user
- Chat handlers use authenticated identity
- Reverse proxy deployment documentation
- Comprehensive integration tests for auth, SSE, rate limiting, and job isolation
2026-03-23 18:01:00 +00:00
43 changed files with 5000 additions and 1138 deletions
+22 -12
View File
@@ -44,7 +44,7 @@ pub struct JobMonitorRoute {
/// the main agent's context window). /// the main agent's context window).
pub fn spawn_job_monitor( pub fn spawn_job_monitor(
job_id: Uuid, job_id: Uuid,
event_rx: broadcast::Receiver<(Uuid, SseEvent)>, event_rx: broadcast::Receiver<(Uuid, String, SseEvent)>,
inject_tx: mpsc::Sender<IncomingMessage>, inject_tx: mpsc::Sender<IncomingMessage>,
route: JobMonitorRoute, route: JobMonitorRoute,
) -> JoinHandle<()> { ) -> JoinHandle<()> {
@@ -56,7 +56,7 @@ pub fn spawn_job_monitor(
/// jobs don't stay `InProgress` forever in the `ContextManager`. /// jobs don't stay `InProgress` forever in the `ContextManager`.
pub fn spawn_job_monitor_with_context( pub fn spawn_job_monitor_with_context(
job_id: Uuid, job_id: Uuid,
mut event_rx: broadcast::Receiver<(Uuid, SseEvent)>, mut event_rx: broadcast::Receiver<(Uuid, String, SseEvent)>,
inject_tx: mpsc::Sender<IncomingMessage>, inject_tx: mpsc::Sender<IncomingMessage>,
route: JobMonitorRoute, route: JobMonitorRoute,
context_manager: Option<Arc<ContextManager>>, context_manager: Option<Arc<ContextManager>>,
@@ -68,7 +68,7 @@ pub fn spawn_job_monitor_with_context(
loop { loop {
match event_rx.recv().await { match event_rx.recv().await {
Ok((ev_job_id, event)) => { Ok((ev_job_id, _user_id, event)) => {
if ev_job_id != job_id { if ev_job_id != job_id {
continue; 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. /// inject messages into) but we still need to free the `max_jobs` slot.
pub fn spawn_completion_watcher( pub fn spawn_completion_watcher(
job_id: Uuid, job_id: Uuid,
mut event_rx: broadcast::Receiver<(Uuid, SseEvent)>, mut event_rx: broadcast::Receiver<(Uuid, String, SseEvent)>,
context_manager: Arc<ContextManager>, context_manager: Arc<ContextManager>,
) -> JoinHandle<()> { ) -> JoinHandle<()> {
let short_id = job_id.to_string()[..8].to_string(); let short_id = job_id.to_string()[..8].to_string();
@@ -170,7 +170,9 @@ pub fn spawn_completion_watcher(
tokio::spawn(async move { tokio::spawn(async move {
loop { loop {
match event_rx.recv().await { 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" { let target = if status == "completed" {
JobState::Completed JobState::Completed
} else { } else {
@@ -227,7 +229,7 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn test_monitor_forwards_assistant_messages() { 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 (inject_tx, mut inject_rx) = mpsc::channel::<IncomingMessage>(16);
let job_id = Uuid::new_v4(); let job_id = Uuid::new_v4();
@@ -237,6 +239,7 @@ mod tests {
event_tx event_tx
.send(( .send((
job_id, job_id,
"test-user".to_string(),
SseEvent::JobMessage { SseEvent::JobMessage {
job_id: job_id.to_string(), job_id: job_id.to_string(),
role: "assistant".to_string(), role: "assistant".to_string(),
@@ -259,7 +262,7 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn test_monitor_ignores_other_jobs() { 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 (inject_tx, mut inject_rx) = mpsc::channel::<IncomingMessage>(16);
let job_id = Uuid::new_v4(); let job_id = Uuid::new_v4();
@@ -270,6 +273,7 @@ mod tests {
event_tx event_tx
.send(( .send((
other_job_id, other_job_id,
"test-user".to_string(),
SseEvent::JobMessage { SseEvent::JobMessage {
job_id: other_job_id.to_string(), job_id: other_job_id.to_string(),
role: "assistant".to_string(), role: "assistant".to_string(),
@@ -289,7 +293,7 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn test_monitor_exits_on_job_result() { 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 (inject_tx, mut inject_rx) = mpsc::channel::<IncomingMessage>(16);
let job_id = Uuid::new_v4(); let job_id = Uuid::new_v4();
@@ -299,6 +303,7 @@ mod tests {
event_tx event_tx
.send(( .send((
job_id, job_id,
"test-user".to_string(),
SseEvent::JobResult { SseEvent::JobResult {
job_id: job_id.to_string(), job_id: job_id.to_string(),
status: "completed".to_string(), status: "completed".to_string(),
@@ -324,7 +329,7 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn test_monitor_skips_tool_events() { 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 (inject_tx, mut inject_rx) = mpsc::channel::<IncomingMessage>(16);
let job_id = Uuid::new_v4(); let job_id = Uuid::new_v4();
@@ -334,6 +339,7 @@ mod tests {
event_tx event_tx
.send(( .send((
job_id, job_id,
"test-user".to_string(),
SseEvent::JobToolUse { SseEvent::JobToolUse {
job_id: job_id.to_string(), job_id: job_id.to_string(),
tool_name: "shell".to_string(), tool_name: "shell".to_string(),
@@ -346,6 +352,7 @@ mod tests {
event_tx event_tx
.send(( .send((
job_id, job_id,
"test-user".to_string(),
SseEvent::JobMessage { SseEvent::JobMessage {
job_id: job_id.to_string(), job_id: job_id.to_string(),
role: "user".to_string(), role: "user".to_string(),
@@ -402,7 +409,7 @@ mod tests {
.await .await
.unwrap(); .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 (inject_tx, mut inject_rx) = mpsc::channel::<IncomingMessage>(16);
let handle = spawn_job_monitor_with_context( let handle = spawn_job_monitor_with_context(
@@ -417,6 +424,7 @@ mod tests {
event_tx event_tx
.send(( .send((
job_id, job_id,
"test-user".to_string(),
SseEvent::JobResult { SseEvent::JobResult {
job_id: job_id.to_string(), job_id: job_id.to_string(),
status: "completed".to_string(), status: "completed".to_string(),
@@ -450,7 +458,7 @@ mod tests {
.await .await
.unwrap(); .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 (inject_tx, mut inject_rx) = mpsc::channel::<IncomingMessage>(16);
let handle = spawn_job_monitor_with_context( let handle = spawn_job_monitor_with_context(
@@ -465,6 +473,7 @@ mod tests {
event_tx event_tx
.send(( .send((
job_id, job_id,
"test-user".to_string(),
SseEvent::JobResult { SseEvent::JobResult {
job_id: job_id.to_string(), job_id: job_id.to_string(),
status: "failed".to_string(), status: "failed".to_string(),
@@ -498,12 +507,13 @@ mod tests {
.await .await
.unwrap(); .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)); let handle = spawn_completion_watcher(job_id, event_tx.subscribe(), Arc::clone(&cm));
event_tx event_tx
.send(( .send((
job_id, job_id,
"test-user".to_string(),
SseEvent::JobResult { SseEvent::JobResult {
job_id: job_id.to_string(), job_id: job_id.to_string(),
status: "completed".to_string(), status: "completed".to_string(),
+1 -1
View File
@@ -1646,7 +1646,7 @@ impl Agent {
}; };
match ext_mgr match ext_mgr
.configure_token(&pending.extension_name, token) .configure_token(&pending.extension_name, token, &message.user_id)
.await .await
{ {
Ok(result) if result.activated => { Ok(result) if result.activated => {
+31 -2
View File
@@ -327,7 +327,7 @@ impl AppBuilder {
.with_search_config(&self.config.search); .with_search_config(&self.config.search);
if let Some(ref emb) = embeddings { 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) // Wire workspace-level settings (read scopes, memory layers)
@@ -341,7 +341,36 @@ impl AppBuilder {
} }
ws = ws.with_memory_layers(self.config.workspace.memory_layers.clone()); ws = ws.with_memory_layers(self.config.workspace.memory_layers.clone());
let ws = Arc::new(ws); 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
// PerUserWorkspaceResolver 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 resolver = Arc::new(
crate::tools::builtin::memory::PerUserWorkspaceResolver::new(
Arc::clone(db),
embeddings.clone(),
emb_cache_config,
self.config.search.clone(),
self.config.workspace.clone(),
),
);
tools.register_memory_tools_with_resolver(resolver);
tracing::info!(
"Memory tools configured with per-user workspace resolver (multi-tenant mode)"
);
} else {
tools.register_memory_tools(Arc::clone(&ws));
}
Some(ws) Some(ws)
} else { } else {
None None
+383 -22
View File
@@ -1,17 +1,133 @@
//! Bearer token authentication middleware for the web gateway. //! 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::{ use axum::{
extract::{Request, State}, extract::{FromRequestParts, Request, State},
http::{HeaderMap, Method, StatusCode}, http::{HeaderMap, Method, StatusCode, request::Parts},
middleware::Next, middleware::Next,
response::{IntoResponse, Response}, response::{IntoResponse, Response},
}; };
use sha2::{Digest, Sha256};
use subtle::ConstantTimeEq; 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)] #[derive(Clone)]
pub struct AuthState { pub struct MultiAuthState {
pub token: String, /// 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. /// 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. /// Auth middleware that validates bearer token from header or query param.
/// ///
/// SSE connections can't set headers from `EventSource`, so we also accept /// 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( pub async fn auth_middleware(
State(auth): State<AuthState>, State(auth): State<MultiAuthState>,
headers: HeaderMap, headers: HeaderMap,
request: Request, mut request: Request,
next: Next, next: Next,
) -> Response { ) -> Response {
// Try Authorization header first (constant-time comparison). // Try Authorization header first.
// RFC 6750 Section 2.1: auth-scheme comparison is case-insensitive. // RFC 6750 Section 2.1: auth-scheme comparison is case-insensitive.
if let Some(auth_header) = headers.get("authorization") if let Some(auth_header) = headers.get("authorization")
&& let Ok(value) = auth_header.to_str() && let Ok(value) = auth_header.to_str()
&& value.len() > 7 && value.len() > 7
&& value[..7].eq_ignore_ascii_case("Bearer ") && 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; 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) if allows_query_token_auth(&request)
&& let Some(token) = query_token(&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; return next.run(request).await;
} }
@@ -83,15 +204,61 @@ pub async fn auth_middleware(
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use super::*; use super::*;
use crate::testing::credentials::{TEST_AUTH_SECRET_TOKEN, TEST_BEARER_TOKEN}; use crate::testing::credentials::TEST_AUTH_SECRET_TOKEN;
#[test] #[test]
fn test_auth_state_clone() { fn test_multi_auth_state_single() {
let state = AuthState { let state = MultiAuthState::single("tok-123".to_string(), "alice".to_string());
token: TEST_BEARER_TOKEN.to_string(), let identity = state.authenticate("tok-123");
}; assert!(identity.is_some());
let cloned = state.clone(); assert_eq!(identity.unwrap().user_id, "alice");
assert_eq!(cloned.token, TEST_BEARER_TOKEN); }
#[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; use axum::Router;
@@ -107,9 +274,7 @@ mod tests {
/// Router with streaming endpoints (query auth allowed) and regular /// Router with streaming endpoints (query auth allowed) and regular
/// endpoints (query auth rejected). /// endpoints (query auth rejected).
fn test_app(token: &str) -> Router { fn test_app(token: &str) -> Router {
let state = AuthState { let state = MultiAuthState::single(token.to_string(), "test-user".to_string());
token: token.to_string(),
};
Router::new() Router::new()
.route("/api/chat/events", get(dummy_handler)) .route("/api/chat/events", get(dummy_handler))
.route("/api/logs/events", get(dummy_handler)) .route("/api/logs/events", get(dummy_handler))
@@ -306,4 +471,200 @@ mod tests {
let resp = app.oneshot(req).await.unwrap(); let resp = app.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::UNAUTHORIZED); 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 uuid::Uuid;
use crate::channels::IncomingMessage; use crate::channels::IncomingMessage;
use crate::channels::web::auth::AuthenticatedUser;
use crate::channels::web::server::GatewayState; use crate::channels::web::server::GatewayState;
use crate::channels::web::types::*; use crate::channels::web::types::*;
use crate::channels::web::util::{build_turns_from_db_messages, truncate_preview}; use crate::channels::web::util::{build_turns_from_db_messages, truncate_preview};
pub async fn chat_send_handler( pub async fn chat_send_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(identity): AuthenticatedUser,
Json(req): Json<SendMessageRequest>, Json(req): Json<SendMessageRequest>,
) -> Result<(StatusCode, Json<SendMessageResponse>), (StatusCode, String)> { ) -> Result<(StatusCode, Json<SendMessageResponse>), (StatusCode, String)> {
if !state.chat_rate_limiter.check() { if !state.chat_rate_limiter.check(&identity.user_id) {
return Err(( return Err((
StatusCode::TOO_MANY_REQUESTS, StatusCode::TOO_MANY_REQUESTS,
"Rate limit exceeded. Try again shortly.".to_string(), "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 { if let Some(ref thread_id) = req.thread_id {
msg = msg.with_thread(thread_id); msg = msg.with_thread(thread_id);
@@ -74,6 +76,7 @@ pub async fn chat_send_handler(
pub async fn chat_approval_handler( pub async fn chat_approval_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(identity): AuthenticatedUser,
Json(req): Json<ApprovalRequest>, Json(req): Json<ApprovalRequest>,
) -> Result<(StatusCode, Json<SendMessageResponse>), (StatusCode, String)> { ) -> Result<(StatusCode, Json<SendMessageResponse>), (StatusCode, String)> {
let (approved, always) = match req.action.as_str() { 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 { if let Some(ref thread_id) = req.thread_id {
msg = msg.with_thread(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. /// The token never touches the LLM, chat history, or SSE stream.
pub async fn chat_auth_token_handler( pub async fn chat_auth_token_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Json(req): Json<AuthTokenRequest>, Json(req): Json<AuthTokenRequest>,
) -> Result<Json<ActionResponse>, (StatusCode, String)> { ) -> Result<Json<ActionResponse>, (StatusCode, String)> {
let ext_mgr = state.extension_manager.as_ref().ok_or(( let ext_mgr = state.extension_manager.as_ref().ok_or((
@@ -158,7 +162,7 @@ pub async fn chat_auth_token_handler(
))?; ))?;
match ext_mgr match ext_mgr
.configure_token(&req.extension_name, &req.token) .configure_token(&req.extension_name, &req.token, &user.user_id)
.await .await
{ {
Ok(result) => { Ok(result) => {
@@ -169,20 +173,26 @@ pub async fn chat_auth_token_handler(
resp.instructions = result.verification.as_ref().map(|v| v.instructions.clone()); resp.instructions = result.verification.as_ref().map(|v| v.instructions.clone());
if result.verification.is_some() { if result.verification.is_some() {
state.sse.broadcast(SseEvent::AuthRequired { state.sse.broadcast_for_user(
extension_name: req.extension_name.clone(), &user.user_id,
instructions: Some(result.message), SseEvent::AuthRequired {
auth_url: None, extension_name: req.extension_name.clone(),
setup_url: None, instructions: Some(result.message),
}); auth_url: None,
setup_url: None,
},
);
} else { } else {
clear_auth_mode(&state).await; clear_auth_mode(&state, &user.user_id).await;
state.sse.broadcast(SseEvent::AuthCompleted { state.sse.broadcast_for_user(
extension_name: req.extension_name.clone(), &user.user_id,
success: true, SseEvent::AuthCompleted {
message: result.message, extension_name: req.extension_name.clone(),
}); success: true,
message: result.message,
},
);
} }
Ok(Json(resp)) Ok(Json(resp))
@@ -190,12 +200,15 @@ pub async fn chat_auth_token_handler(
Err(e) => { Err(e) => {
let msg = e.to_string(); let msg = e.to_string();
if matches!(e, crate::extensions::ExtensionError::ValidationFailed(_)) { if matches!(e, crate::extensions::ExtensionError::ValidationFailed(_)) {
state.sse.broadcast(SseEvent::AuthRequired { state.sse.broadcast_for_user(
extension_name: req.extension_name.clone(), &user.user_id,
instructions: Some(msg.clone()), SseEvent::AuthRequired {
auth_url: None, extension_name: req.extension_name.clone(),
setup_url: None, instructions: Some(msg.clone()),
}); auth_url: None,
setup_url: None,
},
);
} }
Ok(Json(ActionResponse::fail(msg))) Ok(Json(ActionResponse::fail(msg)))
} }
@@ -205,16 +218,17 @@ pub async fn chat_auth_token_handler(
/// Cancel an in-progress auth flow. /// Cancel an in-progress auth flow.
pub async fn chat_auth_cancel_handler( pub async fn chat_auth_cancel_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(identity): AuthenticatedUser,
Json(_req): Json<AuthCancelRequest>, Json(_req): Json<AuthCancelRequest>,
) -> Result<Json<ActionResponse>, (StatusCode, String)> { ) -> Result<Json<ActionResponse>, (StatusCode, String)> {
clear_auth_mode(&state).await; clear_auth_mode(&state, &identity.user_id).await;
Ok(Json(ActionResponse::ok("Auth cancelled"))) Ok(Json(ActionResponse::ok("Auth cancelled")))
} }
/// Clear pending auth mode on the active thread. /// 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 { 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; let mut sess = session.lock().await;
if let Some(thread_id) = sess.active_thread if let Some(thread_id) = sess.active_thread
&& let Some(thread) = sess.threads.get_mut(&thread_id) && 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( pub async fn chat_events_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
) -> Result<impl IntoResponse, (StatusCode, String)> { ) -> Result<impl IntoResponse, (StatusCode, String)> {
state.sse.subscribe().ok_or(( state.sse.subscribe(Some(user.user_id)).ok_or((
StatusCode::SERVICE_UNAVAILABLE, StatusCode::SERVICE_UNAVAILABLE,
"Too many connections".to_string(), "Too many connections".to_string(),
)) ))
@@ -237,6 +252,7 @@ pub async fn chat_ws_handler(
headers: axum::http::HeaderMap, headers: axum::http::HeaderMap,
ws: WebSocketUpgrade, ws: WebSocketUpgrade,
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(identity): AuthenticatedUser,
) -> Result<impl IntoResponse, (StatusCode, String)> { ) -> Result<impl IntoResponse, (StatusCode, String)> {
// Validate Origin header to prevent cross-site WebSocket hijacking. // Validate Origin header to prevent cross-site WebSocket hijacking.
let origin = headers let origin = headers
@@ -262,7 +278,9 @@ pub async fn chat_ws_handler(
"WebSocket origin not allowed".to_string(), "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)] #[derive(Deserialize)]
@@ -274,6 +292,7 @@ pub struct HistoryQuery {
pub async fn chat_history_handler( pub async fn chat_history_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(identity): AuthenticatedUser,
Query(query): Query<HistoryQuery>, Query(query): Query<HistoryQuery>,
) -> Result<Json<HistoryResponse>, (StatusCode, String)> { ) -> Result<Json<HistoryResponse>, (StatusCode, String)> {
let session_manager = state.session_manager.as_ref().ok_or(( 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(), "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 limit = query.limit.unwrap_or(50);
let before_cursor = query let before_cursor = query
@@ -314,7 +335,7 @@ pub async fn chat_history_handler(
&& let Some(ref store) = state.store && let Some(ref store) = state.store
{ {
let owned = store let owned = store
.conversation_belongs_to_user(thread_id, &state.user_id) .conversation_belongs_to_user(thread_id, &identity.user_id)
.await .await
.unwrap_or(false); .unwrap_or(false);
if !owned { if !owned {
@@ -434,24 +455,27 @@ pub async fn chat_history_handler(
pub async fn chat_threads_handler( pub async fn chat_threads_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(identity): AuthenticatedUser,
) -> Result<Json<ThreadListResponse>, (StatusCode, String)> { ) -> Result<Json<ThreadListResponse>, (StatusCode, String)> {
let session_manager = state.session_manager.as_ref().ok_or(( let session_manager = state.session_manager.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE, StatusCode::SERVICE_UNAVAILABLE,
"Session manager not available".to_string(), "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 // Try DB first for persistent thread list
if let Some(ref store) = state.store { if let Some(ref store) = state.store {
// Auto-create assistant thread if it doesn't exist // Auto-create assistant thread if it doesn't exist
let assistant_id = store let assistant_id = store
.get_or_create_assistant_conversation(&state.user_id, "gateway") .get_or_create_assistant_conversation(&identity.user_id, "gateway")
.await .await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
if let Ok(summaries) = store if let Ok(summaries) = store
.list_conversations_all_channels(&state.user_id, 50) .list_conversations_all_channels(&identity.user_id, 50)
.await .await
{ {
let mut assistant_thread = None; let mut assistant_thread = None;
@@ -534,13 +558,16 @@ pub async fn chat_threads_handler(
pub async fn chat_new_thread_handler( pub async fn chat_new_thread_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(identity): AuthenticatedUser,
) -> Result<Json<ThreadInfo>, (StatusCode, String)> { ) -> Result<Json<ThreadInfo>, (StatusCode, String)> {
let session_manager = state.session_manager.as_ref().ok_or(( let session_manager = state.session_manager.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE, StatusCode::SERVICE_UNAVAILABLE,
"Session manager not available".to_string(), "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 (thread_id, info) = {
let mut sess = session.lock().await; let mut sess = session.lock().await;
let thread = sess.create_thread(); 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. // so that the subsequent loadThreads() call from the frontend sees it.
if let Some(ref store) = state.store { if let Some(ref store) = state.store {
match store match store
.ensure_conversation(thread_id, "gateway", &state.user_id, None) .ensure_conversation(thread_id, "gateway", &identity.user_id, None)
.await .await
{ {
Ok(true) => {} Ok(true) => {}
Ok(false) => tracing::warn!( Ok(false) => tracing::warn!(
user = %state.user_id, user = %identity.user_id,
thread_id = %thread_id, thread_id = %thread_id,
"Skipped persisting new thread due to ownership/channel conflict" "Skipped persisting new thread due to ownership/channel conflict"
), ),
+8 -3
View File
@@ -8,11 +8,13 @@ use axum::{
http::StatusCode, http::StatusCode,
}; };
use crate::channels::web::auth::AuthenticatedUser;
use crate::channels::web::server::GatewayState; use crate::channels::web::server::GatewayState;
use crate::channels::web::types::*; use crate::channels::web::types::*;
pub async fn extensions_list_handler( pub async fn extensions_list_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
) -> Result<Json<ExtensionListResponse>, (StatusCode, String)> { ) -> Result<Json<ExtensionListResponse>, (StatusCode, String)> {
let ext_mgr = state.extension_manager.as_ref().ok_or(( let ext_mgr = state.extension_manager.as_ref().ok_or((
StatusCode::NOT_IMPLEMENTED, StatusCode::NOT_IMPLEMENTED,
@@ -20,7 +22,7 @@ pub async fn extensions_list_handler(
))?; ))?;
let installed = ext_mgr let installed = ext_mgr
.list(None, false) .list(None, false, &user.user_id)
.await .await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; .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( pub async fn extensions_tools_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(_user): AuthenticatedUser,
) -> Result<Json<ToolListResponse>, (StatusCode, String)> { ) -> Result<Json<ToolListResponse>, (StatusCode, String)> {
let registry = state.tool_registry.as_ref().ok_or(( let registry = state.tool_registry.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE, StatusCode::SERVICE_UNAVAILABLE,
@@ -100,6 +103,7 @@ pub async fn extensions_tools_handler(
pub async fn extensions_install_handler( pub async fn extensions_install_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Json(req): Json<InstallExtensionRequest>, Json(req): Json<InstallExtensionRequest>,
) -> Result<Json<ActionResponse>, (StatusCode, String)> { ) -> Result<Json<ActionResponse>, (StatusCode, String)> {
let ext_mgr = state.extension_manager.as_ref().ok_or(( let ext_mgr = state.extension_manager.as_ref().ok_or((
@@ -116,7 +120,7 @@ pub async fn extensions_install_handler(
}); });
match ext_mgr 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 .await
{ {
Ok(result) => Ok(Json(ActionResponse::ok(result.message))), Ok(result) => Ok(Json(ActionResponse::ok(result.message))),
@@ -126,6 +130,7 @@ pub async fn extensions_install_handler(
pub async fn extensions_remove_handler( pub async fn extensions_remove_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Path(name): Path<String>, Path(name): Path<String>,
) -> Result<Json<ActionResponse>, (StatusCode, String)> { ) -> Result<Json<ActionResponse>, (StatusCode, String)> {
let ext_mgr = state.extension_manager.as_ref().ok_or(( 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(), "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))), Ok(message) => Ok(Json(ActionResponse::ok(message))),
Err(e) => Ok(Json(ActionResponse::fail(e.to_string()))), Err(e) => Ok(Json(ActionResponse::fail(e.to_string()))),
} }
+388 -277
View File
@@ -11,11 +11,13 @@ use axum::{
use serde::Deserialize; use serde::Deserialize;
use uuid::Uuid; use uuid::Uuid;
use crate::channels::web::auth::AuthenticatedUser;
use crate::channels::web::server::GatewayState; use crate::channels::web::server::GatewayState;
use crate::channels::web::types::*; use crate::channels::web::types::*;
pub async fn jobs_list_handler( pub async fn jobs_list_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
) -> Result<Json<JobListResponse>, (StatusCode, String)> { ) -> Result<Json<JobListResponse>, (StatusCode, String)> {
let store = state.store.as_ref().ok_or(( let store = state.store.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE, StatusCode::SERVICE_UNAVAILABLE,
@@ -25,8 +27,8 @@ pub async fn jobs_list_handler(
let mut jobs: Vec<JobInfo> = Vec::new(); let mut jobs: Vec<JobInfo> = Vec::new();
let mut seen_ids: HashSet<Uuid> = HashSet::new(); let mut seen_ids: HashSet<Uuid> = HashSet::new();
// Fetch sandbox jobs from database. // Fetch sandbox jobs scoped to this user.
match store.list_sandbox_jobs().await { match store.list_sandbox_jobs_for_user(&user.user_id).await {
Ok(sandbox_jobs) => { Ok(sandbox_jobs) => {
for j in &sandbox_jobs { for j in &sandbox_jobs {
let ui_state = match j.status.as_str() { 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. // Fetch agent (non-sandbox) jobs scoped to this user, deduplicating by ID.
match store.list_agent_jobs().await { match store.list_agent_jobs_for_user(&user.user_id).await {
Ok(agent_jobs) => { Ok(agent_jobs) => {
for j in &agent_jobs { for j in &agent_jobs {
if seen_ids.contains(&j.id) { if seen_ids.contains(&j.id) {
@@ -80,6 +82,7 @@ pub async fn jobs_list_handler(
pub async fn jobs_summary_handler( pub async fn jobs_summary_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
) -> Result<Json<JobSummaryResponse>, (StatusCode, String)> { ) -> Result<Json<JobSummaryResponse>, (StatusCode, String)> {
let store = state.store.as_ref().ok_or(( let store = state.store.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE, StatusCode::SERVICE_UNAVAILABLE,
@@ -93,8 +96,8 @@ pub async fn jobs_summary_handler(
let mut failed = 0; let mut failed = 0;
let mut stuck = 0; let mut stuck = 0;
// Sandbox job counts. // Sandbox job counts scoped to this user.
match store.sandbox_job_summary().await { match store.sandbox_job_summary_for_user(&user.user_id).await {
Ok(s) => { Ok(s) => {
total += s.total; total += s.total;
pending += s.creating; pending += s.creating;
@@ -107,8 +110,8 @@ pub async fn jobs_summary_handler(
} }
} }
// Agent job counts. // Agent job counts scoped to this user.
match store.agent_job_summary().await { match store.agent_job_summary_for_user(&user.user_id).await {
Ok(s) => { Ok(s) => {
total += s.total; total += s.total;
pending += s.pending; pending += s.pending;
@@ -134,6 +137,7 @@ pub async fn jobs_summary_handler(
pub async fn jobs_detail_handler( pub async fn jobs_detail_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Path(id): Path<String>, Path(id): Path<String>,
) -> Result<Json<JobDetailResponse>, (StatusCode, String)> { ) -> Result<Json<JobDetailResponse>, (StatusCode, String)> {
let store = state.store.as_ref().ok_or(( 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()))?; .map_err(|_| (StatusCode::BAD_REQUEST, "Invalid job ID".to_string()))?;
// Try sandbox job from DB first. // Try sandbox job from DB first.
if let Ok(Some(job)) = store.get_sandbox_job(job_id).await { match store.get_sandbox_job(job_id).await {
let browse_id = std::path::Path::new(&job.project_dir) Ok(Some(job)) => {
.file_name() if job.user_id != user.user_id {
.map(|n| n.to_string_lossy().to_string()) return Err((StatusCode::NOT_FOUND, "Job not found".to_string()));
.unwrap_or_else(|| job.id.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() { let ui_state = match job.status.as_str() {
"creating" => "pending", "creating" => "pending",
"running" => "in_progress", "running" => "in_progress",
s => s, s => s,
}; };
let elapsed_secs = job.started_at.map(|start| { let elapsed_secs = job.started_at.map(|start| {
let end = job.completed_at.unwrap_or_else(chrono::Utc::now); let end = job.completed_at.unwrap_or_else(chrono::Utc::now);
(end - start).num_seconds().max(0) as u64 (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,
}); });
}
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(); // Synthesize transitions from timestamps.
let is_claude_code = mode.as_deref() == Some("claude_code"); 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 { let mode = store.get_sandbox_job_mode(job.id).await.ok().flatten();
id: job.id, let is_claude_code = mode.as_deref() == Some("claude_code");
title: job.task.clone(),
description: String::new(), return Ok(Json(JobDetailResponse {
state: ui_state.to_string(), id: job.id,
user_id: job.user_id.clone(), title: job.task.clone(),
created_at: job.created_at.to_rfc3339(), description: String::new(),
started_at: job.started_at.map(|dt| dt.to_rfc3339()), state: ui_state.to_string(),
completed_at: job.completed_at.map(|dt| dt.to_rfc3339()), user_id: job.user_id.clone(),
elapsed_secs, created_at: job.created_at.to_rfc3339(),
project_dir: Some(job.project_dir.clone()), started_at: job.started_at.map(|dt| dt.to_rfc3339()),
browse_url: Some(format!("/projects/{}/", browse_id)), completed_at: job.completed_at.map(|dt| dt.to_rfc3339()),
job_mode: mode.filter(|m| m != "worker"), elapsed_secs,
transitions, project_dir: Some(job.project_dir.clone()),
can_restart: state.job_manager.is_some(), browse_url: Some(format!("/projects/{}/", browse_id)),
can_prompt: is_claude_code && state.prompt_queue.is_some(), job_mode: mode.filter(|m| m != "worker"),
job_kind: Some("sandbox".to_string()), 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. // Fall back to agent job from DB.
if let Ok(Some(ctx)) = store.get_job(job_id).await { match store.get_job(job_id).await {
let elapsed_secs = ctx.started_at.map(|start| { Ok(Some(ctx)) => {
let end = ctx.completed_at.unwrap_or_else(chrono::Utc::now); if ctx.user_id != user.user_id {
(end - start).num_seconds().max(0) as u64 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). // 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. // Stuck jobs have no active worker loop, so messages would be silently dropped.
let is_promptable = matches!( let is_promptable = matches!(
ctx.state, ctx.state,
crate::context::JobState::Pending | crate::context::JobState::InProgress crate::context::JobState::Pending | crate::context::JobState::InProgress
); );
return Ok(Json(JobDetailResponse { Ok(Json(JobDetailResponse {
id: ctx.job_id, id: ctx.job_id,
title: ctx.title.clone(), title: ctx.title.clone(),
description: ctx.description.clone(), description: ctx.description.clone(),
state: ctx.state.to_string(), state: ctx.state.to_string(),
user_id: ctx.user_id.clone(), user_id: ctx.user_id.clone(),
created_at: ctx.created_at.to_rfc3339(), created_at: ctx.created_at.to_rfc3339(),
started_at: ctx.started_at.map(|dt| dt.to_rfc3339()), started_at: ctx.started_at.map(|dt| dt.to_rfc3339()),
completed_at: ctx.completed_at.map(|dt| dt.to_rfc3339()), completed_at: ctx.completed_at.map(|dt| dt.to_rfc3339()),
elapsed_secs, elapsed_secs,
project_dir: None, project_dir: None,
browse_url: None, browse_url: None,
job_mode: None, job_mode: None,
transitions: Vec::new(), transitions: Vec::new(),
can_restart: state.scheduler.is_some(), can_restart: state.scheduler.is_some(),
can_prompt: is_promptable && state.scheduler.is_some(), can_prompt: is_promptable && state.scheduler.is_some(),
job_kind: Some("agent".to_string()), 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( pub async fn jobs_cancel_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Path(id): Path<String>, Path(id): Path<String>,
) -> Result<Json<serde_json::Value>, (StatusCode, String)> { ) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
let job_id = Uuid::parse_str(&id) let job_id = Uuid::parse_str(&id)
.map_err(|_| (StatusCode::BAD_REQUEST, "Invalid job ID".to_string()))?; .map_err(|_| (StatusCode::BAD_REQUEST, "Invalid job ID".to_string()))?;
// Try sandbox job cancellation. // Try sandbox job cancellation.
if let Some(ref store) = state.store if let Some(ref store) = state.store {
&& let Ok(Some(job)) = store.get_sandbox_job(job_id).await match store.get_sandbox_job(job_id).await {
{ Ok(Some(job)) => {
if job.status == "running" || job.status == "creating" { if job.user_id != user.user_id {
// Stop the container if we have a job manager. return Err((StatusCode::NOT_FOUND, "Job not found".to_string()));
if let Some(ref jm) = state.job_manager }
&& let Err(e) = jm.stop_job(job_id).await if job.status == "running" || job.status == "creating" {
{ if let Some(ref jm) = state.job_manager
tracing::warn!(job_id = %job_id, error = %e, "Failed to stop container during cancellation"); && 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 // Fall back to agent job cancellation: stop the worker via the scheduler
// (which updates the in-memory ContextManager AND aborts the task handle), // (which updates the in-memory ContextManager AND aborts the task handle),
// then persist the status to the DB as a fallback. // then persist the status to the DB as a fallback.
if let Some(ref store) = state.store if let Some(ref store) = state.store {
&& let Ok(Some(job)) = store.get_job(job_id).await match store.get_job(job_id).await {
{ Ok(Some(job)) => {
if job.state.is_active() { if job.user_id != user.user_id {
// Try to stop via scheduler (aborts the worker task + updates return Err((StatusCode::NOT_FOUND, "Job not found".to_string()));
// in-memory ContextManager). This is best-effort — the job may }
// not be in the scheduler map if it already finished. if job.state.is_active() {
if let Some(ref slot) = state.scheduler // Try to stop via scheduler (aborts the worker task + updates
&& let Some(ref scheduler) = *slot.read().await // in-memory ContextManager). This is best-effort — the job may
{ // not be in the scheduler map if it already finished.
let _ = scheduler.stop(job_id).await; 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 // Always persist cancellation to the DB so the state is
// consistent even if the scheduler wasn't available or the // consistent even if the scheduler wasn't available or the
// job wasn't in its in-memory map. // job wasn't in its in-memory map.
store store
.update_job_status( .update_job_status(
job_id, job_id,
crate::context::JobState::Cancelled, crate::context::JobState::Cancelled,
Some("Cancelled by user"), Some("Cancelled by user"),
) )
.await .await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; .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())) 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( pub async fn jobs_restart_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Path(id): Path<String>, Path(id): Path<String>,
) -> Result<Json<serde_json::Value>, (StatusCode, String)> { ) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
let store = state.store.as_ref().ok_or(( 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()))?; .map_err(|_| (StatusCode::BAD_REQUEST, "Invalid job ID".to_string()))?;
// Try sandbox job restart first. // Try sandbox job restart first.
if let Ok(Some(old_job)) = store.get_sandbox_job(old_job_id).await { match store.get_sandbox_job(old_job_id).await {
if old_job.status != "interrupted" && old_job.status != "failed" { 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(( return Err((
StatusCode::CONFLICT, StatusCode::INTERNAL_SERVER_ERROR,
format!("Cannot restart job in state '{}'", old_job.status), 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. // Try agent job restart: dispatch a new job via the scheduler.
if let Ok(Some(old_job)) = store.get_job(old_job_id).await { match store.get_job(old_job_id).await {
if old_job.state.is_active() { Ok(Some(old_job)) => {
return Err(( if old_job.user_id != user.user_id {
StatusCode::CONFLICT, return Err((StatusCode::NOT_FOUND, "Job not found".to_string()));
format!("Cannot restart job in state '{}'", old_job.state), }
)); 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,
})))
} }
Ok(None) => Err((StatusCode::NOT_FOUND, "Job not found".to_string())),
let slot = state.scheduler.as_ref().ok_or(( Err(e) => Err((
StatusCode::SERVICE_UNAVAILABLE, StatusCode::INTERNAL_SERVER_ERROR,
"Scheduler not available".to_string(), format!("Database error: {}", e),
))?; )),
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,
})));
} }
Err((StatusCode::NOT_FOUND, "Job not found".to_string()))
} }
/// Submit a follow-up prompt to a running job. /// 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) /// - Worker-mode sandbox jobs → not supported (no mechanism to inject)
pub async fn jobs_prompt_handler( pub async fn jobs_prompt_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Path(id): Path<String>, Path(id): Path<String>,
Json(body): Json<serde_json::Value>, Json(body): Json<serde_json::Value>,
) -> Result<Json<serde_json::Value>, (StatusCode, String)> { ) -> 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); 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 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. // It's a sandbox job. Check if Claude Code mode.
let mode = s.get_sandbox_job_mode(job_id).await.ok().flatten(); let mode = s.get_sandbox_job_mode(job_id).await.ok().flatten();
if mode.as_deref() == Some("claude_code") { if mode.as_deref() == Some("claude_code") {
@@ -522,7 +597,14 @@ 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
&& let Ok(Some(agent_job)) = store.get_job(job_id).await
&& agent_job.user_id != user.user_id
{
return Err((StatusCode::NOT_FOUND, "Job not found".to_string()));
}
let slot = state.scheduler.as_ref().ok_or(( let slot = state.scheduler.as_ref().ok_or((
StatusCode::NOT_IMPLEMENTED, StatusCode::NOT_IMPLEMENTED,
"Agent job prompts require the scheduler to be configured".to_string(), "Agent job prompts require the scheduler to be configured".to_string(),
@@ -550,6 +632,7 @@ pub async fn jobs_prompt_handler(
/// Load persisted job events for a job (for history replay on page open). /// Load persisted job events for a job (for history replay on page open).
pub async fn jobs_events_handler( pub async fn jobs_events_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Path(id): Path<String>, Path(id): Path<String>,
) -> Result<Json<serde_json::Value>, (StatusCode, String)> { ) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
let store = state.store.as_ref().ok_or(( let store = state.store.as_ref().ok_or((
@@ -561,6 +644,24 @@ pub async fn jobs_events_handler(
.parse() .parse()
.map_err(|_| (StatusCode::BAD_REQUEST, "Invalid job ID".to_string()))?; .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 let events = store
.list_job_events(job_id, None) .list_job_events(job_id, None)
.await .await
@@ -593,6 +694,7 @@ pub struct FilePathQuery {
pub async fn job_files_list_handler( pub async fn job_files_list_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Path(id): Path<String>, Path(id): Path<String>,
Query(query): Query<FilePathQuery>, Query(query): Query<FilePathQuery>,
) -> Result<Json<ProjectFilesResponse>, (StatusCode, String)> { ) -> Result<Json<ProjectFilesResponse>, (StatusCode, String)> {
@@ -610,6 +712,10 @@ pub async fn job_files_list_handler(
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))? .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?
.ok_or((StatusCode::NOT_FOUND, "Job not found".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 base = std::path::PathBuf::from(&job.project_dir);
let rel_path = query.path.as_deref().unwrap_or(""); let rel_path = query.path.as_deref().unwrap_or("");
let target = base.join(rel_path); let target = base.join(rel_path);
@@ -656,6 +762,7 @@ pub async fn job_files_list_handler(
pub async fn job_files_read_handler( pub async fn job_files_read_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Path(id): Path<String>, Path(id): Path<String>,
Query(query): Query<FilePathQuery>, Query(query): Query<FilePathQuery>,
) -> Result<Json<ProjectFileReadResponse>, (StatusCode, String)> { ) -> Result<Json<ProjectFileReadResponse>, (StatusCode, String)> {
@@ -673,6 +780,10 @@ pub async fn job_files_read_handler(
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))? .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?
.ok_or((StatusCode::NOT_FOUND, "Job not found".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(( let path = query.path.as_deref().ok_or((
StatusCode::BAD_REQUEST, StatusCode::BAD_REQUEST,
"path parameter required".to_string(), "path parameter required".to_string(),
-154
View File
@@ -1,154 +0,0 @@
//! Memory/workspace API handlers.
use std::sync::Arc;
use axum::{
Json,
extract::{Query, State},
http::StatusCode,
};
use serde::Deserialize;
use crate::channels::web::server::GatewayState;
use crate::channels::web::types::*;
#[derive(Deserialize)]
pub struct TreeQuery {
#[allow(dead_code)]
pub depth: Option<usize>,
}
pub async fn memory_tree_handler(
State(state): State<Arc<GatewayState>>,
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(),
))?;
// Build tree from list_all (flat list of all paths)
let all_paths = workspace
.list_all()
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
// Collect unique directories and files
let mut entries: Vec<TreeEntry> = Vec::new();
let mut seen_dirs: std::collections::HashSet<String> = std::collections::HashSet::new();
for path in &all_paths {
// Add parent directories
let parts: Vec<&str> = path.split('/').collect();
for i in 0..parts.len().saturating_sub(1) {
let dir_path = parts[..=i].join("/");
if seen_dirs.insert(dir_path.clone()) {
entries.push(TreeEntry {
path: dir_path,
is_dir: true,
});
}
}
// Add the file itself
entries.push(TreeEntry {
path: path.clone(),
is_dir: false,
});
}
entries.sort_by(|a, b| a.path.cmp(&b.path));
Ok(Json(MemoryTreeResponse { entries }))
}
#[derive(Deserialize)]
pub struct ListQuery {
pub path: Option<String>,
}
pub async fn memory_list_handler(
State(state): State<Arc<GatewayState>>,
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 path = query.path.as_deref().unwrap_or("");
let entries = workspace
.list(path)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
let list_entries: Vec<ListEntry> = entries
.iter()
.map(|e| ListEntry {
name: e.path.rsplit('/').next().unwrap_or(&e.path).to_string(),
path: e.path.clone(),
is_dir: e.is_directory,
updated_at: e.updated_at.map(|dt| dt.to_rfc3339()),
})
.collect();
Ok(Json(MemoryListResponse {
path: path.to_string(),
entries: list_entries,
}))
}
#[derive(Deserialize)]
pub struct ReadQuery {
pub path: String,
}
pub async fn memory_read_handler(
State(state): State<Arc<GatewayState>>,
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 doc = workspace
.read(&query.path)
.await
.map_err(|e| (StatusCode::NOT_FOUND, e.to_string()))?;
Ok(Json(MemoryReadResponse {
path: query.path,
content: doc.content,
updated_at: Some(doc.updated_at.to_rfc3339()),
}))
}
// memory_write_handler lives in server.rs (layer-aware version with append,
// privacy redirect, and proper error status codes).
pub async fn memory_search_handler(
State(state): State<Arc<GatewayState>>,
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 limit = req.limit.unwrap_or(10);
let results = workspace
.search(&req.query, limit)
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
let hits: Vec<SearchHit> = results
.into_iter()
.map(|r| SearchHit {
path: r.document_path,
content: r.content,
score: r.score as f64,
})
.collect();
Ok(Json(MemorySearchResponse { results: hits }))
}
+2 -12
View File
@@ -1,13 +1,9 @@
//! Handler modules for the web gateway API. //! Handler modules for the web gateway API.
//! //!
//! Each module groups related endpoint handlers by domain. //! 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 routines;
pub mod skills; pub mod skills;
// Modules not yet wired into server.rs router -- suppress dead_code until // Modules not yet wired into server.rs router -- suppress dead_code until
@@ -17,12 +13,6 @@ pub mod chat;
#[allow(dead_code)] #[allow(dead_code)]
pub mod extensions; pub mod extensions;
#[allow(dead_code)] #[allow(dead_code)]
pub mod jobs;
#[allow(dead_code)]
pub mod memory;
#[allow(dead_code)]
pub mod routines;
#[allow(dead_code)]
pub mod settings; pub mod settings;
#[allow(dead_code)] #[allow(dead_code)]
pub mod static_files; pub mod static_files;
+42 -3
View File
@@ -11,12 +11,14 @@ use serde::Deserialize;
use uuid::Uuid; use uuid::Uuid;
use crate::agent::routine::{Trigger, next_cron_fire}; use crate::agent::routine::{Trigger, next_cron_fire};
use crate::channels::web::auth::AuthenticatedUser;
use crate::channels::web::server::GatewayState; use crate::channels::web::server::GatewayState;
use crate::channels::web::types::*; use crate::channels::web::types::*;
use crate::error::RoutineError; use crate::error::RoutineError;
pub async fn routines_list_handler( pub async fn routines_list_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
) -> Result<Json<RoutineListResponse>, (StatusCode, String)> { ) -> Result<Json<RoutineListResponse>, (StatusCode, String)> {
let store = state.store.as_ref().ok_or(( let store = state.store.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE, StatusCode::SERVICE_UNAVAILABLE,
@@ -24,7 +26,7 @@ pub async fn routines_list_handler(
))?; ))?;
let routines = store let routines = store
.list_all_routines() .list_routines(&user.user_id)
.await .await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; .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( pub async fn routines_summary_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
) -> Result<Json<RoutineSummaryResponse>, (StatusCode, String)> { ) -> Result<Json<RoutineSummaryResponse>, (StatusCode, String)> {
let store = state.store.as_ref().ok_or(( let store = state.store.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE, StatusCode::SERVICE_UNAVAILABLE,
@@ -42,7 +45,7 @@ pub async fn routines_summary_handler(
))?; ))?;
let routines = store let routines = store
.list_all_routines() .list_routines(&user.user_id)
.await .await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; .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( pub async fn routines_detail_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Path(id): Path<String>, Path(id): Path<String>,
) -> Result<Json<RoutineDetailResponse>, (StatusCode, String)> { ) -> Result<Json<RoutineDetailResponse>, (StatusCode, String)> {
let store = state.store.as_ref().ok_or(( 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()))? .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?
.ok_or((StatusCode::NOT_FOUND, "Routine not found".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 let runs = store
.list_routine_runs(routine_id, 20) .list_routine_runs(routine_id, 20)
.await .await
@@ -137,6 +145,7 @@ pub async fn routines_detail_handler(
pub async fn routines_trigger_handler( pub async fn routines_trigger_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Path(id): Path<String>, Path(id): Path<String>,
) -> Result<Json<serde_json::Value>, (StatusCode, String)> { ) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
// Clone the Arc out of the lock to avoid holding the RwLock across .await. // 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()))?; .map_err(|_| (StatusCode::BAD_REQUEST, "Invalid routine ID".to_string()))?;
let run_id = engine let run_id = engine
.fire_manual(routine_id, Some(&state.user_id)) .fire_manual(routine_id, Some(&user.user_id))
.await .await
.map_err(|e| (routine_error_status(&e), e.to_string()))?; .map_err(|e| (routine_error_status(&e), e.to_string()))?;
@@ -170,6 +179,7 @@ pub struct ToggleRequest {
pub async fn routines_toggle_handler( pub async fn routines_toggle_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Path(id): Path<String>, Path(id): Path<String>,
body: Option<Json<ToggleRequest>>, body: Option<Json<ToggleRequest>>,
) -> Result<Json<serde_json::Value>, (StatusCode, String)> { ) -> 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()))? .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?
.ok_or((StatusCode::NOT_FOUND, "Routine not found".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; let was_enabled = routine.enabled;
// If a specific value was provided, use it; otherwise toggle. // If a specific value was provided, use it; otherwise toggle.
routine.enabled = match body { routine.enabled = match body {
@@ -230,6 +244,7 @@ pub async fn routines_toggle_handler(
pub async fn routines_delete_handler( pub async fn routines_delete_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Path(id): Path<String>, Path(id): Path<String>,
) -> Result<Json<serde_json::Value>, (StatusCode, String)> { ) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
let store = state.store.as_ref().ok_or(( 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) let routine_id = Uuid::parse_str(&id)
.map_err(|_| (StatusCode::BAD_REQUEST, "Invalid routine ID".to_string()))?; .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 let deleted = store
.delete_routine(routine_id) .delete_routine(routine_id)
.await .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( pub async fn routines_runs_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Path(id): Path<String>, Path(id): Path<String>,
) -> Result<Json<serde_json::Value>, (StatusCode, String)> { ) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
let store = state.store.as_ref().ok_or(( 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) let routine_id = Uuid::parse_str(&id)
.map_err(|_| (StatusCode::BAD_REQUEST, "Invalid routine ID".to_string()))?; .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 let runs = store
.list_routine_runs(routine_id, 50) .list_routine_runs(routine_id, 50)
.await .await
+13 -6
View File
@@ -8,17 +8,19 @@ use axum::{
http::StatusCode, http::StatusCode,
}; };
use crate::channels::web::auth::AuthenticatedUser;
use crate::channels::web::server::GatewayState; use crate::channels::web::server::GatewayState;
use crate::channels::web::types::*; use crate::channels::web::types::*;
pub async fn settings_list_handler( pub async fn settings_list_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
) -> Result<Json<SettingsListResponse>, StatusCode> { ) -> Result<Json<SettingsListResponse>, StatusCode> {
let store = state let store = state
.store .store
.as_ref() .as_ref()
.ok_or(StatusCode::SERVICE_UNAVAILABLE)?; .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); tracing::error!("Failed to list settings: {}", e);
StatusCode::INTERNAL_SERVER_ERROR StatusCode::INTERNAL_SERVER_ERROR
})?; })?;
@@ -37,6 +39,7 @@ pub async fn settings_list_handler(
pub async fn settings_get_handler( pub async fn settings_get_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Path(key): Path<String>, Path(key): Path<String>,
) -> Result<Json<SettingResponse>, StatusCode> { ) -> Result<Json<SettingResponse>, StatusCode> {
let store = state let store = state
@@ -44,7 +47,7 @@ pub async fn settings_get_handler(
.as_ref() .as_ref()
.ok_or(StatusCode::SERVICE_UNAVAILABLE)?; .ok_or(StatusCode::SERVICE_UNAVAILABLE)?;
let row = store let row = store
.get_setting_full(&state.user_id, &key) .get_setting_full(&user.user_id, &key)
.await .await
.map_err(|e| { .map_err(|e| {
tracing::error!("Failed to get setting '{}': {}", key, 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( pub async fn settings_set_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Path(key): Path<String>, Path(key): Path<String>,
Json(body): Json<SettingWriteRequest>, Json(body): Json<SettingWriteRequest>,
) -> Result<StatusCode, StatusCode> { ) -> Result<StatusCode, StatusCode> {
@@ -69,7 +73,7 @@ pub async fn settings_set_handler(
.as_ref() .as_ref()
.ok_or(StatusCode::SERVICE_UNAVAILABLE)?; .ok_or(StatusCode::SERVICE_UNAVAILABLE)?;
store store
.set_setting(&state.user_id, &key, &body.value) .set_setting(&user.user_id, &key, &body.value)
.await .await
.map_err(|e| { .map_err(|e| {
tracing::error!("Failed to set setting '{}': {}", key, 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( pub async fn settings_delete_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Path(key): Path<String>, Path(key): Path<String>,
) -> Result<StatusCode, StatusCode> { ) -> Result<StatusCode, StatusCode> {
let store = state let store = state
@@ -88,7 +93,7 @@ pub async fn settings_delete_handler(
.as_ref() .as_ref()
.ok_or(StatusCode::SERVICE_UNAVAILABLE)?; .ok_or(StatusCode::SERVICE_UNAVAILABLE)?;
store store
.delete_setting(&state.user_id, &key) .delete_setting(&user.user_id, &key)
.await .await
.map_err(|e| { .map_err(|e| {
tracing::error!("Failed to delete setting '{}': {}", key, 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( pub async fn settings_export_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
) -> Result<Json<SettingsExportResponse>, StatusCode> { ) -> Result<Json<SettingsExportResponse>, StatusCode> {
let store = state let store = state
.store .store
.as_ref() .as_ref()
.ok_or(StatusCode::SERVICE_UNAVAILABLE)?; .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); tracing::error!("Failed to export settings: {}", e);
StatusCode::INTERNAL_SERVER_ERROR StatusCode::INTERNAL_SERVER_ERROR
})?; })?;
@@ -115,6 +121,7 @@ pub async fn settings_export_handler(
pub async fn settings_import_handler( pub async fn settings_import_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Json(body): Json<SettingsImportRequest>, Json(body): Json<SettingsImportRequest>,
) -> Result<StatusCode, StatusCode> { ) -> Result<StatusCode, StatusCode> {
let store = state let store = state
@@ -122,7 +129,7 @@ pub async fn settings_import_handler(
.as_ref() .as_ref()
.ok_or(StatusCode::SERVICE_UNAVAILABLE)?; .ok_or(StatusCode::SERVICE_UNAVAILABLE)?;
store store
.set_all_settings(&state.user_id, &body.settings) .set_all_settings(&user.user_id, &body.settings)
.await .await
.map_err(|e| { .map_err(|e| {
tracing::error!("Failed to import settings: {}", e); tracing::error!("Failed to import settings: {}", e);
+9
View File
@@ -8,11 +8,13 @@ use axum::{
http::StatusCode, http::StatusCode,
}; };
use crate::channels::web::auth::AuthenticatedUser;
use crate::channels::web::server::GatewayState; use crate::channels::web::server::GatewayState;
use crate::channels::web::types::*; use crate::channels::web::types::*;
pub async fn skills_list_handler( pub async fn skills_list_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(_user): AuthenticatedUser,
) -> Result<Json<SkillListResponse>, (StatusCode, String)> { ) -> Result<Json<SkillListResponse>, (StatusCode, String)> {
let registry = state.skill_registry.as_ref().ok_or(( let registry = state.skill_registry.as_ref().ok_or((
StatusCode::NOT_IMPLEMENTED, StatusCode::NOT_IMPLEMENTED,
@@ -45,6 +47,7 @@ pub async fn skills_list_handler(
pub async fn skills_search_handler( pub async fn skills_search_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(_user): AuthenticatedUser,
Json(req): Json<SkillSearchRequest>, Json(req): Json<SkillSearchRequest>,
) -> Result<Json<SkillSearchResponse>, (StatusCode, String)> { ) -> Result<Json<SkillSearchResponse>, (StatusCode, String)> {
let registry = state.skill_registry.as_ref().ok_or(( 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( pub async fn skills_install_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
headers: axum::http::HeaderMap, headers: axum::http::HeaderMap,
Json(req): Json<SkillInstallRequest>, Json(req): Json<SkillInstallRequest>,
) -> Result<Json<ActionResponse>, (StatusCode, String)> { ) -> 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(( let registry = state.skill_registry.as_ref().ok_or((
StatusCode::NOT_IMPLEMENTED, StatusCode::NOT_IMPLEMENTED,
"Skills system not enabled".to_string(), "Skills system not enabled".to_string(),
@@ -219,6 +225,7 @@ pub async fn skills_install_handler(
pub async fn skills_remove_handler( pub async fn skills_remove_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
headers: axum::http::HeaderMap, headers: axum::http::HeaderMap,
Path(name): Path<String>, Path(name): Path<String>,
) -> Result<Json<ActionResponse>, (StatusCode, 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(( let registry = state.skill_registry.as_ref().ok_or((
StatusCode::NOT_IMPLEMENTED, StatusCode::NOT_IMPLEMENTED,
"Skills system not enabled".to_string(), "Skills system not enabled".to_string(),
@@ -7,6 +7,7 @@ use axum::{
}; };
use crate::bootstrap::ironclaw_base_dir; use crate::bootstrap::ironclaw_base_dir;
use crate::channels::web::auth::AuthenticatedUser;
use crate::channels::web::types::*; use crate::channels::web::types::*;
// --- Static file handlers --- // --- Static file handlers ---
@@ -113,6 +114,7 @@ use crate::channels::web::server::GatewayState;
pub async fn logs_events_handler( pub async fn logs_events_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(_user): AuthenticatedUser,
) -> Result< ) -> Result<
Sse<impl futures::Stream<Item = Result<Event, Infallible>> + Send + 'static>, Sse<impl futures::Stream<Item = Result<Event, Infallible>> + Send + 'static>,
(StatusCode, String), (StatusCode, String),
@@ -152,6 +154,7 @@ pub async fn logs_events_handler(
pub async fn gateway_status_handler( pub async fn gateway_status_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
AuthenticatedUser(_user): AuthenticatedUser,
) -> Json<GatewayStatusResponse> { ) -> Json<GatewayStatusResponse> {
let sse_connections = state.sse.connection_count(); let sse_connections = state.sse.connection_count();
let ws_connections = state let ws_connections = state
+90 -22
View File
@@ -31,6 +31,9 @@ pub mod ws;
/// [`TestGatewayBuilder`](test_helpers::TestGatewayBuilder). /// [`TestGatewayBuilder`](test_helpers::TestGatewayBuilder).
pub mod test_helpers; pub mod test_helpers;
#[cfg(test)]
mod tests;
use std::net::SocketAddr; use std::net::SocketAddr;
use std::sync::Arc; use std::sync::Arc;
@@ -52,6 +55,7 @@ use crate::workspace::Workspace;
use self::log_layer::{LogBroadcaster, LogLevelHandle}; use self::log_layer::{LogBroadcaster, LogLevelHandle};
use self::auth::MultiAuthState;
use self::server::GatewayState; use self::server::GatewayState;
use self::sse::SseManager; use self::sse::SseManager;
use self::types::SseEvent; use self::types::SseEvent;
@@ -60,14 +64,15 @@ use self::types::SseEvent;
pub struct GatewayChannel { pub struct GatewayChannel {
config: GatewayConfig, config: GatewayConfig,
state: Arc<GatewayState>, state: Arc<GatewayState>,
/// The actual auth token in use (generated or from config). /// Multi-user auth state (replaces bare auth_token).
auth_token: String, auth: MultiAuthState,
} }
impl GatewayChannel { impl GatewayChannel {
/// Create a new gateway channel. /// Create a new gateway channel.
/// ///
/// If no auth token is configured, generates a random one and prints it. /// 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 { pub fn new(config: GatewayConfig) -> Self {
let auth_token = config.auth_token.clone().unwrap_or_else(|| { let auth_token = config.auth_token.clone().unwrap_or_else(|| {
use rand::RngCore; use rand::RngCore;
@@ -77,10 +82,13 @@ impl GatewayChannel {
bytes.iter().map(|b| format!("{b:02x}")).collect() bytes.iter().map(|b| format!("{b:02x}")).collect()
}); });
let auth = MultiAuthState::single(auth_token, config.user_id.clone());
let state = Arc::new(GatewayState { let state = Arc::new(GatewayState {
msg_tx: tokio::sync::RwLock::new(None), msg_tx: tokio::sync::RwLock::new(None),
sse: SseManager::new(), sse: Arc::new(SseManager::new()),
workspace: None, workspace: None,
workspace_pool: None,
session_manager: None, session_manager: None,
log_broadcaster: None, log_broadcaster: None,
log_level_handle: None, log_level_handle: None,
@@ -90,13 +98,13 @@ impl GatewayChannel {
job_manager: None, job_manager: None,
prompt_queue: None, prompt_queue: None,
scheduler: None, scheduler: None,
user_id: config.user_id.clone(), default_user_id: config.user_id.clone(),
shutdown_tx: tokio::sync::RwLock::new(None), shutdown_tx: tokio::sync::RwLock::new(None),
ws_tracker: Some(Arc::new(ws::WsConnectionTracker::new())), ws_tracker: Some(Arc::new(ws::WsConnectionTracker::new())),
llm_provider: None, llm_provider: None,
skill_registry: None, skill_registry: None,
skill_catalog: 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), oauth_rate_limiter: server::RateLimiter::new(10, 60),
webhook_rate_limiter: server::RateLimiter::new(10, 60), webhook_rate_limiter: server::RateLimiter::new(10, 60),
registry_entries: Vec::new(), registry_entries: Vec::new(),
@@ -109,7 +117,46 @@ impl GatewayChannel {
Self { Self {
config, config,
state, 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 { let mut new_state = GatewayState {
msg_tx: tokio::sync::RwLock::new(None), msg_tx: tokio::sync::RwLock::new(None),
// Preserve the existing broadcast channel so sender handles remain valid. // 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: self.state.workspace.clone(),
workspace_pool: self.state.workspace_pool.clone(),
session_manager: self.state.session_manager.clone(), session_manager: self.state.session_manager.clone(),
log_broadcaster: self.state.log_broadcaster.clone(), log_broadcaster: self.state.log_broadcaster.clone(),
log_level_handle: self.state.log_level_handle.clone(), log_level_handle: self.state.log_level_handle.clone(),
@@ -129,13 +177,13 @@ impl GatewayChannel {
job_manager: self.state.job_manager.clone(), job_manager: self.state.job_manager.clone(),
prompt_queue: self.state.prompt_queue.clone(), prompt_queue: self.state.prompt_queue.clone(),
scheduler: self.state.scheduler.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), shutdown_tx: tokio::sync::RwLock::new(None),
ws_tracker: self.state.ws_tracker.clone(), ws_tracker: self.state.ws_tracker.clone(),
llm_provider: self.state.llm_provider.clone(), llm_provider: self.state.llm_provider.clone(),
skill_registry: self.state.skill_registry.clone(), skill_registry: self.state.skill_registry.clone(),
skill_catalog: self.state.skill_catalog.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), oauth_rate_limiter: server::RateLimiter::new(10, 60),
webhook_rate_limiter: server::RateLimiter::new(10, 60), webhook_rate_limiter: server::RateLimiter::new(10, 60),
registry_entries: self.state.registry_entries.clone(), registry_entries: self.state.registry_entries.clone(),
@@ -260,9 +308,15 @@ impl GatewayChannel {
self 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 { 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). /// 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))) Ok(Box::pin(ReceiverStream::new(rx)))
} }
@@ -311,10 +365,13 @@ impl Channel for GatewayChannel {
} }
}; };
self.state.sse.broadcast(SseEvent::Response { self.state.sse.broadcast_for_user(
content: response.content, &msg.user_id,
thread_id, SseEvent::Response {
}); content: response.content,
thread_id,
},
);
Ok(()) 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(()) Ok(())
} }
async fn broadcast( async fn broadcast(
&self, &self,
_user_id: &str, user_id: &str,
response: OutgoingResponse, response: OutgoingResponse,
) -> Result<(), ChannelError> { ) -> Result<(), ChannelError> {
let thread_id = match response.thread_id { let thread_id = match response.thread_id {
@@ -445,10 +510,13 @@ impl Channel for GatewayChannel {
return Ok(()); return Ok(());
} }
}; };
self.state.sse.broadcast(SseEvent::Response { self.state.sse.broadcast_for_user(
content: response.content, user_id,
thread_id, SseEvent::Response {
}); content: response.content,
thread_id,
},
);
Ok(()) Ok(())
} }
+2 -1
View File
@@ -463,9 +463,10 @@ fn build_tool_request(
pub async fn chat_completions_handler( pub async fn chat_completions_handler(
State(state): State<Arc<GatewayState>>, State(state): State<Arc<GatewayState>>,
super::auth::AuthenticatedUser(user): super::auth::AuthenticatedUser,
Json(req): Json<OpenAiChatRequest>, Json(req): Json<OpenAiChatRequest>,
) -> Result<impl IntoResponse, (StatusCode, Json<OpenAiErrorResponse>)> { ) -> 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( return Err(openai_error(
StatusCode::TOO_MANY_REQUESTS, StatusCode::TOO_MANY_REQUESTS,
"Rate limit exceeded. Please try again later.", "Rate limit exceeded. Please try again later.",
+509 -187
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. /// Prevents resource exhaustion from connection flooding.
const MAX_CONNECTIONS: u64 = 100; 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. /// 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 { pub struct SseManager {
tx: broadcast::Sender<SseEvent>, tx: broadcast::Sender<ScopedEvent>,
connection_count: Arc<AtomicU64>, connection_count: Arc<AtomicU64>,
max_connections: u64, max_connections: u64,
} }
@@ -45,7 +61,7 @@ impl SseManager {
/// only be called before the server starts accepting connections (i.e., /// only be called before the server starts accepting connections (i.e.,
/// during startup wiring). Calling it after connections are established /// during startup wiring). Calling it after connections are established
/// will break connection tracking and allow exceeding `MAX_CONNECTIONS`. /// 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 { Self {
tx, tx,
connection_count: Arc::new(AtomicU64::new(0)), connection_count: Arc::new(AtomicU64::new(0)),
@@ -53,15 +69,28 @@ impl SseManager {
} }
} }
/// Broadcast an event to all connected clients. /// Get a clone of the broadcast sender for use by other components.
pub fn broadcast(&self, event: SseEvent) { pub(crate) fn sender(&self) -> broadcast::Sender<ScopedEvent> {
// Ignore send errors (no receivers is fine) self.tx.clone()
let _ = self.tx.send(event);
} }
/// Get a clone of the broadcast sender for use by other components. /// Broadcast an event to all connected clients (global/unscoped).
pub fn sender(&self) -> broadcast::Sender<SseEvent> { pub fn broadcast(&self, event: SseEvent) {
self.tx.clone() 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. /// Get current number of active connections.
@@ -71,11 +100,15 @@ impl SseManager {
/// Create a raw broadcast subscription for non-SSE consumers (e.g. WebSocket). /// Create a raw broadcast subscription for non-SSE consumers (e.g. WebSocket).
/// ///
/// Returns a stream of `SseEvent` values and increments/decrements the /// When `user_id` is `Some`, only events scoped to that user (or global
/// connection counter on creation/drop, just like `subscribe()` does for SSE. /// events) are delivered. When `None`, all events are delivered (single-user
/// backwards compatibility).
/// ///
/// Returns `None` if the maximum connection limit has been reached. /// 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 // Atomically increment only if below the limit. This prevents
// concurrent callers from overshooting max_connections. // concurrent callers from overshooting max_connections.
let counter = Arc::clone(&self.connection_count); let counter = Arc::clone(&self.connection_count);
@@ -91,7 +124,19 @@ impl SseManager {
.ok()?; .ok()?;
let rx = self.tx.subscribe(); 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 { Some(CountedStream {
inner: stream, inner: stream,
@@ -101,9 +146,13 @@ impl SseManager {
/// Create a new SSE stream for a client connection. /// 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. /// Returns `None` if the maximum connection limit has been reached.
pub fn subscribe( pub fn subscribe(
&self, &self,
user_id: Option<String>,
) -> Option<Sse<impl Stream<Item = Result<Event, Infallible>> + Send + 'static + use<>>> { ) -> Option<Sse<impl Stream<Item = Result<Event, Infallible>> + Send + 'static + use<>>> {
// Atomically increment only if below the limit. // Atomically increment only if below the limit.
let counter = Arc::clone(&self.connection_count); let counter = Arc::clone(&self.connection_count);
@@ -120,9 +169,23 @@ impl SseManager {
let rx = self.tx.subscribe(); let rx = self.tx.subscribe();
let stream = BroadcastStream::new(rx) let stream = BroadcastStream::new(rx)
.filter_map(|result| result.ok()) .filter_map(move |result| match result {
.map(|event| { Ok(scoped) => match (&user_id, &scoped.user_id) {
let data = serde_json::to_string(&event).unwrap_or_default(); (_, 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 { let event_type = match &event {
SseEvent::Response { .. } => "response", SseEvent::Response { .. } => "response",
SseEvent::Thinking { .. } => "thinking", SseEvent::Thinking { .. } => "thinking",
@@ -147,7 +210,7 @@ impl SseManager {
SseEvent::TurnCost { .. } => "turn_cost", SseEvent::TurnCost { .. } => "turn_cost",
SseEvent::ExtensionStatus { .. } => "extension_status", 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 // Wrap in a stream that decrements on drop
@@ -215,16 +278,14 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn test_broadcast_to_receiver() { async fn test_broadcast_to_receiver() {
let manager = SseManager::new(); 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 { manager.broadcast(SseEvent::Status {
message: "test".to_string(), message: "test".to_string(),
thread_id: None, thread_id: None,
}); });
let event = rx.next().await; let event = stream.next().await.unwrap();
assert!(event.is_some());
let event = event.unwrap().unwrap();
match event { match event {
SseEvent::Status { message, .. } => assert_eq!(message, "test"), SseEvent::Status { message, .. } => assert_eq!(message, "test"),
_ => panic!("unexpected event type"), _ => panic!("unexpected event type"),
@@ -234,7 +295,7 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn test_subscribe_raw_receives_events() { async fn test_subscribe_raw_receives_events() {
let manager = SseManager::new(); 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); assert_eq!(manager.connection_count(), 1);
@@ -254,7 +315,7 @@ mod tests {
async fn test_subscribe_raw_decrements_on_drop() { async fn test_subscribe_raw_decrements_on_drop() {
let manager = SseManager::new(); 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); assert_eq!(manager.connection_count(), 1);
} }
// Stream dropped, counter should decrement // Stream dropped, counter should decrement
@@ -264,8 +325,8 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn test_subscribe_raw_multiple_subscribers() { async fn test_subscribe_raw_multiple_subscribers() {
let manager = SseManager::new(); let manager = SseManager::new();
let mut s1 = 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().expect("should subscribe")); let mut s2 = Box::pin(manager.subscribe_raw(None).expect("should subscribe"));
assert_eq!(manager.connection_count(), 2); assert_eq!(manager.connection_count(), 2);
manager.broadcast(SseEvent::Heartbeat); manager.broadcast(SseEvent::Heartbeat);
@@ -286,12 +347,51 @@ mod tests {
let mut manager = SseManager::new(); let mut manager = SseManager::new();
manager.max_connections = 2; // Low limit for testing manager.max_connections = 2; // Low limit for testing
let _s1 = Box::pin(manager.subscribe_raw().expect("first should succeed")); let _s1 = Box::pin(manager.subscribe_raw(None).expect("first should succeed"));
let _s2 = Box::pin(manager.subscribe_raw().expect("second should succeed")); let _s2 = Box::pin(manager.subscribe_raw(None).expect("second should succeed"));
assert_eq!(manager.connection_count(), 2); assert_eq!(manager.connection_count(), 2);
// Third should be rejected // Third should be rejected
assert!(manager.subscribe_raw().is_none()); assert!(manager.subscribe_raw(None).is_none());
assert!(manager.subscribe().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();
assert!(matches!(e, SseEvent::Heartbeat));
} }
} }
+23 -6
View File
@@ -10,7 +10,8 @@ use std::sync::Arc;
use tokio::sync::mpsc; use tokio::sync::mpsc;
use crate::channels::IncomingMessage; 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::sse::SseManager;
use crate::channels::web::ws::WsConnectionTracker; use crate::channels::web::ws::WsConnectionTracker;
@@ -64,8 +65,9 @@ impl TestGatewayBuilder {
pub fn build(self) -> Arc<GatewayState> { pub fn build(self) -> Arc<GatewayState> {
Arc::new(GatewayState { Arc::new(GatewayState {
msg_tx: tokio::sync::RwLock::new(self.msg_tx), msg_tx: tokio::sync::RwLock::new(self.msg_tx),
sse: SseManager::new(), sse: Arc::new(SseManager::new()),
workspace: None, workspace: None,
workspace_pool: None,
session_manager: None, session_manager: None,
log_broadcaster: None, log_broadcaster: None,
log_level_handle: None, log_level_handle: None,
@@ -74,14 +76,14 @@ impl TestGatewayBuilder {
store: None, store: None,
job_manager: None, job_manager: None,
prompt_queue: None, prompt_queue: None,
user_id: self.user_id, default_user_id: self.user_id,
shutdown_tx: tokio::sync::RwLock::new(None), shutdown_tx: tokio::sync::RwLock::new(None),
ws_tracker: Some(Arc::new(WsConnectionTracker::new())), ws_tracker: Some(Arc::new(WsConnectionTracker::new())),
llm_provider: self.llm_provider, llm_provider: self.llm_provider,
skill_registry: None, skill_registry: None,
skill_catalog: None, skill_catalog: None,
scheduler: None, scheduler: None,
chat_rate_limiter: RateLimiter::new(30, 60), chat_rate_limiter: PerUserRateLimiter::new(30, 60),
oauth_rate_limiter: RateLimiter::new(10, 60), oauth_rate_limiter: RateLimiter::new(10, 60),
webhook_rate_limiter: RateLimiter::new(10, 60), webhook_rate_limiter: RateLimiter::new(10, 60),
registry_entries: Vec::new(), registry_entries: Vec::new(),
@@ -98,11 +100,26 @@ impl TestGatewayBuilder {
self, self,
auth_token: &str, auth_token: &str,
) -> Result<(SocketAddr, Arc<GatewayState>), crate::error::ChannelError> { ) -> Result<(SocketAddr, Arc<GatewayState>), crate::error::ChannelError> {
let auth = MultiAuthState::single(auth_token.to_string(), "test-user".to_string());
let state = self.build(); let state = self.build();
let addr: SocketAddr = "127.0.0.1:0" let addr: SocketAddr = "127.0.0.1:0"
.parse() .parse()
.expect("hard-coded address must parse"); .expect("hard-coded address must parse"); // safety: constant literal
let bound = start_server(addr, state.clone(), auth_token.to_string()).await?; 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)) 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 /// When either task ends (client disconnect or broadcast closed), both are
/// cleaned up. /// 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(); let (mut ws_sink, mut ws_stream) = socket.split();
// Track connection // 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(); 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. // 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"); tracing::warn!("WebSocket rejected: too many connections");
// Decrement the WS tracker we already incremented above. // Decrement the WS tracker we already incremented above.
if let Some(ref tracker) = tracker_for_drop { 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 // 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 { while let Some(Ok(frame)) = ws_stream.next().await {
match frame { match frame {
Message::Text(text) => { Message::Text(text) => {
@@ -263,10 +267,14 @@ async fn handle_client_message(
token, token,
} => { } => {
if let Some(ref ext_mgr) = state.extension_manager { 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) => { Ok(result) => {
if result.verification.is_some() { if result.verification.is_some() {
state.sse.broadcast( state.sse.broadcast_for_user(
user_id,
crate::channels::web::types::SseEvent::AuthRequired { crate::channels::web::types::SseEvent::AuthRequired {
extension_name: extension_name.clone(), extension_name: extension_name.clone(),
instructions: Some(result.message), instructions: Some(result.message),
@@ -275,8 +283,9 @@ async fn handle_client_message(
}, },
); );
} else { } else {
crate::channels::web::server::clear_auth_mode(state).await; crate::channels::web::server::clear_auth_mode(state, user_id).await;
state.sse.broadcast( state.sse.broadcast_for_user(
user_id,
crate::channels::web::types::SseEvent::AuthCompleted { crate::channels::web::types::SseEvent::AuthCompleted {
extension_name, extension_name,
success: true, success: true,
@@ -288,7 +297,8 @@ async fn handle_client_message(
Err(e) => { Err(e) => {
let msg = format!("Auth failed: {}", e); let msg = format!("Auth failed: {}", e);
if matches!(e, crate::extensions::ExtensionError::ValidationFailed(_)) { if matches!(e, crate::extensions::ExtensionError::ValidationFailed(_)) {
state.sse.broadcast( state.sse.broadcast_for_user(
user_id,
crate::channels::web::types::SseEvent::AuthRequired { crate::channels::web::types::SseEvent::AuthRequired {
extension_name: extension_name.clone(), extension_name: extension_name.clone(),
instructions: Some(msg.clone()), instructions: Some(msg.clone()),
@@ -311,7 +321,7 @@ async fn handle_client_message(
} }
} }
WsClientMessage::AuthCancel { .. } => { WsClientMessage::AuthCancel { .. } => {
crate::channels::web::server::clear_auth_mode(state).await; crate::channels::web::server::clear_auth_mode(state, user_id).await;
} }
WsClientMessage::Ping => { WsClientMessage::Ping => {
let _ = direct_tx.send(WsServerMessage::Pong).await; let _ = direct_tx.send(WsServerMessage::Pong).await;
@@ -498,8 +508,9 @@ mod tests {
GatewayState { GatewayState {
msg_tx: tokio::sync::RwLock::new(msg_tx), msg_tx: tokio::sync::RwLock::new(msg_tx),
sse: SseManager::new(), sse: Arc::new(SseManager::new()),
workspace: None, workspace: None,
workspace_pool: None,
session_manager: None, session_manager: None,
log_broadcaster: None, log_broadcaster: None,
log_level_handle: None, log_level_handle: None,
@@ -509,13 +520,13 @@ mod tests {
job_manager: None, job_manager: None,
prompt_queue: None, prompt_queue: None,
scheduler: None, scheduler: None,
user_id: "test".to_string(), default_user_id: "test".to_string(),
shutdown_tx: tokio::sync::RwLock::new(None), shutdown_tx: tokio::sync::RwLock::new(None),
ws_tracker: Some(Arc::new(WsConnectionTracker::new())), ws_tracker: Some(Arc::new(WsConnectionTracker::new())),
llm_provider: None, llm_provider: None,
skill_registry: None, skill_registry: None,
skill_catalog: 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), oauth_rate_limiter: crate::channels::web::server::RateLimiter::new(10, 60),
webhook_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(), registry_entries: Vec::new(),
+2 -2
View File
@@ -447,8 +447,8 @@ pub struct PendingOAuthFlow {
pub user_id: String, pub user_id: String,
/// Secrets store reference for token persistence. /// Secrets store reference for token persistence.
pub secrets: Arc<dyn SecretsStore + Send + Sync>, pub secrets: Arc<dyn SecretsStore + Send + Sync>,
/// SSE broadcast sender for notifying the web UI. /// SSE broadcast manager for notifying the web UI.
pub sse_sender: Option<tokio::sync::broadcast::Sender<crate::channels::web::types::SseEvent>>, pub sse_manager: Option<Arc<crate::channels::web::sse::SseManager>>,
/// Gateway auth token for authenticating with the platform token exchange proxy. /// Gateway auth token for authenticating with the platform token exchange proxy.
pub gateway_token: Option<String>, pub gateway_token: Option<String>,
/// Additional form params for the token exchange request. /// Additional form params for the token exchange request.
+142
View File
@@ -2,6 +2,7 @@ use std::collections::HashMap;
use std::path::PathBuf; use std::path::PathBuf;
use secrecy::SecretString; use secrecy::SecretString;
use serde::Deserialize;
use crate::bootstrap::ironclaw_base_dir; use crate::bootstrap::ironclaw_base_dir;
use crate::config::helpers::{optional_env, parse_bool_env, parse_optional_env}; 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. /// Bearer token for authentication. Random hex generated at startup if unset.
pub auth_token: Option<String>, pub auth_token: Option<String>,
pub user_id: 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). /// Signal channel configuration (signal-cli daemon HTTP/JSON-RPC).
@@ -115,6 +136,118 @@ impl ChannelsConfig {
.or_else(|| cs.gateway_user_id.clone()) .or_else(|| cs.gateway_user_id.clone())
.unwrap_or_else(|| owner_id.to_string()); .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 { Some(GatewayConfig {
host: optional_env("GATEWAY_HOST")? host: optional_env("GATEWAY_HOST")?
.or_else(|| cs.gateway_host.clone()) .or_else(|| cs.gateway_host.clone())
@@ -126,6 +259,9 @@ impl ChannelsConfig {
auth_token: optional_env("GATEWAY_AUTH_TOKEN")? auth_token: optional_env("GATEWAY_AUTH_TOKEN")?
.or_else(|| cs.gateway_auth_token.clone()), .or_else(|| cs.gateway_auth_token.clone()),
user_id, user_id,
workspace_read_scopes,
memory_layers,
user_tokens,
}) })
} else { } else {
None None
@@ -281,6 +417,9 @@ mod tests {
port: 3000, port: 3000,
auth_token: Some("tok-abc".to_string()), auth_token: Some("tok-abc".to_string()),
user_id: "default".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.host, "127.0.0.1");
assert_eq!(cfg.port, 3000); assert_eq!(cfg.port, 3000);
@@ -295,6 +434,9 @@ mod tests {
port: 3001, port: 3001,
auth_token: None, auth_token: None,
user_id: "anon".to_string(), user_id: "anon".to_string(),
workspace_read_scopes: vec![],
memory_layers: vec![],
user_tokens: None,
}; };
assert!(cfg.auth_token.is_none()); assert!(cfg.auth_token.is_none());
} }
+69
View File
@@ -230,6 +230,49 @@ impl JobStore for LibSqlBackend {
Ok(jobs) 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( async fn get_agent_job_failure_reason(
&self, &self,
id: Uuid, id: Uuid,
@@ -277,6 +320,32 @@ impl JobStore for LibSqlBackend {
Ok(summary) 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> { async fn save_action(&self, job_id: Uuid, action: &ActionRecord) -> Result<(), DatabaseError> {
let conn = self.connect().await?; let conn = self.connect().await?;
let duration_ms = action.duration.as_millis() as i64; let duration_ms = action.duration.as_millis() as i64;
+8
View File
@@ -409,7 +409,15 @@ pub trait JobStore: Send + Sync {
async fn mark_job_stuck(&self, id: Uuid) -> Result<(), DatabaseError>; async fn mark_job_stuck(&self, id: Uuid) -> Result<(), DatabaseError>;
async fn get_stuck_jobs(&self) -> Result<Vec<Uuid>, 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(&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(&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). /// Get the failure reason for a single agent job (O(1) lookup).
async fn get_agent_job_failure_reason(&self, id: Uuid) async fn get_agent_job_failure_reason(&self, id: Uuid)
-> Result<Option<String>, DatabaseError>; -> Result<Option<String>, DatabaseError>;
+14
View File
@@ -249,10 +249,24 @@ impl JobStore for PgBackend {
self.store.list_agent_jobs().await 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> { async fn agent_job_summary(&self) -> Result<AgentJobSummary, DatabaseError> {
self.store.agent_job_summary().await 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( async fn get_agent_job_failure_reason(
&self, &self,
id: Uuid, id: Uuid,
+326 -233
View File
File diff suppressed because it is too large Load Diff
+53
View File
@@ -842,6 +842,38 @@ impl Store {
.collect()) .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. /// Get the failure reason for a single agent job.
pub async fn get_agent_job_failure_reason( pub async fn get_agent_job_failure_reason(
&self, &self,
@@ -875,6 +907,27 @@ impl Store {
} }
Ok(summary) 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 ==================== // ==================== Job Events ====================
+67 -15
View File
@@ -589,15 +589,48 @@ async fn async_main() -> anyhow::Result<()> {
// ── Gateway channel ──────────────────────────────────────────────── // ── Gateway channel ────────────────────────────────────────────────
let mut gateway_url: Option<String> = None; let mut gateway_url: Option<String> = None;
let mut sse_sender: Option< let mut sse_manager: Option<std::sync::Arc<ironclaw::channels::web::sse::SseManager>> = None;
tokio::sync::broadcast::Sender<ironclaw::channels::web::types::SseEvent>, let mut _gateway_state: Option<std::sync::Arc<ironclaw::channels::web::server::GatewayState>> =
> = None; None;
if let Some(ref gw_config) = config.channels.gateway { if let Some(ref gw_config) = config.channels.gateway {
let mut gw = // Build multi-user auth state if user_tokens is configured, else single-user.
GatewayChannel::new(gw_config.clone()).with_llm_provider(Arc::clone(&components.llm)); 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 { if let Some(ref ws) = components.workspace {
gw = gw.with_workspace(Arc::clone(ws)); 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_session_manager(Arc::clone(&session_manager));
gw = gw.with_log_broadcaster(Arc::clone(&log_broadcaster)); gw = gw.with_log_broadcaster(Arc::clone(&log_broadcaster));
gw = gw.with_log_level_handle(Arc::clone(&log_level_handle)); gw = gw.with_log_level_handle(Arc::clone(&log_level_handle));
@@ -648,8 +681,12 @@ async fn async_main() -> anyhow::Result<()> {
let mut rx = tx.subscribe(); let mut rx = tx.subscribe();
let gw_state = Arc::clone(gw.state()); let gw_state = Arc::clone(gw.state());
tokio::spawn(async move { tokio::spawn(async move {
while let Ok((_job_id, event)) = rx.recv().await { while let Ok((_job_id, user_id, event)) = rx.recv().await {
gw_state.sse.broadcast(event); if user_id.is_empty() {
gw_state.sse.broadcast(event);
} else {
gw_state.sse.broadcast_for_user(&user_id, event);
}
} }
}); });
} }
@@ -691,7 +728,8 @@ async fn async_main() -> anyhow::Result<()> {
// Capture SSE sender and routine engine slot before moving gw into channels. // Capture SSE sender and routine engine slot before moving gw into channels.
// IMPORTANT: This must come after all `with_*` calls since `rebuild_state` // IMPORTANT: This must come after all `with_*` calls since `rebuild_state`
// creates a new SseManager, which would orphan this sender. // creates a new SseManager, which would orphan this sender.
sse_sender = Some(gw.state().sse.sender()); sse_manager = Some(Arc::clone(&gw.state().sse));
_gateway_state = Some(Arc::clone(gw.state()));
channel_names.push("gateway".to_string()); channel_names.push("gateway".to_string());
channels.add(Box::new(gw)).await; channels.add(Box::new(gw)).await;
} }
@@ -774,12 +812,20 @@ async fn async_main() -> anyhow::Result<()> {
// Auto-activate WASM channels that were active in a previous session. // Auto-activate WASM channels that were active in a previous session.
// Relay channels are handled separately below via restore_relay_channels(). // Relay channels are handled separately below via restore_relay_channels().
let persisted = ext_mgr.load_persisted_active_channels().await; let ext_user_id = config
.channels
.gateway
.as_ref()
.map(|g| g.user_id.clone())
.unwrap_or_else(|| "default".to_string());
let persisted = ext_mgr.load_persisted_active_channels(&ext_user_id).await;
for name in &persisted { 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; continue;
} }
match ext_mgr.activate(name).await { match ext_mgr.activate(name, &ext_user_id).await {
Ok(result) => { Ok(result) => {
tracing::debug!( tracing::debug!(
channel = %name, channel = %name,
@@ -804,14 +850,20 @@ async fn async_main() -> anyhow::Result<()> {
ext_mgr ext_mgr
.set_relay_channel_manager(Arc::clone(&channels)) .set_relay_channel_manager(Arc::clone(&channels))
.await; .await;
ext_mgr.restore_relay_channels().await; let ext_user_id = config
.channels
.gateway
.as_ref()
.map(|g| g.user_id.clone())
.unwrap_or_else(|| "default".to_string());
ext_mgr.restore_relay_channels(&ext_user_id).await;
} }
// Wire SSE sender into extension manager for broadcasting status events. // Wire SSE sender into extension manager for broadcasting status events.
if let Some(ref ext_mgr) = components.extension_manager if let Some(ref ext_mgr) = components.extension_manager
&& let Some(ref sender) = sse_sender && let Some(sse) = sse_manager
{ {
ext_mgr.set_sse_sender(sender.clone()).await; ext_mgr.set_sse_sender(sse).await;
} }
// Snapshot memory for trace recording before the agent starts // Snapshot memory for trace recording before the agent starts
@@ -849,7 +901,7 @@ async fn async_main() -> anyhow::Result<()> {
skills_config: config.skills.clone(), skills_config: config.skills.clone(),
hooks: components.hooks, hooks: components.hooks,
cost_guard: components.cost_guard, cost_guard: components.cost_guard,
sse_tx: sse_sender, sse_tx: None, // TODO: wire SseManager into scheduler (needs Sender<SseEvent> → Arc<SseManager> refactor)
http_interceptor, http_interceptor,
transcription: config.transcription.create_provider().map(|p| { transcription: config.transcription.create_provider().map(|p| {
Arc::new(ironclaw::llm::transcription::TranscriptionMiddleware::new( Arc::new(ironclaw::llm::transcription::TranscriptionMiddleware::new(
+24 -6
View File
@@ -40,7 +40,8 @@ pub struct OrchestratorState {
pub job_manager: Arc<ContainerJobManager>, pub job_manager: Arc<ContainerJobManager>,
pub token_store: TokenStore, pub token_store: TokenStore,
/// Broadcast channel for job events (consumed by the web gateway SSE). /// 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. /// Buffered follow-up prompts for sandbox jobs, keyed by job_id.
pub prompt_queue: Arc<Mutex<HashMap<Uuid, VecDeque<PendingPrompt>>>>, pub prompt_queue: Arc<Mutex<HashMap<Uuid, VecDeque<PendingPrompt>>>>,
/// Database handle for persisting job events. /// Database handle for persisting job events.
@@ -351,9 +352,24 @@ async fn job_event_handler(
}, },
}; };
// Broadcast via the channel (if configured) // Broadcast via the channel (if configured).
// Look up the job owner so the gateway can scope delivery per-user.
if let Some(ref tx) = state.job_event_tx { if let Some(ref tx) = state.job_event_tx {
let _ = tx.send((job_id, sse_event)); let user_id = 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(uid) = user_id {
let _ = tx.send((job_id, uid, sse_event));
} else {
// Fallback: broadcast globally (single-user mode or job not found).
let _ = tx.send((job_id, String::new(), sse_event));
}
} }
Ok(StatusCode::OK) Ok(StatusCode::OK)
@@ -769,8 +785,10 @@ mod tests {
let resp = router.oneshot(req).await.unwrap(); let resp = router.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK); 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); assert_eq!(recv_id, job_id);
// No store configured, so user_id falls back to empty string.
assert_eq!(recv_uid, "");
match event { match event {
SseEvent::JobMessage { SseEvent::JobMessage {
job_id: jid, job_id: jid,
@@ -824,7 +842,7 @@ mod tests {
let resp = router.oneshot(req).await.unwrap(); let resp = router.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK); 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 { match event {
SseEvent::JobToolUse { tool_name, .. } => { SseEvent::JobToolUse { tool_name, .. } => {
assert_eq!(tool_name, "shell"); assert_eq!(tool_name, "shell");
@@ -869,7 +887,7 @@ mod tests {
let resp = router.oneshot(req).await.unwrap(); let resp = router.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK); 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 // Unknown event types fall through to JobStatus
assert!(matches!(event, SseEvent::JobStatus { .. })); assert!(matches!(event, SseEvent::JobStatus { .. }));
} }
+1 -1
View File
@@ -63,7 +63,7 @@ fn resolve_orchestrator_port() -> u16 {
/// Result of orchestrator setup, containing all handles needed by the agent. /// Result of orchestrator setup, containing all handles needed by the agent.
pub struct OrchestratorSetup { pub struct OrchestratorSetup {
pub container_job_manager: Option<Arc<ContainerJobManager>>, 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 prompt_queue: Arc<Mutex<HashMap<Uuid, VecDeque<api::PendingPrompt>>>>,
pub docker_status: crate::sandbox::DockerStatus, pub docker_status: crate::sandbox::DockerStatus,
} }
+17 -17
View File
@@ -130,7 +130,7 @@ impl Tool for ToolInstallTool {
async fn execute( async fn execute(
&self, &self,
params: serde_json::Value, params: serde_json::Value,
_ctx: &JobContext, ctx: &JobContext,
) -> Result<ToolOutput, ToolError> { ) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now(); let start = std::time::Instant::now();
@@ -150,7 +150,7 @@ impl Tool for ToolInstallTool {
let result = self let result = self
.manager .manager
.install(name, url, kind_hint) .install(name, url, kind_hint, &ctx.user_id)
.await .await
.map_err(|e| ToolError::ExecutionFailed(e.to_string()))?; .map_err(|e| ToolError::ExecutionFailed(e.to_string()))?;
@@ -205,7 +205,7 @@ impl Tool for ToolAuthTool {
async fn execute( async fn execute(
&self, &self,
params: serde_json::Value, params: serde_json::Value,
_ctx: &JobContext, ctx: &JobContext,
) -> Result<ToolOutput, ToolError> { ) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now(); let start = std::time::Instant::now();
@@ -213,13 +213,13 @@ impl Tool for ToolAuthTool {
let result = self let result = self
.manager .manager
.auth(name) .auth(name, &ctx.user_id)
.await .await
.map_err(|e| ToolError::ExecutionFailed(e.to_string()))?; .map_err(|e| ToolError::ExecutionFailed(e.to_string()))?;
// Auto-activate after successful auth so tools are available immediately // Auto-activate after successful auth so tools are available immediately
if result.is_authenticated() { if result.is_authenticated() {
match self.manager.activate(name).await { match self.manager.activate(name, &ctx.user_id).await {
Ok(activate_result) => { Ok(activate_result) => {
let output = serde_json::json!({ let output = serde_json::json!({
"status": "authenticated_and_activated", "status": "authenticated_and_activated",
@@ -304,13 +304,13 @@ impl Tool for ToolActivateTool {
async fn execute( async fn execute(
&self, &self,
params: serde_json::Value, params: serde_json::Value,
_ctx: &JobContext, ctx: &JobContext,
) -> Result<ToolOutput, ToolError> { ) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now(); let start = std::time::Instant::now();
let name = require_str(&params, "name")?; let name = require_str(&params, "name")?;
match self.manager.activate(name).await { match self.manager.activate(name, &ctx.user_id).await {
Ok(result) => { Ok(result) => {
let output = serde_json::to_value(&result) let output = serde_json::to_value(&result)
.unwrap_or_else(|_| serde_json::json!({"error": "serialization failed"})); .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 // Activation failed due to missing auth; initiate auth flow
// so the agent loop can show the auth card. // 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() => { Ok(auth_result) if auth_result.is_authenticated() => {
// Auth succeeded (e.g. env var was set); retry activation. // Auth succeeded (e.g. env var was set); retry activation.
let result = self let result = self
.manager .manager
.activate(name) .activate(name, &ctx.user_id)
.await .await
.map_err(|e| ToolError::ExecutionFailed(e.to_string()))?; .map_err(|e| ToolError::ExecutionFailed(e.to_string()))?;
let output = serde_json::to_value(&result).unwrap_or_else( let output = serde_json::to_value(&result).unwrap_or_else(
@@ -404,7 +404,7 @@ impl Tool for ToolListTool {
async fn execute( async fn execute(
&self, &self,
params: serde_json::Value, params: serde_json::Value,
_ctx: &JobContext, ctx: &JobContext,
) -> Result<ToolOutput, ToolError> { ) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now(); let start = std::time::Instant::now();
@@ -425,7 +425,7 @@ impl Tool for ToolListTool {
let extensions = self let extensions = self
.manager .manager
.list(kind_filter, include_available) .list(kind_filter, include_available, &ctx.user_id)
.await .await
.map_err(|e| ToolError::ExecutionFailed(e.to_string()))?; .map_err(|e| ToolError::ExecutionFailed(e.to_string()))?;
@@ -477,7 +477,7 @@ impl Tool for ToolRemoveTool {
async fn execute( async fn execute(
&self, &self,
params: serde_json::Value, params: serde_json::Value,
_ctx: &JobContext, ctx: &JobContext,
) -> Result<ToolOutput, ToolError> { ) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now(); let start = std::time::Instant::now();
@@ -485,7 +485,7 @@ impl Tool for ToolRemoveTool {
let message = self let message = self
.manager .manager
.remove(name) .remove(name, &ctx.user_id)
.await .await
.map_err(|e| ToolError::ExecutionFailed(e.to_string()))?; .map_err(|e| ToolError::ExecutionFailed(e.to_string()))?;
@@ -541,7 +541,7 @@ impl Tool for ToolUpgradeTool {
async fn execute( async fn execute(
&self, &self,
params: serde_json::Value, params: serde_json::Value,
_ctx: &JobContext, ctx: &JobContext,
) -> Result<ToolOutput, ToolError> { ) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now(); let start = std::time::Instant::now();
@@ -549,7 +549,7 @@ impl Tool for ToolUpgradeTool {
let result = self let result = self
.manager .manager
.upgrade(name) .upgrade(name, &ctx.user_id)
.await .await
.map_err(|e| ToolError::ExecutionFailed(e.to_string()))?; .map_err(|e| ToolError::ExecutionFailed(e.to_string()))?;
@@ -603,7 +603,7 @@ impl Tool for ExtensionInfoTool {
async fn execute( async fn execute(
&self, &self,
params: serde_json::Value, params: serde_json::Value,
_ctx: &JobContext, ctx: &JobContext,
) -> Result<ToolOutput, ToolError> { ) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now(); let start = std::time::Instant::now();
@@ -611,7 +611,7 @@ impl Tool for ExtensionInfoTool {
let info = self let info = self
.manager .manager
.extension_info(name) .extension_info(name, &ctx.user_id)
.await .await
.map_err(|e| ToolError::ExecutionFailed(e.to_string()))?; .map_err(|e| ToolError::ExecutionFailed(e.to_string()))?;
+2 -2
View File
@@ -85,7 +85,7 @@ pub struct CreateJobTool {
job_manager: Option<Arc<ContainerJobManager>>, job_manager: Option<Arc<ContainerJobManager>>,
store: Option<Arc<dyn Database>>, store: Option<Arc<dyn Database>>,
/// Broadcast sender for job events (used to subscribe a monitor). /// 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. /// Injection channel for pushing messages into the agent loop.
inject_tx: Option<tokio::sync::mpsc::Sender<IncomingMessage>>, inject_tx: Option<tokio::sync::mpsc::Sender<IncomingMessage>>,
/// Encrypted secrets store for validating credential grants. /// Encrypted secrets store for validating credential grants.
@@ -120,7 +120,7 @@ impl CreateJobTool {
/// monitor that forwards Claude Code output to the main agent loop. /// monitor that forwards Claude Code output to the main agent loop.
pub fn with_monitor_deps( pub fn with_monitor_deps(
mut self, 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>, inject_tx: tokio::sync::mpsc::Sender<IncomingMessage>,
) -> Self { ) -> Self {
self.event_tx = Some(event_tx); self.event_tx = Some(event_tx);
+358 -45
View File
@@ -12,15 +12,119 @@
//! Use `memory_write` to persist important facts that should be remembered //! Use `memory_write` to persist important facts that should be remembered
//! across sessions. //! across sessions.
use std::collections::HashMap;
use std::path::Path; use std::path::Path;
use std::sync::Arc; use std::sync::Arc;
use async_trait::async_trait; use async_trait::async_trait;
use tokio::sync::RwLock;
use crate::context::JobContext; use crate::context::JobContext;
use crate::tools::tool::{Tool, ToolError, ToolOutput, require_str}; use crate::tools::tool::{Tool, ToolError, ToolOutput, require_str};
use crate::workspace::{Workspace, paths}; 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)
}
}
/// Creates per-user workspaces on demand, caching them for reuse.
///
/// Used in multi-tenant mode where each authenticated user gets their own
/// workspace scope. The workspace is constructed with the same configuration
/// (embeddings, search config, memory layers) as the startup workspace.
pub struct PerUserWorkspaceResolver {
db: Arc<dyn crate::db::Database>,
embeddings: Option<Arc<dyn crate::workspace::EmbeddingProvider>>,
embedding_cache_config: crate::workspace::EmbeddingCacheConfig,
search_config: crate::config::WorkspaceSearchConfig,
workspace_config: crate::config::WorkspaceConfig,
cache: RwLock<HashMap<String, Arc<Workspace>>>,
}
impl PerUserWorkspaceResolver {
pub fn new(
db: Arc<dyn crate::db::Database>,
embeddings: Option<Arc<dyn crate::workspace::EmbeddingProvider>>,
embedding_cache_config: crate::workspace::EmbeddingCacheConfig,
search_config: crate::config::WorkspaceSearchConfig,
workspace_config: crate::config::WorkspaceConfig,
) -> Self {
Self {
db,
embeddings,
embedding_cache_config,
search_config,
workspace_config,
cache: RwLock::new(HashMap::new()),
}
}
fn build_workspace(&self, user_id: &str) -> Arc<Workspace> {
let mut ws = Workspace::new_with_db(user_id, Arc::clone(&self.db))
.with_search_config(&self.search_config);
if let Some(ref emb) = self.embeddings {
ws = ws.with_embeddings_cached(Arc::clone(emb), self.embedding_cache_config.clone());
}
if !self.workspace_config.read_scopes.is_empty() {
ws = ws.with_additional_read_scopes(self.workspace_config.read_scopes.clone());
}
ws = ws.with_memory_layers(self.workspace_config.memory_layers.clone());
Arc::new(ws)
}
}
#[async_trait]
impl WorkspaceResolver for PerUserWorkspaceResolver {
async fn resolve(&self, user_id: &str) -> Arc<Workspace> {
// Fast path: read lock
{
let cache = self.cache.read().await;
if let Some(ws) = cache.get(user_id) {
return Arc::clone(ws);
}
}
// Slow path: write lock, double-check
let mut cache = self.cache.write().await;
if let Some(ws) = cache.get(user_id) {
return Arc::clone(ws);
}
let ws = self.build_workspace(user_id);
cache.insert(user_id.to_string(), Arc::clone(&ws));
tracing::debug!(user_id = user_id, "Created per-user workspace");
ws
}
}
/// Detect paths that are clearly local filesystem references, not workspace-memory docs. /// Detect paths that are clearly local filesystem references, not workspace-memory docs.
/// ///
/// Examples: /// Examples:
@@ -62,13 +166,20 @@ fn map_write_err(e: crate::error::WorkspaceError) -> ToolError {
/// The agent should call this tool before answering questions about /// The agent should call this tool before answering questions about
/// prior work, decisions, preferences, or any historical context. /// prior work, decisions, preferences, or any historical context.
pub struct MemorySearchTool { pub struct MemorySearchTool {
workspace: Arc<Workspace>, resolver: Arc<dyn WorkspaceResolver>,
} }
impl MemorySearchTool { impl MemorySearchTool {
/// Create a new memory search tool. /// Create a new memory search tool with a workspace resolver.
pub fn new(workspace: Arc<Workspace>) -> Self { pub fn new(resolver: Arc<dyn WorkspaceResolver>) -> Self {
Self { workspace } 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 +218,7 @@ impl Tool for MemorySearchTool {
async fn execute( async fn execute(
&self, &self,
params: serde_json::Value, params: serde_json::Value,
_ctx: &JobContext, ctx: &JobContext,
) -> Result<ToolOutput, ToolError> { ) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now(); let start = std::time::Instant::now();
@@ -119,8 +230,8 @@ impl Tool for MemorySearchTool {
.unwrap_or(5) .unwrap_or(5)
.min(20) as usize; .min(20) as usize;
let results = self let workspace = self.resolver.resolve(&ctx.user_id).await;
.workspace let results = workspace
.search(query, limit) .search(query, limit)
.await .await
.map_err(|e| ToolError::ExecutionFailed(format!("Search failed: {}", e)))?; .map_err(|e| ToolError::ExecutionFailed(format!("Search failed: {}", e)))?;
@@ -151,13 +262,20 @@ impl Tool for MemorySearchTool {
/// Use this to persist important information that should be remembered /// Use this to persist important information that should be remembered
/// across sessions: decisions, preferences, facts, lessons learned. /// across sessions: decisions, preferences, facts, lessons learned.
pub struct MemoryWriteTool { pub struct MemoryWriteTool {
workspace: Arc<Workspace>, resolver: Arc<dyn WorkspaceResolver>,
} }
impl MemoryWriteTool { impl MemoryWriteTool {
/// Create a new memory write tool. /// Create a new memory write tool with a workspace resolver.
pub fn new(workspace: Arc<Workspace>) -> Self { pub fn new(resolver: Arc<dyn WorkspaceResolver>) -> Self {
Self { workspace } 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 +349,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. // Bootstrap target: clear BOOTSTRAP.md to mark first-run ritual complete.
// Handled early because it accepts empty content (unlike other targets). // Handled early because it accepts empty content (unlike other targets).
if target == "bootstrap" { if target == "bootstrap" {
// Write empty content to effectively disable the bootstrap injection. // Write empty content to effectively disable the bootstrap injection.
// system_prompt_for_context() skips empty files. // system_prompt_for_context() skips empty files.
self.workspace workspace
.write(paths::BOOTSTRAP, "") .write(paths::BOOTSTRAP, "")
.await .await
.map_err(map_write_err)?; .map_err(map_write_err)?;
// Also set the in-memory flag so BOOTSTRAP.md injection stops // Also set the in-memory flag so BOOTSTRAP.md injection stops
// immediately without waiting for a restart. // immediately without waiting for a restart.
self.workspace.mark_bootstrap_completed(); workspace.mark_bootstrap_completed();
let output = serde_json::json!({ let output = serde_json::json!({
"status": "cleared", "status": "cleared",
@@ -289,12 +409,12 @@ impl Tool for MemoryWriteTool {
// Otherwise, use default workspace methods (which include injection scanning). // Otherwise, use default workspace methods (which include injection scanning).
let layer_result = if let Some(layer_name) = layer { let layer_result = if let Some(layer_name) = layer {
let result = if append { let result = if append {
self.workspace workspace
.append_to_layer(layer_name, &resolved_path, content, force) .append_to_layer(layer_name, &resolved_path, content, force)
.await .await
.map_err(map_write_err)? .map_err(map_write_err)?
} else { } else {
self.workspace workspace
.write_to_layer(layer_name, &resolved_path, content, force) .write_to_layer(layer_name, &resolved_path, content, force)
.await .await
.map_err(map_write_err)? .map_err(map_write_err)?
@@ -307,31 +427,33 @@ impl Tool for MemoryWriteTool {
match target { match target {
"memory" => { "memory" => {
if append { if append {
self.workspace workspace
.append_memory(content) .append_memory(content)
.await .await
.map_err(map_write_err)?; .map_err(map_write_err)?;
} else { } else {
self.workspace workspace
.write(paths::MEMORY, content) .write(paths::MEMORY, content)
.await .await
.map_err(map_write_err)?; .map_err(map_write_err)?;
} }
} }
"daily_log" => { "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) .append_daily_log_tz(content, tz)
.await .await
.map_err(map_write_err)?; .map_err(map_write_err)?;
} }
_ => { _ => {
if append { if append {
self.workspace workspace
.append(&resolved_path, content) .append(&resolved_path, content)
.await .await
.map_err(map_write_err)?; .map_err(map_write_err)?;
} else { } else {
self.workspace workspace
.write(&resolved_path, content) .write(&resolved_path, content)
.await .await
.map_err(map_write_err)?; .map_err(map_write_err)?;
@@ -361,12 +483,12 @@ impl Tool for MemoryWriteTool {
}; };
let mut synced_docs: Vec<&str> = Vec::new(); let mut synced_docs: Vec<&str> = Vec::new();
if normalized_path == paths::PROFILE { if normalized_path == paths::PROFILE {
match self.workspace.sync_profile_documents().await { match workspace.sync_profile_documents().await {
Ok(true) => { Ok(true) => {
tracing::info!("profile write: synced USER.md + assistant-directives.md"); tracing::info!("profile write: synced USER.md + assistant-directives.md");
synced_docs.extend_from_slice(&[paths::USER, paths::ASSISTANT_DIRECTIVES]); 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(); let toml_path = crate::settings::Settings::default_toml_path();
if let Ok(Some(mut settings)) = crate::settings::Settings::load_toml(&toml_path) if let Ok(Some(mut settings)) = crate::settings::Settings::load_toml(&toml_path)
&& !settings.profile_onboarding_completed && !settings.profile_onboarding_completed
@@ -416,13 +538,20 @@ impl Tool for MemoryWriteTool {
/// ///
/// Use this to read the full content of any file in the workspace. /// Use this to read the full content of any file in the workspace.
pub struct MemoryReadTool { pub struct MemoryReadTool {
workspace: Arc<Workspace>, resolver: Arc<dyn WorkspaceResolver>,
} }
impl MemoryReadTool { impl MemoryReadTool {
/// Create a new memory read tool. /// Create a new memory read tool with a workspace resolver.
pub fn new(workspace: Arc<Workspace>) -> Self { pub fn new(resolver: Arc<dyn WorkspaceResolver>) -> Self {
Self { workspace } 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 +585,7 @@ impl Tool for MemoryReadTool {
async fn execute( async fn execute(
&self, &self,
params: serde_json::Value, params: serde_json::Value,
_ctx: &JobContext, ctx: &JobContext,
) -> Result<ToolOutput, ToolError> { ) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now(); let start = std::time::Instant::now();
@@ -470,8 +599,8 @@ impl Tool for MemoryReadTool {
))); )));
} }
let doc = self let workspace = self.resolver.resolve(&ctx.user_id).await;
.workspace let doc = workspace
.read(path) .read(path)
.await .await
.map_err(|e| ToolError::ExecutionFailed(format!("Read failed: {}", e)))?; .map_err(|e| ToolError::ExecutionFailed(format!("Read failed: {}", e)))?;
@@ -495,20 +624,27 @@ impl Tool for MemoryReadTool {
/// ///
/// Returns a hierarchical view of files and directories with configurable depth. /// Returns a hierarchical view of files and directories with configurable depth.
pub struct MemoryTreeTool { pub struct MemoryTreeTool {
workspace: Arc<Workspace>, resolver: Arc<dyn WorkspaceResolver>,
} }
impl MemoryTreeTool { impl MemoryTreeTool {
/// Create a new memory tree tool. /// Create a new memory tree tool with a workspace resolver.
pub fn new(workspace: Arc<Workspace>) -> Self { pub fn new(resolver: Arc<dyn WorkspaceResolver>) -> Self {
Self { workspace } 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. /// Recursively build tree structure.
/// ///
/// Returns a compact format where directories end with `/` and may have children. /// Returns a compact format where directories end with `/` and may have children.
async fn build_tree( async fn build_tree(
&self, workspace: &Arc<Workspace>,
path: &str, path: &str,
current_depth: usize, current_depth: usize,
max_depth: usize, max_depth: usize,
@@ -517,8 +653,7 @@ impl MemoryTreeTool {
return Ok(Vec::new()); return Ok(Vec::new());
} }
let entries = self let entries = workspace
.workspace
.list(path) .list(path)
.await .await
.map_err(|e| ToolError::ExecutionFailed(format!("Tree failed: {}", e)))?; .map_err(|e| ToolError::ExecutionFailed(format!("Tree failed: {}", e)))?;
@@ -533,8 +668,13 @@ impl MemoryTreeTool {
}; };
if entry.is_directory && current_depth < max_depth { if entry.is_directory && current_depth < max_depth {
let children = let children = Box::pin(Self::build_tree(
Box::pin(self.build_tree(&entry.path, current_depth + 1, max_depth)).await?; workspace,
&entry.path,
current_depth + 1,
max_depth,
))
.await?;
if children.is_empty() { if children.is_empty() {
result.push(serde_json::Value::String(display_path)); result.push(serde_json::Value::String(display_path));
} else { } else {
@@ -584,7 +724,7 @@ impl Tool for MemoryTreeTool {
async fn execute( async fn execute(
&self, &self,
params: serde_json::Value, params: serde_json::Value,
_ctx: &JobContext, ctx: &JobContext,
) -> Result<ToolOutput, ToolError> { ) -> Result<ToolOutput, ToolError> {
let start = std::time::Instant::now(); let start = std::time::Instant::now();
@@ -596,7 +736,8 @@ impl Tool for MemoryTreeTool {
.unwrap_or(1) .unwrap_or(1)
.clamp(1, 10) as usize; .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 // Compact output: just the tree array
Ok(ToolOutput::success( Ok(ToolOutput::success(
@@ -650,7 +791,7 @@ mod tests {
#[test] #[test]
fn test_memory_search_schema() { fn test_memory_search_schema() {
let workspace = make_test_workspace(); let workspace = make_test_workspace();
let tool = MemorySearchTool::new(workspace); let tool = MemorySearchTool::from_workspace(workspace);
assert_eq!(tool.name(), "memory_search"); assert_eq!(tool.name(), "memory_search");
assert!(!tool.requires_sanitization()); assert!(!tool.requires_sanitization());
@@ -668,7 +809,7 @@ mod tests {
#[test] #[test]
fn test_memory_write_schema() { fn test_memory_write_schema() {
let workspace = make_test_workspace(); let workspace = make_test_workspace();
let tool = MemoryWriteTool::new(workspace); let tool = MemoryWriteTool::from_workspace(workspace);
assert_eq!(tool.name(), "memory_write"); assert_eq!(tool.name(), "memory_write");
@@ -681,7 +822,7 @@ mod tests {
#[test] #[test]
fn test_memory_read_schema() { fn test_memory_read_schema() {
let workspace = make_test_workspace(); let workspace = make_test_workspace();
let tool = MemoryReadTool::new(workspace); let tool = MemoryReadTool::from_workspace(workspace);
assert_eq!(tool.name(), "memory_read"); assert_eq!(tool.name(), "memory_read");
@@ -698,7 +839,7 @@ mod tests {
#[test] #[test]
fn test_memory_tree_schema() { fn test_memory_tree_schema() {
let workspace = make_test_workspace(); let workspace = make_test_workspace();
let tool = MemoryTreeTool::new(workspace); let tool = MemoryTreeTool::from_workspace(workspace);
assert_eq!(tool.name(), "memory_tree"); assert_eq!(tool.name(), "memory_tree");
@@ -711,7 +852,7 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn test_memory_write_rejects_injection_to_identity_file() { async fn test_memory_write_rejects_injection_to_identity_file() {
let workspace = make_test_workspace(); let workspace = make_test_workspace();
let tool = MemoryWriteTool::new(workspace); let tool = MemoryWriteTool::from_workspace(workspace);
let ctx = JobContext::default(); let ctx = JobContext::default();
let params = serde_json::json!({ let params = serde_json::json!({
@@ -733,4 +874,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_per_user_workspace_resolver_returns_different_workspaces() {
let db = make_test_db().await;
let resolver = PerUserWorkspaceResolver::new(
db,
None,
crate::workspace::EmbeddingCacheConfig::default(),
crate::config::WorkspaceSearchConfig::default(),
crate::config::WorkspaceConfig::default(),
);
let ws_alice = resolver.resolve("alice").await;
let ws_bob = resolver.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_per_user_workspace_resolver_caches_workspace() {
let db = make_test_db().await;
let resolver = PerUserWorkspaceResolver::new(
db,
None,
crate::workspace::EmbeddingCacheConfig::default(),
crate::config::WorkspaceSearchConfig::default(),
crate::config::WorkspaceConfig::default(),
);
let ws1 = resolver.resolve("alice").await;
let ws2 = resolver.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 http;
mod job; mod job;
mod json; mod json;
mod memory; pub mod memory;
mod message; mod message;
pub mod path_utils; pub mod path_utils;
mod restart; mod restart;
+32 -6
View File
@@ -334,15 +334,37 @@ impl ToolRegistry {
tracing::debug!("Registered 5 development tools"); 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 /// Memory tools require a workspace for persistence. Call this after
/// `register_builtin_tools()` if you have a workspace available. /// `register_builtin_tools()` if you have a workspace available.
pub fn register_memory_tools(&self, workspace: Arc<Workspace>) { pub fn register_memory_tools(&self, workspace: Arc<Workspace>) {
self.register_sync(Arc::new(MemorySearchTool::new(Arc::clone(&workspace)))); self.register_sync(Arc::new(MemorySearchTool::from_workspace(Arc::clone(
self.register_sync(Arc::new(MemoryWriteTool::new(Arc::clone(&workspace)))); &workspace,
self.register_sync(Arc::new(MemoryReadTool::new(Arc::clone(&workspace)))); ))));
self.register_sync(Arc::new(MemoryTreeTool::new(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"); tracing::debug!("Registered 4 memory tools");
} }
@@ -361,7 +383,11 @@ impl ToolRegistry {
job_manager: Option<Arc<ContainerJobManager>>, job_manager: Option<Arc<ContainerJobManager>>,
store: Option<Arc<dyn Database>>, store: Option<Arc<dyn Database>>,
job_event_tx: Option< 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>>, inject_tx: Option<tokio::sync::mpsc::Sender<crate::channels::IncomingMessage>>,
prompt_queue: Option<PromptQueue>, prompt_queue: Option<PromptQueue>,
+1 -1
View File
@@ -661,7 +661,7 @@ mod advanced {
.await .await
.expect("failed to inject test token"); .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!( assert!(
activate_result.is_ok(), activate_result.is_ok(),
"activation failed: {:?}", "activation failed: {:?}",
+1 -1
View File
@@ -216,7 +216,7 @@ async fn extension_manager_with_process_manager_constructs() {
); );
// Verify the manager is functional — list returns Ok. // 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.is_ok(), "list should succeed on empty manager");
assert!(result.unwrap().is_empty()); 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>) { ) -> (SocketAddr, Arc<GatewayState>) {
let state = Arc::new(GatewayState { let state = Arc::new(GatewayState {
msg_tx: tokio::sync::RwLock::new(None), msg_tx: tokio::sync::RwLock::new(None),
sse: SseManager::new(), sse: Arc::new(SseManager::new()),
workspace: None, workspace: None,
workspace_pool: None,
session_manager: None, session_manager: None,
log_broadcaster: None, log_broadcaster: None,
log_level_handle: None, log_level_handle: None,
@@ -202,13 +203,13 @@ async fn start_test_server_with_provider(
job_manager: None, job_manager: None,
prompt_queue: None, prompt_queue: None,
scheduler: None, scheduler: None,
user_id: "test-user".to_string(), default_user_id: "test-user".to_string(),
shutdown_tx: tokio::sync::RwLock::new(None), shutdown_tx: tokio::sync::RwLock::new(None),
ws_tracker: Some(Arc::new(WsConnectionTracker::new())), ws_tracker: Some(Arc::new(WsConnectionTracker::new())),
llm_provider: Some(llm_provider), llm_provider: Some(llm_provider),
skill_registry: None, skill_registry: None,
skill_catalog: 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), oauth_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(10, 60),
webhook_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(), registry_entries: Vec::new(),
@@ -218,8 +219,12 @@ async fn start_test_server_with_provider(
active_config: ironclaw::channels::web::server::ActiveConfigSnapshot::default(), 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 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 .await
.expect("Failed to start test server"); .expect("Failed to start test server");
@@ -684,8 +689,9 @@ async fn test_no_llm_provider_returns_503() {
// Create state WITHOUT llm_provider // Create state WITHOUT llm_provider
let state = Arc::new(GatewayState { let state = Arc::new(GatewayState {
msg_tx: tokio::sync::RwLock::new(None), msg_tx: tokio::sync::RwLock::new(None),
sse: SseManager::new(), sse: Arc::new(SseManager::new()),
workspace: None, workspace: None,
workspace_pool: None,
session_manager: None, session_manager: None,
log_broadcaster: None, log_broadcaster: None,
log_level_handle: None, log_level_handle: None,
@@ -695,13 +701,13 @@ async fn test_no_llm_provider_returns_503() {
job_manager: None, job_manager: None,
prompt_queue: None, prompt_queue: None,
scheduler: None, scheduler: None,
user_id: "test-user".to_string(), default_user_id: "test-user".to_string(),
shutdown_tx: tokio::sync::RwLock::new(None), shutdown_tx: tokio::sync::RwLock::new(None),
ws_tracker: Some(Arc::new(WsConnectionTracker::new())), ws_tracker: Some(Arc::new(WsConnectionTracker::new())),
llm_provider: None, // No LLM! llm_provider: None, // No LLM!
skill_registry: None, skill_registry: None,
skill_catalog: 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), oauth_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(10, 60),
webhook_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(), registry_entries: Vec::new(),
@@ -711,10 +717,12 @@ async fn test_no_llm_provider_returns_503() {
active_config: ironclaw::channels::web::server::ActiveConfigSnapshot::default(), 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 addr: SocketAddr = "127.0.0.1:0".parse().unwrap();
let bound_addr = start_server(addr, state, AUTH_TOKEN.to_string()) let bound_addr = start_server(addr, state, auth).await.unwrap();
.await
.unwrap();
let url = format!("http://{}/v1/chat/completions", bound_addr); let url = format!("http://{}/v1/chat/completions", bound_addr);
let resp = client() let resp = client()
@@ -741,9 +749,10 @@ async fn test_chat_completions_body_too_large() {
let state = ironclaw::channels::web::test_helpers::TestGatewayBuilder::new() let state = ironclaw::channels::web::test_helpers::TestGatewayBuilder::new()
.llm_provider(llm_provider) .llm_provider(llm_provider)
.build(); .build();
let auth_state = ironclaw::channels::web::auth::AuthState { let auth_state = ironclaw::channels::web::auth::MultiAuthState::single(
token: AUTH_TOKEN.to_string(), AUTH_TOKEN.to_string(),
}; "test-user".to_string(),
);
let app = Router::new() let app = Router::new()
.route( .route(
+11 -6
View File
@@ -13,8 +13,11 @@ use ironclaw::agent::routine_engine::RoutineEngine;
use ironclaw::agent::{Agent, AgentDeps, SessionManager as AgentSessionManager}; use ironclaw::agent::{Agent, AgentDeps, SessionManager as AgentSessionManager};
use ironclaw::app::{AppBuilder, AppBuilderFlags}; use ironclaw::app::{AppBuilder, AppBuilderFlags};
use ironclaw::channels::IncomingMessage; use ironclaw::channels::IncomingMessage;
use ironclaw::channels::web::auth::MultiAuthState;
use ironclaw::channels::web::log_layer::LogBroadcaster; 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::sse::SseManager;
use ironclaw::channels::web::ws::WsConnectionTracker; use ironclaw::channels::web::ws::WsConnectionTracker;
use ironclaw::config::{Config, RegistryProviderConfig, RoutineConfig}; use ironclaw::config::{Config, RegistryProviderConfig, RoutineConfig};
@@ -211,8 +214,9 @@ impl GatewayWorkflowHarness {
let gateway_state = Arc::new(GatewayState { let gateway_state = Arc::new(GatewayState {
msg_tx: tokio::sync::RwLock::new(Some(gw_tx)), msg_tx: tokio::sync::RwLock::new(Some(gw_tx)),
sse: SseManager::new(), sse: Arc::new(SseManager::new()),
workspace: components.workspace.clone(), workspace: components.workspace.clone(),
workspace_pool: None,
session_manager: Some(Arc::clone(&agent_session_manager)), session_manager: Some(Arc::clone(&agent_session_manager)),
log_broadcaster: None, log_broadcaster: None,
log_level_handle: None, log_level_handle: None,
@@ -222,13 +226,13 @@ impl GatewayWorkflowHarness {
job_manager: None, job_manager: None,
prompt_queue: None, prompt_queue: None,
scheduler: Some(scheduler_slot.clone()), scheduler: Some(scheduler_slot.clone()),
user_id: user_id.clone(), default_user_id: user_id.clone(),
shutdown_tx: tokio::sync::RwLock::new(None), shutdown_tx: tokio::sync::RwLock::new(None),
ws_tracker: Some(Arc::new(WsConnectionTracker::new())), ws_tracker: Some(Arc::new(WsConnectionTracker::new())),
llm_provider: Some(Arc::clone(&components.llm)), llm_provider: Some(Arc::clone(&components.llm)),
skill_registry: components.skill_registry.clone(), skill_registry: components.skill_registry.clone(),
skill_catalog: components.skill_catalog.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), oauth_rate_limiter: RateLimiter::new(10, 60),
webhook_rate_limiter: RateLimiter::new(10, 60), webhook_rate_limiter: RateLimiter::new(10, 60),
registry_entries: Vec::new(), registry_entries: Vec::new(),
@@ -254,7 +258,7 @@ impl GatewayWorkflowHarness {
skills_config: components.config.skills.clone(), skills_config: components.config.skills.clone(),
hooks: components.hooks, hooks: components.hooks,
cost_guard: components.cost_guard, cost_guard: components.cost_guard,
sse_tx: Some(gateway_state.sse.sender()), sse_tx: None,
http_interceptor: None, http_interceptor: None,
transcription: None, transcription: None,
document_extraction: None, document_extraction: None,
@@ -288,10 +292,11 @@ impl GatewayWorkflowHarness {
} }
let auth_token = "gateway-test-token".to_string(); let auth_token = "gateway-test-token".to_string();
let auth = MultiAuthState::single(auth_token.clone(), user_id.clone());
let addr = start_server( let addr = start_server(
"127.0.0.1:0".parse().expect("valid localhost addr"), "127.0.0.1:0".parse().expect("valid localhost addr"),
Arc::clone(&gateway_state), Arc::clone(&gateway_state),
auth_token.clone(), auth,
) )
.await .await
.expect("failed to start gateway server"); .expect("failed to start gateway server");
+9 -4
View File
@@ -39,8 +39,9 @@ async fn start_test_server() -> (
let state = Arc::new(GatewayState { let state = Arc::new(GatewayState {
msg_tx: tokio::sync::RwLock::new(Some(agent_tx)), msg_tx: tokio::sync::RwLock::new(Some(agent_tx)),
sse: SseManager::new(), sse: Arc::new(SseManager::new()),
workspace: None, workspace: None,
workspace_pool: None,
session_manager: None, session_manager: None,
log_broadcaster: None, log_broadcaster: None,
log_level_handle: None, log_level_handle: None,
@@ -50,13 +51,13 @@ async fn start_test_server() -> (
job_manager: None, job_manager: None,
prompt_queue: None, prompt_queue: None,
scheduler: None, scheduler: None,
user_id: "test-user".to_string(), default_user_id: "test-user".to_string(),
shutdown_tx: tokio::sync::RwLock::new(None), shutdown_tx: tokio::sync::RwLock::new(None),
ws_tracker: Some(Arc::new(WsConnectionTracker::new())), ws_tracker: Some(Arc::new(WsConnectionTracker::new())),
llm_provider: None, llm_provider: None,
skill_registry: None, skill_registry: None,
skill_catalog: 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), oauth_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(10, 60),
webhook_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(), registry_entries: Vec::new(),
@@ -66,8 +67,12 @@ async fn start_test_server() -> (
active_config: ironclaw::channels::web::server::ActiveConfigSnapshot::default(), 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 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 .await
.expect("Failed to start test server"); .expect("Failed to start test server");