diff --git a/src/agent/job_monitor.rs b/src/agent/job_monitor.rs index 675d0426..8d302c1d 100644 --- a/src/agent/job_monitor.rs +++ b/src/agent/job_monitor.rs @@ -44,7 +44,7 @@ pub struct JobMonitorRoute { /// the main agent's context window). pub fn spawn_job_monitor( job_id: Uuid, - event_rx: broadcast::Receiver<(Uuid, SseEvent)>, + event_rx: broadcast::Receiver<(Uuid, String, SseEvent)>, inject_tx: mpsc::Sender, route: JobMonitorRoute, ) -> JoinHandle<()> { @@ -56,7 +56,7 @@ pub fn spawn_job_monitor( /// jobs don't stay `InProgress` forever in the `ContextManager`. pub fn spawn_job_monitor_with_context( job_id: Uuid, - mut event_rx: broadcast::Receiver<(Uuid, SseEvent)>, + mut event_rx: broadcast::Receiver<(Uuid, String, SseEvent)>, inject_tx: mpsc::Sender, route: JobMonitorRoute, context_manager: Option>, @@ -68,7 +68,7 @@ pub fn spawn_job_monitor_with_context( loop { match event_rx.recv().await { - Ok((ev_job_id, event)) => { + Ok((ev_job_id, _user_id, event)) => { if ev_job_id != job_id { continue; } @@ -162,7 +162,7 @@ pub fn spawn_job_monitor_with_context( /// inject messages into) but we still need to free the `max_jobs` slot. pub fn spawn_completion_watcher( job_id: Uuid, - mut event_rx: broadcast::Receiver<(Uuid, SseEvent)>, + mut event_rx: broadcast::Receiver<(Uuid, String, SseEvent)>, context_manager: Arc, ) -> JoinHandle<()> { let short_id = job_id.to_string()[..8].to_string(); @@ -170,7 +170,7 @@ pub fn spawn_completion_watcher( tokio::spawn(async move { loop { match event_rx.recv().await { - Ok((ev_job_id, SseEvent::JobResult { status, .. })) if ev_job_id == job_id => { + Ok((ev_job_id, _user_id, SseEvent::JobResult { status, .. })) if ev_job_id == job_id => { let target = if status == "completed" { JobState::Completed } else { @@ -227,7 +227,7 @@ mod tests { #[tokio::test] async fn test_monitor_forwards_assistant_messages() { - let (event_tx, _) = broadcast::channel::<(Uuid, SseEvent)>(16); + let (event_tx, _) = broadcast::channel::<(Uuid, String, SseEvent)>(16); let (inject_tx, mut inject_rx) = mpsc::channel::(16); let job_id = Uuid::new_v4(); @@ -237,6 +237,7 @@ mod tests { event_tx .send(( job_id, + "test-user".to_string(), SseEvent::JobMessage { job_id: job_id.to_string(), role: "assistant".to_string(), @@ -259,7 +260,7 @@ mod tests { #[tokio::test] async fn test_monitor_ignores_other_jobs() { - let (event_tx, _) = broadcast::channel::<(Uuid, SseEvent)>(16); + let (event_tx, _) = broadcast::channel::<(Uuid, String, SseEvent)>(16); let (inject_tx, mut inject_rx) = mpsc::channel::(16); let job_id = Uuid::new_v4(); @@ -270,6 +271,7 @@ mod tests { event_tx .send(( other_job_id, + "test-user".to_string(), SseEvent::JobMessage { job_id: other_job_id.to_string(), role: "assistant".to_string(), @@ -289,7 +291,7 @@ mod tests { #[tokio::test] async fn test_monitor_exits_on_job_result() { - let (event_tx, _) = broadcast::channel::<(Uuid, SseEvent)>(16); + let (event_tx, _) = broadcast::channel::<(Uuid, String, SseEvent)>(16); let (inject_tx, mut inject_rx) = mpsc::channel::(16); let job_id = Uuid::new_v4(); @@ -299,6 +301,7 @@ mod tests { event_tx .send(( job_id, + "test-user".to_string(), SseEvent::JobResult { job_id: job_id.to_string(), status: "completed".to_string(), @@ -324,7 +327,7 @@ mod tests { #[tokio::test] async fn test_monitor_skips_tool_events() { - let (event_tx, _) = broadcast::channel::<(Uuid, SseEvent)>(16); + let (event_tx, _) = broadcast::channel::<(Uuid, String, SseEvent)>(16); let (inject_tx, mut inject_rx) = mpsc::channel::(16); let job_id = Uuid::new_v4(); @@ -334,6 +337,7 @@ mod tests { event_tx .send(( job_id, + "test-user".to_string(), SseEvent::JobToolUse { job_id: job_id.to_string(), tool_name: "shell".to_string(), @@ -346,6 +350,7 @@ mod tests { event_tx .send(( job_id, + "test-user".to_string(), SseEvent::JobMessage { job_id: job_id.to_string(), role: "user".to_string(), @@ -402,7 +407,7 @@ mod tests { .await .unwrap(); - let (event_tx, _) = broadcast::channel::<(Uuid, SseEvent)>(16); + let (event_tx, _) = broadcast::channel::<(Uuid, String, SseEvent)>(16); let (inject_tx, mut inject_rx) = mpsc::channel::(16); let handle = spawn_job_monitor_with_context( @@ -417,6 +422,7 @@ mod tests { event_tx .send(( job_id, + "test-user".to_string(), SseEvent::JobResult { job_id: job_id.to_string(), status: "completed".to_string(), @@ -450,7 +456,7 @@ mod tests { .await .unwrap(); - let (event_tx, _) = broadcast::channel::<(Uuid, SseEvent)>(16); + let (event_tx, _) = broadcast::channel::<(Uuid, String, SseEvent)>(16); let (inject_tx, mut inject_rx) = mpsc::channel::(16); let handle = spawn_job_monitor_with_context( @@ -465,6 +471,7 @@ mod tests { event_tx .send(( job_id, + "test-user".to_string(), SseEvent::JobResult { job_id: job_id.to_string(), status: "failed".to_string(), @@ -498,12 +505,13 @@ mod tests { .await .unwrap(); - let (event_tx, _) = broadcast::channel::<(Uuid, SseEvent)>(16); + let (event_tx, _) = broadcast::channel::<(Uuid, String, SseEvent)>(16); let handle = spawn_completion_watcher(job_id, event_tx.subscribe(), Arc::clone(&cm)); event_tx .send(( job_id, + "test-user".to_string(), SseEvent::JobResult { job_id: job_id.to_string(), status: "completed".to_string(), diff --git a/src/agent/thread_ops.rs b/src/agent/thread_ops.rs index eec29099..ddfd0c0f 100644 --- a/src/agent/thread_ops.rs +++ b/src/agent/thread_ops.rs @@ -1646,7 +1646,7 @@ impl Agent { }; match ext_mgr - .configure_token(&pending.extension_name, token) + .configure_token(&pending.extension_name, token, &message.user_id) .await { Ok(result) if result.activated => { diff --git a/src/channels/web/auth.rs b/src/channels/web/auth.rs index b2fa4e4f..49688bf8 100644 --- a/src/channels/web/auth.rs +++ b/src/channels/web/auth.rs @@ -1,17 +1,107 @@ //! Bearer token authentication middleware for the web gateway. +//! +//! Supports multi-user mode: each token maps to a `UserIdentity` that carries +//! the user_id. The identity is inserted into request extensions so downstream +//! handlers can extract it via `AuthenticatedUser`. + +use std::collections::HashMap; use axum::{ - extract::{Request, State}, - http::{HeaderMap, Method, StatusCode}, + extract::{FromRequestParts, Request, State}, + http::{HeaderMap, Method, StatusCode, request::Parts}, middleware::Next, response::{IntoResponse, Response}, }; use 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, +} + +/// Multi-user auth state: maps tokens to user identities. +/// +/// In single-user mode (the default), contains exactly one entry. #[derive(Clone)] -pub struct AuthState { - pub token: String, +pub struct MultiAuthState { + tokens: HashMap, +} + +impl MultiAuthState { + /// Create a single-user auth state (backwards compatible). + pub fn single(token: String, user_id: String) -> Self { + let mut tokens = HashMap::new(); + tokens.insert( + token, + UserIdentity { + user_id, + workspace_read_scopes: Vec::new(), + }, + ); + Self { tokens } + } + + /// Create a multi-user auth state from a map of tokens to identities. + pub fn multi(tokens: HashMap) -> Self { + Self { tokens } + } + + /// Authenticate a token, returning the associated identity if valid. + /// + /// Uses constant-time comparison (`subtle::ConstantTimeEq`) to prevent + /// timing side-channels that could leak token information. 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_bytes = candidate.as_bytes(); + let mut matched: Option<&UserIdentity> = None; + for (token, identity) in &self.tokens { + let token_bytes = token.as_bytes(); + // ct_eq requires equal lengths; pad comparison to avoid length leak + if candidate_bytes.len() == token_bytes.len() + && bool::from(candidate_bytes.ct_eq(token_bytes)) + { + matched = Some(identity); + } + } + matched + } + + /// Get the first token (for backwards-compatible printing at startup). + pub fn first_token(&self) -> Option<&str> { + self.tokens.keys().next().map(|s| s.as_str()) + } + + /// Get the first user identity (for single-user fallback). + pub fn first_identity(&self) -> Option<&UserIdentity> { + self.tokens.values().next() + } +} + +/// 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 FromRequestParts for AuthenticatedUser +where + S: Send + Sync, +{ + type Rejection = (StatusCode, &'static str); + + async fn from_request_parts(parts: &mut Parts, _state: &S) -> Result { + parts + .extensions + .get::() + .cloned() + .map(AuthenticatedUser) + .ok_or((StatusCode::UNAUTHORIZED, "Not authenticated")) + } } /// Whether query-string token auth is allowed for this request. @@ -51,47 +141,101 @@ fn query_token(request: &Request) -> Option { /// Auth middleware that validates bearer token from header or query param. /// /// SSE connections can't set headers from `EventSource`, so we also accept -/// `?token=xxx` as a query parameter, but only on SSE endpoints. +/// `?token=xxx` as a query parameter, but only on SSE/WS endpoints. +/// +/// On successful authentication, inserts the matching `UserIdentity` into +/// request extensions for downstream extraction via `AuthenticatedUser`. pub async fn auth_middleware( - State(auth): State, + State(auth): State, headers: HeaderMap, - request: Request, + mut request: Request, next: Next, ) -> Response { - // Try Authorization header first (constant-time comparison). + // Try Authorization header first. // RFC 6750 Section 2.1: auth-scheme comparison is case-insensitive. if let Some(auth_header) = headers.get("authorization") && let Ok(value) = auth_header.to_str() && value.len() > 7 && value[..7].eq_ignore_ascii_case("Bearer ") - && bool::from(value.as_bytes()[7..].ct_eq(auth.token.as_bytes())) + && let Some(identity) = auth.authenticate(&value[7..]) { + request.extensions_mut().insert(identity.clone()); return next.run(request).await; } - // Fall back to query parameter, but only for SSE endpoints (constant-time comparison). + // Fall back to query parameter, but only for SSE/WS endpoints. if allows_query_token_auth(&request) && let Some(token) = query_token(&request) - && bool::from(token.as_bytes().ct_eq(auth.token.as_bytes())) + && let Some(identity) = auth.authenticate(&token) { + request.extensions_mut().insert(identity.clone()); return next.run(request).await; } (StatusCode::UNAUTHORIZED, "Invalid or missing auth token").into_response() } +// Keep the old type as an alias for any external references during migration. +pub type AuthState = MultiAuthState; + #[cfg(test)] mod tests { use super::*; - use crate::testing::credentials::{TEST_AUTH_SECRET_TOKEN, TEST_BEARER_TOKEN}; + use crate::testing::credentials::TEST_AUTH_SECRET_TOKEN; #[test] - fn test_auth_state_clone() { - let state = AuthState { - token: TEST_BEARER_TOKEN.to_string(), - }; - let cloned = state.clone(); - assert_eq!(cloned.token, TEST_BEARER_TOKEN); + fn test_multi_auth_state_single() { + let state = MultiAuthState::single("tok-123".to_string(), "alice".to_string()); + let identity = state.authenticate("tok-123"); + assert!(identity.is_some()); + assert_eq!(identity.unwrap().user_id, "alice"); + } + + #[test] + fn test_multi_auth_state_reject_wrong_token() { + let state = MultiAuthState::single("tok-123".to_string(), "alice".to_string()); + assert!(state.authenticate("wrong-token").is_none()); + } + + #[test] + fn test_multi_auth_state_multi_users() { + let mut tokens = HashMap::new(); + tokens.insert( + "tok-alice".to_string(), + UserIdentity { + user_id: "alice".to_string(), + workspace_read_scopes: Vec::new(), + }, + ); + tokens.insert( + "tok-bob".to_string(), + UserIdentity { + user_id: "bob".to_string(), + workspace_read_scopes: Vec::new(), + }, + ); + let state = MultiAuthState::multi(tokens); + + let alice = state.authenticate("tok-alice").unwrap(); + assert_eq!(alice.user_id, "alice"); + + let bob = state.authenticate("tok-bob").unwrap(); + assert_eq!(bob.user_id, "bob"); + + assert!(state.authenticate("tok-charlie").is_none()); + } + + #[test] + fn test_multi_auth_state_first_token() { + let state = MultiAuthState::single("my-token".to_string(), "user1".to_string()); + assert_eq!(state.first_token(), Some("my-token")); + } + + #[test] + fn test_multi_auth_state_first_identity() { + let state = MultiAuthState::single("my-token".to_string(), "user1".to_string()); + let identity = state.first_identity().unwrap(); + assert_eq!(identity.user_id, "user1"); } use axum::Router; @@ -107,9 +251,7 @@ mod tests { /// Router with streaming endpoints (query auth allowed) and regular /// endpoints (query auth rejected). fn test_app(token: &str) -> Router { - let state = AuthState { - token: token.to_string(), - }; + let state = MultiAuthState::single(token.to_string(), "test-user".to_string()); Router::new() .route("/api/chat/events", get(dummy_handler)) .route("/api/logs/events", get(dummy_handler)) @@ -306,4 +448,200 @@ mod tests { let resp = app.oneshot(req).await.unwrap(); assert_eq!(resp.status(), StatusCode::UNAUTHORIZED); } + + // --- Multi-tenant auth integration tests --- + + /// Handler that extracts `AuthenticatedUser` and returns the resolved user_id. + async fn identity_handler(AuthenticatedUser(identity): AuthenticatedUser) -> String { + identity.user_id + } + + /// Handler that extracts `AuthenticatedUser` and returns workspace_read_scopes as JSON. + async fn scopes_handler(AuthenticatedUser(identity): AuthenticatedUser) -> String { + serde_json::to_string(&identity.workspace_read_scopes).unwrap() + } + + /// Build a multi-user router where each token maps to a distinct identity. + fn multi_user_app(tokens: HashMap) -> 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 { + 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 = 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 = 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 = 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()); + } } diff --git a/src/channels/web/handlers/chat.rs b/src/channels/web/handlers/chat.rs index 5cb2b9ea..640ee260 100644 --- a/src/channels/web/handlers/chat.rs +++ b/src/channels/web/handlers/chat.rs @@ -12,22 +12,24 @@ use serde::Deserialize; use uuid::Uuid; use crate::channels::IncomingMessage; +use crate::channels::web::auth::AuthenticatedUser; use crate::channels::web::server::GatewayState; use crate::channels::web::types::*; use crate::channels::web::util::{build_turns_from_db_messages, truncate_preview}; pub async fn chat_send_handler( State(state): State>, + AuthenticatedUser(identity): AuthenticatedUser, Json(req): Json, ) -> Result<(StatusCode, Json), (StatusCode, String)> { - if !state.chat_rate_limiter.check() { + if !state.chat_rate_limiter.check(&identity.user_id) { return Err(( StatusCode::TOO_MANY_REQUESTS, "Rate limit exceeded. Try again shortly.".to_string(), )); } - let mut msg = IncomingMessage::new("gateway", &state.user_id, &req.content); + let mut msg = IncomingMessage::new("gateway", &identity.user_id, &req.content); if let Some(ref thread_id) = req.thread_id { msg = msg.with_thread(thread_id); @@ -74,6 +76,7 @@ pub async fn chat_send_handler( pub async fn chat_approval_handler( State(state): State>, + AuthenticatedUser(identity): AuthenticatedUser, Json(req): Json, ) -> Result<(StatusCode, Json), (StatusCode, String)> { let (approved, always) = match req.action.as_str() { @@ -109,7 +112,7 @@ pub async fn chat_approval_handler( ) })?; - let mut msg = IncomingMessage::new("gateway", &state.user_id, content); + let mut msg = IncomingMessage::new("gateway", &identity.user_id, content); if let Some(ref thread_id) = req.thread_id { msg = msg.with_thread(thread_id); @@ -150,6 +153,7 @@ pub async fn chat_approval_handler( /// The token never touches the LLM, chat history, or SSE stream. pub async fn chat_auth_token_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Json(req): Json, ) -> Result, (StatusCode, String)> { let ext_mgr = state.extension_manager.as_ref().ok_or(( @@ -158,7 +162,7 @@ pub async fn chat_auth_token_handler( ))?; match ext_mgr - .configure_token(&req.extension_name, &req.token) + .configure_token(&req.extension_name, &req.token, &user.user_id) .await { Ok(result) => { @@ -176,7 +180,7 @@ pub async fn chat_auth_token_handler( setup_url: None, }); } else { - clear_auth_mode(&state).await; + clear_auth_mode(&state, &user.user_id).await; state.sse.broadcast(SseEvent::AuthCompleted { extension_name: req.extension_name.clone(), @@ -205,16 +209,17 @@ pub async fn chat_auth_token_handler( /// Cancel an in-progress auth flow. pub async fn chat_auth_cancel_handler( State(state): State>, + AuthenticatedUser(identity): AuthenticatedUser, Json(_req): Json, ) -> Result, (StatusCode, String)> { - clear_auth_mode(&state).await; + clear_auth_mode(&state, &identity.user_id).await; Ok(Json(ActionResponse::ok("Auth cancelled"))) } /// Clear pending auth mode on the active thread. -pub async fn clear_auth_mode(state: &GatewayState) { +pub async fn clear_auth_mode(state: &GatewayState, user_id: &str) { if let Some(ref sm) = state.session_manager { - let session = sm.get_or_create_session(&state.user_id).await; + let session = sm.get_or_create_session(user_id).await; let mut sess = session.lock().await; if let Some(thread_id) = sess.active_thread && let Some(thread) = sess.threads.get_mut(&thread_id) @@ -226,8 +231,9 @@ pub async fn clear_auth_mode(state: &GatewayState) { pub async fn chat_events_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, ) -> Result { - state.sse.subscribe().ok_or(( + state.sse.subscribe(Some(user.user_id)).ok_or(( StatusCode::SERVICE_UNAVAILABLE, "Too many connections".to_string(), )) @@ -237,6 +243,7 @@ pub async fn chat_ws_handler( headers: axum::http::HeaderMap, ws: WebSocketUpgrade, State(state): State>, + AuthenticatedUser(identity): AuthenticatedUser, ) -> Result { // Validate Origin header to prevent cross-site WebSocket hijacking. let origin = headers @@ -262,7 +269,9 @@ pub async fn chat_ws_handler( "WebSocket origin not allowed".to_string(), )); } - Ok(ws.on_upgrade(move |socket| crate::channels::web::ws::handle_ws_connection(socket, state))) + Ok(ws.on_upgrade(move |socket| { + crate::channels::web::ws::handle_ws_connection(socket, state, identity) + })) } #[derive(Deserialize)] @@ -274,6 +283,7 @@ pub struct HistoryQuery { pub async fn chat_history_handler( State(state): State>, + AuthenticatedUser(identity): AuthenticatedUser, Query(query): Query, ) -> Result, (StatusCode, String)> { let session_manager = state.session_manager.as_ref().ok_or(( @@ -281,7 +291,7 @@ pub async fn chat_history_handler( "Session manager not available".to_string(), ))?; - let session = session_manager.get_or_create_session(&state.user_id).await; + let session = session_manager.get_or_create_session(&identity.user_id).await; let limit = query.limit.unwrap_or(50); let before_cursor = query @@ -314,7 +324,7 @@ pub async fn chat_history_handler( && let Some(ref store) = state.store { let owned = store - .conversation_belongs_to_user(thread_id, &state.user_id) + .conversation_belongs_to_user(thread_id, &identity.user_id) .await .unwrap_or(false); if !owned { @@ -434,24 +444,25 @@ pub async fn chat_history_handler( pub async fn chat_threads_handler( State(state): State>, + AuthenticatedUser(identity): AuthenticatedUser, ) -> Result, (StatusCode, String)> { let session_manager = state.session_manager.as_ref().ok_or(( StatusCode::SERVICE_UNAVAILABLE, "Session manager not available".to_string(), ))?; - let session = session_manager.get_or_create_session(&state.user_id).await; + let session = session_manager.get_or_create_session(&identity.user_id).await; // Try DB first for persistent thread list if let Some(ref store) = state.store { // Auto-create assistant thread if it doesn't exist let assistant_id = store - .get_or_create_assistant_conversation(&state.user_id, "gateway") + .get_or_create_assistant_conversation(&identity.user_id, "gateway") .await .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; if let Ok(summaries) = store - .list_conversations_all_channels(&state.user_id, 50) + .list_conversations_all_channels(&identity.user_id, 50) .await { let mut assistant_thread = None; @@ -534,13 +545,14 @@ pub async fn chat_threads_handler( pub async fn chat_new_thread_handler( State(state): State>, + AuthenticatedUser(identity): AuthenticatedUser, ) -> Result, (StatusCode, String)> { let session_manager = state.session_manager.as_ref().ok_or(( StatusCode::SERVICE_UNAVAILABLE, "Session manager not available".to_string(), ))?; - let session = session_manager.get_or_create_session(&state.user_id).await; + let session = session_manager.get_or_create_session(&identity.user_id).await; let (thread_id, info) = { let mut sess = session.lock().await; let thread = sess.create_thread(); @@ -562,12 +574,12 @@ pub async fn chat_new_thread_handler( // so that the subsequent loadThreads() call from the frontend sees it. if let Some(ref store) = state.store { match store - .ensure_conversation(thread_id, "gateway", &state.user_id, None) + .ensure_conversation(thread_id, "gateway", &identity.user_id, None) .await { Ok(true) => {} Ok(false) => tracing::warn!( - user = %state.user_id, + user = %identity.user_id, thread_id = %thread_id, "Skipped persisting new thread due to ownership/channel conflict" ), diff --git a/src/channels/web/handlers/extensions.rs b/src/channels/web/handlers/extensions.rs index 855fba3e..f855aa43 100644 --- a/src/channels/web/handlers/extensions.rs +++ b/src/channels/web/handlers/extensions.rs @@ -8,11 +8,13 @@ use axum::{ http::StatusCode, }; +use crate::channels::web::auth::AuthenticatedUser; use crate::channels::web::server::GatewayState; use crate::channels::web::types::*; pub async fn extensions_list_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, ) -> Result, (StatusCode, String)> { let ext_mgr = state.extension_manager.as_ref().ok_or(( StatusCode::NOT_IMPLEMENTED, @@ -20,7 +22,7 @@ pub async fn extensions_list_handler( ))?; let installed = ext_mgr - .list(None, false) + .list(None, false, &user.user_id) .await .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; @@ -100,6 +102,7 @@ pub async fn extensions_tools_handler( pub async fn extensions_install_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Json(req): Json, ) -> Result, (StatusCode, String)> { let ext_mgr = state.extension_manager.as_ref().ok_or(( @@ -116,7 +119,7 @@ pub async fn extensions_install_handler( }); match ext_mgr - .install(&req.name, req.url.as_deref(), kind_hint) + .install(&req.name, req.url.as_deref(), kind_hint, &user.user_id) .await { Ok(result) => Ok(Json(ActionResponse::ok(result.message))), @@ -126,6 +129,7 @@ pub async fn extensions_install_handler( pub async fn extensions_remove_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Path(name): Path, ) -> Result, (StatusCode, String)> { let ext_mgr = state.extension_manager.as_ref().ok_or(( @@ -133,7 +137,7 @@ pub async fn extensions_remove_handler( "Extension manager not available (secrets store required)".to_string(), ))?; - match ext_mgr.remove(&name).await { + match ext_mgr.remove(&name, &user.user_id).await { Ok(message) => Ok(Json(ActionResponse::ok(message))), Err(e) => Ok(Json(ActionResponse::fail(e.to_string()))), } diff --git a/src/channels/web/handlers/jobs.rs b/src/channels/web/handlers/jobs.rs index 5a94e055..2e313c81 100644 --- a/src/channels/web/handlers/jobs.rs +++ b/src/channels/web/handlers/jobs.rs @@ -11,11 +11,13 @@ use axum::{ use serde::Deserialize; use uuid::Uuid; +use crate::channels::web::auth::AuthenticatedUser; use crate::channels::web::server::GatewayState; use crate::channels::web::types::*; pub async fn jobs_list_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, ) -> Result, (StatusCode, String)> { let store = state.store.as_ref().ok_or(( StatusCode::SERVICE_UNAVAILABLE, @@ -25,8 +27,8 @@ pub async fn jobs_list_handler( let mut jobs: Vec = Vec::new(); let mut seen_ids: HashSet = HashSet::new(); - // Fetch sandbox jobs from database. - match store.list_sandbox_jobs().await { + // Fetch sandbox jobs scoped to this user. + match store.list_sandbox_jobs_for_user(&user.user_id).await { Ok(sandbox_jobs) => { for j in &sandbox_jobs { let ui_state = match j.status.as_str() { @@ -54,6 +56,9 @@ pub async fn jobs_list_handler( match store.list_agent_jobs().await { Ok(agent_jobs) => { for j in &agent_jobs { + if j.user_id != user.user_id { + continue; + } if seen_ids.contains(&j.id) { continue; } @@ -80,6 +85,7 @@ pub async fn jobs_list_handler( pub async fn jobs_summary_handler( State(state): State>, + AuthenticatedUser(_user): AuthenticatedUser, ) -> Result, (StatusCode, String)> { let store = state.store.as_ref().ok_or(( StatusCode::SERVICE_UNAVAILABLE, @@ -134,6 +140,7 @@ pub async fn jobs_summary_handler( pub async fn jobs_detail_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Path(id): Path, ) -> Result, (StatusCode, String)> { let store = state.store.as_ref().ok_or(( @@ -145,169 +152,213 @@ pub async fn jobs_detail_handler( .map_err(|_| (StatusCode::BAD_REQUEST, "Invalid job ID".to_string()))?; // Try sandbox job from DB first. - if let Ok(Some(job)) = store.get_sandbox_job(job_id).await { - let browse_id = std::path::Path::new(&job.project_dir) - .file_name() - .map(|n| n.to_string_lossy().to_string()) - .unwrap_or_else(|| job.id.to_string()); + match store.get_sandbox_job(job_id).await { + Ok(Some(job)) => { + if job.user_id != user.user_id { + return Err((StatusCode::NOT_FOUND, "Job not found".to_string())); + } + let browse_id = std::path::Path::new(&job.project_dir) + .file_name() + .map(|n| n.to_string_lossy().to_string()) + .unwrap_or_else(|| job.id.to_string()); - let ui_state = match job.status.as_str() { - "creating" => "pending", - "running" => "in_progress", - s => s, - }; + let ui_state = match job.status.as_str() { + "creating" => "pending", + "running" => "in_progress", + s => s, + }; - let elapsed_secs = job.started_at.map(|start| { - let end = job.completed_at.unwrap_or_else(chrono::Utc::now); - (end - start).num_seconds().max(0) as u64 - }); - - // Synthesize transitions from timestamps. - let mut transitions = Vec::new(); - if let Some(started) = job.started_at { - transitions.push(TransitionInfo { - from: "creating".to_string(), - to: "running".to_string(), - timestamp: started.to_rfc3339(), - reason: None, + let elapsed_secs = job.started_at.map(|start| { + let end = job.completed_at.unwrap_or_else(chrono::Utc::now); + (end - start).num_seconds().max(0) as u64 }); - } - if let Some(completed) = job.completed_at { - transitions.push(TransitionInfo { - from: "running".to_string(), - to: job.status.clone(), - timestamp: completed.to_rfc3339(), - reason: job.failure_reason.clone(), - }); - } - let mode = store.get_sandbox_job_mode(job.id).await.ok().flatten(); - let is_claude_code = mode.as_deref() == Some("claude_code"); + // Synthesize transitions from timestamps. + let mut transitions = Vec::new(); + if let Some(started) = job.started_at { + transitions.push(TransitionInfo { + from: "creating".to_string(), + to: "running".to_string(), + timestamp: started.to_rfc3339(), + reason: None, + }); + } + if let Some(completed) = job.completed_at { + transitions.push(TransitionInfo { + from: "running".to_string(), + to: job.status.clone(), + timestamp: completed.to_rfc3339(), + reason: job.failure_reason.clone(), + }); + } - return Ok(Json(JobDetailResponse { - id: job.id, - title: job.task.clone(), - description: String::new(), - state: ui_state.to_string(), - user_id: job.user_id.clone(), - created_at: job.created_at.to_rfc3339(), - started_at: job.started_at.map(|dt| dt.to_rfc3339()), - completed_at: job.completed_at.map(|dt| dt.to_rfc3339()), - elapsed_secs, - project_dir: Some(job.project_dir.clone()), - browse_url: Some(format!("/projects/{}/", browse_id)), - job_mode: mode.filter(|m| m != "worker"), - transitions, - can_restart: state.job_manager.is_some(), - can_prompt: is_claude_code && state.prompt_queue.is_some(), - job_kind: Some("sandbox".to_string()), - })); + let mode = store.get_sandbox_job_mode(job.id).await.ok().flatten(); + let is_claude_code = mode.as_deref() == Some("claude_code"); + + return Ok(Json(JobDetailResponse { + id: job.id, + title: job.task.clone(), + description: String::new(), + state: ui_state.to_string(), + user_id: job.user_id.clone(), + created_at: job.created_at.to_rfc3339(), + started_at: job.started_at.map(|dt| dt.to_rfc3339()), + completed_at: job.completed_at.map(|dt| dt.to_rfc3339()), + elapsed_secs, + project_dir: Some(job.project_dir.clone()), + browse_url: Some(format!("/projects/{}/", browse_id)), + job_mode: mode.filter(|m| m != "worker"), + transitions, + can_restart: state.job_manager.is_some(), + can_prompt: is_claude_code && state.prompt_queue.is_some(), + job_kind: Some("sandbox".to_string()), + })); + } + Ok(None) => {} + Err(e) => { + return Err(( + StatusCode::INTERNAL_SERVER_ERROR, + format!("Database error: {}", e), + )); + } } // Fall back to agent job from DB. - if let Ok(Some(ctx)) = store.get_job(job_id).await { - let elapsed_secs = ctx.started_at.map(|start| { - let end = ctx.completed_at.unwrap_or_else(chrono::Utc::now); - (end - start).num_seconds().max(0) as u64 - }); + match store.get_job(job_id).await { + Ok(Some(ctx)) => { + if ctx.user_id != user.user_id { + return Err((StatusCode::NOT_FOUND, "Job not found".to_string())); + } + let elapsed_secs = ctx.started_at.map(|start| { + let end = ctx.completed_at.unwrap_or_else(chrono::Utc::now); + (end - start).num_seconds().max(0) as u64 + }); - // Only show prompt bar for jobs that have a running worker (Pending/InProgress). - // Stuck jobs have no active worker loop, so messages would be silently dropped. - let is_promptable = matches!( - ctx.state, - crate::context::JobState::Pending | crate::context::JobState::InProgress - ); - return Ok(Json(JobDetailResponse { - id: ctx.job_id, - title: ctx.title.clone(), - description: ctx.description.clone(), - state: ctx.state.to_string(), - user_id: ctx.user_id.clone(), - created_at: ctx.created_at.to_rfc3339(), - started_at: ctx.started_at.map(|dt| dt.to_rfc3339()), - completed_at: ctx.completed_at.map(|dt| dt.to_rfc3339()), - elapsed_secs, - project_dir: None, - browse_url: None, - job_mode: None, - transitions: Vec::new(), - can_restart: state.scheduler.is_some(), - can_prompt: is_promptable && state.scheduler.is_some(), - job_kind: Some("agent".to_string()), - })); + // Only show prompt bar for jobs that have a running worker (Pending/InProgress). + // Stuck jobs have no active worker loop, so messages would be silently dropped. + let is_promptable = matches!( + ctx.state, + crate::context::JobState::Pending | crate::context::JobState::InProgress + ); + Ok(Json(JobDetailResponse { + id: ctx.job_id, + title: ctx.title.clone(), + description: ctx.description.clone(), + state: ctx.state.to_string(), + user_id: ctx.user_id.clone(), + created_at: ctx.created_at.to_rfc3339(), + started_at: ctx.started_at.map(|dt| dt.to_rfc3339()), + completed_at: ctx.completed_at.map(|dt| dt.to_rfc3339()), + elapsed_secs, + project_dir: None, + browse_url: None, + job_mode: None, + transitions: Vec::new(), + can_restart: state.scheduler.is_some(), + can_prompt: is_promptable && state.scheduler.is_some(), + job_kind: Some("agent".to_string()), + })) + } + Ok(None) => Err((StatusCode::NOT_FOUND, "Job not found".to_string())), + Err(e) => Err(( + StatusCode::INTERNAL_SERVER_ERROR, + format!("Database error: {}", e), + )), } - - Err((StatusCode::NOT_FOUND, "Job not found".to_string())) } pub async fn jobs_cancel_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Path(id): Path, ) -> Result, (StatusCode, String)> { let job_id = Uuid::parse_str(&id) .map_err(|_| (StatusCode::BAD_REQUEST, "Invalid job ID".to_string()))?; // Try sandbox job cancellation. - if let Some(ref store) = state.store - && let Ok(Some(job)) = store.get_sandbox_job(job_id).await - { - if job.status == "running" || job.status == "creating" { - // Stop the container if we have a job manager. - if let Some(ref jm) = state.job_manager - && let Err(e) = jm.stop_job(job_id).await - { - tracing::warn!(job_id = %job_id, error = %e, "Failed to stop container during cancellation"); + if let Some(ref store) = state.store { + match store.get_sandbox_job(job_id).await { + Ok(Some(job)) => { + if job.user_id != user.user_id { + return Err((StatusCode::NOT_FOUND, "Job not found".to_string())); + } + if job.status == "running" || job.status == "creating" { + if let Some(ref jm) = state.job_manager + && let Err(e) = jm.stop_job(job_id).await + { + tracing::warn!(job_id = %job_id, error = %e, "Failed to stop container during cancellation"); + } + store + .update_sandbox_job_status( + job_id, + "failed", + Some(false), + Some("Cancelled by user"), + None, + Some(chrono::Utc::now()), + ) + .await + .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; + } + return Ok(Json(serde_json::json!({ + "status": "cancelled", + "job_id": job_id, + }))); + } + Ok(None) => {} + Err(e) => { + return Err(( + StatusCode::INTERNAL_SERVER_ERROR, + format!("Database error: {}", e), + )); } - store - .update_sandbox_job_status( - job_id, - "failed", - Some(false), - Some("Cancelled by user"), - None, - Some(chrono::Utc::now()), - ) - .await - .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; } - return Ok(Json(serde_json::json!({ - "status": "cancelled", - "job_id": job_id, - }))); } // Fall back to agent job cancellation: stop the worker via the scheduler // (which updates the in-memory ContextManager AND aborts the task handle), // then persist the status to the DB as a fallback. - if let Some(ref store) = state.store - && let Ok(Some(job)) = store.get_job(job_id).await - { - if job.state.is_active() { - // Try to stop via scheduler (aborts the worker task + updates - // in-memory ContextManager). This is best-effort — the job may - // not be in the scheduler map if it already finished. - if let Some(ref slot) = state.scheduler - && let Some(ref scheduler) = *slot.read().await - { - let _ = scheduler.stop(job_id).await; - } + if let Some(ref store) = state.store { + match store.get_job(job_id).await { + Ok(Some(job)) => { + if job.user_id != user.user_id { + return Err((StatusCode::NOT_FOUND, "Job not found".to_string())); + } + if job.state.is_active() { + // Try to stop via scheduler (aborts the worker task + updates + // in-memory ContextManager). This is best-effort — the job may + // not be in the scheduler map if it already finished. + if let Some(ref slot) = state.scheduler + && let Some(ref scheduler) = *slot.read().await + { + let _ = scheduler.stop(job_id).await; + } - // Always persist cancellation to the DB so the state is - // consistent even if the scheduler wasn't available or the - // job wasn't in its in-memory map. - store - .update_job_status( - job_id, - crate::context::JobState::Cancelled, - Some("Cancelled by user"), - ) - .await - .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; + // Always persist cancellation to the DB so the state is + // consistent even if the scheduler wasn't available or the + // job wasn't in its in-memory map. + store + .update_job_status( + job_id, + crate::context::JobState::Cancelled, + Some("Cancelled by user"), + ) + .await + .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; + } + return Ok(Json(serde_json::json!({ + "status": "cancelled", + "job_id": job_id, + }))); + } + Ok(None) => {} + Err(e) => { + return Err(( + StatusCode::INTERNAL_SERVER_ERROR, + format!("Database error: {}", e), + )); + } } - return Ok(Json(serde_json::json!({ - "status": "cancelled", - "job_id": job_id, - }))); } Err((StatusCode::NOT_FOUND, "Job not found".to_string())) @@ -315,6 +366,7 @@ pub async fn jobs_cancel_handler( pub async fn jobs_restart_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Path(id): Path, ) -> Result, (StatusCode, String)> { let store = state.store.as_ref().ok_or(( @@ -327,6 +379,9 @@ pub async fn jobs_restart_handler( // Try sandbox job restart first. if let Ok(Some(old_job)) = store.get_sandbox_job(old_job_id).await { + 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, @@ -476,6 +531,7 @@ pub async fn jobs_restart_handler( /// - Worker-mode sandbox jobs → not supported (no mechanism to inject) pub async fn jobs_prompt_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Path(id): Path, Json(body): Json, ) -> Result, (StatusCode, String)> { @@ -483,6 +539,26 @@ pub async fn jobs_prompt_handler( .parse() .map_err(|_| (StatusCode::BAD_REQUEST, "Invalid job ID".to_string()))?; + // Verify ownership before queuing a prompt. + if let Some(ref store) = state.store { + match store.get_sandbox_job(job_id).await { + Ok(Some(job)) => { + if job.user_id != user.user_id { + return Err((StatusCode::NOT_FOUND, "Job not found".to_string())); + } + } + 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 content = body .get("content") .and_then(|v| v.as_str()) @@ -550,6 +626,7 @@ pub async fn jobs_prompt_handler( /// Load persisted job events for a job (for history replay on page open). pub async fn jobs_events_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Path(id): Path, ) -> Result, (StatusCode, String)> { let store = state.store.as_ref().ok_or(( @@ -561,6 +638,24 @@ pub async fn jobs_events_handler( .parse() .map_err(|_| (StatusCode::BAD_REQUEST, "Invalid job ID".to_string()))?; + // Verify ownership before returning events. + match store.get_sandbox_job(job_id).await { + Ok(Some(job)) => { + if job.user_id != user.user_id { + return Err((StatusCode::NOT_FOUND, "Job not found".to_string())); + } + } + Ok(None) => { + return Err((StatusCode::NOT_FOUND, "Job not found".to_string())); + } + Err(e) => { + return Err(( + StatusCode::INTERNAL_SERVER_ERROR, + format!("Database error: {}", e), + )); + } + } + let events = store .list_job_events(job_id, None) .await @@ -593,6 +688,7 @@ pub struct FilePathQuery { pub async fn job_files_list_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Path(id): Path, Query(query): Query, ) -> Result, (StatusCode, String)> { @@ -610,6 +706,10 @@ pub async fn job_files_list_handler( .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))? .ok_or((StatusCode::NOT_FOUND, "Job not found".to_string()))?; + if job.user_id != user.user_id { + return Err((StatusCode::NOT_FOUND, "Job not found".to_string())); + } + let base = std::path::PathBuf::from(&job.project_dir); let rel_path = query.path.as_deref().unwrap_or(""); let target = base.join(rel_path); @@ -656,6 +756,7 @@ pub async fn job_files_list_handler( pub async fn job_files_read_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Path(id): Path, Query(query): Query, ) -> Result, (StatusCode, String)> { @@ -673,6 +774,10 @@ pub async fn job_files_read_handler( .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))? .ok_or((StatusCode::NOT_FOUND, "Job not found".to_string()))?; + if job.user_id != user.user_id { + return Err((StatusCode::NOT_FOUND, "Job not found".to_string())); + } + let path = query.path.as_deref().ok_or(( StatusCode::BAD_REQUEST, "path parameter required".to_string(), diff --git a/src/channels/web/handlers/mod.rs b/src/channels/web/handlers/mod.rs index 2f942058..82187e21 100644 --- a/src/channels/web/handlers/mod.rs +++ b/src/channels/web/handlers/mod.rs @@ -8,6 +8,7 @@ //! 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 skills; // Modules not yet wired into server.rs router -- suppress dead_code until @@ -17,8 +18,6 @@ pub mod chat; #[allow(dead_code)] pub mod extensions; #[allow(dead_code)] -pub mod jobs; -#[allow(dead_code)] pub mod memory; #[allow(dead_code)] pub mod routines; diff --git a/src/channels/web/handlers/routines.rs b/src/channels/web/handlers/routines.rs index 368a28ae..19e219b4 100644 --- a/src/channels/web/handlers/routines.rs +++ b/src/channels/web/handlers/routines.rs @@ -137,6 +137,7 @@ pub async fn routines_detail_handler( pub async fn routines_trigger_handler( State(state): State>, + crate::channels::web::auth::AuthenticatedUser(user): crate::channels::web::auth::AuthenticatedUser, Path(id): Path, ) -> Result, (StatusCode, String)> { // Clone the Arc out of the lock to avoid holding the RwLock across .await. @@ -152,7 +153,7 @@ pub async fn routines_trigger_handler( .map_err(|_| (StatusCode::BAD_REQUEST, "Invalid routine ID".to_string()))?; let run_id = engine - .fire_manual(routine_id, Some(&state.user_id)) + .fire_manual(routine_id, Some(&user.user_id)) .await .map_err(|e| (routine_error_status(&e), e.to_string()))?; diff --git a/src/channels/web/handlers/settings.rs b/src/channels/web/handlers/settings.rs index dd66027b..4dd7299a 100644 --- a/src/channels/web/handlers/settings.rs +++ b/src/channels/web/handlers/settings.rs @@ -8,17 +8,19 @@ use axum::{ http::StatusCode, }; +use crate::channels::web::auth::AuthenticatedUser; use crate::channels::web::server::GatewayState; use crate::channels::web::types::*; pub async fn settings_list_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, ) -> Result, StatusCode> { let store = state .store .as_ref() .ok_or(StatusCode::SERVICE_UNAVAILABLE)?; - let rows = store.list_settings(&state.user_id).await.map_err(|e| { + let rows = store.list_settings(&user.user_id).await.map_err(|e| { tracing::error!("Failed to list settings: {}", e); StatusCode::INTERNAL_SERVER_ERROR })?; @@ -37,6 +39,7 @@ pub async fn settings_list_handler( pub async fn settings_get_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Path(key): Path, ) -> Result, StatusCode> { let store = state @@ -44,7 +47,7 @@ pub async fn settings_get_handler( .as_ref() .ok_or(StatusCode::SERVICE_UNAVAILABLE)?; let row = store - .get_setting_full(&state.user_id, &key) + .get_setting_full(&user.user_id, &key) .await .map_err(|e| { tracing::error!("Failed to get setting '{}': {}", key, e); @@ -61,6 +64,7 @@ pub async fn settings_get_handler( pub async fn settings_set_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Path(key): Path, Json(body): Json, ) -> Result { @@ -69,7 +73,7 @@ pub async fn settings_set_handler( .as_ref() .ok_or(StatusCode::SERVICE_UNAVAILABLE)?; store - .set_setting(&state.user_id, &key, &body.value) + .set_setting(&user.user_id, &key, &body.value) .await .map_err(|e| { tracing::error!("Failed to set setting '{}': {}", key, e); @@ -81,6 +85,7 @@ pub async fn settings_set_handler( pub async fn settings_delete_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Path(key): Path, ) -> Result { let store = state @@ -88,7 +93,7 @@ pub async fn settings_delete_handler( .as_ref() .ok_or(StatusCode::SERVICE_UNAVAILABLE)?; store - .delete_setting(&state.user_id, &key) + .delete_setting(&user.user_id, &key) .await .map_err(|e| { tracing::error!("Failed to delete setting '{}': {}", key, e); @@ -100,12 +105,13 @@ pub async fn settings_delete_handler( pub async fn settings_export_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, ) -> Result, StatusCode> { let store = state .store .as_ref() .ok_or(StatusCode::SERVICE_UNAVAILABLE)?; - let settings = store.get_all_settings(&state.user_id).await.map_err(|e| { + let settings = store.get_all_settings(&user.user_id).await.map_err(|e| { tracing::error!("Failed to export settings: {}", e); StatusCode::INTERNAL_SERVER_ERROR })?; @@ -115,6 +121,7 @@ pub async fn settings_export_handler( pub async fn settings_import_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Json(body): Json, ) -> Result { let store = state @@ -122,7 +129,7 @@ pub async fn settings_import_handler( .as_ref() .ok_or(StatusCode::SERVICE_UNAVAILABLE)?; store - .set_all_settings(&state.user_id, &body.settings) + .set_all_settings(&user.user_id, &body.settings) .await .map_err(|e| { tracing::error!("Failed to import settings: {}", e); diff --git a/src/channels/web/mod.rs b/src/channels/web/mod.rs index f40834cb..a03120ce 100644 --- a/src/channels/web/mod.rs +++ b/src/channels/web/mod.rs @@ -52,6 +52,7 @@ use crate::workspace::Workspace; use self::log_layer::{LogBroadcaster, LogLevelHandle}; +use self::auth::MultiAuthState; use self::server::GatewayState; use self::sse::SseManager; use self::types::SseEvent; @@ -60,14 +61,15 @@ use self::types::SseEvent; pub struct GatewayChannel { config: GatewayConfig, state: Arc, - /// The actual auth token in use (generated or from config). - auth_token: String, + /// Multi-user auth state (replaces bare auth_token). + auth: MultiAuthState, } impl GatewayChannel { /// Create a new gateway channel. /// /// If no auth token is configured, generates a random one and prints it. + /// Builds a single-user `MultiAuthState` from the config. pub fn new(config: GatewayConfig) -> Self { let auth_token = config.auth_token.clone().unwrap_or_else(|| { use rand::RngCore; @@ -77,10 +79,13 @@ impl GatewayChannel { bytes.iter().map(|b| format!("{b:02x}")).collect() }); + let auth = MultiAuthState::single(auth_token, config.user_id.clone()); + let state = Arc::new(GatewayState { msg_tx: tokio::sync::RwLock::new(None), - sse: SseManager::new(), + sse: Arc::new(SseManager::new()), workspace: None, + workspace_pool: None, session_manager: None, log_broadcaster: None, log_level_handle: None, @@ -90,13 +95,13 @@ impl GatewayChannel { job_manager: None, prompt_queue: None, scheduler: None, - user_id: config.user_id.clone(), + default_user_id: config.user_id.clone(), shutdown_tx: tokio::sync::RwLock::new(None), ws_tracker: Some(Arc::new(ws::WsConnectionTracker::new())), llm_provider: None, skill_registry: None, skill_catalog: None, - chat_rate_limiter: server::RateLimiter::new(30, 60), + chat_rate_limiter: server::PerUserRateLimiter::new(30, 60), oauth_rate_limiter: server::RateLimiter::new(10, 60), webhook_rate_limiter: server::RateLimiter::new(10, 60), registry_entries: Vec::new(), @@ -109,7 +114,46 @@ impl GatewayChannel { Self { config, state, - auth_token, + auth, + } + } + + /// Create a gateway channel with a pre-built multi-user auth state. + pub fn new_multi_auth(config: GatewayConfig, auth: MultiAuthState) -> Self { + let state = Arc::new(GatewayState { + msg_tx: tokio::sync::RwLock::new(None), + sse: Arc::new(SseManager::new()), + workspace: None, + workspace_pool: None, + session_manager: None, + log_broadcaster: None, + log_level_handle: None, + extension_manager: None, + tool_registry: None, + store: None, + job_manager: None, + prompt_queue: None, + scheduler: None, + default_user_id: config.user_id.clone(), + shutdown_tx: tokio::sync::RwLock::new(None), + ws_tracker: Some(Arc::new(ws::WsConnectionTracker::new())), + llm_provider: None, + skill_registry: None, + skill_catalog: None, + chat_rate_limiter: server::PerUserRateLimiter::new(30, 60), + oauth_rate_limiter: server::RateLimiter::new(10, 60), + registry_entries: Vec::new(), + cost_guard: None, + routine_engine: Arc::new(tokio::sync::RwLock::new(None)), + startup_time: std::time::Instant::now(), + webhook_rate_limiter: server::RateLimiter::new(10, 60), + active_config: server::ActiveConfigSnapshot::default(), + }); + + Self { + config, + state, + auth, } } @@ -118,8 +162,9 @@ impl GatewayChannel { let mut new_state = GatewayState { msg_tx: tokio::sync::RwLock::new(None), // Preserve the existing broadcast channel so sender handles remain valid. - sse: SseManager::from_sender(self.state.sse.sender()), + sse: Arc::new(SseManager::from_sender(self.state.sse.sender())), workspace: self.state.workspace.clone(), + workspace_pool: self.state.workspace_pool.clone(), session_manager: self.state.session_manager.clone(), log_broadcaster: self.state.log_broadcaster.clone(), log_level_handle: self.state.log_level_handle.clone(), @@ -129,13 +174,13 @@ impl GatewayChannel { job_manager: self.state.job_manager.clone(), prompt_queue: self.state.prompt_queue.clone(), scheduler: self.state.scheduler.clone(), - user_id: self.state.user_id.clone(), + default_user_id: self.state.default_user_id.clone(), shutdown_tx: tokio::sync::RwLock::new(None), ws_tracker: self.state.ws_tracker.clone(), llm_provider: self.state.llm_provider.clone(), skill_registry: self.state.skill_registry.clone(), skill_catalog: self.state.skill_catalog.clone(), - chat_rate_limiter: server::RateLimiter::new(30, 60), + chat_rate_limiter: server::PerUserRateLimiter::new(30, 60), oauth_rate_limiter: server::RateLimiter::new(10, 60), webhook_rate_limiter: server::RateLimiter::new(10, 60), registry_entries: self.state.registry_entries.clone(), @@ -260,9 +305,15 @@ impl GatewayChannel { self } - /// Get the auth token (for printing to console on startup). + /// Inject the per-user workspace pool for multi-user mode. + pub fn with_workspace_pool(mut self, pool: Arc) -> Self { + self.rebuild_state(|s| s.workspace_pool = Some(pool)); + self + } + + /// Get the first auth token (for printing to console on startup). pub fn auth_token(&self) -> &str { - &self.auth_token + self.auth.first_token().unwrap_or("") } /// Get a reference to the shared gateway state (for the agent to push SSE events). @@ -291,7 +342,7 @@ impl Channel for GatewayChannel { ), })?; - server::start_server(addr, self.state.clone(), self.auth_token.clone()).await?; + server::start_server(addr, self.state.clone(), self.auth.clone()).await?; Ok(Box::pin(ReceiverStream::new(rx))) } @@ -311,10 +362,13 @@ impl Channel for GatewayChannel { } }; - self.state.sse.broadcast(SseEvent::Response { - content: response.content, - thread_id, - }); + self.state.sse.broadcast_for_user( + &msg.user_id, + SseEvent::Response { + content: response.content, + thread_id, + }, + ); Ok(()) } @@ -427,13 +481,21 @@ impl Channel for GatewayChannel { }, }; - self.state.sse.broadcast(event); + // Scope events to the user when user_id is available in metadata. + // When user_id is missing (heartbeat, routines), events go to all + // subscribers. In multi-tenant mode this leaks status across users. + if let Some(uid) = metadata.get("user_id").and_then(|v| v.as_str()) { + self.state.sse.broadcast_for_user(uid, event); + } else { + tracing::debug!("Status event missing user_id in metadata; broadcasting globally"); + self.state.sse.broadcast(event); + } Ok(()) } async fn broadcast( &self, - _user_id: &str, + user_id: &str, response: OutgoingResponse, ) -> Result<(), ChannelError> { let thread_id = match response.thread_id { @@ -445,10 +507,13 @@ impl Channel for GatewayChannel { return Ok(()); } }; - self.state.sse.broadcast(SseEvent::Response { - content: response.content, - thread_id, - }); + self.state.sse.broadcast_for_user( + user_id, + SseEvent::Response { + content: response.content, + thread_id, + }, + ); Ok(()) } diff --git a/src/channels/web/openai_compat.rs b/src/channels/web/openai_compat.rs index 51577e06..55b7c854 100644 --- a/src/channels/web/openai_compat.rs +++ b/src/channels/web/openai_compat.rs @@ -463,9 +463,10 @@ fn build_tool_request( pub async fn chat_completions_handler( State(state): State>, + super::auth::AuthenticatedUser(user): super::auth::AuthenticatedUser, Json(req): Json, ) -> Result)> { - if !state.chat_rate_limiter.check() { + if !state.chat_rate_limiter.check(&user.user_id) { return Err(openai_error( StatusCode::TOO_MANY_REQUESTS, "Rate limit exceeded. Please try again later.", diff --git a/src/channels/web/server.rs b/src/channels/web/server.rs index 7edaad67..bc3c78a7 100644 --- a/src/channels/web/server.rs +++ b/src/channels/web/server.rs @@ -30,7 +30,7 @@ use crate::agent::SessionManager; use crate::bootstrap::ironclaw_base_dir; use crate::channels::IncomingMessage; use crate::channels::relay::DEFAULT_RELAY_NAME; -use crate::channels::web::auth::{AuthState, auth_middleware}; +use crate::channels::web::auth::{AuthenticatedUser, MultiAuthState, UserIdentity, auth_middleware}; use crate::channels::web::handlers::jobs::{ job_files_list_handler, job_files_read_handler, jobs_cancel_handler, jobs_detail_handler, jobs_events_handler, jobs_list_handler, jobs_prompt_handler, jobs_restart_handler, @@ -80,7 +80,6 @@ fn redact_oauth_state_for_logs(state: &str) -> String { /// Simple sliding-window rate limiter. /// /// Tracks the number of requests in the current window. Resets when the window expires. -/// Not per-IP (since this is a single-user gateway with auth), but prevents flooding. pub struct RateLimiter { /// Requests remaining in the current window. remaining: AtomicU64, @@ -108,6 +107,12 @@ impl RateLimiter { } /// Try to consume one request. Returns `true` if allowed, `false` if rate limited. + /// + /// Note: There is a benign TOCTOU race between checking `window_start` and + /// resetting it — two concurrent threads may both see an expired window + /// and reset it, granting a few extra requests at the window boundary. + /// This is acceptable for chat rate limiting where approximate enforcement + /// is sufficient, and avoids the cost of a Mutex. pub fn check(&self) -> bool { let now = std::time::SystemTime::now() .duration_since(std::time::UNIX_EPOCH) @@ -148,14 +153,116 @@ pub struct ActiveConfigSnapshot { pub enabled_channels: Vec, } +/// Per-user rate limiter that maintains a separate sliding window per user_id. +/// +/// Prevents one user from exhausting the rate limit for all users in multi-tenant mode. +pub struct PerUserRateLimiter { + limiters: std::sync::RwLock>, + max_requests: u64, + window_secs: u64, +} + +impl PerUserRateLimiter { + pub fn new(max_requests: u64, window_secs: u64) -> Self { + Self { + limiters: std::sync::RwLock::new(std::collections::HashMap::new()), + max_requests, + window_secs, + } + } + + /// Try to consume one request for the given user. Returns `true` if allowed. + pub fn check(&self, user_id: &str) -> bool { + // Fast path: check existing limiter under read lock. + // On lock poisoning (another thread panicked while holding the lock), + // allow the request rather than crashing the server. + { + let map = match self.limiters.read() { + Ok(m) => m, + Err(e) => { + tracing::warn!("PerUserRateLimiter read lock poisoned; recovering"); + e.into_inner() + } + }; + if let Some(limiter) = map.get(user_id) { + return limiter.check(); + } + } + // Slow path: create limiter under write lock. + let mut map = match self.limiters.write() { + Ok(m) => m, + Err(e) => { + tracing::warn!("PerUserRateLimiter write lock poisoned; recovering"); + e.into_inner() + } + }; + let limiter = map + .entry(user_id.to_string()) + .or_insert_with(|| RateLimiter::new(self.max_requests, self.window_secs)); + limiter.check() + } +} + +/// Per-user workspace pool: lazily creates and caches workspaces keyed by user_id. +/// +/// In single-user mode, exactly one workspace is cached. In multi-user mode, +/// each authenticated user gets their own workspace with appropriate scopes. +pub struct WorkspacePool { + db: Arc, + embeddings: Option>, + cache: tokio::sync::RwLock>>, +} + +impl WorkspacePool { + pub fn new( + db: Arc, + embeddings: Option>, + ) -> Self { + Self { + db, + embeddings, + cache: tokio::sync::RwLock::new(std::collections::HashMap::new()), + } + } + + /// Get or create a workspace for the given user identity. + pub async fn get_or_create(&self, identity: &UserIdentity) -> Arc { + // Fast path: check read lock + { + let cache = self.cache.read().await; + if let Some(ws) = cache.get(&identity.user_id) { + return Arc::clone(ws); + } + } + + // Slow path: create workspace under write lock + let mut cache = self.cache.write().await; + // Double-check after acquiring write lock + if let Some(ws) = cache.get(&identity.user_id) { + return Arc::clone(ws); + } + + let mut ws = Workspace::new_with_db(&identity.user_id, Arc::clone(&self.db)); + if let Some(ref emb) = self.embeddings { + ws = ws.with_embeddings(Arc::clone(emb)); + } + + let ws = Arc::new(ws); + cache.insert(identity.user_id.clone(), Arc::clone(&ws)); + ws + } +} + /// Shared state for all gateway handlers. pub struct GatewayState { /// Channel to send messages to the agent loop. pub msg_tx: tokio::sync::RwLock>>, - /// SSE broadcast manager. - pub sse: SseManager, - /// Workspace for memory API. + /// SSE broadcast manager (Arc-wrapped so extension manager can hold a reference). + pub sse: Arc, + /// Workspace for memory API (single-user fallback). pub workspace: Option>, + /// Per-user workspace pool for multi-user mode. + pub workspace_pool: Option>, /// Session manager for thread info. pub session_manager: Option>, /// Log broadcaster for the logs SSE endpoint. @@ -172,8 +279,8 @@ pub struct GatewayState { pub job_manager: Option>, /// Prompt queue for Claude Code follow-up prompts. pub prompt_queue: Option, - /// User ID for this gateway. - pub user_id: String, + /// Default user ID (fallback for non-request contexts like heartbeat/routines). + pub default_user_id: String, /// Shutdown signal sender. pub shutdown_tx: tokio::sync::RwLock>>, /// WebSocket connection tracker. @@ -186,8 +293,8 @@ pub struct GatewayState { pub skill_catalog: Option>, /// Scheduler for sending follow-up messages to running agent jobs. pub scheduler: Option, - /// Rate limiter for chat endpoints (30 messages per 60 seconds). - pub chat_rate_limiter: RateLimiter, + /// Per-user rate limiter for chat endpoints (30 messages per 60 seconds per user). + pub chat_rate_limiter: PerUserRateLimiter, /// Rate limiter for OAuth callback endpoints (10 requests per 60 seconds). pub oauth_rate_limiter: RateLimiter, /// Rate limiter for webhook trigger endpoints (10 requests per 60 seconds). @@ -211,7 +318,7 @@ pub struct GatewayState { pub async fn start_server( addr: SocketAddr, state: Arc, - auth_token: String, + auth: MultiAuthState, ) -> Result { let listener = tokio::net::TcpListener::bind(addr).await.map_err(|e| { crate::error::ChannelError::StartupFailed { @@ -242,7 +349,7 @@ pub async fn start_server( ); // Protected routes (require auth) - let auth_state = AuthState { token: auth_token }; + let auth_state = auth; let protected = Router::new() // Chat .route("/api/chat/send", post(chat_send_handler)) @@ -568,14 +675,12 @@ async fn oauth_callback_handler( .get("error_description") .cloned() .unwrap_or_else(|| error.clone()); - clear_auth_mode(&state).await; return oauth_error_page(&description); } let state_param = match params.get("state") { Some(s) if !s.is_empty() => s.clone(), _ => { - clear_auth_mode(&state).await; return oauth_error_page("IronClaw"); } }; @@ -583,7 +688,6 @@ async fn oauth_callback_handler( let code = match params.get("code") { Some(c) if !c.is_empty() => c.clone(), _ => { - clear_auth_mode(&state).await; return oauth_error_page("IronClaw"); } }; @@ -592,7 +696,6 @@ async fn oauth_callback_handler( let ext_mgr = match state.extension_manager.as_ref() { Some(mgr) => mgr, None => { - clear_auth_mode(&state).await; return oauth_error_page("IronClaw"); } }; @@ -606,7 +709,7 @@ async fn oauth_callback_handler( error = %error, "OAuth callback received with malformed state" ); - clear_auth_mode(&state).await; + clear_auth_mode(&state, &state.default_user_id).await; return oauth_error_page("IronClaw"); } }; @@ -628,7 +731,6 @@ async fn oauth_callback_handler( lookup_key = %redacted_lookup_key, "OAuth callback received with unknown or expired state" ); - clear_auth_mode(&state).await; return oauth_error_page("IronClaw"); } }; @@ -640,14 +742,14 @@ async fn oauth_callback_handler( "OAuth flow expired" ); // Notify UI so auth card can show error instead of staying stuck - if let Some(ref sender) = flow.sse_sender { - let _ = sender.send(SseEvent::AuthCompleted { + if let Some(ref sse) = flow.sse_manager { + sse.broadcast(SseEvent::AuthCompleted { extension_name: flow.extension_name.clone(), success: false, message: "OAuth flow expired. Please try again.".to_string(), }); } - clear_auth_mode(&state).await; + clear_auth_mode(&state, &flow.user_id).await; return oauth_error_page(&flow.display_name); } @@ -753,14 +855,14 @@ async fn oauth_callback_handler( // Clear auth mode regardless of outcome so the next user message goes // through to the LLM instead of being intercepted as a token. - clear_auth_mode(&state).await; + clear_auth_mode(&state, &flow.user_id).await; // After successful OAuth, auto-activate the extension so it moves // from "Installed (Authenticate)" → "Active" without a second click. // OAuth success is independent of activation — tokens are already stored. // Report auth as successful and attempt activation as a bonus step. let final_message = if success { - match ext_mgr.activate(&flow.extension_name).await { + match ext_mgr.activate(&flow.extension_name, &flow.user_id).await { Ok(result) => result.message, Err(e) => { tracing::warn!( @@ -779,8 +881,8 @@ async fn oauth_callback_handler( }; // Broadcast SSE event to notify the web UI - if let Some(ref sender) = flow.sse_sender { - let _ = sender.send(SseEvent::AuthCompleted { + if let Some(ref sse) = flow.sse_manager { + sse.broadcast(SseEvent::AuthCompleted { extension_name: flow.extension_name, success, message: final_message.clone(), @@ -962,7 +1064,7 @@ async fn slack_relay_oauth_callback_handler( let state_key = format!("relay:{}:oauth_state", DEFAULT_RELAY_NAME); let stored_state = match ext_mgr .secrets() - .get_decrypted(&state.user_id, &state_key) + .get_decrypted(&state.default_user_id, &state_key) .await { Ok(secret) => secret.expose().to_string(), @@ -986,7 +1088,7 @@ async fn slack_relay_oauth_callback_handler( } // Delete the nonce (one-time use) - let _ = ext_mgr.secrets().delete(&state.user_id, &state_key).await; + let _ = ext_mgr.secrets().delete(&state.default_user_id, &state_key).await; let result: Result<(), String> = async { let store = state.store.as_ref().ok_or_else(|| { @@ -997,12 +1099,12 @@ async fn slack_relay_oauth_callback_handler( // Store team_id in settings let team_id_key = format!("relay:{}:team_id", DEFAULT_RELAY_NAME); let _ = store - .set_setting(&state.user_id, &team_id_key, &serde_json::json!(team_id)) + .set_setting(&state.default_user_id, &team_id_key, &serde_json::json!(team_id)) .await; // Activate the relay channel ext_mgr - .activate_stored_relay(DEFAULT_RELAY_NAME) + .activate_stored_relay(DEFAULT_RELAY_NAME, &state.default_user_id) .await .map_err(|e| format!("Failed to activate relay channel: {}", e))?; @@ -1104,6 +1206,7 @@ fn mime_to_ext(mime: &str) -> &str { async fn chat_send_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, headers: axum::http::HeaderMap, Json(req): Json, ) -> Result<(StatusCode, Json), (StatusCode, String)> { @@ -1113,14 +1216,14 @@ async fn chat_send_handler( req.thread_id ); - if !state.chat_rate_limiter.check() { + if !state.chat_rate_limiter.check(&user.user_id) { return Err(( StatusCode::TOO_MANY_REQUESTS, "Rate limit exceeded. Try again shortly.".to_string(), )); } - let mut msg = IncomingMessage::new("gateway", &state.user_id, &req.content); + let mut msg = IncomingMessage::new("gateway", &user.user_id, &req.content); // Prefer timezone from JSON body, fall back to X-Timezone header let tz = req .timezone @@ -1130,10 +1233,13 @@ async fn chat_send_handler( msg = msg.with_timezone(tz); } + // Always include user_id in metadata so downstream SSE broadcasts can scope events. + let mut meta = serde_json::json!({"user_id": &user.user_id}); if let Some(ref thread_id) = req.thread_id { msg = msg.with_thread(thread_id); - msg = msg.with_metadata(serde_json::json!({"thread_id": thread_id})); + meta["thread_id"] = serde_json::json!(thread_id); } + msg = msg.with_metadata(meta); // Convert uploaded images to IncomingAttachments if !req.images.is_empty() { @@ -1182,6 +1288,7 @@ async fn chat_send_handler( async fn chat_approval_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Json(req): Json, ) -> Result<(StatusCode, Json), (StatusCode, String)> { let (approved, always) = match req.action.as_str() { @@ -1217,7 +1324,7 @@ async fn chat_approval_handler( ) })?; - let mut msg = IncomingMessage::new("gateway", &state.user_id, content); + let mut msg = IncomingMessage::new("gateway", &user.user_id, content); if let Some(ref thread_id) = req.thread_id { msg = msg.with_thread(thread_id); @@ -1258,6 +1365,7 @@ async fn chat_approval_handler( /// The token never touches the LLM, chat history, or SSE stream. async fn chat_auth_token_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Json(req): Json, ) -> Result, (StatusCode, String)> { let ext_mgr = state.extension_manager.as_ref().ok_or(( @@ -1266,7 +1374,7 @@ async fn chat_auth_token_handler( ))?; match ext_mgr - .configure_token(&req.extension_name, &req.token) + .configure_token(&req.extension_name, &req.token, &user.user_id) .await { Ok(result) => { @@ -1281,27 +1389,36 @@ async fn chat_auth_token_handler( resp.instructions = result.verification.as_ref().map(|v| v.instructions.clone()); if result.verification.is_some() { - state.sse.broadcast(SseEvent::AuthRequired { - extension_name: req.extension_name.clone(), - instructions: Some(result.message), - auth_url: None, - setup_url: None, - }); + state.sse.broadcast_for_user( + &user.user_id, + SseEvent::AuthRequired { + extension_name: req.extension_name.clone(), + instructions: Some(result.message), + auth_url: None, + setup_url: None, + }, + ); } else if result.activated { // Clear auth mode on the active thread - clear_auth_mode(&state).await; + clear_auth_mode(&state, &user.user_id).await; - state.sse.broadcast(SseEvent::AuthCompleted { - extension_name: req.extension_name.clone(), - success: true, - message: result.message, - }); + state.sse.broadcast_for_user( + &user.user_id, + SseEvent::AuthCompleted { + extension_name: req.extension_name.clone(), + success: true, + message: result.message, + }, + ); } else { - state.sse.broadcast(SseEvent::AuthCompleted { - extension_name: req.extension_name.clone(), - success: false, - message: result.message, - }); + state.sse.broadcast_for_user( + &user.user_id, + SseEvent::AuthCompleted { + extension_name: req.extension_name.clone(), + success: false, + message: result.message, + }, + ); } Ok(Json(resp)) @@ -1310,12 +1427,15 @@ async fn chat_auth_token_handler( let msg = e.to_string(); // Re-emit auth_required for retry on validation errors if matches!(e, crate::extensions::ExtensionError::ValidationFailed(_)) { - state.sse.broadcast(SseEvent::AuthRequired { - extension_name: req.extension_name.clone(), - instructions: Some(msg.clone()), - auth_url: None, - setup_url: None, - }); + state.sse.broadcast_for_user( + &user.user_id, + SseEvent::AuthRequired { + extension_name: req.extension_name.clone(), + instructions: Some(msg.clone()), + auth_url: None, + setup_url: None, + }, + ); } Ok(Json(ActionResponse::fail(msg))) } @@ -1325,16 +1445,17 @@ async fn chat_auth_token_handler( /// Cancel an in-progress auth flow. async fn chat_auth_cancel_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Json(_req): Json, ) -> Result, (StatusCode, String)> { - clear_auth_mode(&state).await; + clear_auth_mode(&state, &user.user_id).await; Ok(Json(ActionResponse::ok("Auth cancelled"))) } /// Clear pending auth mode on the active thread. -pub async fn clear_auth_mode(state: &GatewayState) { +pub async fn clear_auth_mode(state: &GatewayState, user_id: &str) { if let Some(ref sm) = state.session_manager { - let session = sm.get_or_create_session(&state.user_id).await; + let session = sm.get_or_create_session(user_id).await; let mut sess = session.lock().await; if let Some(thread_id) = sess.active_thread && let Some(thread) = sess.threads.get_mut(&thread_id) @@ -1346,8 +1467,9 @@ pub async fn clear_auth_mode(state: &GatewayState) { async fn chat_events_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, ) -> Result { - let sse = state.sse.subscribe().ok_or(( + let sse = state.sse.subscribe(Some(user.user_id)).ok_or(( StatusCode::SERVICE_UNAVAILABLE, "Too many connections".to_string(), ))?; @@ -1357,7 +1479,31 @@ async fn chat_events_handler( )) } +/// Check whether an Origin header value points to a local address. +/// +/// Extracts the host from the origin (handling both IPv4/hostname and IPv6 +/// literal formats) and compares it against known local addresses. Used to +/// prevent cross-site WebSocket hijacking while allowing localhost access. +fn is_local_origin(origin: &str) -> bool { + let host = origin + .strip_prefix("http://") + .or_else(|| origin.strip_prefix("https://")) + .and_then(|rest| { + if rest.starts_with('[') { + // IPv6 literal: extract "[::1]" up to and including ']' + rest.find(']').map(|i| &rest[..=i]) + } else { + // IPv4 or hostname: take up to the first ':' (port) or '/' (path) + rest.split(':').next()?.split('/').next() + } + }) + .unwrap_or(""); + + matches!(host, "localhost" | "127.0.0.1" | "[::1]") +} + async fn chat_ws_handler( + AuthenticatedUser(user): AuthenticatedUser, headers: axum::http::HeaderMap, ws: WebSocketUpgrade, State(state): State>, @@ -1375,23 +1521,16 @@ async fn chat_ws_handler( ) })?; - // Extract the host from the origin and compare exactly, so that - // crafted origins like "http://localhost.evil.com" are rejected. - // Origin format is "scheme://host[:port]". - let host = origin - .strip_prefix("http://") - .or_else(|| origin.strip_prefix("https://")) - .and_then(|rest| rest.split(':').next()?.split('/').next()) - .unwrap_or(""); - - let is_local = matches!(host, "localhost" | "127.0.0.1" | "[::1]"); + let is_local = is_local_origin(origin); if !is_local { return Err(( StatusCode::FORBIDDEN, "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, user) + })) } #[derive(Deserialize)] @@ -1403,6 +1542,7 @@ struct HistoryQuery { async fn chat_history_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Query(query): Query, ) -> Result, (StatusCode, String)> { let session_manager = state.session_manager.as_ref().ok_or(( @@ -1410,7 +1550,7 @@ async fn chat_history_handler( "Session manager not available".to_string(), ))?; - let session = session_manager.get_or_create_session(&state.user_id).await; + let session = session_manager.get_or_create_session(&user.user_id).await; let sess = session.lock().await; let limit = query.limit.unwrap_or(50); @@ -1445,9 +1585,12 @@ async fn chat_history_handler( && let Some(ref store) = state.store { let owned = store - .conversation_belongs_to_user(thread_id, &state.user_id) + .conversation_belongs_to_user(thread_id, &user.user_id) .await - .unwrap_or(false); + .map_err(|e| { + tracing::error!(thread_id = %thread_id, error = %e, "DB error during thread ownership check"); + (StatusCode::INTERNAL_SERVER_ERROR, "Database error".to_string()) + })?; if !owned && !sess.threads.contains_key(&thread_id) { return Err((StatusCode::NOT_FOUND, "Thread not found".to_string())); } @@ -1558,27 +1701,29 @@ async fn chat_history_handler( async fn chat_threads_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, ) -> Result, (StatusCode, String)> { let session_manager = state.session_manager.as_ref().ok_or(( StatusCode::SERVICE_UNAVAILABLE, "Session manager not available".to_string(), ))?; - let session = session_manager.get_or_create_session(&state.user_id).await; + let session = session_manager.get_or_create_session(&user.user_id).await; let sess = session.lock().await; // Try DB first for persistent thread list if let Some(ref store) = state.store { // Auto-create assistant thread if it doesn't exist let assistant_id = store - .get_or_create_assistant_conversation(&state.user_id, "gateway") + .get_or_create_assistant_conversation(&user.user_id, "gateway") .await .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; - if let Ok(summaries) = store - .list_conversations_all_channels(&state.user_id, 50) + match store + .list_conversations_all_channels(&user.user_id, 50) .await { + Ok(summaries) => { let mut assistant_thread = None; let mut threads = Vec::new(); @@ -1621,6 +1766,10 @@ async fn chat_threads_handler( active_thread: sess.active_thread, })); } + Err(e) => { + tracing::error!(user_id = %user.user_id, error = %e, "DB error listing threads; falling back to in-memory"); + } + } } // Fallback: in-memory only (no assistant thread without DB) @@ -1649,13 +1798,14 @@ async fn chat_threads_handler( async fn chat_new_thread_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, ) -> Result, (StatusCode, String)> { let session_manager = state.session_manager.as_ref().ok_or(( StatusCode::SERVICE_UNAVAILABLE, "Session manager not available".to_string(), ))?; - let session = session_manager.get_or_create_session(&state.user_id).await; + let session = session_manager.get_or_create_session(&user.user_id).await; let (thread_id, info) = { let mut sess = session.lock().await; let thread = sess.create_thread(); @@ -1677,12 +1827,12 @@ async fn chat_new_thread_handler( // so that the subsequent loadThreads() call from the frontend sees it. if let Some(ref store) = state.store { match store - .ensure_conversation(thread_id, "gateway", &state.user_id, None) + .ensure_conversation(thread_id, "gateway", &user.user_id, None) .await { Ok(true) => {} Ok(false) => tracing::warn!( - user = %state.user_id, + user = %user.user_id, thread_id = %thread_id, "Skipped persisting new thread due to ownership/channel conflict" ), @@ -1702,6 +1852,27 @@ async fn chat_new_thread_handler( // --- Memory handlers --- +/// Resolve the workspace for the authenticated user. +/// +/// Prefers `workspace_pool` (multi-user mode) when available, falling back +/// to the single-user `state.workspace`. +async fn resolve_workspace( + state: &GatewayState, + user: &UserIdentity, +) -> Result, (StatusCode, String)> { + if let Some(ref pool) = state.workspace_pool { + return Ok(pool.get_or_create(user).await); + } + state + .workspace + .as_ref() + .cloned() + .ok_or(( + StatusCode::SERVICE_UNAVAILABLE, + "Workspace not available".to_string(), + )) +} + #[derive(Deserialize)] struct TreeQuery { #[allow(dead_code)] @@ -1710,12 +1881,10 @@ struct TreeQuery { async fn memory_tree_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Query(_query): Query, ) -> Result, (StatusCode, String)> { - let workspace = state.workspace.as_ref().ok_or(( - StatusCode::SERVICE_UNAVAILABLE, - "Workspace not available".to_string(), - ))?; + let workspace = resolve_workspace(&state, &user).await?; // Build tree from list_all (flat list of all paths) let all_paths = workspace @@ -1758,12 +1927,10 @@ struct ListQuery { async fn memory_list_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Query(query): Query, ) -> Result, (StatusCode, String)> { - let workspace = state.workspace.as_ref().ok_or(( - StatusCode::SERVICE_UNAVAILABLE, - "Workspace not available".to_string(), - ))?; + let workspace = resolve_workspace(&state, &user).await?; let path = query.path.as_deref().unwrap_or(""); let entries = workspace @@ -1794,12 +1961,10 @@ struct ReadQuery { async fn memory_read_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Query(query): Query, ) -> Result, (StatusCode, String)> { - let workspace = state.workspace.as_ref().ok_or(( - StatusCode::SERVICE_UNAVAILABLE, - "Workspace not available".to_string(), - ))?; + let workspace = resolve_workspace(&state, &user).await?; let doc = workspace .read(&query.path) @@ -1815,12 +1980,10 @@ async fn memory_read_handler( async fn memory_write_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Json(req): Json, ) -> Result, (StatusCode, String)> { - let workspace = state.workspace.as_ref().ok_or(( - StatusCode::SERVICE_UNAVAILABLE, - "Workspace not available".to_string(), - ))?; + let workspace = resolve_workspace(&state, &user).await?; // Route through layer-aware methods when a layer is specified. // @@ -1880,12 +2043,10 @@ async fn memory_write_handler( async fn memory_search_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Json(req): Json, ) -> Result, (StatusCode, String)> { - let workspace = state.workspace.as_ref().ok_or(( - StatusCode::SERVICE_UNAVAILABLE, - "Workspace not available".to_string(), - ))?; + let workspace = resolve_workspace(&state, &user).await?; let limit = req.limit.unwrap_or(10); let results = workspace @@ -1981,6 +2142,7 @@ async fn logs_level_set_handler( async fn extensions_list_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, ) -> Result, (StatusCode, String)> { let ext_mgr = state.extension_manager.as_ref().ok_or(( StatusCode::NOT_IMPLEMENTED, @@ -1988,7 +2150,7 @@ async fn extensions_list_handler( ))?; let installed = ext_mgr - .list(None, false) + .list(None, false, &user.user_id) .await .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; @@ -2068,6 +2230,7 @@ async fn extensions_tools_handler( async fn extensions_install_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Json(req): Json, ) -> Result, (StatusCode, String)> { // When extension manager isn't available, check registry entries for a helpful message @@ -2103,7 +2266,7 @@ async fn extensions_install_handler( }); match ext_mgr - .install(&req.name, req.url.as_deref(), kind_hint) + .install(&req.name, req.url.as_deref(), kind_hint, &user.user_id) .await { Ok(result) => { @@ -2111,7 +2274,7 @@ async fn extensions_install_handler( // Auto-activate WASM tools after install (install = active). if result.kind == crate::extensions::ExtensionKind::WasmTool { - if let Err(e) = ext_mgr.activate(&req.name).await { + if let Err(e) = ext_mgr.activate(&req.name, &user.user_id).await { tracing::debug!( extension = %req.name, error = %e, @@ -2123,7 +2286,7 @@ async fn extensions_install_handler( // expansion and for first-time auth when credentials are already // configured (e.g., built-in providers). We only surface an auth_url // when the extension reports it is awaiting authorization. - match ext_mgr.auth(&req.name).await { + match ext_mgr.auth(&req.name, &user.user_id).await { Ok(auth_result) if auth_result.auth_url().is_some() => { // Scope expansion or initial OAuth: user needs to authorize resp.auth_url = auth_result.auth_url().map(String::from); @@ -2140,6 +2303,7 @@ async fn extensions_install_handler( async fn extensions_activate_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Path(name): Path, ) -> Result, (StatusCode, String)> { let ext_mgr = state.extension_manager.as_ref().ok_or(( @@ -2147,14 +2311,14 @@ async fn extensions_activate_handler( "Extension manager not available (secrets store required)".to_string(), ))?; - match ext_mgr.activate(&name).await { + match ext_mgr.activate(&name, &user.user_id).await { Ok(result) => { // Activation loaded the WASM module. Check if the tool needs // OAuth scope expansion (e.g., adding google-docs when gmail // already has a token but missing the documents scope). // Initial OAuth setup is triggered via configure. let mut resp = ActionResponse::ok(result.message); - if let Ok(auth_result) = ext_mgr.auth(&name).await + if let Ok(auth_result) = ext_mgr.auth(&name, &user.user_id).await && auth_result.auth_url().is_some() { resp.auth_url = auth_result.auth_url().map(String::from); @@ -2172,10 +2336,10 @@ async fn extensions_activate_handler( } // Activation failed due to auth; try authenticating first. - match ext_mgr.auth(&name).await { + match ext_mgr.auth(&name, &user.user_id).await { Ok(auth_result) if auth_result.is_authenticated() => { // Auth succeeded, retry activation. - match ext_mgr.activate(&name).await { + match ext_mgr.activate(&name, &user.user_id).await { Ok(result) => Ok(Json(ActionResponse::ok(result.message))), Err(e) => Ok(Json(ActionResponse::fail(e.to_string()))), } @@ -2206,22 +2370,57 @@ async fn extensions_activate_handler( /// Redirect `/projects/{id}` to `/projects/{id}/` so relative paths in /// the served HTML resolve within the project namespace. -async fn project_redirect_handler(Path(project_id): Path) -> impl IntoResponse { - axum::response::Redirect::permanent(&format!("/projects/{project_id}/")) +async fn project_redirect_handler( + State(state): State>, + super::auth::AuthenticatedUser(user): super::auth::AuthenticatedUser, + Path(project_id): Path, +) -> impl IntoResponse { + if !verify_project_ownership(&state, &project_id, &user.user_id).await { + return (StatusCode::NOT_FOUND, "Not found").into_response(); + } + axum::response::Redirect::permanent(&format!("/projects/{project_id}/")).into_response() } /// Serve `index.html` when hitting `/projects/{project_id}/`. -async fn project_index_handler(Path(project_id): Path) -> impl IntoResponse { +async fn project_index_handler( + State(state): State>, + super::auth::AuthenticatedUser(user): super::auth::AuthenticatedUser, + Path(project_id): Path, +) -> impl IntoResponse { + if !verify_project_ownership(&state, &project_id, &user.user_id).await { + return (StatusCode::NOT_FOUND, "Not found").into_response(); + } serve_project_file(&project_id, "index.html").await } /// Serve any file under `/projects/{project_id}/{path}`. async fn project_file_handler( + State(state): State>, + super::auth::AuthenticatedUser(user): super::auth::AuthenticatedUser, Path((project_id, path)): Path<(String, String)>, ) -> impl IntoResponse { + if !verify_project_ownership(&state, &project_id, &user.user_id).await { + return (StatusCode::NOT_FOUND, "Not found").into_response(); + } serve_project_file(&project_id, &path).await } +/// Check that a project directory belongs to a job owned by the given user. +/// Returns false if the store is unavailable or the project is not found. +async fn verify_project_ownership(state: &GatewayState, project_id: &str, user_id: &str) -> bool { + let Some(ref store) = state.store else { + return false; + }; + // The project_id is a sandbox job UUID used as the directory name. + let Ok(job_id) = project_id.parse::() else { + return false; + }; + match store.get_sandbox_job(job_id).await { + Ok(Some(job)) => job.user_id == user_id, + _ => false, + } +} + /// Shared logic: resolve the file inside `~/.ironclaw/projects/{project_id}/`, /// guard against path traversal, and stream the content with the right MIME type. async fn serve_project_file(project_id: &str, path: &str) -> axum::response::Response { @@ -2264,6 +2463,7 @@ async fn serve_project_file(project_id: &str, path: &str) -> axum::response::Res async fn extensions_remove_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Path(name): Path, ) -> Result, (StatusCode, String)> { let ext_mgr = state.extension_manager.as_ref().ok_or(( @@ -2271,7 +2471,7 @@ async fn extensions_remove_handler( "Extension manager not available (secrets store required)".to_string(), ))?; - match ext_mgr.remove(&name).await { + match ext_mgr.remove(&name, &user.user_id).await { Ok(message) => Ok(Json(ActionResponse::ok(message))), Err(e) => Ok(Json(ActionResponse::fail(e.to_string()))), } @@ -2279,6 +2479,7 @@ async fn extensions_remove_handler( async fn extensions_registry_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Query(params): Query, ) -> Json { let query = params.query.unwrap_or_default(); @@ -2311,7 +2512,7 @@ async fn extensions_registry_handler( let installed: std::collections::HashSet<(String, String)> = if let Some(ext_mgr) = state.extension_manager.as_ref() { ext_mgr - .list(None, false) + .list(None, false, &user.user_id) .await .unwrap_or_default() .into_iter() @@ -2342,6 +2543,7 @@ async fn extensions_registry_handler( async fn extensions_setup_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Path(name): Path, ) -> Result, (StatusCode, String)> { let ext_mgr = state.extension_manager.as_ref().ok_or(( @@ -2350,12 +2552,12 @@ async fn extensions_setup_handler( ))?; let setup = ext_mgr - .get_setup_schema(&name) + .get_setup_schema(&name, &user.user_id) .await .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; let kind = ext_mgr - .list(None, false) + .list(None, false, &user.user_id) .await .ok() .and_then(|list| list.into_iter().find(|e| e.name == name)) @@ -2372,6 +2574,7 @@ async fn extensions_setup_handler( async fn extensions_setup_submit_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Path(name): Path, Json(req): Json, ) -> Result, (StatusCode, String)> { @@ -2382,9 +2585,9 @@ async fn extensions_setup_submit_handler( // Clear auth mode regardless of outcome so the next user message goes // through to the LLM instead of being intercepted as a token. - clear_auth_mode(&state).await; + clear_auth_mode(&state, &user.user_id).await; - match ext_mgr.configure(&name, &req.secrets, &req.fields).await { + match ext_mgr.configure(&name, &req.secrets, &req.fields, &user.user_id).await { Ok(result) => { let mut resp = if result.verification.is_some() || result.activated { ActionResponse::ok(result.message) @@ -2462,6 +2665,7 @@ async fn pairing_approve_handler( async fn routines_runs_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Path(id): Path, ) -> Result, (StatusCode, String)> { let store = state.store.as_ref().ok_or(( @@ -2472,6 +2676,17 @@ async fn routines_runs_handler( let routine_id = Uuid::parse_str(&id) .map_err(|_| (StatusCode::BAD_REQUEST, "Invalid routine ID".to_string()))?; + // Verify ownership before listing runs. + let routine = store + .get_routine(routine_id) + .await + .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))? + .ok_or((StatusCode::NOT_FOUND, "Routine not found".to_string()))?; + + if routine.user_id != user.user_id { + return Err((StatusCode::NOT_FOUND, "Routine not found".to_string())); + } + let runs = store .list_routine_runs(routine_id, 50) .await @@ -2501,12 +2716,13 @@ async fn routines_runs_handler( async fn settings_list_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, ) -> Result, StatusCode> { let store = state .store .as_ref() .ok_or(StatusCode::SERVICE_UNAVAILABLE)?; - let rows = store.list_settings(&state.user_id).await.map_err(|e| { + let rows = store.list_settings(&user.user_id).await.map_err(|e| { tracing::error!("Failed to list settings: {}", e); StatusCode::INTERNAL_SERVER_ERROR })?; @@ -2525,6 +2741,7 @@ async fn settings_list_handler( async fn settings_get_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Path(key): Path, ) -> Result, StatusCode> { let store = state @@ -2532,7 +2749,7 @@ async fn settings_get_handler( .as_ref() .ok_or(StatusCode::SERVICE_UNAVAILABLE)?; let row = store - .get_setting_full(&state.user_id, &key) + .get_setting_full(&user.user_id, &key) .await .map_err(|e| { tracing::error!("Failed to get setting '{}': {}", key, e); @@ -2549,6 +2766,7 @@ async fn settings_get_handler( async fn settings_set_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Path(key): Path, Json(body): Json, ) -> Result { @@ -2557,7 +2775,7 @@ async fn settings_set_handler( .as_ref() .ok_or(StatusCode::SERVICE_UNAVAILABLE)?; store - .set_setting(&state.user_id, &key, &body.value) + .set_setting(&user.user_id, &key, &body.value) .await .map_err(|e| { tracing::error!("Failed to set setting '{}': {}", key, e); @@ -2569,6 +2787,7 @@ async fn settings_set_handler( async fn settings_delete_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Path(key): Path, ) -> Result { let store = state @@ -2576,7 +2795,7 @@ async fn settings_delete_handler( .as_ref() .ok_or(StatusCode::SERVICE_UNAVAILABLE)?; store - .delete_setting(&state.user_id, &key) + .delete_setting(&user.user_id, &key) .await .map_err(|e| { tracing::error!("Failed to delete setting '{}': {}", key, e); @@ -2588,12 +2807,13 @@ async fn settings_delete_handler( async fn settings_export_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, ) -> Result, StatusCode> { let store = state .store .as_ref() .ok_or(StatusCode::SERVICE_UNAVAILABLE)?; - let settings = store.get_all_settings(&state.user_id).await.map_err(|e| { + let settings = store.get_all_settings(&user.user_id).await.map_err(|e| { tracing::error!("Failed to export settings: {}", e); StatusCode::INTERNAL_SERVER_ERROR })?; @@ -2603,6 +2823,7 @@ async fn settings_export_handler( async fn settings_import_handler( State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, Json(body): Json, ) -> Result { let store = state @@ -2610,7 +2831,7 @@ async fn settings_import_handler( .as_ref() .ok_or(StatusCode::SERVICE_UNAVAILABLE)?; store - .set_all_settings(&state.user_id, &body.settings) + .set_all_settings(&user.user_id, &body.settings) .await .map_err(|e| { tracing::error!("Failed to import settings: {}", e); @@ -2870,8 +3091,9 @@ mod tests { fn test_gateway_state(ext_mgr: Option>) -> Arc { Arc::new(GatewayState { msg_tx: tokio::sync::RwLock::new(None), - sse: SseManager::new(), + sse: Arc::new(SseManager::new()), workspace: None, + workspace_pool: None, session_manager: None, log_broadcaster: None, log_level_handle: None, @@ -2880,14 +3102,14 @@ mod tests { store: None, job_manager: None, prompt_queue: None, - user_id: "test".to_string(), + 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: RateLimiter::new(30, 60), + chat_rate_limiter: PerUserRateLimiter::new(30, 60), oauth_rate_limiter: RateLimiter::new(10, 60), webhook_rate_limiter: RateLimiter::new(10, 60), registry_entries: vec![], @@ -2951,12 +3173,18 @@ mod tests { "BOT_TOKEN": "dummy-token" } }); - let req = axum::http::Request::builder() + let mut req = axum::http::Request::builder() .method("POST") .uri(format!("/api/extensions/{channel_name}/setup")) .header("content-type", "application/json") .body(Body::from(req_body.to_string())) .expect("request"); + // Inject AuthenticatedUser so the handler's extractor succeeds + // without needing the full auth middleware layer. + req.extensions_mut().insert(UserIdentity { + user_id: "test".to_string(), + workspace_read_scopes: Vec::new(), + }); let resp = ServiceExt::>::oneshot(app, req) .await @@ -3056,7 +3284,7 @@ mod tests { break; } match timeout(remaining, receiver.recv()).await { - Ok(Ok(crate::channels::web::types::SseEvent::AuthRequired { .. })) => { + Ok(Ok(scoped)) if matches!(scoped.event, crate::channels::web::types::SseEvent::AuthRequired { .. }) => { panic!("verification responses should not emit auth_required SSE events") } Ok(Ok(_)) => continue, @@ -3077,7 +3305,8 @@ mod tests { let state = test_gateway_state(None); let addr: SocketAddr = "127.0.0.1:0".parse().unwrap(); - let bound = start_server(addr, state.clone(), "test-token".to_string()) + let auth = MultiAuthState::single("test-token".to_string(), "test".to_string()); + let bound = start_server(addr, state.clone(), auth) .await .expect("server should start"); @@ -3239,7 +3468,7 @@ mod tests { scopes: vec![], user_id: "test".to_string(), secrets, - sse_sender: None, + sse_manager: None, gateway_token: None, token_exchange_extra_params: std::collections::HashMap::new(), client_id_secret_name: None, @@ -3287,7 +3516,8 @@ mod tests { ))); let (ext_mgr, _wasm_tools_dir, _wasm_channels_dir) = test_ext_mgr(secrets.clone()); - let (sender, mut receiver) = tokio::sync::broadcast::channel(4); + let sse_mgr = Arc::new(SseManager::new()); + let mut receiver = sse_mgr.sender().subscribe(); let Some(created_at) = expired_flow_created_at() else { eprintln!("Skipping expired OAuth flow SSE test: monotonic uptime below expiry window"); return; @@ -3307,7 +3537,7 @@ mod tests { scopes: vec![], user_id: "test".to_string(), secrets, - sse_sender: Some(sender), + sse_manager: Some(sse_mgr), gateway_token: None, token_exchange_extra_params: std::collections::HashMap::new(), client_id_secret_name: None, @@ -3333,7 +3563,7 @@ mod tests { .expect("response"); assert_eq!(resp.status(), StatusCode::OK); - match receiver.recv().await.expect("auth_completed event") { + match receiver.recv().await.expect("auth_completed event").event { crate::channels::web::types::SseEvent::AuthCompleted { extension_name, success, @@ -3410,7 +3640,7 @@ mod tests { scopes: vec![], user_id: "test".to_string(), secrets, - sse_sender: None, + sse_manager: None, gateway_token: None, token_exchange_extra_params: std::collections::HashMap::new(), client_id_secret_name: None, @@ -3497,7 +3727,7 @@ mod tests { scopes: vec![], user_id: "test".to_string(), secrets, - sse_sender: None, + sse_manager: None, gateway_token: None, token_exchange_extra_params: std::collections::HashMap::new(), client_id_secret_name: None, @@ -3718,4 +3948,36 @@ mod tests { let exists = secrets.exists("test", &state_key).await.unwrap_or(true); assert!(!exists, "CSRF nonce should be deleted after use"); } + + #[test] + fn test_is_local_origin_localhost() { + assert!(is_local_origin("http://localhost:3001")); + assert!(is_local_origin("http://localhost")); + assert!(is_local_origin("https://localhost:3001")); + } + + #[test] + fn test_is_local_origin_ipv4() { + assert!(is_local_origin("http://127.0.0.1:3001")); + assert!(is_local_origin("http://127.0.0.1")); + } + + #[test] + fn test_is_local_origin_ipv6() { + assert!(is_local_origin("http://[::1]:3001")); + assert!(is_local_origin("http://[::1]")); + } + + #[test] + fn test_is_local_origin_rejects_remote() { + assert!(!is_local_origin("http://evil.com")); + assert!(!is_local_origin("http://localhost.evil.com")); + assert!(!is_local_origin("http://192.168.1.1:3001")); + } + + #[test] + fn test_is_local_origin_rejects_garbage() { + assert!(!is_local_origin("not-a-url")); + assert!(!is_local_origin("")); + } } diff --git a/src/channels/web/sse.rs b/src/channels/web/sse.rs index 7b952346..8654f7bd 100644 --- a/src/channels/web/sse.rs +++ b/src/channels/web/sse.rs @@ -17,9 +17,25 @@ use crate::channels::web::types::SseEvent; /// Prevents resource exhaustion from connection flooding. const MAX_CONNECTIONS: u64 = 100; +/// Envelope for broadcast events: carries an optional user scope. +/// +/// `user_id = None` means the event is global (e.g. Heartbeat) and delivered +/// to all subscribers. `user_id = Some(id)` means the event is only delivered +/// to subscribers that match that user_id. +#[derive(Debug, Clone)] +pub(crate) struct ScopedEvent { + pub(crate) user_id: Option, + pub(crate) event: SseEvent, +} + /// Manages SSE broadcast to all connected browser tabs. +/// +/// In multi-user mode, events are scoped by user_id so that each subscriber +/// only receives events intended for their user (plus global events like +/// Heartbeat). In single-user mode, all events are delivered to all subscribers +/// (backwards compatible). pub struct SseManager { - tx: broadcast::Sender, + tx: broadcast::Sender, connection_count: Arc, max_connections: u64, } @@ -45,7 +61,7 @@ impl SseManager { /// only be called before the server starts accepting connections (i.e., /// during startup wiring). Calling it after connections are established /// will break connection tracking and allow exceeding `MAX_CONNECTIONS`. - pub fn from_sender(tx: broadcast::Sender) -> Self { + pub(crate) fn from_sender(tx: broadcast::Sender) -> Self { Self { tx, connection_count: Arc::new(AtomicU64::new(0)), @@ -53,15 +69,28 @@ impl SseManager { } } - /// Broadcast an event to all connected clients. - pub fn broadcast(&self, event: SseEvent) { - // Ignore send errors (no receivers is fine) - let _ = self.tx.send(event); + /// Get a clone of the broadcast sender for use by other components. + pub(crate) fn sender(&self) -> broadcast::Sender { + self.tx.clone() } - /// Get a clone of the broadcast sender for use by other components. - pub fn sender(&self) -> broadcast::Sender { - self.tx.clone() + /// Broadcast an event to all connected clients (global/unscoped). + pub fn broadcast(&self, event: SseEvent) { + let _ = self.tx.send(ScopedEvent { + user_id: None, + event, + }); + } + + /// Broadcast an event scoped to a specific user. + /// + /// Only subscribers for this user_id (or unscoped subscribers) will + /// receive the event. + pub fn broadcast_for_user(&self, user_id: &str, event: SseEvent) { + let _ = self.tx.send(ScopedEvent { + user_id: Some(user_id.to_string()), + event, + }); } /// Get current number of active connections. @@ -71,11 +100,15 @@ impl SseManager { /// Create a raw broadcast subscription for non-SSE consumers (e.g. WebSocket). /// - /// Returns a stream of `SseEvent` values and increments/decrements the - /// connection counter on creation/drop, just like `subscribe()` does for SSE. + /// When `user_id` is `Some`, only events scoped to that user (or global + /// events) are delivered. When `None`, all events are delivered (single-user + /// backwards compatibility). /// /// Returns `None` if the maximum connection limit has been reached. - pub fn subscribe_raw(&self) -> Option + Send + 'static + use<>> { + pub fn subscribe_raw( + &self, + user_id: Option, + ) -> Option + Send + 'static + use<>> { // Atomically increment only if below the limit. This prevents // concurrent callers from overshooting max_connections. let counter = Arc::clone(&self.connection_count); @@ -91,7 +124,19 @@ impl SseManager { .ok()?; let rx = self.tx.subscribe(); - let stream = BroadcastStream::new(rx).filter_map(|result| result.ok()); + let stream = BroadcastStream::new(rx).filter_map(move |result| match result { + Ok(scoped) => { + // Global events (user_id=None) always pass through. + // Scoped events only pass if the subscriber matches (or subscriber is unscoped). + match (&user_id, &scoped.user_id) { + (_, None) => Some(scoped.event), // global -> all + (None, _) => Some(scoped.event), // unscoped subscriber -> all + (Some(sub), Some(ev)) if sub == ev => Some(scoped.event), // match + _ => None, // different user -> skip + } + } + Err(_) => None, + }); Some(CountedStream { inner: stream, @@ -101,9 +146,13 @@ impl SseManager { /// Create a new SSE stream for a client connection. /// + /// When `user_id` is `Some`, only events for that user (or global events) + /// are delivered. When `None`, all events are delivered. + /// /// Returns `None` if the maximum connection limit has been reached. pub fn subscribe( &self, + user_id: Option, ) -> Option> + Send + 'static + use<>>> { // Atomically increment only if below the limit. let counter = Arc::clone(&self.connection_count); @@ -120,7 +169,15 @@ impl SseManager { let rx = self.tx.subscribe(); let stream = BroadcastStream::new(rx) - .filter_map(|result| result.ok()) + .filter_map(move |result| match result { + Ok(scoped) => match (&user_id, &scoped.user_id) { + (_, None) => Some(scoped.event), + (None, _) => Some(scoped.event), + (Some(sub), Some(ev)) if sub == ev => Some(scoped.event), + _ => None, + }, + Err(_) => None, + }) .map(|event| { let data = serde_json::to_string(&event).unwrap_or_default(); let event_type = match &event { @@ -215,16 +272,14 @@ mod tests { #[tokio::test] async fn test_broadcast_to_receiver() { let manager = SseManager::new(); - let mut rx = BroadcastStream::new(manager.tx.subscribe()); + let mut stream = Box::pin(manager.subscribe_raw(None).expect("should subscribe")); manager.broadcast(SseEvent::Status { message: "test".to_string(), thread_id: None, }); - let event = rx.next().await; - assert!(event.is_some()); - let event = event.unwrap().unwrap(); + let event = stream.next().await.unwrap(); match event { SseEvent::Status { message, .. } => assert_eq!(message, "test"), _ => panic!("unexpected event type"), @@ -234,7 +289,7 @@ mod tests { #[tokio::test] async fn test_subscribe_raw_receives_events() { let manager = SseManager::new(); - let mut stream = Box::pin(manager.subscribe_raw().expect("should subscribe")); + let mut stream = Box::pin(manager.subscribe_raw(None).expect("should subscribe")); assert_eq!(manager.connection_count(), 1); @@ -254,7 +309,7 @@ mod tests { async fn test_subscribe_raw_decrements_on_drop() { let manager = SseManager::new(); { - let _stream = Box::pin(manager.subscribe_raw().expect("should subscribe")); + let _stream = Box::pin(manager.subscribe_raw(None).expect("should subscribe")); assert_eq!(manager.connection_count(), 1); } // Stream dropped, counter should decrement @@ -264,8 +319,8 @@ mod tests { #[tokio::test] async fn test_subscribe_raw_multiple_subscribers() { let manager = SseManager::new(); - let mut s1 = Box::pin(manager.subscribe_raw().expect("should subscribe")); - let mut s2 = Box::pin(manager.subscribe_raw().expect("should subscribe")); + let mut s1 = Box::pin(manager.subscribe_raw(None).expect("should subscribe")); + let mut s2 = Box::pin(manager.subscribe_raw(None).expect("should subscribe")); assert_eq!(manager.connection_count(), 2); manager.broadcast(SseEvent::Heartbeat); @@ -286,12 +341,51 @@ mod tests { let mut manager = SseManager::new(); manager.max_connections = 2; // Low limit for testing - let _s1 = Box::pin(manager.subscribe_raw().expect("first should succeed")); - let _s2 = Box::pin(manager.subscribe_raw().expect("second should succeed")); + let _s1 = Box::pin(manager.subscribe_raw(None).expect("first should succeed")); + let _s2 = Box::pin(manager.subscribe_raw(None).expect("second should succeed")); assert_eq!(manager.connection_count(), 2); // Third should be rejected - assert!(manager.subscribe_raw().is_none()); - assert!(manager.subscribe().is_none()); + assert!(manager.subscribe_raw(None).is_none()); + assert!(manager.subscribe(None).is_none()); + } + + #[tokio::test] + async fn test_scoped_events_filtered_by_user() { + let manager = SseManager::new(); + let mut alice = Box::pin( + manager + .subscribe_raw(Some("alice".to_string())) + .expect("subscribe"), + ); + let mut bob = Box::pin( + manager + .subscribe_raw(Some("bob".to_string())) + .expect("subscribe"), + ); + + // Send event scoped to alice + manager.broadcast_for_user( + "alice", + SseEvent::Status { + message: "alice only".to_string(), + thread_id: None, + }, + ); + + // Send global event + manager.broadcast(SseEvent::Heartbeat); + + // Alice gets her scoped event + let e = alice.next().await.unwrap(); + assert!(matches!(e, SseEvent::Status { .. })); + + // Alice also gets the global heartbeat + let e = alice.next().await.unwrap(); + assert!(matches!(e, SseEvent::Heartbeat)); + + // Bob only gets the global heartbeat (alice's event was filtered) + let e = bob.next().await.unwrap(); + assert!(matches!(e, SseEvent::Heartbeat)); } } diff --git a/src/channels/web/test_helpers.rs b/src/channels/web/test_helpers.rs index 8751be6a..87b1c86d 100644 --- a/src/channels/web/test_helpers.rs +++ b/src/channels/web/test_helpers.rs @@ -10,7 +10,8 @@ use std::sync::Arc; use tokio::sync::mpsc; use crate::channels::IncomingMessage; -use crate::channels::web::server::{GatewayState, RateLimiter, start_server}; +use crate::channels::web::auth::MultiAuthState; +use crate::channels::web::server::{GatewayState, PerUserRateLimiter, RateLimiter, start_server}; use crate::channels::web::sse::SseManager; use crate::channels::web::ws::WsConnectionTracker; @@ -64,8 +65,9 @@ impl TestGatewayBuilder { pub fn build(self) -> Arc { Arc::new(GatewayState { msg_tx: tokio::sync::RwLock::new(self.msg_tx), - sse: SseManager::new(), + sse: Arc::new(SseManager::new()), workspace: None, + workspace_pool: None, session_manager: None, log_broadcaster: None, log_level_handle: None, @@ -74,14 +76,14 @@ impl TestGatewayBuilder { store: None, job_manager: None, prompt_queue: None, - user_id: self.user_id, + default_user_id: self.user_id, shutdown_tx: tokio::sync::RwLock::new(None), ws_tracker: Some(Arc::new(WsConnectionTracker::new())), llm_provider: self.llm_provider, skill_registry: None, skill_catalog: None, scheduler: None, - chat_rate_limiter: RateLimiter::new(30, 60), + chat_rate_limiter: PerUserRateLimiter::new(30, 60), oauth_rate_limiter: RateLimiter::new(10, 60), webhook_rate_limiter: RateLimiter::new(10, 60), registry_entries: Vec::new(), @@ -98,11 +100,26 @@ impl TestGatewayBuilder { self, auth_token: &str, ) -> Result<(SocketAddr, Arc), crate::error::ChannelError> { + let auth = MultiAuthState::single(auth_token.to_string(), "test-user".to_string()); let state = self.build(); let addr: SocketAddr = "127.0.0.1:0" .parse() .expect("hard-coded address must parse"); - let bound = start_server(addr, state.clone(), auth_token.to_string()).await?; + 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), crate::error::ChannelError> { + let state = self.build(); + let addr: SocketAddr = "127.0.0.1:0" + .parse() + .expect("hard-coded address must parse"); + let bound = start_server(addr, state.clone(), auth).await?; Ok((bound, state)) } } diff --git a/src/channels/web/ws.rs b/src/channels/web/ws.rs index 470c3422..4ee6002a 100644 --- a/src/channels/web/ws.rs +++ b/src/channels/web/ws.rs @@ -62,7 +62,11 @@ impl Default for WsConnectionTracker { /// /// When either task ends (client disconnect or broadcast closed), both are /// cleaned up. -pub async fn handle_ws_connection(socket: WebSocket, state: Arc) { +pub async fn handle_ws_connection( + socket: WebSocket, + state: Arc, + user: crate::channels::web::auth::UserIdentity, +) { let (mut ws_sink, mut ws_stream) = socket.split(); // Track connection @@ -71,9 +75,9 @@ pub async fn handle_ws_connection(socket: WebSocket, state: Arc) { } let tracker_for_drop = state.ws_tracker.clone(); - // Subscribe to broadcast events (same source as SSE). + // Subscribe to broadcast events (same source as SSE), scoped to this user. // Reject if we've hit the connection limit. - let Some(raw_stream) = state.sse.subscribe_raw() else { + let Some(raw_stream) = state.sse.subscribe_raw(Some(user.user_id.clone())) else { tracing::warn!("WebSocket rejected: too many connections"); // Decrement the WS tracker we already incremented above. if let Some(ref tracker) = tracker_for_drop { @@ -117,7 +121,7 @@ pub async fn handle_ws_connection(socket: WebSocket, state: Arc) { }); // Receiver task: read client frames and route to agent - let user_id = state.user_id.clone(); + let user_id = user.user_id; while let Some(Ok(frame)) = ws_stream.next().await { match frame { Message::Text(text) => { @@ -263,10 +267,11 @@ async fn handle_client_message( token, } => { if let Some(ref ext_mgr) = state.extension_manager { - match ext_mgr.configure_token(&extension_name, &token).await { + match ext_mgr.configure_token(&extension_name, &token, user_id).await { Ok(result) => { if result.verification.is_some() { - state.sse.broadcast( + state.sse.broadcast_for_user( + user_id, crate::channels::web::types::SseEvent::AuthRequired { extension_name: extension_name.clone(), instructions: Some(result.message), @@ -275,8 +280,9 @@ async fn handle_client_message( }, ); } else { - crate::channels::web::server::clear_auth_mode(state).await; - state.sse.broadcast( + crate::channels::web::server::clear_auth_mode(state, user_id).await; + state.sse.broadcast_for_user( + user_id, crate::channels::web::types::SseEvent::AuthCompleted { extension_name, success: true, @@ -288,7 +294,8 @@ async fn handle_client_message( Err(e) => { let msg = format!("Auth failed: {}", e); if matches!(e, crate::extensions::ExtensionError::ValidationFailed(_)) { - state.sse.broadcast( + state.sse.broadcast_for_user( + user_id, crate::channels::web::types::SseEvent::AuthRequired { extension_name: extension_name.clone(), instructions: Some(msg.clone()), @@ -311,7 +318,7 @@ async fn handle_client_message( } } WsClientMessage::AuthCancel { .. } => { - crate::channels::web::server::clear_auth_mode(state).await; + crate::channels::web::server::clear_auth_mode(state, user_id).await; } WsClientMessage::Ping => { let _ = direct_tx.send(WsServerMessage::Pong).await; @@ -498,8 +505,9 @@ mod tests { GatewayState { msg_tx: tokio::sync::RwLock::new(msg_tx), - sse: SseManager::new(), + sse: Arc::new(SseManager::new()), workspace: None, + workspace_pool: None, session_manager: None, log_broadcaster: None, log_level_handle: None, @@ -509,13 +517,13 @@ mod tests { job_manager: None, prompt_queue: None, scheduler: None, - user_id: "test".to_string(), + default_user_id: "test".to_string(), shutdown_tx: tokio::sync::RwLock::new(None), ws_tracker: Some(Arc::new(WsConnectionTracker::new())), llm_provider: None, skill_registry: None, skill_catalog: None, - chat_rate_limiter: crate::channels::web::server::RateLimiter::new(30, 60), + chat_rate_limiter: crate::channels::web::server::PerUserRateLimiter::new(30, 60), oauth_rate_limiter: crate::channels::web::server::RateLimiter::new(10, 60), webhook_rate_limiter: crate::channels::web::server::RateLimiter::new(10, 60), registry_entries: Vec::new(), diff --git a/src/cli/oauth_defaults.rs b/src/cli/oauth_defaults.rs index 531d474e..3b57872f 100644 --- a/src/cli/oauth_defaults.rs +++ b/src/cli/oauth_defaults.rs @@ -447,8 +447,8 @@ pub struct PendingOAuthFlow { pub user_id: String, /// Secrets store reference for token persistence. pub secrets: Arc, - /// SSE broadcast sender for notifying the web UI. - pub sse_sender: Option>, + /// SSE broadcast manager for notifying the web UI. + pub sse_manager: Option>, /// Gateway auth token for authenticating with the platform token exchange proxy. pub gateway_token: Option, /// Additional form params for the token exchange request. diff --git a/src/config/channels.rs b/src/config/channels.rs index d249dd18..45c180df 100644 --- a/src/config/channels.rs +++ b/src/config/channels.rs @@ -2,6 +2,7 @@ use std::collections::HashMap; use std::path::PathBuf; use secrecy::SecretString; +use serde::Deserialize; use crate::bootstrap::ironclaw_base_dir; use crate::config::helpers::{optional_env, parse_bool_env, parse_optional_env}; @@ -45,6 +46,26 @@ pub struct GatewayConfig { /// Bearer token for authentication. Random hex generated at startup if unset. pub auth_token: Option, 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, + /// Memory layer definitions (JSON in env var, or from external config). + pub memory_layers: Vec, + /// 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>, +} + +/// 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, } /// Signal channel configuration (signal-cli daemon HTTP/JSON-RPC). @@ -115,6 +136,125 @@ impl ChannelsConfig { .or_else(|| cs.gateway_user_id.clone()) .unwrap_or_else(|| owner_id.to_string()); + let memory_layers: Vec = + 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> = + match optional_env("GATEWAY_USER_TOKENS")? { + Some(json_str) => { + let tokens: HashMap = + 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 = optional_env("WORKSPACE_READ_SCOPES")? + .map(|s| { + s.split(',') + .map(|s| s.trim().to_string()) + .filter(|s| !s.is_empty()) + .collect() + }) + .unwrap_or_default(); + + for scope in &workspace_read_scopes { + if scope.len() > 128 { + return Err(ConfigError::InvalidValue { + key: "WORKSPACE_READ_SCOPES".to_string(), + message: format!( + "scope '{}...' exceeds 128 characters", + &scope[..32] + ), + }); + } + } Some(GatewayConfig { host: optional_env("GATEWAY_HOST")? .or_else(|| cs.gateway_host.clone()) @@ -126,6 +266,9 @@ impl ChannelsConfig { auth_token: optional_env("GATEWAY_AUTH_TOKEN")? .or_else(|| cs.gateway_auth_token.clone()), user_id, + workspace_read_scopes, + memory_layers, + user_tokens, }) } else { None @@ -281,6 +424,9 @@ mod tests { port: 3000, auth_token: Some("tok-abc".to_string()), user_id: "default".to_string(), + workspace_read_scopes: vec![], + memory_layers: vec![], + user_tokens: None, }; assert_eq!(cfg.host, "127.0.0.1"); assert_eq!(cfg.port, 3000); @@ -295,6 +441,9 @@ mod tests { port: 3001, auth_token: None, user_id: "anon".to_string(), + workspace_read_scopes: vec![], + memory_layers: vec![], + user_tokens: None, }; assert!(cfg.auth_token.is_none()); } diff --git a/src/extensions/manager.rs b/src/extensions/manager.rs index df5de72d..c61c5a2d 100644 --- a/src/extensions/manager.rs +++ b/src/extensions/manager.rs @@ -411,9 +411,8 @@ pub struct ExtensionManager { installed_relay_extensions: RwLock>, /// Last activation error for each WASM channel (ephemeral, cleared on success). activation_errors: RwLock>, - /// SSE broadcast sender (set post-construction via `set_sse_sender()`). - sse_sender: - RwLock>>, + /// SSE broadcast manager (set post-construction via `set_sse_sender()`). + sse_manager: RwLock>>, /// Shared registry of pending OAuth flows for gateway-routed callbacks. /// /// Keyed by CSRF `state` parameter. Populated in `start_wasm_oauth()` @@ -484,7 +483,7 @@ impl ExtensionManager { pub async fn active_tool_names(&self) -> HashSet { let mut names = HashSet::new(); - match self.list(None, false).await { + match self.list(None, false, &self.user_id).await { Ok(extensions) => { for extension in extensions { match extension.kind { @@ -550,7 +549,7 @@ impl ExtensionManager { active_channel_names: RwLock::new(HashSet::new()), installed_relay_extensions: RwLock::new(HashSet::new()), activation_errors: RwLock::new(HashMap::new()), - sse_sender: RwLock::new(None), + sse_manager: RwLock::new(None), pending_oauth_flows: crate::cli::oauth_defaults::new_pending_oauth_registry(), gateway_token: std::env::var("GATEWAY_AUTH_TOKEN").ok(), relay_config: crate::config::RelayConfig::from_env(), @@ -892,25 +891,18 @@ impl ExtensionManager { *self.relay_channel_manager.write().await = Some(channel_manager); } - /// Check if a channel name corresponds to a relay extension (has stored team_id + /// Check if a channel name corresponds to a relay extension (has stored stream token /// or is tracked in the installed relay extensions set). - pub async fn is_relay_channel(&self, name: &str) -> bool { + pub async fn is_relay_channel(&self, name: &str, user_id: &str) -> bool { // Check in-memory installed set first (supports no-store mode) if self.installed_relay_extensions.read().await.contains(name) { return true; } - // Then check persistent settings - if let Some(ref store) = self.store { - let team_id_key = format!("relay:{}:team_id", name); - store - .get_setting(&self.user_id, &team_id_key) - .await - .ok() - .flatten() - .is_some() - } else { - false - } + // Then check for stored stream token + self.secrets + .exists(user_id, &format!("relay:{}:stream_token", name)) + .await + .unwrap_or(false) } /// Restore persisted relay channels after startup. @@ -921,18 +913,18 @@ impl ExtensionManager { /// /// Call this only after `set_relay_channel_manager()` or `set_channel_runtime()`. /// Otherwise, each activation attempt fails with "Channel manager not initialized". - pub async fn restore_relay_channels(&self) { - let persisted = self.load_persisted_active_channels().await; + pub async fn restore_relay_channels(&self, user_id: &str) { + let persisted = self.load_persisted_active_channels(user_id).await; let already_active = self.active_channel_names.read().await.clone(); for name in &persisted { if already_active.contains(name) { continue; } - if !self.is_relay_channel(name).await { + if !self.is_relay_channel(name, user_id).await { continue; } - match self.activate_stored_relay(name).await { + match self.activate_stored_relay(name, user_id).await { Ok(_) => { tracing::debug!(channel = %name, "Restored persisted relay channel"); } @@ -987,7 +979,7 @@ impl ExtensionManager { /// Persist the set of active channel names to the settings store. /// /// Saved under key `activated_channels` so channels auto-activate on restart. - async fn persist_active_channels(&self) { + async fn persist_active_channels(&self, user_id: &str) { let Some(ref store) = self.store else { return; }; @@ -1000,7 +992,7 @@ impl ExtensionManager { .collect(); let value = serde_json::json!(names); if let Err(e) = store - .set_setting(&self.user_id, "activated_channels", &value) + .set_setting(user_id, "activated_channels", &value) .await { tracing::warn!(error = %e, "Failed to persist activated_channels setting"); @@ -1011,11 +1003,11 @@ impl ExtensionManager { /// /// Returns channel names that were activated in a prior session so they can /// be auto-activated at startup. - pub async fn load_persisted_active_channels(&self) -> Vec { + pub async fn load_persisted_active_channels(&self, user_id: &str) -> Vec { let Some(ref store) = self.store else { return Vec::new(); }; - match store.get_setting(&self.user_id, "activated_channels").await { + match store.get_setting(user_id, "activated_channels").await { Ok(Some(value)) => match serde_json::from_value(value) { Ok(names) => names, Err(e) => { @@ -1034,9 +1026,9 @@ impl ExtensionManager { /// Set the SSE broadcast sender for pushing extension status events to the web UI. pub async fn set_sse_sender( &self, - sender: tokio::sync::broadcast::Sender, + sse: Arc, ) { - *self.sse_sender.write().await = Some(sender); + *self.sse_manager.write().await = Some(sse); } /// Returns the pending OAuth flow registry for sharing with the web gateway. @@ -1141,8 +1133,8 @@ impl ExtensionManager { /// Broadcast an extension status change to the web UI via SSE. async fn broadcast_extension_status(&self, name: &str, status: &str, message: Option<&str>) { - if let Some(ref sender) = *self.sse_sender.read().await { - let _ = sender.send(crate::channels::web::types::SseEvent::ExtensionStatus { + if let Some(ref sse) = *self.sse_manager.read().await { + sse.broadcast(crate::channels::web::types::SseEvent::ExtensionStatus { extension_name: name.to_string(), status: status.to_string(), message: message.map(|m| m.to_string()), @@ -1185,15 +1177,14 @@ impl ExtensionManager { &self, name: &str, url: Option<&str>, - kind_hint: Option, - ) -> Result { + kind_hint: Option, user_id: &str) -> Result { let sanitized_url = url.map(sanitize_url_for_logging); tracing::info!(extension = %name, url = ?sanitized_url, kind = ?kind_hint, "Installing extension"); Self::validate_extension_name(name)?; // If we have a registry entry, use it (prefer kind_hint to resolve collisions) if let Some(entry) = self.registry.get_with_kind(name, kind_hint).await { - return self.install_from_entry(&entry).await.map_err(|e| { + return self.install_from_entry(&entry, user_id).await.map_err(|e| { tracing::error!(extension = %name, error = %e, "Extension install failed"); e }); @@ -1203,7 +1194,7 @@ impl ExtensionManager { if let Some(url) = url { let kind = kind_hint.unwrap_or_else(|| infer_kind_from_url(url)); return match kind { - ExtensionKind::McpServer => self.install_mcp_from_url(name, url).await, + ExtensionKind::McpServer => self.install_mcp_from_url(name, url, user_id).await, ExtensionKind::WasmTool => self.install_wasm_tool_from_url(name, url).await, ExtensionKind::WasmChannel => { self.install_wasm_channel_from_url(name, url, None).await @@ -1234,31 +1225,31 @@ impl ExtensionManager { /// /// Read-only for WASM extensions; may initiate OAuth for MCP servers. /// To provide secrets, use [`configure()`] instead. - pub async fn auth(&self, name: &str) -> Result { + pub async fn auth(&self, name: &str, user_id: &str) -> Result { // Clean up expired pending auths self.cleanup_expired_auths().await; // Determine what kind of extension this is - let kind = self.determine_installed_kind(name).await?; + let kind = self.determine_installed_kind(name, user_id).await?; match kind { - ExtensionKind::McpServer => self.auth_mcp(name).await, - ExtensionKind::WasmTool => self.auth_wasm_tool(name).await, - ExtensionKind::WasmChannel => self.auth_wasm_channel_status(name).await, - ExtensionKind::ChannelRelay => self.auth_channel_relay(name).await, + ExtensionKind::McpServer => self.auth_mcp(name, user_id).await, + ExtensionKind::WasmTool => self.auth_wasm_tool(name, user_id).await, + ExtensionKind::WasmChannel => self.auth_wasm_channel_status(name, user_id).await, + ExtensionKind::ChannelRelay => self.auth_channel_relay(name, user_id).await, } } /// Activate an installed (and optionally authenticated) extension. - pub async fn activate(&self, name: &str) -> Result { + pub async fn activate(&self, name: &str, user_id: &str) -> Result { Self::validate_extension_name(name)?; - let kind = self.determine_installed_kind(name).await?; + let kind = self.determine_installed_kind(name, user_id).await?; match kind { - ExtensionKind::McpServer => self.activate_mcp(name).await, - ExtensionKind::WasmTool => self.activate_wasm_tool(name).await, - ExtensionKind::WasmChannel => self.activate_wasm_channel(name).await, - ExtensionKind::ChannelRelay => self.activate_channel_relay(name).await, + ExtensionKind::McpServer => self.activate_mcp(name, user_id).await, + ExtensionKind::WasmTool => self.activate_wasm_tool(name, user_id).await, + ExtensionKind::WasmChannel => self.activate_wasm_channel(name, user_id).await, + ExtensionKind::ChannelRelay => self.activate_channel_relay(name, user_id).await, } } @@ -1270,16 +1261,17 @@ impl ExtensionManager { &self, kind_filter: Option, include_available: bool, + user_id: &str, ) -> Result, ExtensionError> { let mut extensions = Vec::new(); // List MCP servers if kind_filter.is_none() || kind_filter == Some(ExtensionKind::McpServer) { - match self.load_mcp_servers().await { + match self.load_mcp_servers(user_id).await { Ok(servers) => { for server in &servers.servers { let authenticated = - is_authenticated(server, &self.secrets, &self.user_id).await; + is_authenticated(server, &self.secrets, user_id).await; let clients = self.mcp_clients.read().await; let active = clients.contains_key(&server.name); @@ -1337,7 +1329,7 @@ impl ExtensionManager { .get_with_kind(&name, Some(ExtensionKind::WasmTool)) .await; let display_name = registry_entry.as_ref().map(|e| e.display_name.clone()); - let auth_state = self.check_tool_auth_status(&name).await; + let auth_state = self.check_tool_auth_status(&name, user_id).await; let version = if let Some(ref cap_path) = discovered.capabilities_path { tokio::fs::read(cap_path) .await @@ -1384,7 +1376,7 @@ impl ExtensionManager { let errors = self.activation_errors.read().await; for (name, discovered) in channels { let active = active_names.contains(&name); - let auth_state = self.check_channel_auth_status(&name).await; + let auth_state = self.check_channel_auth_status(&name, user_id).await; let activation_error = errors.get(&name).cloned(); let registry_entry = self .registry @@ -1436,7 +1428,7 @@ impl ExtensionManager { let active_names = self.active_channel_names.read().await; for name in installed.iter() { let active = active_names.contains(name); - let has_token = self.is_relay_channel(name).await; + let has_token = self.is_relay_channel(name, user_id).await; let registry_entry = self .registry .get_with_kind(name, Some(ExtensionKind::ChannelRelay)) @@ -1499,9 +1491,9 @@ impl ExtensionManager { } /// Remove an installed extension. - pub async fn remove(&self, name: &str) -> Result { + pub async fn remove(&self, name: &str, user_id: &str) -> Result { Self::validate_extension_name(name)?; - let kind = self.determine_installed_kind(name).await?; + let kind = self.determine_installed_kind(name, user_id).await?; // Clean up any in-progress OAuth flows for this extension. // TCP mode: abort the listener task so port 9876 is freed immediately. @@ -1535,7 +1527,7 @@ impl ExtensionManager { self.mcp_clients.write().await.remove(name); // Remove from config - self.remove_mcp_server(name) + self.remove_mcp_server(name, user_id) .await .map_err(|e| ExtensionError::Config(e.to_string()))?; @@ -1595,7 +1587,7 @@ impl ExtensionManager { ExtensionKind::WasmChannel => { // Remove from active set and persist self.active_channel_names.write().await.remove(name); - self.persist_active_channels().await; + self.persist_active_channels(user_id).await; // Clear stale activation errors so reinstall starts clean self.activation_errors.write().await.remove(name); @@ -1629,15 +1621,14 @@ impl ExtensionManager { // Remove from active channels self.active_channel_names.write().await.remove(name); - self.persist_active_channels().await; + self.persist_active_channels(user_id).await; self.activation_errors.write().await.remove(name); - // Remove stored team_id - if let Some(ref store) = self.store { - let _ = store - .delete_setting(&self.user_id, &format!("relay:{}:team_id", name)) - .await; - } + // Remove stored stream token + let _ = self + .secrets + .delete(user_id, &format!("relay:{}:stream_token", name)) + .await; // Stop webhook traffic before removing the channel from the managers. self.clear_relay_webhook_state().await; @@ -1672,13 +1663,13 @@ impl ExtensionManager { /// /// The upgrade preserves authentication secrets — only the `.wasm` binary /// (and `.capabilities.json`) are replaced. - pub async fn upgrade(&self, name: Option<&str>) -> Result { + pub async fn upgrade(&self, name: Option<&str>, user_id: &str) -> Result { // Collect extensions to check let mut candidates: Vec<(String, ExtensionKind)> = Vec::new(); if let Some(name) = name { Self::validate_extension_name(name)?; - let kind = self.determine_installed_kind(name).await?; + let kind = self.determine_installed_kind(name, user_id).await?; if kind == ExtensionKind::McpServer { return Err(ExtensionError::Other( "MCP servers don't have WIT versions and cannot be upgraded this way" @@ -1716,7 +1707,7 @@ impl ExtensionManager { let mut outcomes = Vec::new(); for (ext_name, kind) in &candidates { - let outcome = self.upgrade_one(ext_name, *kind).await; + let outcome = self.upgrade_one(ext_name, *kind, user_id).await; outcomes.push(outcome); } @@ -1742,7 +1733,7 @@ impl ExtensionManager { } /// Upgrade a single WASM extension if its WIT version is outdated. - async fn upgrade_one(&self, name: &str, kind: ExtensionKind) -> UpgradeOutcome { + async fn upgrade_one(&self, name: &str, kind: ExtensionKind, user_id: &str) -> UpgradeOutcome { let (cap_dir, host_wit) = match kind { ExtensionKind::WasmTool => (&self.wasm_tools_dir, crate::tools::wasm::WIT_TOOL_VERSION), ExtensionKind::WasmChannel => ( @@ -1838,7 +1829,7 @@ impl ExtensionManager { } // Reinstall from registry - match self.install_from_entry(&entry).await { + match self.install_from_entry(&entry, user_id).await { Ok(_) => { tracing::info!( extension = %name, @@ -1867,9 +1858,9 @@ impl ExtensionManager { } /// Get detailed info about an installed extension (version, wit_version, host compatibility). - pub async fn extension_info(&self, name: &str) -> Result { + pub async fn extension_info(&self, name: &str, user_id: &str) -> Result { Self::validate_extension_name(name)?; - let kind = self.determine_installed_kind(name).await?; + let kind = self.determine_installed_kind(name, user_id).await?; match kind { ExtensionKind::WasmTool => { @@ -1949,11 +1940,10 @@ impl ExtensionManager { // ── MCP config helpers (DB with disk fallback) ───────────────────── async fn load_mcp_servers( - &self, - ) -> Result + &self, user_id: &str) -> Result { if let Some(ref store) = self.store { - crate::tools::mcp::config::load_mcp_servers_from_db(store.as_ref(), &self.user_id).await + crate::tools::mcp::config::load_mcp_servers_from_db(store.as_ref(), user_id).await } else { crate::tools::mcp::config::load_mcp_servers().await } @@ -1961,9 +1951,8 @@ impl ExtensionManager { async fn get_mcp_server( &self, - name: &str, - ) -> Result { - let servers = self.load_mcp_servers().await?; + name: &str, user_id: &str) -> Result { + let servers = self.load_mcp_servers(user_id).await?; servers.get(name).cloned().ok_or_else(|| { crate::tools::mcp::config::ConfigError::ServerNotFound { name: name.to_string(), @@ -1973,11 +1962,10 @@ impl ExtensionManager { async fn add_mcp_server( &self, - config: McpServerConfig, - ) -> Result<(), crate::tools::mcp::config::ConfigError> { + config: McpServerConfig, user_id: &str) -> Result<(), crate::tools::mcp::config::ConfigError> { config.validate()?; if let Some(ref store) = self.store { - crate::tools::mcp::config::add_mcp_server_db(store.as_ref(), &self.user_id, config) + crate::tools::mcp::config::add_mcp_server_db(store.as_ref(), user_id, config) .await } else { crate::tools::mcp::config::add_mcp_server(config).await @@ -1986,10 +1974,9 @@ impl ExtensionManager { async fn remove_mcp_server( &self, - name: &str, - ) -> Result<(), crate::tools::mcp::config::ConfigError> { + name: &str, user_id: &str) -> Result<(), crate::tools::mcp::config::ConfigError> { if let Some(ref store) = self.store { - crate::tools::mcp::config::remove_mcp_server_db(store.as_ref(), &self.user_id, name) + crate::tools::mcp::config::remove_mcp_server_db(store.as_ref(), user_id, name) .await } else { crate::tools::mcp::config::remove_mcp_server(name).await @@ -2001,8 +1988,9 @@ impl ExtensionManager { async fn install_from_entry( &self, entry: &RegistryEntry, + user_id: &str, ) -> Result { - let primary_result = self.try_install_from_source(entry, &entry.source).await; + let primary_result = self.try_install_from_source(entry, &entry.source, user_id).await; match fallback_decision(&primary_result, &entry.fallback_source) { FallbackDecision::Return => primary_result, FallbackDecision::TryFallback => { @@ -2017,7 +2005,7 @@ impl ExtensionManager { primary_error = %primary_err, "Primary install failed, trying fallback source" ); - match self.try_install_from_source(entry, fallback).await { + match self.try_install_from_source(entry, fallback, user_id).await { Ok(result) => Ok(result), Err(fallback_err) => { tracing::error!( @@ -2037,6 +2025,7 @@ impl ExtensionManager { &self, entry: &RegistryEntry, source: &ExtensionSource, + user_id: &str, ) -> Result { match entry.kind { ExtensionKind::McpServer => { @@ -2049,7 +2038,7 @@ impl ExtensionManager { )); } }; - self.install_mcp_from_url(&entry.name, &url).await + self.install_mcp_from_url(&entry.name, &url, user_id).await } ExtensionKind::WasmTool => match source { ExtensionSource::WasmDownload { @@ -2133,9 +2122,10 @@ impl ExtensionManager { &self, name: &str, url: &str, + user_id: &str, ) -> Result { // Check if already installed - if self.get_mcp_server(name).await.is_ok() { + if self.get_mcp_server(name, user_id).await.is_ok() { return Err(ExtensionError::AlreadyInstalled(name.to_string())); } @@ -2144,7 +2134,7 @@ impl ExtensionManager { .validate() .map_err(|e| ExtensionError::InvalidUrl(e.to_string()))?; - self.add_mcp_server(config) + self.add_mcp_server(config, user_id) .await .map_err(|e| ExtensionError::Config(e.to_string()))?; @@ -2505,14 +2495,14 @@ impl ExtensionManager { }) } - async fn auth_mcp(&self, name: &str) -> Result { + async fn auth_mcp(&self, name: &str, user_id: &str) -> Result { let server = self - .get_mcp_server(name) + .get_mcp_server(name, user_id) .await .map_err(|e| ExtensionError::NotInstalled(e.to_string()))?; // Check if already authenticated - if is_authenticated(&server, &self.secrets, &self.user_id).await { + if is_authenticated(&server, &self.secrets, user_id).await { return Ok(AuthResult::authenticated(name, ExtensionKind::McpServer)); } @@ -2520,7 +2510,7 @@ impl ExtensionManager { // open in the same browser. The gateway's /oauth/callback handler will // complete the token exchange. if self.should_use_gateway_mode() { - return match self.auth_mcp_build_url(name, &server).await { + return match self.auth_mcp_build_url(name, &server, user_id).await { Ok(result) => Ok(result), Err(ExtensionError::AuthNotSupported(_)) => Ok(AuthResult::awaiting_token( name, @@ -2537,14 +2527,14 @@ impl ExtensionManager { } // CLI/local mode: run the full blocking OAuth flow (opens browser, waits for callback) - match authorize_mcp_server(&server, &self.secrets, &self.user_id).await { + match authorize_mcp_server(&server, &self.secrets, user_id).await { Ok(_token) => { tracing::info!("MCP server '{}' authenticated via OAuth", name); Ok(AuthResult::authenticated(name, ExtensionKind::McpServer)) } Err(crate::tools::mcp::auth::AuthError::NotSupported) => { // Server doesn't support OAuth, try building a URL - match self.auth_mcp_build_url(name, &server).await { + match self.auth_mcp_build_url(name, &server, user_id).await { Ok(result) => Ok(result), Err(_) => Ok(AuthResult::awaiting_token( name, @@ -2583,8 +2573,7 @@ impl ExtensionManager { async fn auth_mcp_build_url( &self, name: &str, - server: &McpServerConfig, - ) -> Result { + server: &McpServerConfig, user_id: &str) -> Result { // Try to discover OAuth metadata and build a URL the user can open manually let metadata = discover_full_oauth_metadata(&server.url) .await @@ -2672,9 +2661,9 @@ impl ExtensionManager { provider: Some(format!("mcp:{}", name)), validation_endpoint: None, scopes, - user_id: self.user_id.clone(), + user_id: user_id.to_string(), secrets: Arc::clone(&self.secrets), - sse_sender: self.sse_sender.read().await.clone(), + sse_manager: self.sse_manager.read().await.clone(), gateway_token: self.gateway_token.clone(), token_exchange_extra_params, client_id_secret_name: if server.oauth.is_none() { @@ -2715,7 +2704,7 @@ impl ExtensionManager { } } - async fn auth_wasm_tool(&self, name: &str) -> Result { + async fn auth_wasm_tool(&self, name: &str, user_id: &str) -> Result { // Read the capabilities file to get auth config let cap_path = self .wasm_tools_dir @@ -2747,7 +2736,7 @@ impl ExtensionManager { let params = CreateSecretParams::new(&auth.secret_name, &value).with_provider(name.to_string()); self.secrets - .create(&self.user_id, params) + .create(user_id, params) .await .map_err(|e| ExtensionError::AuthFailed(e.to_string()))?; @@ -2757,7 +2746,7 @@ impl ExtensionManager { // Check if already authenticated (with scope expansion detection) let token_exists = self .secrets - .exists(&self.user_id, &auth.secret_name) + .exists(user_id, &auth.secret_name) .await .unwrap_or(false); @@ -2765,9 +2754,9 @@ impl ExtensionManager { // If this tool has OAuth config, check whether new scopes are needed let needs_reauth = if let Some(ref oauth) = auth.oauth { let merged = self - .collect_shared_scopes(&auth.secret_name, &oauth.scopes) + .collect_shared_scopes(&auth.secret_name, &oauth.scopes, user_id) .await; - let needs = self.needs_scope_expansion(&auth.secret_name, &merged).await; + let needs = self.needs_scope_expansion(&auth.secret_name, &merged, user_id).await; tracing::debug!( tool = name, secret_name = %auth.secret_name, @@ -2790,7 +2779,7 @@ impl ExtensionManager { // But only if credentials are available — if the tool has setup secrets // for client_id/secret that aren't configured yet, return needs_setup. if let Some(ref oauth) = auth.oauth { - if self.needs_setup_credentials(name, &auth, oauth).await { + if self.needs_setup_credentials(name, &auth, oauth, user_id).await { let display = auth.display_name.as_deref().unwrap_or(name); return Ok(AuthResult::needs_setup( name, @@ -2804,7 +2793,7 @@ impl ExtensionManager { } return self - .start_wasm_oauth(name, &auth, oauth) + .start_wasm_oauth(name, &auth, oauth, user_id) .await .map_err(|e| ExtensionError::AuthFailed(e.to_string())); } @@ -2824,7 +2813,7 @@ impl ExtensionManager { } /// Determine the auth readiness of a WASM channel. - async fn check_channel_auth_status(&self, name: &str) -> ToolAuthState { + async fn check_channel_auth_status(&self, name: &str, user_id: &str) -> ToolAuthState { let cap_path = self .wasm_channels_dir .join(format!("{}.capabilities.json", name)); @@ -2849,7 +2838,7 @@ impl ExtensionManager { let all_provided = futures::future::join_all( required .iter() - .map(|s| self.secrets.exists(&self.user_id, &s.name)), + .map(|s| self.secrets.exists(user_id, &s.name)), ) .await .into_iter() @@ -2884,8 +2873,7 @@ impl ExtensionManager { async fn collect_shared_scopes( &self, secret_name: &str, - base_scopes: &[String], - ) -> Vec { + base_scopes: &[String], _user_id: &str) -> Vec { let mut all_scopes: std::collections::BTreeSet = base_scopes.iter().cloned().collect(); @@ -2905,14 +2893,14 @@ impl ExtensionManager { } /// Check whether the stored scopes are insufficient for the merged scopes. - async fn needs_scope_expansion(&self, secret_name: &str, merged_scopes: &[String]) -> bool { + async fn needs_scope_expansion(&self, secret_name: &str, merged_scopes: &[String], user_id: &str) -> bool { if merged_scopes.is_empty() { return false; } let scopes_key = format!("{}_scopes", secret_name); let stored_scopes: std::collections::HashSet = - match self.secrets.get_decrypted(&self.user_id, &scopes_key).await { + match self.secrets.get_decrypted(user_id, &scopes_key).await { Ok(secret) => { let scopes: std::collections::HashSet = secret .expose() @@ -2979,8 +2967,7 @@ impl ExtensionManager { &self, name: &str, auth: &crate::tools::wasm::AuthCapabilitySchema, - oauth: &crate::tools::wasm::OAuthConfigSchema, - ) -> bool { + oauth: &crate::tools::wasm::OAuthConfigSchema, user_id: &str) -> bool { let builtin = crate::cli::oauth_defaults::builtin_credentials(&auth.secret_name); let (id_entry, secret_entry) = self.find_setup_credential_names(name).await; @@ -3005,7 +2992,7 @@ impl ExtensionManager { continue; } let resolved = self - .resolve_oauth_credential(inline, env, fallback, Some(setup_name)) + .resolve_oauth_credential(inline, env, fallback, Some(setup_name), user_id) .await .is_some(); if !resolved { @@ -3024,11 +3011,10 @@ impl ExtensionManager { inline_value: &Option, env_var_name: &Option, builtin_value: Option<&str>, - setup_secret_name: Option<&str>, - ) -> Option { + setup_secret_name: Option<&str>, user_id: &str) -> Option { // 1. Check secrets store (entered via Setup tab) if let Some(secret_name) = setup_secret_name - && let Ok(secret) = self.secrets.get_decrypted(&self.user_id, secret_name).await + && let Ok(secret) = self.secrets.get_decrypted(user_id, secret_name).await { let val = secret.expose(); if !val.is_empty() { @@ -3061,8 +3047,7 @@ impl ExtensionManager { &self, name: &str, auth: &crate::tools::wasm::AuthCapabilitySchema, - oauth: &crate::tools::wasm::OAuthConfigSchema, - ) -> Result { + oauth: &crate::tools::wasm::OAuthConfigSchema, user_id: &str) -> Result { use crate::cli::oauth_defaults; let builtin = oauth_defaults::builtin_credentials(&auth.secret_name); @@ -3081,8 +3066,7 @@ impl ExtensionManager { &oauth.client_id, &oauth.client_id_env, builtin.as_ref().map(|c| c.client_id), - setup_client_id_name.as_deref(), - ) + setup_client_id_name.as_deref(), user_id) .await .ok_or_else(|| { let env_name = oauth @@ -3109,8 +3093,7 @@ impl ExtensionManager { &oauth.client_secret, &oauth.client_secret_env, builtin.as_ref().map(|c| c.client_secret), - setup_client_secret_name.as_deref(), - ) + setup_client_secret_name.as_deref(), user_id) .await; self.clear_pending_extension_auth(name).await; @@ -3122,7 +3105,7 @@ impl ExtensionManager { // Merge scopes from all tools sharing this provider let merged_scopes = self - .collect_shared_scopes(&auth.secret_name, &oauth.scopes) + .collect_shared_scopes(&auth.secret_name, &oauth.scopes, user_id) .await; // Build authorization URL with CSRF state @@ -3169,9 +3152,9 @@ impl ExtensionManager { provider: auth.provider.clone(), validation_endpoint: auth.validation_endpoint.clone(), scopes: merged_scopes, - user_id: self.user_id.clone(), + user_id: user_id.to_string(), secrets: Arc::clone(&self.secrets), - sse_sender: self.sse_sender.read().await.clone(), + sse_manager: self.sse_manager.read().await.clone(), gateway_token: self.gateway_token.clone(), token_exchange_extra_params: std::collections::HashMap::new(), client_id_secret_name: None, @@ -3199,9 +3182,9 @@ impl ExtensionManager { let secret_name = auth.secret_name.clone(); let provider = auth.provider.clone(); let validation_endpoint = auth.validation_endpoint.clone(); - let user_id = self.user_id.clone(); + let user_id = user_id.to_string(); let secrets = Arc::clone(&self.secrets); - let sse_sender = self.sse_sender.read().await.clone(); + let sse_manager = self.sse_manager.read().await.clone(); let ext_name = name.to_string(); let task_handle = tokio::spawn(async move { @@ -3280,8 +3263,8 @@ impl ExtensionManager { } } - if let Some(ref sender) = sse_sender { - let _ = sender.send(crate::channels::web::types::SseEvent::AuthCompleted { + if let Some(ref sse) = sse_manager { + sse.broadcast(crate::channels::web::types::SseEvent::AuthCompleted { extension_name: ext_name, success, message, @@ -3351,7 +3334,7 @@ impl ExtensionManager { } /// Determine the auth readiness of a WASM tool. - async fn check_tool_auth_status(&self, name: &str) -> ToolAuthState { + async fn check_tool_auth_status(&self, name: &str, user_id: &str) -> ToolAuthState { let Some(cap_file) = self.load_tool_capabilities(name).await else { return ToolAuthState::NoAuth; }; @@ -3402,7 +3385,7 @@ impl ExtensionManager { if let Some(ref auth) = cap_file.auth { let has_token = self .secrets - .exists(&self.user_id, &auth.secret_name) + .exists(user_id, &auth.secret_name) .await .unwrap_or(false) || auth @@ -3424,11 +3407,27 @@ impl ExtensionManager { return ToolAuthState::NoAuth; } - ToolAuthState::Ready + let all_provided = futures::future::join_all( + setup + .required_secrets + .iter() + .filter(|s| !s.optional) + .filter(|s| !Self::is_auto_resolved_oauth_field(&s.name, &cap_file)) + .map(|s| self.secrets.exists(user_id, &s.name)), + ) + .await + .into_iter() + .all(|r| r.unwrap_or(false)); + + if all_provided { + ToolAuthState::Ready + } else { + ToolAuthState::NeedsSetup + } } /// Check auth status for a WASM channel (read-only). - async fn auth_wasm_channel_status(&self, name: &str) -> Result { + async fn auth_wasm_channel_status(&self, name: &str, user_id: &str) -> Result { let cap_path = self .wasm_channels_dir .join(format!("{}.capabilities.json", name)); @@ -3463,7 +3462,7 @@ impl ExtensionManager { } if !self .secrets - .exists(&self.user_id, &secret.name) + .exists(user_id, &secret.name) .await .unwrap_or(false) { @@ -3485,7 +3484,7 @@ impl ExtensionManager { )) } - async fn activate_mcp(&self, name: &str) -> Result { + async fn activate_mcp(&self, name: &str, user_id: &str) -> Result { // Check if already activated { let clients = self.mcp_clients.read().await; @@ -3509,7 +3508,7 @@ impl ExtensionManager { } let server = self - .get_mcp_server(name) + .get_mcp_server(name, user_id) .await .map_err(|e| ExtensionError::NotInstalled(e.to_string()))?; @@ -3518,7 +3517,7 @@ impl ExtensionManager { &self.mcp_session_manager, &self.mcp_process_manager, Some(Arc::clone(&self.secrets)), - &self.user_id, + user_id, ) .await .map_err(|e| ExtensionError::ActivationFailed(e.to_string()))?; @@ -3576,7 +3575,7 @@ impl ExtensionManager { }) } - async fn activate_wasm_tool(&self, name: &str) -> Result { + async fn activate_wasm_tool(&self, name: &str, user_id: &str) -> Result { // Check if already active if self.tool_registry.has(name).await { return Ok(ActivateResult { @@ -3590,7 +3589,7 @@ impl ExtensionManager { // Check auth status — block activation if required secrets are missing. // NeedsAuth (OAuth not yet completed) is allowed because configure() loads // the tool first, then starts the OAuth flow to obtain the token. - let auth_state = self.check_tool_auth_status(name).await; + let auth_state = self.check_tool_auth_status(name, user_id).await; if auth_state == ToolAuthState::NeedsSetup { return Err(ExtensionError::ActivationFailed(format!( "Tool '{}' requires configuration. Use the setup form to provide credentials.", @@ -3670,14 +3669,14 @@ impl ExtensionManager { /// Loads the channel from its WASM file, injects credentials and config, /// registers it with the webhook router, and hot-adds it to the channel manager /// so its stream feeds into the agent loop. - async fn activate_wasm_channel(&self, name: &str) -> Result { + async fn activate_wasm_channel(&self, name: &str, user_id: &str) -> Result { // If already active, re-inject credentials and refresh webhook secret. // Handles the case where a channel was loaded at startup before the // user saved secrets via the web UI. { let active = self.active_channel_names.read().await; if active.contains(name) { - return self.refresh_active_channel(name).await; + return self.refresh_active_channel(name, user_id).await; } } @@ -3704,7 +3703,7 @@ impl ExtensionManager { }; // Check auth status first - let auth_state = self.check_channel_auth_status(name).await; + let auth_state = self.check_channel_auth_status(name, user_id).await; if auth_state != ToolAuthState::Ready && auth_state != ToolAuthState::NoAuth { return Err(ExtensionError::ActivationFailed(format!( "Channel '{}' requires configuration. Use the setup form to provide credentials.", @@ -3914,7 +3913,7 @@ impl ExtensionManager { .insert(channel_name.clone()); // Persist activation state so the channel auto-activates on restart - self.persist_active_channels().await; + self.persist_active_channels(&self.user_id).await; tracing::info!(channel = %channel_name, "Hot-activated WASM channel"); @@ -3930,7 +3929,7 @@ impl ExtensionManager { /// /// Called when the user saves new secrets via the setup form for a channel /// that was loaded at startup (possibly without credentials). - async fn refresh_active_channel(&self, name: &str) -> Result { + async fn refresh_active_channel(&self, name: &str, user_id: &str) -> Result { let router = { let rt_guard = self.channel_runtime.read().await; match rt_guard.as_ref() { @@ -3964,7 +3963,7 @@ impl ExtensionManager { &existing_channel, Some(self.secrets.as_ref()), name, - &self.user_id, + user_id, ) .await { @@ -4013,7 +4012,7 @@ impl ExtensionManager { // Refresh webhook secret if let Ok(secret) = self .secrets - .get_decrypted(&self.user_id, &webhook_secret_name) + .get_decrypted(user_id, &webhook_secret_name) .await { router @@ -4030,7 +4029,7 @@ impl ExtensionManager { if let Some(ref sig_key_name) = sig_key_secret_name && let Ok(key_secret) = self .secrets - .get_decrypted(&self.user_id, sig_key_name) + .get_decrypted(user_id, sig_key_name) .await { match router @@ -4050,7 +4049,7 @@ impl ExtensionManager { if let Some(ref hmac_secret_name_ref) = hmac_secret_name { match self .secrets - .get_decrypted(&self.user_id, hmac_secret_name_ref) + .get_decrypted(user_id, hmac_secret_name_ref) .await { Ok(secret) => { @@ -4108,9 +4107,9 @@ impl ExtensionManager { // ── Channel-relay extension methods ────────────────────────────────── /// Derive a stable instance ID from the relay config and user_id. - fn relay_instance_id(&self, config: &crate::config::RelayConfig) -> String { + fn relay_instance_id(&self, config: &crate::config::RelayConfig, user_id: &str) -> String { config.instance_id.clone().unwrap_or_else(|| { - uuid::Uuid::new_v5(&uuid::Uuid::NAMESPACE_DNS, self.user_id.as_bytes()).to_string() + uuid::Uuid::new_v5(&uuid::Uuid::NAMESPACE_DNS, user_id.as_bytes()).to_string() }) } @@ -4119,9 +4118,9 @@ impl ExtensionManager { /// For Slack: initiates OAuth flow (redirect-based). /// For Telegram: accepts a bot token, registers it with channel-relay, /// and stores the returned stream token. - async fn auth_channel_relay(&self, name: &str) -> Result { - // Check if already authenticated (has stored team_id) - if self.is_relay_channel(name).await { + async fn auth_channel_relay(&self, name: &str, user_id: &str) -> Result { + // Check if already authenticated (stream token exists) + if self.is_relay_channel(name, user_id).await { return Ok(AuthResult::authenticated(name, ExtensionKind::ChannelRelay)); } @@ -4140,10 +4139,11 @@ impl ExtensionManager { // state and appends it to the post-OAuth redirect URL. let state_nonce = uuid::Uuid::new_v4().to_string(); let state_key = format!("relay:{}:oauth_state", name); - let _ = self.secrets.delete(&self.user_id, &state_key).await; + // Delete any stale nonce before storing the new one + let _ = self.secrets.delete(user_id, &state_key).await; self.secrets .create( - &self.user_id, + user_id, CreateSecretParams::new(&state_key, &state_nonce), ) .await @@ -4163,23 +4163,36 @@ impl ExtensionManager { } /// Activate a channel-relay extension. - async fn activate_channel_relay(&self, name: &str) -> Result { + async fn activate_channel_relay(&self, name: &str, user_id: &str) -> Result { + let token_key = format!("relay:{}:stream_token", name); let team_id_key = format!("relay:{}:team_id", name); - let store = self.store.as_ref().ok_or(ExtensionError::AuthRequired)?; - let team_id = store - .get_setting(&self.user_id, &team_id_key) - .await - .ok() - .flatten() - .and_then(|v| v.as_str().map(|s| s.to_string())) - .filter(|s| !s.is_empty()) - .ok_or(ExtensionError::AuthRequired)?; + // Check if we have a stream token + // Verify auth: stream token must exist (even though we don't use it in this constructor path) + let _stream_token = match self.secrets.get_decrypted(user_id, &token_key).await { + Ok(secret) => secret.expose().to_string(), + Err(_) => { + return Err(ExtensionError::AuthRequired); + } + }; + + // Get team_id from settings + let team_id = if let Some(ref store) = self.store { + store + .get_setting(user_id, &team_id_key) + .await + .ok() + .flatten() + .and_then(|v| v.as_str().map(|s| s.to_string())) + .unwrap_or_default() + } else { + String::new() + }; // Use relay config captured at startup let relay_config = self.relay_config()?; - let instance_id = self.relay_instance_id(relay_config); + let instance_id = self.relay_instance_id(relay_config, user_id); let client = crate::channels::relay::RelayClient::new( relay_config.url.clone(), @@ -4206,13 +4219,6 @@ impl ExtensionManager { event_rx, ); - // Callback URL is now set during OAuth flow, not via PUT /callbacks. - // The relay webhook endpoint path is still needed for the web gateway. - tracing::info!( - webhook_path = %relay_config.webhook_path, - "Relay channel activated (callback URL set during OAuth)" - ); - // Hot-add to channel manager let cm_guard = self.relay_channel_manager.read().await; let channel_mgr = cm_guard.as_ref().ok_or_else(|| { @@ -4236,7 +4242,7 @@ impl ExtensionManager { .write() .await .insert(name.to_string()); - self.persist_active_channels().await; + self.persist_active_channels(user_id).await; // Broadcast status let status_msg = "Slack connected via channel relay".to_string(); @@ -4252,12 +4258,12 @@ impl ExtensionManager { } /// Activate a channel-relay extension from stored credentials (for startup reconnect). - pub async fn activate_stored_relay(&self, name: &str) -> Result<(), ExtensionError> { - self.activate_channel_relay(name).await?; + pub async fn activate_stored_relay(&self, name: &str, user_id: &str) -> Result<(), ExtensionError> { self.installed_relay_extensions .write() .await .insert(name.to_string()); + self.activate_channel_relay(name, user_id).await?; Ok(()) } @@ -4266,9 +4272,9 @@ impl ExtensionManager { /// This is a read-only check — it never modifies `installed_relay_extensions`. /// To mark a relay extension as installed, use `activate_stored_relay()` or /// the explicit install flow. - async fn determine_installed_kind(&self, name: &str) -> Result { + async fn determine_installed_kind(&self, name: &str, user_id: &str) -> Result { // Check MCP servers first - if self.get_mcp_server(name).await.is_ok() { + if self.get_mcp_server(name, user_id).await.is_ok() { return Ok(ExtensionKind::McpServer); } @@ -4288,8 +4294,8 @@ impl ExtensionManager { if self.installed_relay_extensions.read().await.contains(name) { return Ok(ExtensionKind::ChannelRelay); } - // Also check if there's a stored team_id (persisted across restarts) - if self.is_relay_channel(name).await { + // Also check if there's a stored stream token (persisted across restarts) + if self.is_relay_channel(name, user_id).await { return Ok(ExtensionKind::ChannelRelay); } @@ -4424,9 +4430,10 @@ impl ExtensionManager { pub async fn get_setup_schema( &self, name: &str, + user_id: &str, ) -> Result { Self::validate_extension_name(name)?; - let kind = self.determine_installed_kind(name).await?; + let kind = self.determine_installed_kind(name, user_id).await?; match kind { ExtensionKind::WasmChannel => { let cap_path = self @@ -4449,7 +4456,7 @@ impl ExtensionManager { for secret in &cap_file.setup.required_secrets { let provided = self .secrets - .exists(&self.user_id, &secret.name) + .exists(user_id, &secret.name) .await .unwrap_or(false); secrets.push(crate::channels::web::types::SecretFieldInfo { @@ -4486,7 +4493,7 @@ impl ExtensionManager { } let provided = self .secrets - .exists(&self.user_id, &secret.name) + .exists(user_id, &secret.name) .await .unwrap_or(false); secrets.push(crate::channels::web::types::SecretFieldInfo { @@ -4849,9 +4856,10 @@ impl ExtensionManager { name: &str, secrets: &std::collections::HashMap, fields: &std::collections::HashMap, + user_id: &str, ) -> Result { Self::validate_extension_name(name)?; - let kind = self.determine_installed_kind(name).await?; + let kind = self.determine_installed_kind(name, user_id).await?; // Load allowed secret names and tool setup field definitions from capabilities. let mut channel_cap_file: Option = None; @@ -4907,7 +4915,7 @@ impl ExtensionManager { } ExtensionKind::McpServer => { let server = self - .get_mcp_server(name) + .get_mcp_server(name, user_id) .await .map_err(|e| ExtensionError::NotInstalled(e.to_string()))?; let mut names = std::collections::HashSet::new(); @@ -4993,7 +5001,7 @@ impl ExtensionManager { let params = CreateSecretParams::new(secret_name, trimmed_value).with_provider(name.to_string()); self.secrets - .create(&self.user_id, params) + .create(user_id, params) .await .map_err(|e| ExtensionError::AuthFailed(e.to_string()))?; } @@ -5071,7 +5079,7 @@ impl ExtensionManager { .is_some_and(|v| !v.trim().is_empty()); let already_stored = self .secrets - .exists(&self.user_id, &secret_def.name) + .exists(user_id, &secret_def.name) .await .unwrap_or(false); if !already_provided && !already_stored { @@ -5083,7 +5091,7 @@ impl ExtensionManager { let params = CreateSecretParams::new(&secret_def.name, &hex_value) .with_provider(name.to_string()); self.secrets - .create(&self.user_id, params) + .create(user_id, params) .await .map_err(|e| ExtensionError::AuthFailed(e.to_string()))?; tracing::info!( @@ -5119,7 +5127,7 @@ impl ExtensionManager { // For tools, save and attempt auto-activation, then check auth. if kind == ExtensionKind::WasmTool { - match self.activate_wasm_tool(name).await { + match self.activate_wasm_tool(name, user_id).await { Ok(result) => { // Delete existing OAuth token so auth() starts a fresh flow. // Done AFTER activation succeeds to avoid losing tokens on failure. @@ -5130,16 +5138,16 @@ impl ExtensionManager { { let _ = self .secrets - .delete(&self.user_id, &auth_cfg.secret_name) + .delete(user_id, &auth_cfg.secret_name) .await; let _ = self .secrets - .delete(&self.user_id, &format!("{}_scopes", auth_cfg.secret_name)) + .delete(user_id, &format!("{}_scopes", auth_cfg.secret_name)) .await; let _ = self .secrets .delete( - &self.user_id, + user_id, &format!("{}_refresh_token", auth_cfg.secret_name), ) .await; @@ -5150,7 +5158,7 @@ impl ExtensionManager { let mut auth_url = None; // Box::pin breaks the async recursion cycle: // auth() → auth_wasm_tool() → (OAuth) → configure() → auth() - if let Ok(auth_result) = Box::pin(self.auth(name)).await { + if let Ok(auth_result) = Box::pin(self.auth(name, user_id)).await { auth_url = auth_result.auth_url().map(String::from); } let message = if auth_url.is_some() { @@ -5192,9 +5200,9 @@ impl ExtensionManager { // Activate the extension now that secrets are saved. // Dispatch by kind — WasmTool was already handled above with an early return. let activate_result = match kind { - ExtensionKind::WasmChannel => self.activate_wasm_channel(name).await, - ExtensionKind::McpServer => self.activate_mcp(name).await, - ExtensionKind::ChannelRelay => self.activate_channel_relay(name).await, + ExtensionKind::WasmChannel => self.activate_wasm_channel(name, user_id).await, + ExtensionKind::McpServer => self.activate_mcp(name, user_id).await, + ExtensionKind::ChannelRelay => self.activate_channel_relay(name, user_id).await, ExtensionKind::WasmTool => { return Ok(ConfigureResult { message: format!("Configuration saved for '{}'.", name), @@ -5268,9 +5276,8 @@ impl ExtensionManager { pub async fn configure_token( &self, name: &str, - token: &str, - ) -> Result { - let kind = self.determine_installed_kind(name).await?; + token: &str, user_id: &str) -> Result { + let kind = self.determine_installed_kind(name, user_id).await?; let secret_name = match kind { ExtensionKind::WasmChannel => { let cap_path = self @@ -5291,7 +5298,7 @@ impl ExtensionManager { } if !self .secrets - .exists(&self.user_id, &s.name) + .exists(user_id, &s.name) .await .unwrap_or(false) { @@ -5321,7 +5328,7 @@ impl ExtensionManager { if let Some(ref auth) = cap.auth { if !self .secrets - .exists(&self.user_id, &auth.secret_name) + .exists(user_id, &auth.secret_name) .await .unwrap_or(false) { @@ -5332,7 +5339,7 @@ impl ExtensionManager { for s in &setup.required_secrets { if !self .secrets - .exists(&self.user_id, &s.name) + .exists(user_id, &s.name) .await .unwrap_or(false) { @@ -5359,7 +5366,7 @@ impl ExtensionManager { } ExtensionKind::McpServer => { let server = self - .get_mcp_server(name) + .get_mcp_server(name, user_id) .await .map_err(|e| ExtensionError::NotInstalled(e.to_string()))?; server.token_secret_name() @@ -5369,7 +5376,7 @@ impl ExtensionManager { let mut secrets = std::collections::HashMap::new(); secrets.insert(secret_name, token.to_string()); - self.configure(name, &secrets, &std::collections::HashMap::new()) + self.configure(name, &secrets, &std::collections::HashMap::new(), user_id) .await } @@ -5931,7 +5938,7 @@ mod tests { tools_dir, channels_dir, None, // tunnel_url - "test".to_string(), + "test".to_string(), // user_id store, vec![], ) @@ -6128,7 +6135,7 @@ mod tests { let runtime = Arc::new(crate::tools::wasm::WasmToolRuntime::new(config).expect("runtime")); let mgr = make_test_manager(Some(runtime), dir.path().to_path_buf()); - let err = mgr.activate("nonexistent").await.unwrap_err(); + let err = mgr.activate("nonexistent", "test").await.unwrap_err(); let msg = err.to_string(); assert!( !msg.contains("WASM runtime not available"), @@ -6152,7 +6159,7 @@ mod tests { let mgr = make_test_manager(None, dir.path().to_path_buf()); - let err = mgr.activate("fake").await.unwrap_err(); + let err = mgr.activate("fake", "test").await.unwrap_err(); let msg = err.to_string(); assert!( msg.contains("WASM runtime not available"), @@ -6187,7 +6194,7 @@ mod tests { #[tokio::test] async fn test_upgrade_no_installed_extensions() { let manager = make_manager_with_temp_dirs(); - let result = manager.upgrade(None).await.unwrap(); + let result = manager.upgrade(None, "test").await.unwrap(); assert!(result.results.is_empty()); assert!(result.message.contains("No WASM extensions installed")); } @@ -6196,7 +6203,7 @@ mod tests { async fn test_upgrade_mcp_server_rejected() { let manager = make_manager_with_temp_dirs(); // MCP servers can't be upgraded via tool_upgrade - let err = manager.upgrade(Some("some-mcp")).await; + let err = manager.upgrade(Some("some-mcp"), "test").await; // It will fail with NotInstalled because there's no MCP server named "some-mcp", // but if it were installed, the MCP code path would be rejected. assert!(err.is_err()); @@ -6222,7 +6229,7 @@ mod tests { let manager = make_manager_custom_dirs(dir.path().join("tools"), channels_dir); - let result = manager.upgrade(Some("test-channel")).await.unwrap(); + let result = manager.upgrade(Some("test-channel"), "test").await.unwrap(); assert_eq!(result.results.len(), 1); assert_eq!(result.results[0].status, "already_up_to_date"); } @@ -6247,7 +6254,7 @@ mod tests { let manager = make_manager_custom_dirs(dir.path().join("tools"), channels_dir); - let result = manager.upgrade(Some("custom-channel")).await.unwrap(); + let result = manager.upgrade(Some("custom-channel"), "test").await.unwrap(); assert_eq!(result.results.len(), 1); assert_eq!(result.results[0].status, "not_in_registry"); } @@ -6502,6 +6509,7 @@ mod tests { "123456789:ABCdefGhI".to_string(), )]), &std::collections::HashMap::new(), + "test", ) .await .map_err(|err| format!("configure succeeds: {err}"))?; @@ -6532,7 +6540,7 @@ mod tests { "telegram should be hot-added to the running channel manager", )?; require_eq( - manager.load_persisted_active_channels().await, + manager.load_persisted_active_channels("test").await, vec!["telegram".to_string()], "persisted active channels", )?; @@ -6630,6 +6638,7 @@ mod tests { "123456789:ABCdefGhI".to_string(), )]), &std::collections::HashMap::new(), + "test", ) .await .map_err(|err| format!("configure returned challenge: {err}"))?; @@ -6943,7 +6952,7 @@ mod tests { ); // Calling determine_installed_kind for a non-installed name returns NotInstalled - let result = mgr.determine_installed_kind("slack-relay").await; + let result = mgr.determine_installed_kind("slack-relay", "test").await; assert!(result.is_err(), "Should return NotInstalled"); // Crucially: installed_relay_extensions must still be empty @@ -6958,8 +6967,8 @@ mod tests { let dir = tempfile::tempdir().expect("temp dir"); let mgr = make_test_manager(None, dir.path().to_path_buf()); - // With no DB store, is_relay_channel always returns false - assert!(!mgr.is_relay_channel("slack-relay").await); + // No token stored → not a relay channel + assert!(!mgr.is_relay_channel("slack-relay", "test").await); } #[tokio::test] @@ -6967,7 +6976,7 @@ mod tests { let dir = tempfile::tempdir().expect("temp dir"); let mgr = make_test_manager(None, dir.path().to_path_buf()); - let err = mgr.activate_channel_relay("slack-relay").await.unwrap_err(); + let err = mgr.activate_channel_relay("slack-relay", "test").await.unwrap_err(); assert!( matches!(err, ExtensionError::AuthRequired), "expected AuthRequired, got: {err:?}" @@ -7011,7 +7020,7 @@ mod tests { assert!(cm.get_channel("slack-relay").await.is_some()); // Remove should succeed and shut down the channel - let result = mgr.remove("slack-relay").await; + let result = mgr.remove("slack-relay", "test").await; assert!(result.is_ok(), "remove should succeed: {:?}", result.err()); // installed_relay_extensions should be cleared @@ -7080,7 +7089,7 @@ mod tests { scopes: vec![], user_id: "test".to_string(), secrets: Arc::clone(&secrets), - sse_sender: None, + sse_manager: None, gateway_token: None, token_exchange_extra_params: std::collections::HashMap::new(), client_id_secret_name: None, @@ -7104,7 +7113,7 @@ mod tests { scopes: vec![], user_id: "test".to_string(), secrets, - sse_sender: None, + sse_manager: None, gateway_token: None, token_exchange_extra_params: std::collections::HashMap::new(), client_id_secret_name: None, @@ -7112,7 +7121,7 @@ mod tests { }, ); - let result = mgr.remove("gmail").await; + let result = mgr.remove("gmail", "test").await; assert!(result.is_ok(), "remove should succeed: {:?}", result.err()); tokio::task::yield_now().await; @@ -7158,7 +7167,7 @@ mod tests { .await .insert("telegram".to_string(), "channel failed".to_string()); - let result = mgr.remove("telegram").await; + let result = mgr.remove("telegram", "test").await; assert!(result.is_ok(), "remove should succeed: {:?}", result.err()); assert!( @@ -7568,7 +7577,7 @@ mod tests { .expect("store SECRET_A"); // configure_token should target SECRET_B (the first missing one) - let _result = mgr.configure_token("multi", "value-b").await; + let _result = mgr.configure_token("multi", "value-b", "test").await; // configure will fail at activation (no real WASM runtime), but the // secret should still have been stored before activation was attempted. // Check that SECRET_B was stored. @@ -7608,7 +7617,7 @@ mod tests { let mgr = make_manager_custom_dirs(dir.path().join("tools"), channels_dir); // auth() should return a result without storing anything - let result = mgr.auth("test-ch").await; + let result = mgr.auth("test-ch", "test").await; assert!(result.is_ok(), "auth should succeed: {:?}", result.err()); // No secrets should have been created @@ -7651,7 +7660,7 @@ mod tests { let mgr = make_manager_custom_dirs(dir.path().join("tools"), channels_dir); let result = mgr - .auth("telegram") + .auth("telegram", "test") .await .map_err(|err| format!("telegram auth status: {err}"))?; let instructions = result @@ -7784,7 +7793,7 @@ mod tests { ); let result = mgr - .configure("test-relay", &secrets, &std::collections::HashMap::new()) + .configure("test-relay", &secrets, &std::collections::HashMap::new(), "test") .await; assert!( result.is_ok(), diff --git a/src/main.rs b/src/main.rs index 2cf8fd53..09a16bdb 100644 --- a/src/main.rs +++ b/src/main.rs @@ -589,15 +589,42 @@ async fn async_main() -> anyhow::Result<()> { // ── Gateway channel ──────────────────────────────────────────────── let mut gateway_url: Option = None; - let mut sse_sender: Option< - tokio::sync::broadcast::Sender, - > = None; + let mut sse_manager: Option> = None; + let mut _gateway_state: Option> = + None; if let Some(ref gw_config) = config.channels.gateway { - let mut gw = - GatewayChannel::new(gw_config.clone()).with_llm_provider(Arc::clone(&components.llm)); + // Build multi-user auth state if user_tokens is configured, else single-user. + let mut gw = if let Some(ref user_tokens) = gw_config.user_tokens { + use ironclaw::channels::web::auth::{MultiAuthState, UserIdentity}; + let tokens = user_tokens + .iter() + .map(|(token, cfg)| { + ( + token.clone(), + UserIdentity { + user_id: cfg.user_id.clone(), + workspace_read_scopes: cfg.workspace_read_scopes.clone(), + }, + ) + }) + .collect(); + let auth = MultiAuthState::multi(tokens); + GatewayChannel::new_multi_auth(gw_config.clone(), auth) + } else { + GatewayChannel::new(gw_config.clone()) + }; + gw = gw.with_llm_provider(Arc::clone(&components.llm)); if let Some(ref ws) = components.workspace { gw = gw.with_workspace(Arc::clone(ws)); } + // Create per-user workspace pool for multi-user mode. + if let Some(ref db) = components.db { + let pool = Arc::new(ironclaw::channels::web::server::WorkspacePool::new( + Arc::clone(db), + components.embeddings.clone(), + )); + gw = gw.with_workspace_pool(pool); + } gw = gw.with_session_manager(Arc::clone(&session_manager)); gw = gw.with_log_broadcaster(Arc::clone(&log_broadcaster)); gw = gw.with_log_level_handle(Arc::clone(&log_level_handle)); @@ -648,8 +675,12 @@ async fn async_main() -> anyhow::Result<()> { let mut rx = tx.subscribe(); let gw_state = Arc::clone(gw.state()); tokio::spawn(async move { - while let Ok((_job_id, event)) = rx.recv().await { - gw_state.sse.broadcast(event); + while let Ok((_job_id, user_id, event)) = rx.recv().await { + if user_id.is_empty() { + gw_state.sse.broadcast(event); + } else { + gw_state.sse.broadcast_for_user(&user_id, event); + } } }); } @@ -691,7 +722,8 @@ async fn async_main() -> anyhow::Result<()> { // Capture SSE sender and routine engine slot before moving gw into channels. // IMPORTANT: This must come after all `with_*` calls since `rebuild_state` // creates a new SseManager, which would orphan this sender. - sse_sender = Some(gw.state().sse.sender()); + sse_manager = Some(Arc::clone(&gw.state().sse)); + _gateway_state = Some(Arc::clone(gw.state())); channel_names.push("gateway".to_string()); channels.add(Box::new(gw)).await; } @@ -774,12 +806,18 @@ async fn async_main() -> anyhow::Result<()> { // Auto-activate WASM channels that were active in a previous session. // Relay channels are handled separately below via restore_relay_channels(). - let persisted = ext_mgr.load_persisted_active_channels().await; + let 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 { - if active_at_startup.contains(name) || ext_mgr.is_relay_channel(name).await { + if active_at_startup.contains(name) || ext_mgr.is_relay_channel(name, &ext_user_id).await { continue; } - match ext_mgr.activate(name).await { + match ext_mgr.activate(name, &ext_user_id).await { Ok(result) => { tracing::debug!( channel = %name, @@ -804,14 +842,20 @@ async fn async_main() -> anyhow::Result<()> { ext_mgr .set_relay_channel_manager(Arc::clone(&channels)) .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. 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 @@ -849,7 +893,7 @@ async fn async_main() -> anyhow::Result<()> { skills_config: config.skills.clone(), hooks: components.hooks, cost_guard: components.cost_guard, - sse_tx: sse_sender, + sse_tx: None, // TODO: wire SseManager into scheduler (needs Sender → Arc refactor) http_interceptor, transcription: config.transcription.create_provider().map(|p| { Arc::new(ironclaw::llm::transcription::TranscriptionMiddleware::new( diff --git a/src/orchestrator/api.rs b/src/orchestrator/api.rs index 8d77c581..0ebfd415 100644 --- a/src/orchestrator/api.rs +++ b/src/orchestrator/api.rs @@ -40,7 +40,8 @@ pub struct OrchestratorState { pub job_manager: Arc, pub token_store: TokenStore, /// Broadcast channel for job events (consumed by the web gateway SSE). - pub job_event_tx: Option>, + /// Tuple: (job_id, user_id, event). + pub job_event_tx: Option>, /// Buffered follow-up prompts for sandbox jobs, keyed by job_id. pub prompt_queue: Arc>>>, /// 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 { - 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) @@ -769,8 +785,10 @@ mod tests { let resp = router.oneshot(req).await.unwrap(); assert_eq!(resp.status(), StatusCode::OK); - let (recv_id, event) = rx.recv().await.unwrap(); + let (recv_id, recv_uid, event) = rx.recv().await.unwrap(); assert_eq!(recv_id, job_id); + // No store configured, so user_id falls back to empty string. + assert_eq!(recv_uid, ""); match event { SseEvent::JobMessage { job_id: jid, @@ -824,7 +842,7 @@ mod tests { let resp = router.oneshot(req).await.unwrap(); assert_eq!(resp.status(), StatusCode::OK); - let (_recv_id, event) = rx.recv().await.unwrap(); + let (_recv_id, _recv_uid, event) = rx.recv().await.unwrap(); match event { SseEvent::JobToolUse { tool_name, .. } => { assert_eq!(tool_name, "shell"); @@ -869,7 +887,7 @@ mod tests { let resp = router.oneshot(req).await.unwrap(); assert_eq!(resp.status(), StatusCode::OK); - let (_recv_id, event) = rx.recv().await.unwrap(); + let (_recv_id, _recv_uid, event) = rx.recv().await.unwrap(); // Unknown event types fall through to JobStatus assert!(matches!(event, SseEvent::JobStatus { .. })); } diff --git a/src/orchestrator/mod.rs b/src/orchestrator/mod.rs index d6e028a5..e68051dc 100644 --- a/src/orchestrator/mod.rs +++ b/src/orchestrator/mod.rs @@ -63,7 +63,7 @@ fn resolve_orchestrator_port() -> u16 { /// Result of orchestrator setup, containing all handles needed by the agent. pub struct OrchestratorSetup { pub container_job_manager: Option>, - pub job_event_tx: Option>, + pub job_event_tx: Option>, pub prompt_queue: Arc>>>, pub docker_status: crate::sandbox::DockerStatus, } diff --git a/src/tools/builtin/extension_tools.rs b/src/tools/builtin/extension_tools.rs index cb0f71dd..fba61613 100644 --- a/src/tools/builtin/extension_tools.rs +++ b/src/tools/builtin/extension_tools.rs @@ -130,7 +130,7 @@ impl Tool for ToolInstallTool { async fn execute( &self, params: serde_json::Value, - _ctx: &JobContext, + ctx: &JobContext, ) -> Result { let start = std::time::Instant::now(); @@ -150,7 +150,7 @@ impl Tool for ToolInstallTool { let result = self .manager - .install(name, url, kind_hint) + .install(name, url, kind_hint, &ctx.user_id) .await .map_err(|e| ToolError::ExecutionFailed(e.to_string()))?; @@ -205,7 +205,7 @@ impl Tool for ToolAuthTool { async fn execute( &self, params: serde_json::Value, - _ctx: &JobContext, + ctx: &JobContext, ) -> Result { let start = std::time::Instant::now(); @@ -213,13 +213,13 @@ impl Tool for ToolAuthTool { let result = self .manager - .auth(name) + .auth(name, &ctx.user_id) .await .map_err(|e| ToolError::ExecutionFailed(e.to_string()))?; // Auto-activate after successful auth so tools are available immediately if result.is_authenticated() { - match self.manager.activate(name).await { + match self.manager.activate(name, &ctx.user_id).await { Ok(activate_result) => { let output = serde_json::json!({ "status": "authenticated_and_activated", @@ -304,13 +304,13 @@ impl Tool for ToolActivateTool { async fn execute( &self, params: serde_json::Value, - _ctx: &JobContext, + ctx: &JobContext, ) -> Result { let start = std::time::Instant::now(); let name = require_str(¶ms, "name")?; - match self.manager.activate(name).await { + match self.manager.activate(name, &ctx.user_id).await { Ok(result) => { let output = serde_json::to_value(&result) .unwrap_or_else(|_| serde_json::json!({"error": "serialization failed"})); @@ -329,12 +329,12 @@ impl Tool for ToolActivateTool { // Activation failed due to missing auth; initiate auth flow // so the agent loop can show the auth card. - match self.manager.auth(name).await { + match self.manager.auth(name, &ctx.user_id).await { Ok(auth_result) if auth_result.is_authenticated() => { // Auth succeeded (e.g. env var was set); retry activation. let result = self .manager - .activate(name) + .activate(name, &ctx.user_id) .await .map_err(|e| ToolError::ExecutionFailed(e.to_string()))?; let output = serde_json::to_value(&result).unwrap_or_else( @@ -404,7 +404,7 @@ impl Tool for ToolListTool { async fn execute( &self, params: serde_json::Value, - _ctx: &JobContext, + ctx: &JobContext, ) -> Result { let start = std::time::Instant::now(); @@ -425,7 +425,7 @@ impl Tool for ToolListTool { let extensions = self .manager - .list(kind_filter, include_available) + .list(kind_filter, include_available, &ctx.user_id) .await .map_err(|e| ToolError::ExecutionFailed(e.to_string()))?; @@ -477,7 +477,7 @@ impl Tool for ToolRemoveTool { async fn execute( &self, params: serde_json::Value, - _ctx: &JobContext, + ctx: &JobContext, ) -> Result { let start = std::time::Instant::now(); @@ -485,7 +485,7 @@ impl Tool for ToolRemoveTool { let message = self .manager - .remove(name) + .remove(name, &ctx.user_id) .await .map_err(|e| ToolError::ExecutionFailed(e.to_string()))?; @@ -541,7 +541,7 @@ impl Tool for ToolUpgradeTool { async fn execute( &self, params: serde_json::Value, - _ctx: &JobContext, + ctx: &JobContext, ) -> Result { let start = std::time::Instant::now(); @@ -549,7 +549,7 @@ impl Tool for ToolUpgradeTool { let result = self .manager - .upgrade(name) + .upgrade(name, &ctx.user_id) .await .map_err(|e| ToolError::ExecutionFailed(e.to_string()))?; @@ -603,7 +603,7 @@ impl Tool for ExtensionInfoTool { async fn execute( &self, params: serde_json::Value, - _ctx: &JobContext, + ctx: &JobContext, ) -> Result { let start = std::time::Instant::now(); @@ -611,7 +611,7 @@ impl Tool for ExtensionInfoTool { let info = self .manager - .extension_info(name) + .extension_info(name, &ctx.user_id) .await .map_err(|e| ToolError::ExecutionFailed(e.to_string()))?; diff --git a/src/tools/builtin/job.rs b/src/tools/builtin/job.rs index 0933ee40..86d7e44d 100644 --- a/src/tools/builtin/job.rs +++ b/src/tools/builtin/job.rs @@ -85,7 +85,7 @@ pub struct CreateJobTool { job_manager: Option>, store: Option>, /// Broadcast sender for job events (used to subscribe a monitor). - event_tx: Option>, + event_tx: Option>, /// Injection channel for pushing messages into the agent loop. inject_tx: Option>, /// Encrypted secrets store for validating credential grants. @@ -120,7 +120,7 @@ impl CreateJobTool { /// monitor that forwards Claude Code output to the main agent loop. pub fn with_monitor_deps( mut self, - event_tx: tokio::sync::broadcast::Sender<(Uuid, SseEvent)>, + event_tx: tokio::sync::broadcast::Sender<(Uuid, String, SseEvent)>, inject_tx: tokio::sync::mpsc::Sender, ) -> Self { self.event_tx = Some(event_tx); diff --git a/src/tools/registry.rs b/src/tools/registry.rs index dff09a5c..9dfee0a3 100644 --- a/src/tools/registry.rs +++ b/src/tools/registry.rs @@ -361,7 +361,7 @@ impl ToolRegistry { job_manager: Option>, store: 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>, prompt_queue: Option, diff --git a/tests/e2e_advanced_traces.rs b/tests/e2e_advanced_traces.rs index 2b9fac29..b3efc8d9 100644 --- a/tests/e2e_advanced_traces.rs +++ b/tests/e2e_advanced_traces.rs @@ -661,7 +661,7 @@ mod advanced { .await .expect("failed to inject test token"); - let activate_result = ext_mgr.activate("mock-notion").await; + let activate_result = ext_mgr.activate("mock-notion", "default").await; assert!( activate_result.is_ok(), "activation failed: {:?}", diff --git a/tests/module_init_integration.rs b/tests/module_init_integration.rs index c75ccc6f..3aea7984 100644 --- a/tests/module_init_integration.rs +++ b/tests/module_init_integration.rs @@ -216,7 +216,7 @@ async fn extension_manager_with_process_manager_constructs() { ); // Verify the manager is functional — list returns Ok. - let result = manager.list(None, false).await; + let result = manager.list(None, false, "test").await; assert!(result.is_ok(), "list should succeed on empty manager"); assert!(result.unwrap().is_empty()); } diff --git a/tests/multi_tenant_integration.rs b/tests/multi_tenant_integration.rs new file mode 100644 index 00000000..6924a1a7 --- /dev/null +++ b/tests/multi_tenant_integration.rs @@ -0,0 +1,1057 @@ +//! Integration tests for multi-tenant auth, isolation, and per-user scoping. +//! +//! These tests verify that multi-tenant infrastructure works correctly: +//! - Token-to-identity mapping via MultiAuthState +//! - Per-user SSE event scoping (user A doesn't see user B's events) +//! - Per-user rate limiting (user A exhausting limit doesn't block user B) +//! - Auth middleware inserts correct UserIdentity into request extensions +//! - WebSocket connections are scoped to the authenticated user + +use std::collections::HashMap; +use std::net::SocketAddr; +use std::sync::Arc; +use std::time::Duration; + +use axum::body::Body; +use axum::http::{Request, StatusCode}; +use axum::middleware; +use axum::routing::{get, post}; +use axum::Router; +use tower::ServiceExt; + +use ironclaw::channels::web::auth::{ + AuthenticatedUser, MultiAuthState, UserIdentity, auth_middleware, +}; +use ironclaw::channels::web::server::{GatewayState, PerUserRateLimiter, RateLimiter}; +use ironclaw::channels::web::sse::SseManager; +use ironclaw::channels::web::test_helpers::TestGatewayBuilder; +use ironclaw::channels::web::ws::WsConnectionTracker; +use ironclaw::context::JobContext; +use ironclaw::db::Database; + +// --------------------------------------------------------------------------- +// Helpers +// --------------------------------------------------------------------------- + +const ALICE_TOKEN: &str = "tok-alice-secret"; +const BOB_TOKEN: &str = "tok-bob-secret"; +const ALICE_USER_ID: &str = "alice"; +const BOB_USER_ID: &str = "bob"; + +/// Build a MultiAuthState with two users. +fn two_user_auth() -> MultiAuthState { + let mut tokens = HashMap::new(); + tokens.insert( + ALICE_TOKEN.to_string(), + UserIdentity { + user_id: ALICE_USER_ID.to_string(), + workspace_read_scopes: Vec::new(), + }, + ); + tokens.insert( + BOB_TOKEN.to_string(), + UserIdentity { + user_id: BOB_USER_ID.to_string(), + workspace_read_scopes: vec!["shared".to_string()], + }, + ); + MultiAuthState::multi(tokens) +} + +/// Build a test Router that echoes the authenticated user_id back. +fn user_echo_app(auth: MultiAuthState) -> Router { + async fn echo_user( + AuthenticatedUser(user): AuthenticatedUser, + ) -> String { + user.user_id + } + + async fn echo_user_with_scopes( + AuthenticatedUser(user): AuthenticatedUser, + ) -> String { + format!("{}:{}", user.user_id, user.workspace_read_scopes.join(",")) + } + + Router::new() + .route("/api/whoami", get(echo_user)) + .route("/api/whoami/scopes", get(echo_user_with_scopes)) + .route("/api/action", post(echo_user)) + .route("/api/chat/events", get(echo_user)) // SSE endpoint (allows query token) + .layer(middleware::from_fn_with_state(auth, auth_middleware)) +} + +// =========================================================================== +// Auth: token-to-identity mapping +// =========================================================================== + +#[tokio::test] +async fn alice_token_resolves_to_alice_identity() { + let app = user_echo_app(two_user_auth()); + let resp = app + .oneshot( + Request::builder() + .uri("/api/whoami") + .header("Authorization", format!("Bearer {ALICE_TOKEN}")) + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + + assert_eq!(resp.status(), StatusCode::OK); + let body = axum::body::to_bytes(resp.into_body(), 1024).await.unwrap(); + assert_eq!(std::str::from_utf8(&body).unwrap(), ALICE_USER_ID); +} + +#[tokio::test] +async fn bob_token_resolves_to_bob_identity() { + let app = user_echo_app(two_user_auth()); + let resp = app + .oneshot( + Request::builder() + .uri("/api/whoami") + .header("Authorization", format!("Bearer {BOB_TOKEN}")) + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + + assert_eq!(resp.status(), StatusCode::OK); + let body = axum::body::to_bytes(resp.into_body(), 1024).await.unwrap(); + assert_eq!(std::str::from_utf8(&body).unwrap(), BOB_USER_ID); +} + +#[tokio::test] +async fn bob_identity_carries_workspace_read_scopes() { + let app = user_echo_app(two_user_auth()); + let resp = app + .oneshot( + Request::builder() + .uri("/api/whoami/scopes") + .header("Authorization", format!("Bearer {BOB_TOKEN}")) + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + + assert_eq!(resp.status(), StatusCode::OK); + let body = axum::body::to_bytes(resp.into_body(), 1024).await.unwrap(); + assert_eq!(std::str::from_utf8(&body).unwrap(), "bob:shared"); +} + +#[tokio::test] +async fn unknown_token_rejected() { + let app = user_echo_app(two_user_auth()); + let resp = app + .oneshot( + Request::builder() + .uri("/api/whoami") + .header("Authorization", "Bearer unknown-token") + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + + assert_eq!(resp.status(), StatusCode::UNAUTHORIZED); +} + +#[tokio::test] +async fn no_token_rejected() { + let app = user_echo_app(two_user_auth()); + let resp = app + .oneshot( + Request::builder() + .uri("/api/whoami") + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + + assert_eq!(resp.status(), StatusCode::UNAUTHORIZED); +} + +#[tokio::test] +async fn alice_token_does_not_authenticate_as_bob() { + let app = user_echo_app(two_user_auth()); + let resp = app + .oneshot( + Request::builder() + .uri("/api/whoami") + .header("Authorization", format!("Bearer {ALICE_TOKEN}")) + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + + let body = axum::body::to_bytes(resp.into_body(), 1024).await.unwrap(); + let user_id = std::str::from_utf8(&body).unwrap(); + assert_eq!(user_id, ALICE_USER_ID); + assert_ne!(user_id, BOB_USER_ID); +} + +// =========================================================================== +// Auth: query token on SSE/WS endpoints +// =========================================================================== + +#[tokio::test] +async fn query_token_works_for_sse_endpoint_multi_user() { + let app = user_echo_app(two_user_auth()); + let resp = app + .oneshot( + Request::builder() + .uri(format!("/api/chat/events?token={ALICE_TOKEN}")) + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + + assert_eq!(resp.status(), StatusCode::OK); + let body = axum::body::to_bytes(resp.into_body(), 1024).await.unwrap(); + assert_eq!(std::str::from_utf8(&body).unwrap(), ALICE_USER_ID); +} + +#[tokio::test] +async fn query_token_rejected_for_non_sse_endpoint_multi_user() { + let app = user_echo_app(two_user_auth()); + let resp = app + .oneshot( + Request::builder() + .uri(format!("/api/whoami?token={ALICE_TOKEN}")) + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + + assert_eq!(resp.status(), StatusCode::UNAUTHORIZED); +} + +#[tokio::test] +async fn query_token_rejected_for_post_multi_user() { + let app = user_echo_app(two_user_auth()); + let resp = app + .oneshot( + Request::builder() + .method("POST") + .uri(format!("/api/action?token={ALICE_TOKEN}")) + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + + assert_eq!(resp.status(), StatusCode::UNAUTHORIZED); +} + +// =========================================================================== +// Per-user rate limiting +// =========================================================================== + +#[test] +fn per_user_rate_limiter_isolates_users() { + let limiter = PerUserRateLimiter::new(3, 60); + + // Alice uses all 3 requests + assert!(limiter.check("alice")); + assert!(limiter.check("alice")); + assert!(limiter.check("alice")); + // Alice is now rate-limited + assert!(!limiter.check("alice")); + + // Bob is unaffected — gets his own 3 requests + assert!(limiter.check("bob")); + assert!(limiter.check("bob")); + assert!(limiter.check("bob")); + assert!(!limiter.check("bob")); +} + +#[test] +fn per_user_rate_limiter_different_users_independent() { + let limiter = PerUserRateLimiter::new(2, 60); + + // Interleave requests from different users + assert!(limiter.check("alice")); + assert!(limiter.check("bob")); + assert!(limiter.check("alice")); + assert!(limiter.check("bob")); + + // Both exhausted independently + assert!(!limiter.check("alice")); + assert!(!limiter.check("bob")); + + // Charlie is fresh + assert!(limiter.check("charlie")); +} + +#[test] +fn per_user_rate_limiter_single_user_mode() { + // In single-user mode, only one user_id is used + let limiter = PerUserRateLimiter::new(5, 60); + for _ in 0..5 { + assert!(limiter.check("default")); + } + assert!(!limiter.check("default")); +} + +// =========================================================================== +// SSE event scoping +// =========================================================================== + +#[tokio::test] +async fn sse_scoped_event_only_delivered_to_target_user() { + use ironclaw::channels::web::types::SseEvent; + use tokio_stream::StreamExt; + + let manager = SseManager::new(); + let mut alice_stream = Box::pin( + manager + .subscribe_raw(Some(ALICE_USER_ID.to_string())) + .expect("subscribe"), + ); + let mut bob_stream = Box::pin( + manager + .subscribe_raw(Some(BOB_USER_ID.to_string())) + .expect("subscribe"), + ); + + // Send event scoped to alice + manager.broadcast_for_user( + ALICE_USER_ID, + SseEvent::Status { + message: "alice's event".to_string(), + thread_id: None, + }, + ); + + // Send global heartbeat (both should get it) + manager.broadcast(SseEvent::Heartbeat); + + // Alice gets her scoped event first + let e = alice_stream.next().await.unwrap(); + match &e { + SseEvent::Status { message, .. } => assert_eq!(message, "alice's event"), + _ => panic!("Expected Status, got {:?}", e), + } + + // Alice also gets heartbeat + let e = alice_stream.next().await.unwrap(); + assert!(matches!(e, SseEvent::Heartbeat)); + + // Bob only gets the heartbeat (alice's event was filtered) + let e = bob_stream.next().await.unwrap(); + assert!(matches!(e, SseEvent::Heartbeat)); +} + +#[tokio::test] +async fn sse_global_event_delivered_to_all_users() { + use ironclaw::channels::web::types::SseEvent; + use tokio_stream::StreamExt; + + let manager = SseManager::new(); + let mut alice = Box::pin( + manager + .subscribe_raw(Some(ALICE_USER_ID.to_string())) + .expect("subscribe"), + ); + let mut bob = Box::pin( + manager + .subscribe_raw(Some(BOB_USER_ID.to_string())) + .expect("subscribe"), + ); + + manager.broadcast(SseEvent::Status { + message: "global announcement".to_string(), + thread_id: None, + }); + + let ea = alice.next().await.unwrap(); + let eb = bob.next().await.unwrap(); + match (&ea, &eb) { + (SseEvent::Status { message: a, .. }, SseEvent::Status { message: b, .. }) => { + assert_eq!(a, "global announcement"); + assert_eq!(b, "global announcement"); + } + _ => panic!("Expected Status events"), + } +} + +#[tokio::test] +async fn sse_user_b_event_not_visible_to_user_a() { + use ironclaw::channels::web::types::SseEvent; + use tokio_stream::StreamExt; + + let manager = SseManager::new(); + let mut alice = Box::pin( + manager + .subscribe_raw(Some(ALICE_USER_ID.to_string())) + .expect("subscribe"), + ); + + // Send event for bob only + manager.broadcast_for_user( + BOB_USER_ID, + SseEvent::Response { + content: "bob's secret".to_string(), + thread_id: "t1".to_string(), + }, + ); + + // Send heartbeat so alice has something to receive + manager.broadcast(SseEvent::Heartbeat); + + // Alice should only get heartbeat, not bob's response + let e = alice.next().await.unwrap(); + assert!( + matches!(e, SseEvent::Heartbeat), + "Expected Heartbeat, got {:?}", + e + ); +} + +#[tokio::test] +async fn sse_unscoped_subscriber_receives_all_events() { + use ironclaw::channels::web::types::SseEvent; + use tokio_stream::StreamExt; + + let manager = SseManager::new(); + // Unscoped subscriber (None user_id) — backwards-compatible single-user mode + let mut stream = Box::pin(manager.subscribe_raw(None).expect("subscribe")); + + manager.broadcast_for_user( + ALICE_USER_ID, + SseEvent::Status { + message: "alice only".to_string(), + thread_id: None, + }, + ); + manager.broadcast_for_user( + BOB_USER_ID, + SseEvent::Status { + message: "bob only".to_string(), + thread_id: None, + }, + ); + manager.broadcast(SseEvent::Heartbeat); + + // Unscoped subscriber gets ALL three events + let e1 = stream.next().await.unwrap(); + let e2 = stream.next().await.unwrap(); + let e3 = stream.next().await.unwrap(); + + match &e1 { + SseEvent::Status { message, .. } => assert_eq!(message, "alice only"), + _ => panic!("Expected alice's Status"), + } + match &e2 { + SseEvent::Status { message, .. } => assert_eq!(message, "bob only"), + _ => panic!("Expected bob's Status"), + } + assert!(matches!(e3, SseEvent::Heartbeat)); +} + +// =========================================================================== +// MultiAuthState: edge cases +// =========================================================================== + +#[test] +fn multi_auth_state_empty_token_not_valid() { + let state = MultiAuthState::single("real-token".to_string(), "user1".to_string()); + assert!(state.authenticate("").is_none()); +} + +#[test] +fn multi_auth_state_first_token_returns_any_token() { + let auth = two_user_auth(); + let first = auth.first_token().unwrap(); + // Should be one of the two tokens (HashMap ordering is non-deterministic) + assert!(first == ALICE_TOKEN || first == BOB_TOKEN); +} + +#[test] +fn multi_auth_state_first_identity_returns_valid_user() { + let auth = two_user_auth(); + let identity = auth.first_identity().unwrap(); + assert!(identity.user_id == ALICE_USER_ID || identity.user_id == BOB_USER_ID); +} + +#[test] +fn multi_auth_state_token_prefix_not_valid() { + // Ensure partial token matches don't authenticate + let state = MultiAuthState::single("secret-token-123".to_string(), "user1".to_string()); + assert!(state.authenticate("secret-token").is_none()); + assert!(state.authenticate("secret-token-1234").is_none()); + assert!(state.authenticate("secret-token-123").is_some()); +} + +// =========================================================================== +// Connection counting with user scoping +// =========================================================================== + +#[tokio::test] +async fn sse_connection_count_tracks_scoped_subscribers() { + let manager = SseManager::new(); + assert_eq!(manager.connection_count(), 0); + + let _alice = Box::pin( + manager + .subscribe_raw(Some(ALICE_USER_ID.to_string())) + .expect("subscribe"), + ); + assert_eq!(manager.connection_count(), 1); + + let _bob = Box::pin( + manager + .subscribe_raw(Some(BOB_USER_ID.to_string())) + .expect("subscribe"), + ); + assert_eq!(manager.connection_count(), 2); + + drop(_alice); + assert_eq!(manager.connection_count(), 1); + + drop(_bob); + assert_eq!(manager.connection_count(), 0); +} + +// =========================================================================== +// GatewayState construction: multi-user fields +// =========================================================================== + +#[test] +fn gateway_state_has_multi_tenant_fields() { + // Verify the GatewayState struct accepts all multi-tenant fields. + // This is a compile-time check that the conflict resolution didn't + // drop any fields. + let state = GatewayState { + msg_tx: tokio::sync::RwLock::new(None), + sse: Arc::new(SseManager::new()), + workspace: None, + workspace_pool: None, // Multi-tenant: per-user workspace pool + 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: "fallback".to_string(), // Multi-tenant: renamed from user_id + shutdown_tx: tokio::sync::RwLock::new(None), + ws_tracker: Some(Arc::new(WsConnectionTracker::new())), + llm_provider: None, + skill_registry: None, + skill_catalog: None, + chat_rate_limiter: PerUserRateLimiter::new(30, 60), // Multi-tenant: per-user + oauth_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(), + webhook_rate_limiter: RateLimiter::new(10, 60), + active_config: Default::default(), + }; + + assert_eq!(state.default_user_id, "fallback"); + assert!(state.workspace_pool.is_none()); +} + +// =========================================================================== +// Full-server handler-level tests (real HTTP through auth middleware) +// =========================================================================== + +/// Build a MultiAuthState with two users and start a real server. +async fn start_multi_user_server() -> (SocketAddr, Arc) { + let (agent_tx, _agent_rx) = tokio::sync::mpsc::channel(64); + let auth = two_user_auth(); + TestGatewayBuilder::new() + .msg_tx(agent_tx) + .start_multi(auth) + .await + .expect("Failed to start multi-user test server") +} + +#[tokio::test] +async fn full_server_alice_can_access_protected_endpoint() { + let (addr, _state) = start_multi_user_server().await; + + let client = reqwest::Client::new(); + let resp = client + .get(format!("http://{}/api/gateway/status", addr)) + .header("Authorization", format!("Bearer {}", ALICE_TOKEN)) + .send() + .await + .unwrap(); + + assert_eq!(resp.status(), 200); +} + +#[tokio::test] +async fn full_server_bob_can_access_protected_endpoint() { + let (addr, _state) = start_multi_user_server().await; + + let client = reqwest::Client::new(); + let resp = client + .get(format!("http://{}/api/gateway/status", addr)) + .header("Authorization", format!("Bearer {}", BOB_TOKEN)) + .send() + .await + .unwrap(); + + assert_eq!(resp.status(), 200); +} + +#[tokio::test] +async fn full_server_unknown_token_returns_401() { + let (addr, _state) = start_multi_user_server().await; + + let client = reqwest::Client::new(); + let resp = client + .get(format!("http://{}/api/gateway/status", addr)) + .header("Authorization", "Bearer wrong-token") + .send() + .await + .unwrap(); + + assert_eq!(resp.status(), 401); +} + +#[tokio::test] +async fn full_server_no_auth_header_returns_401() { + let (addr, _state) = start_multi_user_server().await; + + let client = reqwest::Client::new(); + let resp = client + .get(format!("http://{}/api/gateway/status", addr)) + .send() + .await + .unwrap(); + + assert_eq!(resp.status(), 401); +} + +#[tokio::test] +async fn full_server_health_is_public() { + let (addr, _state) = start_multi_user_server().await; + + let client = reqwest::Client::new(); + let resp = client + .get(format!("http://{}/api/health", addr)) + .send() + .await + .unwrap(); + + assert_eq!(resp.status(), 200); +} + +#[tokio::test] +async fn full_server_chat_send_accepted_for_alice() { + let (agent_tx, mut agent_rx) = tokio::sync::mpsc::channel(64); + let auth = two_user_auth(); + let (addr, _state) = TestGatewayBuilder::new() + .msg_tx(agent_tx) + .start_multi(auth) + .await + .expect("Failed to start server"); + + let client = reqwest::Client::new(); + let resp = client + .post(format!("http://{}/api/chat/send", addr)) + .header("Authorization", format!("Bearer {}", ALICE_TOKEN)) + .header("Content-Type", "application/json") + .body(r#"{"content":"hello from alice"}"#) + .send() + .await + .unwrap(); + + assert_eq!(resp.status(), 202); // ACCEPTED + + // Verify the message reached the agent channel + let msg = tokio::time::timeout(Duration::from_secs(2), agent_rx.recv()) + .await + .expect("Timed out waiting for agent message") + .expect("Agent channel closed"); + + assert_eq!(msg.content, "hello from alice"); + assert_eq!(msg.channel, "gateway"); +} + +#[tokio::test] +async fn full_server_chat_send_rejected_without_auth() { + let (addr, _state) = start_multi_user_server().await; + + let client = reqwest::Client::new(); + let resp = client + .post(format!("http://{}/api/chat/send", addr)) + .header("Content-Type", "application/json") + .body(r#"{"content":"unauthorized message"}"#) + .send() + .await + .unwrap(); + + assert_eq!(resp.status(), 401); +} + +#[tokio::test] +async fn full_server_query_token_works_for_sse() { + let (addr, _state) = start_multi_user_server().await; + + let client = reqwest::Client::new(); + // SSE endpoint should accept query token + let resp = client + .get(format!( + "http://{}/api/chat/events?token={}", + addr, ALICE_TOKEN + )) + .send() + .await + .unwrap(); + + // Should get 200 (SSE stream starts) + assert_eq!(resp.status(), 200); +} + +#[tokio::test] +async fn full_server_query_token_rejected_for_non_sse() { + let (addr, _state) = start_multi_user_server().await; + + let client = reqwest::Client::new(); + // Non-SSE endpoint should NOT accept query token + let resp = client + .get(format!( + "http://{}/api/gateway/status?token={}", + addr, ALICE_TOKEN + )) + .send() + .await + .unwrap(); + + assert_eq!(resp.status(), 401); +} + +#[tokio::test] +async fn full_server_jobs_endpoint_returns_503_without_db() { + let (addr, _state) = start_multi_user_server().await; + + let client = reqwest::Client::new(); + // Jobs endpoint requires database — should return 503 (no DB configured) + // but NOT 401 (auth should pass) + let resp = client + .get(format!("http://{}/api/jobs", addr)) + .header("Authorization", format!("Bearer {}", ALICE_TOKEN)) + .send() + .await + .unwrap(); + + // Without a database, this should return a server error, not an auth error + let status = resp.status().as_u16(); + assert_ne!(status, 401, "Should not be auth error — token is valid"); + assert_ne!(status, 403, "Should not be forbidden — token is valid"); +} + +#[tokio::test] +async fn full_server_jobs_endpoint_rejected_without_auth() { + let (addr, _state) = start_multi_user_server().await; + + let client = reqwest::Client::new(); + let resp = client + .get(format!("http://{}/api/jobs", addr)) + .send() + .await + .unwrap(); + + assert_eq!(resp.status(), 401); +} + +#[tokio::test] +async fn full_server_ws_multi_user_event_isolation() { + use futures::StreamExt; + use ironclaw::channels::web::types::SseEvent; + use tokio_tungstenite::tungstenite::Message; + use tokio_tungstenite::tungstenite::client::IntoClientRequest; + + let (addr, state) = start_multi_user_server().await; + + // Connect Alice's WS + let alice_url = format!("ws://{}/api/chat/ws?token={}", addr, ALICE_TOKEN); + let mut alice_req = alice_url.into_client_request().unwrap(); + alice_req.headers_mut().insert( + "Origin", + format!("http://127.0.0.1:{}", addr.port()).parse().unwrap(), + ); + let (mut alice_ws, _) = tokio_tungstenite::connect_async(alice_req) + .await + .expect("Alice WS connect failed"); + + // Connect Bob's WS + let bob_url = format!("ws://{}/api/chat/ws?token={}", addr, BOB_TOKEN); + let mut bob_req = bob_url.into_client_request().unwrap(); + bob_req.headers_mut().insert( + "Origin", + format!("http://127.0.0.1:{}", addr.port()).parse().unwrap(), + ); + let (mut bob_ws, _) = tokio_tungstenite::connect_async(bob_req) + .await + .expect("Bob WS connect failed"); + + tokio::time::sleep(Duration::from_millis(100)).await; + + // Broadcast an event scoped to Alice only + state.sse.broadcast_for_user( + ALICE_USER_ID, + SseEvent::Status { + message: "alice-only-event".to_string(), + thread_id: None, + }, + ); + + // Broadcast a global heartbeat so Bob has something to receive + state.sse.broadcast(SseEvent::Heartbeat); + + // Alice should get her scoped event + let alice_msg = tokio::time::timeout(Duration::from_secs(2), alice_ws.next()) + .await + .expect("Alice WS timed out") + .expect("Alice stream ended") + .expect("Alice WS error"); + + if let Message::Text(text) = alice_msg { + let parsed: serde_json::Value = serde_json::from_str(&text).unwrap(); + assert_eq!(parsed["type"], "event"); + assert_eq!(parsed["event_type"], "status"); + assert_eq!(parsed["data"]["message"], "alice-only-event"); + } else { + panic!("Expected Text frame from Alice WS, got {:?}", alice_msg); + } + + // Bob should only get the heartbeat, NOT alice's event + let bob_msg = tokio::time::timeout(Duration::from_secs(2), bob_ws.next()) + .await + .expect("Bob WS timed out") + .expect("Bob stream ended") + .expect("Bob WS error"); + + if let Message::Text(text) = bob_msg { + let parsed: serde_json::Value = serde_json::from_str(&text).unwrap(); + assert_eq!(parsed["type"], "event"); + assert_eq!( + parsed["event_type"], "heartbeat", + "Bob should only see heartbeat, not alice's event. Got: {}", + text + ); + } else { + panic!("Expected Text frame from Bob WS, got {:?}", bob_msg); + } + + alice_ws.close(None).await.ok(); + bob_ws.close(None).await.ok(); +} + +// =========================================================================== +// DB-backed job ownership tests (libSQL in-memory) +// =========================================================================== + +/// Start a multi-user server with a real (in-memory) database. +#[cfg(feature = "libsql")] +async fn start_multi_user_server_with_db() -> ( + SocketAddr, + Arc, + Arc, + tempfile::TempDir, +) { + let temp_dir = tempfile::tempdir().expect("failed to create temp dir"); + let path = temp_dir.path().join("test.db"); + let backend = ironclaw::db::libsql::LibSqlBackend::new_local(&path) + .await + .expect("failed to create test DB"); + backend.run_migrations().await.expect("failed to run migrations"); + let db: Arc = Arc::new(backend); + let (agent_tx, _agent_rx) = tokio::sync::mpsc::channel(64); + let auth = two_user_auth(); + + // Build state manually so we can inject the DB + let state = Arc::new(GatewayState { + msg_tx: tokio::sync::RwLock::new(Some(agent_tx)), + 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: Some(Arc::clone(&db)), + job_manager: None, + prompt_queue: None, + scheduler: None, + default_user_id: ALICE_USER_ID.to_string(), + shutdown_tx: tokio::sync::RwLock::new(None), + ws_tracker: Some(Arc::new(WsConnectionTracker::new())), + llm_provider: None, + skill_registry: None, + skill_catalog: None, + chat_rate_limiter: PerUserRateLimiter::new(30, 60), + oauth_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(), + webhook_rate_limiter: RateLimiter::new(10, 60), + active_config: Default::default(), + }); + + let addr: SocketAddr = "127.0.0.1:0".parse().unwrap(); + let bound = ironclaw::channels::web::server::start_server(addr, state.clone(), auth) + .await + .expect("Failed to start server with DB"); + + (bound, state, db, temp_dir) +} + +#[cfg(feature = "libsql")] +#[tokio::test] +async fn full_server_alice_sees_own_jobs_only() { + let (addr, _state, db, _tmp) = start_multi_user_server_with_db().await; + + // Create jobs owned by Alice and Bob + let alice_job = JobContext::with_user(ALICE_USER_ID, "Alice's job", "Alice's work"); + let bob_job = JobContext::with_user(BOB_USER_ID, "Bob's job", "Bob's work"); + let alice_job_id = alice_job.job_id; + + db.save_job(&alice_job).await.unwrap(); + db.save_job(&bob_job).await.unwrap(); + + let client = reqwest::Client::new(); + + // Alice lists jobs — should only see her own + let resp = client + .get(format!("http://{}/api/jobs", addr)) + .header("Authorization", format!("Bearer {}", ALICE_TOKEN)) + .send() + .await + .unwrap(); + + assert_eq!(resp.status(), 200); + let body: serde_json::Value = resp.json().await.unwrap(); + let jobs = body["jobs"].as_array().unwrap(); + + // Alice should see exactly 1 job + assert_eq!(jobs.len(), 1, "Alice should see only her own job"); + assert_eq!(jobs[0]["id"], alice_job_id.to_string()); + assert_eq!(jobs[0]["title"], "Alice's job"); +} + +#[cfg(feature = "libsql")] +#[tokio::test] +async fn full_server_bob_cannot_see_alice_job_detail() { + let (addr, _state, db, _tmp) = start_multi_user_server_with_db().await; + + // Create a job owned by Alice + let alice_job = JobContext::with_user(ALICE_USER_ID, "Alice's secret job", "Private"); + let alice_job_id = alice_job.job_id; + db.save_job(&alice_job).await.unwrap(); + + let client = reqwest::Client::new(); + + // Bob tries to access Alice's job by ID — should get 404 (not 403, to prevent enumeration) + let resp = client + .get(format!("http://{}/api/jobs/{}", addr, alice_job_id)) + .header("Authorization", format!("Bearer {}", BOB_TOKEN)) + .send() + .await + .unwrap(); + + assert_eq!( + resp.status(), + 404, + "Bob should not be able to see Alice's job" + ); +} + +#[cfg(feature = "libsql")] +#[tokio::test] +async fn full_server_alice_can_see_own_job_detail() { + let (addr, _state, db, _tmp) = start_multi_user_server_with_db().await; + + let alice_job = JobContext::with_user(ALICE_USER_ID, "Alice's visible job", "Details here"); + let alice_job_id = alice_job.job_id; + db.save_job(&alice_job).await.unwrap(); + + let client = reqwest::Client::new(); + + let resp = client + .get(format!("http://{}/api/jobs/{}", addr, alice_job_id)) + .header("Authorization", format!("Bearer {}", ALICE_TOKEN)) + .send() + .await + .unwrap(); + + assert_eq!(resp.status(), 200); + let body: serde_json::Value = resp.json().await.unwrap(); + assert_eq!(body["id"], alice_job_id.to_string()); + assert_eq!(body["title"], "Alice's visible job"); +} + +#[cfg(feature = "libsql")] +#[tokio::test] +async fn full_server_bob_sees_own_jobs_only() { + let (addr, _state, db, _tmp) = start_multi_user_server_with_db().await; + + // Create multiple jobs for each user + for i in 0..3 { + let aj = JobContext::with_user(ALICE_USER_ID, format!("Alice job {}", i), ""); + db.save_job(&aj).await.unwrap(); + } + for i in 0..2 { + let bj = JobContext::with_user(BOB_USER_ID, format!("Bob job {}", i), ""); + db.save_job(&bj).await.unwrap(); + } + + let client = reqwest::Client::new(); + + // Bob lists jobs + let resp = client + .get(format!("http://{}/api/jobs", addr)) + .header("Authorization", format!("Bearer {}", BOB_TOKEN)) + .send() + .await + .unwrap(); + + assert_eq!(resp.status(), 200); + let body: serde_json::Value = resp.json().await.unwrap(); + let jobs = body["jobs"].as_array().unwrap(); + + assert_eq!(jobs.len(), 2, "Bob should see only his 2 jobs, not Alice's 3"); + for job in jobs { + let title = job["title"].as_str().unwrap(); + assert!( + title.starts_with("Bob job"), + "Bob should only see his own jobs, got: {}", + title + ); + } +} + +#[cfg(feature = "libsql")] +#[tokio::test] +async fn full_server_nonexistent_job_returns_404() { + let (addr, _state, _db, _tmp) = start_multi_user_server_with_db().await; + + let client = reqwest::Client::new(); + let fake_id = uuid::Uuid::new_v4(); + + let resp = client + .get(format!("http://{}/api/jobs/{}", addr, fake_id)) + .header("Authorization", format!("Bearer {}", ALICE_TOKEN)) + .send() + .await + .unwrap(); + + assert_eq!(resp.status(), 404); +} diff --git a/tests/multi_tenant_system_prompt.rs b/tests/multi_tenant_system_prompt.rs new file mode 100644 index 00000000..ce6a5e94 --- /dev/null +++ b/tests/multi_tenant_system_prompt.rs @@ -0,0 +1,254 @@ +//! 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 = (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 = 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, 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]) -> Option { + 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(); + } +} diff --git a/tests/openai_compat_integration.rs b/tests/openai_compat_integration.rs index 2a472d00..c3892ace 100644 --- a/tests/openai_compat_integration.rs +++ b/tests/openai_compat_integration.rs @@ -191,8 +191,9 @@ async fn start_test_server_with_provider( ) -> (SocketAddr, Arc) { let state = Arc::new(GatewayState { msg_tx: tokio::sync::RwLock::new(None), - sse: SseManager::new(), + sse: Arc::new(SseManager::new()), workspace: None, + workspace_pool: None, session_manager: None, log_broadcaster: None, log_level_handle: None, @@ -202,13 +203,13 @@ async fn start_test_server_with_provider( job_manager: None, prompt_queue: None, scheduler: None, - user_id: "test-user".to_string(), + default_user_id: "test-user".to_string(), shutdown_tx: tokio::sync::RwLock::new(None), ws_tracker: Some(Arc::new(WsConnectionTracker::new())), llm_provider: Some(llm_provider), skill_registry: None, skill_catalog: None, - chat_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(30, 60), + chat_rate_limiter: ironclaw::channels::web::server::PerUserRateLimiter::new(30, 60), oauth_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(10, 60), webhook_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(10, 60), registry_entries: Vec::new(), @@ -218,8 +219,12 @@ async fn start_test_server_with_provider( active_config: ironclaw::channels::web::server::ActiveConfigSnapshot::default(), }); + let auth = ironclaw::channels::web::auth::MultiAuthState::single( + AUTH_TOKEN.to_string(), + "test-user".to_string(), + ); let addr: SocketAddr = "127.0.0.1:0".parse().unwrap(); - let bound_addr = start_server(addr, state.clone(), AUTH_TOKEN.to_string()) + let bound_addr = start_server(addr, state.clone(), auth) .await .expect("Failed to start test server"); @@ -684,8 +689,9 @@ async fn test_no_llm_provider_returns_503() { // Create state WITHOUT llm_provider let state = Arc::new(GatewayState { msg_tx: tokio::sync::RwLock::new(None), - sse: SseManager::new(), + sse: Arc::new(SseManager::new()), workspace: None, + workspace_pool: None, session_manager: None, log_broadcaster: None, log_level_handle: None, @@ -695,13 +701,13 @@ async fn test_no_llm_provider_returns_503() { job_manager: None, prompt_queue: None, scheduler: None, - user_id: "test-user".to_string(), + default_user_id: "test-user".to_string(), shutdown_tx: tokio::sync::RwLock::new(None), ws_tracker: Some(Arc::new(WsConnectionTracker::new())), llm_provider: None, // No LLM! skill_registry: None, skill_catalog: None, - chat_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(30, 60), + chat_rate_limiter: ironclaw::channels::web::server::PerUserRateLimiter::new(30, 60), oauth_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(10, 60), webhook_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(10, 60), registry_entries: Vec::new(), @@ -711,8 +717,12 @@ async fn test_no_llm_provider_returns_503() { active_config: ironclaw::channels::web::server::ActiveConfigSnapshot::default(), }); + let auth = ironclaw::channels::web::auth::MultiAuthState::single( + AUTH_TOKEN.to_string(), + "test-user".to_string(), + ); let addr: SocketAddr = "127.0.0.1:0".parse().unwrap(); - let bound_addr = start_server(addr, state, AUTH_TOKEN.to_string()) + let bound_addr = start_server(addr, state, auth) .await .unwrap(); @@ -741,9 +751,10 @@ async fn test_chat_completions_body_too_large() { let state = ironclaw::channels::web::test_helpers::TestGatewayBuilder::new() .llm_provider(llm_provider) .build(); - let auth_state = ironclaw::channels::web::auth::AuthState { - token: AUTH_TOKEN.to_string(), - }; + let auth_state = ironclaw::channels::web::auth::MultiAuthState::single( + AUTH_TOKEN.to_string(), + "test-user".to_string(), + ); let app = Router::new() .route( diff --git a/tests/support/gateway_workflow_harness.rs b/tests/support/gateway_workflow_harness.rs index d33c6fe0..a6353f27 100644 --- a/tests/support/gateway_workflow_harness.rs +++ b/tests/support/gateway_workflow_harness.rs @@ -14,7 +14,8 @@ use ironclaw::agent::{Agent, AgentDeps, SessionManager as AgentSessionManager}; use ironclaw::app::{AppBuilder, AppBuilderFlags}; use ironclaw::channels::IncomingMessage; use ironclaw::channels::web::log_layer::LogBroadcaster; -use ironclaw::channels::web::server::{GatewayState, RateLimiter, start_server}; +use ironclaw::channels::web::auth::MultiAuthState; +use ironclaw::channels::web::server::{GatewayState, PerUserRateLimiter, RateLimiter, start_server}; use ironclaw::channels::web::sse::SseManager; use ironclaw::channels::web::ws::WsConnectionTracker; use ironclaw::config::{Config, RegistryProviderConfig, RoutineConfig}; @@ -211,8 +212,9 @@ impl GatewayWorkflowHarness { let gateway_state = Arc::new(GatewayState { msg_tx: tokio::sync::RwLock::new(Some(gw_tx)), - sse: SseManager::new(), + sse: Arc::new(SseManager::new()), workspace: components.workspace.clone(), + workspace_pool: None, session_manager: Some(Arc::clone(&agent_session_manager)), log_broadcaster: None, log_level_handle: None, @@ -222,13 +224,13 @@ impl GatewayWorkflowHarness { job_manager: None, prompt_queue: None, scheduler: Some(scheduler_slot.clone()), - user_id: user_id.clone(), + default_user_id: user_id.clone(), shutdown_tx: tokio::sync::RwLock::new(None), ws_tracker: Some(Arc::new(WsConnectionTracker::new())), llm_provider: Some(Arc::clone(&components.llm)), skill_registry: components.skill_registry.clone(), skill_catalog: components.skill_catalog.clone(), - chat_rate_limiter: RateLimiter::new(120, 60), + chat_rate_limiter: PerUserRateLimiter::new(120, 60), oauth_rate_limiter: RateLimiter::new(10, 60), webhook_rate_limiter: RateLimiter::new(10, 60), registry_entries: Vec::new(), @@ -254,7 +256,7 @@ impl GatewayWorkflowHarness { skills_config: components.config.skills.clone(), hooks: components.hooks, cost_guard: components.cost_guard, - sse_tx: Some(gateway_state.sse.sender()), + sse_tx: None, http_interceptor: None, transcription: None, document_extraction: None, @@ -288,10 +290,11 @@ impl GatewayWorkflowHarness { } let auth_token = "gateway-test-token".to_string(); + let auth = MultiAuthState::single(auth_token.clone(), user_id.clone()); let addr = start_server( "127.0.0.1:0".parse().expect("valid localhost addr"), Arc::clone(&gateway_state), - auth_token.clone(), + auth, ) .await .expect("failed to start gateway server"); diff --git a/tests/ws_gateway_integration.rs b/tests/ws_gateway_integration.rs index 556c5dcc..43277389 100644 --- a/tests/ws_gateway_integration.rs +++ b/tests/ws_gateway_integration.rs @@ -39,8 +39,9 @@ async fn start_test_server() -> ( let state = Arc::new(GatewayState { msg_tx: tokio::sync::RwLock::new(Some(agent_tx)), - sse: SseManager::new(), + sse: Arc::new(SseManager::new()), workspace: None, + workspace_pool: None, session_manager: None, log_broadcaster: None, log_level_handle: None, @@ -50,13 +51,13 @@ async fn start_test_server() -> ( job_manager: None, prompt_queue: None, scheduler: None, - user_id: "test-user".to_string(), + default_user_id: "test-user".to_string(), shutdown_tx: tokio::sync::RwLock::new(None), ws_tracker: Some(Arc::new(WsConnectionTracker::new())), llm_provider: None, skill_registry: None, skill_catalog: None, - chat_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(30, 60), + chat_rate_limiter: ironclaw::channels::web::server::PerUserRateLimiter::new(30, 60), oauth_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(10, 60), webhook_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(10, 60), registry_entries: Vec::new(), @@ -66,8 +67,12 @@ async fn start_test_server() -> ( active_config: ironclaw::channels::web::server::ActiveConfigSnapshot::default(), }); + let auth = ironclaw::channels::web::auth::MultiAuthState::single( + AUTH_TOKEN.to_string(), + "test-user".to_string(), + ); let addr: SocketAddr = "127.0.0.1:0".parse().unwrap(); - let bound_addr = start_server(addr, state.clone(), AUTH_TOKEN.to_string()) + let bound_addr = start_server(addr, state.clone(), auth) .await .expect("Failed to start test server");