From 27c43e185fb4d5045117b4051069c1184dd3e539 Mon Sep 17 00:00:00 2001 From: "ilblackdragon@gmail.com" Date: Tue, 24 Mar 2026 14:31:02 -0700 Subject: [PATCH] feat(web): DB-backed auth, user/token/invitation API handlers MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Adds the web gateway layer for DB-backed user management (#1605): Auth refactor: - CombinedAuthState wraps env-var tokens (MultiAuthState) + optional DbAuthenticator for DB-backed token lookup with LRU cache (60s TTL, 1024 max entries) - auth_middleware tries env-var tokens first, then DB fallback - From impl for backward compatibility - main.rs wires with_db_auth when database is available API handlers (12 new endpoints): - /api/admin/users — CRUD: create, list, detail, update, suspend, activate - /api/tokens — create (returns plaintext once), list, revoke - /api/invitations — create, list, accept (creates user + first token) Token creation: 32 random bytes → hex plaintext, SHA-256 hash stored. Invitation accept: validates hash + pending + not expired, creates user record and first API token atomically. All test files updated for CombinedAuthState type change. Co-Authored-By: Claude Opus 4.6 (1M context) --- src/channels/web/auth.rs | 158 ++++++++++++-- src/channels/web/handlers/invitations.rs | 229 ++++++++++++++++++++ src/channels/web/handlers/mod.rs | 3 + src/channels/web/handlers/tokens.rs | 132 ++++++++++++ src/channels/web/handlers/users.rs | 245 ++++++++++++++++++++++ src/channels/web/mod.rs | 23 +- src/channels/web/server.rs | 48 ++++- src/channels/web/test_helpers.rs | 4 +- src/channels/web/tests/multi_tenant.rs | 15 +- src/main.rs | 1 + tests/multi_tenant_integration.rs | 7 +- tests/openai_compat_integration.rs | 6 +- tests/support/gateway_workflow_harness.rs | 2 +- tests/ws_gateway_integration.rs | 2 +- 14 files changed, 836 insertions(+), 39 deletions(-) create mode 100644 src/channels/web/handlers/invitations.rs create mode 100644 src/channels/web/handlers/tokens.rs create mode 100644 src/channels/web/handlers/users.rs diff --git a/src/channels/web/auth.rs b/src/channels/web/auth.rs index 7dc8adb4..f7a79f32 100644 --- a/src/channels/web/auth.rs +++ b/src/channels/web/auth.rs @@ -13,7 +13,12 @@ use axum::{ response::{IntoResponse, Response}, }; use sha2::{Digest, Sha256}; +use std::sync::Arc; +use std::time::Instant; use subtle::ConstantTimeEq; +use tokio::sync::RwLock; + +use crate::db::Database; /// Identity resolved from a bearer token. #[derive(Debug, Clone)] @@ -108,6 +113,96 @@ impl MultiAuthState { } } +/// DB-backed token authenticator with an in-memory LRU cache. +/// +/// Checks an LRU cache first (TTL 60s), then falls back to a DB query. +/// Cache entries expire naturally — revoking a token or suspending a user +/// has at most 60s of stale authentication before the cache entry expires. +#[derive(Clone)] +#[allow(clippy::type_complexity)] +pub struct DbAuthenticator { + store: Arc, + /// LRU cache: token_hash → (identity, inserted_at). + cache: Arc>>, +} + +impl DbAuthenticator { + /// Cache TTL — how long a successful auth is cached before re-querying the DB. + const CACHE_TTL_SECS: u64 = 60; + /// Maximum cache entries to prevent unbounded growth. + const MAX_CACHE_ENTRIES: usize = 1024; + + pub fn new(store: Arc) -> Self { + Self { + store, + cache: Arc::new(RwLock::new(HashMap::new())), + } + } + + /// Authenticate a token against the database, using cache when possible. + pub async fn authenticate(&self, candidate: &str) -> Option { + let hash = hash_token(candidate); + + // Check cache first + { + let cache = self.cache.read().await; + if let Some((identity, inserted_at)) = cache.get(&hash) + && inserted_at.elapsed().as_secs() < Self::CACHE_TTL_SECS + { + return Some(identity.clone()); + } + } + + // Cache miss or expired — query DB + let (token_record, user_record) = self.store.authenticate_token(&hash).await.ok()??; + + let identity = UserIdentity { + user_id: user_record.id.clone(), + workspace_read_scopes: Vec::new(), // DB-backed users don't have static scopes yet + }; + + // Record token usage (best-effort, don't block auth) + let store = self.store.clone(); + let token_id = token_record.id; + let user_id = user_record.id; + tokio::spawn(async move { + let _ = store.record_token_usage(token_id).await; + let _ = store.record_login(&user_id).await; + }); + + // Update cache + { + let mut cache = self.cache.write().await; + // Evict stale entries if cache is full + if cache.len() >= Self::MAX_CACHE_ENTRIES { + let now = Instant::now(); + cache.retain(|_, (_, ts)| now.duration_since(*ts).as_secs() < Self::CACHE_TTL_SECS); + } + cache.insert(hash, (identity.clone(), Instant::now())); + } + + Some(identity) + } +} + +/// Combined auth state: tries env-var tokens first, then DB-backed tokens. +#[derive(Clone)] +pub struct CombinedAuthState { + /// In-memory tokens from GATEWAY_USER_TOKENS or GATEWAY_AUTH_TOKEN. + pub env_auth: MultiAuthState, + /// DB-backed token authenticator (optional — only when a database is available). + pub db_auth: Option, +} + +impl From for CombinedAuthState { + fn from(env_auth: MultiAuthState) -> Self { + Self { + env_auth, + db_auth: None, + } + } +} + /// Axum extractor that provides the authenticated user identity. /// /// Only available on routes behind `auth_middleware`. Extracts the @@ -166,39 +261,58 @@ 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/WS endpoints. +/// Tries env-var tokens first (constant-time, in-memory), then falls back +/// to DB-backed token lookup if configured. SSE connections can't set +/// headers from `EventSource`, so we also accept `?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, mut request: Request, next: Next, ) -> Response { - // Try Authorization header first. - // RFC 6750 Section 2.1: auth-scheme comparison is case-insensitive. + // Extract the candidate token from header or query param. + let token = extract_token(&headers, &request); + + if let Some(ref tok) = token { + // 1. Try env-var tokens first (fast, constant-time, in-memory). + if let Some(identity) = auth.env_auth.authenticate(tok) { + request.extensions_mut().insert(identity.clone()); + return next.run(request).await; + } + + // 2. Fall back to DB-backed token lookup. + if let Some(ref db_auth) = auth.db_auth + && let Some(identity) = db_auth.authenticate(tok).await + { + request.extensions_mut().insert(identity); + return next.run(request).await; + } + } + + (StatusCode::UNAUTHORIZED, "Invalid or missing auth token").into_response() +} + +/// Extract a bearer token from the Authorization header or query parameter. +fn extract_token(headers: &HeaderMap, request: &Request) -> Option { + // Try Authorization header first (RFC 6750). 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 ") - && let Some(identity) = auth.authenticate(&value[7..]) { - request.extensions_mut().insert(identity.clone()); - return next.run(request).await; + return Some(value[7..].to_string()); } - // Fall back to query parameter, but only for SSE/WS endpoints. - if allows_query_token_auth(&request) - && let Some(token) = query_token(&request) - && let Some(identity) = auth.authenticate(&token) - { - request.extensions_mut().insert(identity.clone()); - return next.run(request).await; + // Fall back to query parameter for SSE/WS endpoints. + if allows_query_token_auth(request) { + return query_token(request); } - (StatusCode::UNAUTHORIZED, "Invalid or missing auth token").into_response() + None } #[cfg(test)] @@ -274,7 +388,10 @@ mod tests { /// Router with streaming endpoints (query auth allowed) and regular /// endpoints (query auth rejected). fn test_app(token: &str) -> Router { - let state = MultiAuthState::single(token.to_string(), "test-user".to_string()); + let state = CombinedAuthState::from(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)) @@ -486,7 +603,7 @@ mod tests { /// 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); + let state = CombinedAuthState::from(MultiAuthState::multi(tokens)); Router::new() .route("/api/chat/events", get(identity_handler)) .route("/api/chat/send", post(identity_handler)) @@ -643,7 +760,10 @@ mod tests { #[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 state = CombinedAuthState::from(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)); diff --git a/src/channels/web/handlers/invitations.rs b/src/channels/web/handlers/invitations.rs new file mode 100644 index 00000000..74ea96dd --- /dev/null +++ b/src/channels/web/handlers/invitations.rs @@ -0,0 +1,229 @@ +//! Invitation management API handlers. + +use std::sync::Arc; + +use axum::{Json, extract::State, http::StatusCode}; +use rand::RngCore; +use rand::rngs::OsRng; +use sha2::{Digest, Sha256}; +use uuid::Uuid; + +use crate::channels::web::auth::AuthenticatedUser; +use crate::channels::web::server::GatewayState; +use crate::db::{InvitationRecord, UserRecord}; + +/// POST /api/invitations — create an invitation. +pub async fn invitations_create_handler( + State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, + Json(body): Json, +) -> Result, (StatusCode, String)> { + let store = state.store.as_ref().ok_or(( + StatusCode::SERVICE_UNAVAILABLE, + "Database not available".to_string(), + ))?; + + let email = body.get("email").and_then(|v| v.as_str()).map(String::from); + + let expires_in_days = body + .get("expires_in_days") + .and_then(|v| v.as_u64()) + .unwrap_or(7); + + let now = chrono::Utc::now(); + let expires_at = now + chrono::Duration::days(expires_in_days as i64); + + // Generate 32 random bytes for the invite token. + let mut token_bytes = [0u8; 32]; + OsRng.fill_bytes(&mut token_bytes); + let plaintext_token = hex::encode(token_bytes); + + // SHA-256 hash for storage — plaintext is never persisted. + let mut hasher = Sha256::new(); + hasher.update(token_bytes); + let hash: [u8; 32] = hasher.finalize().into(); + + let invitation_id = Uuid::new_v4(); + let invitation = InvitationRecord { + id: invitation_id, + email: email.clone(), + invited_by: user.user_id.clone(), + status: "pending".to_string(), + expires_at, + accepted_at: None, + accepted_by: None, + created_at: now, + }; + + store + .create_invitation(&invitation, &hash) + .await + .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; + + // Return the plaintext token — this is the ONLY time it is shown. + Ok(Json(serde_json::json!({ + "invite_token": plaintext_token, + "id": invitation_id.to_string(), + "email": email, + "expires_at": expires_at.to_rfc3339(), + "created_at": now.to_rfc3339(), + }))) +} + +/// GET /api/invitations — list invitations created by the current user. +pub async fn invitations_list_handler( + State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, +) -> Result, (StatusCode, String)> { + let store = state.store.as_ref().ok_or(( + StatusCode::SERVICE_UNAVAILABLE, + "Database not available".to_string(), + ))?; + + let invitations = store + .list_invitations(Some(&user.user_id)) + .await + .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; + + let invitations_json: Vec = invitations + .into_iter() + .map(|inv| { + serde_json::json!({ + "id": inv.id.to_string(), + "email": inv.email, + "invited_by": inv.invited_by, + "status": inv.status, + "expires_at": inv.expires_at.to_rfc3339(), + "accepted_at": inv.accepted_at.map(|dt| dt.to_rfc3339()), + "accepted_by": inv.accepted_by, + "created_at": inv.created_at.to_rfc3339(), + }) + }) + .collect(); + + Ok(Json(serde_json::json!({ "invitations": invitations_json }))) +} + +/// POST /api/invitations/accept — accept an invitation and create a user account. +pub async fn invitations_accept_handler( + State(state): State>, + Json(body): Json, +) -> Result, (StatusCode, String)> { + let store = state.store.as_ref().ok_or(( + StatusCode::SERVICE_UNAVAILABLE, + "Database not available".to_string(), + ))?; + + let invite_token = body.get("invite_token").and_then(|v| v.as_str()).ok_or(( + StatusCode::BAD_REQUEST, + "Missing required field 'invite_token'".to_string(), + ))?; + + let display_name = body + .get("display_name") + .and_then(|v| v.as_str()) + .ok_or(( + StatusCode::BAD_REQUEST, + "Missing required field 'display_name'".to_string(), + ))? + .to_string(); + + // Hash the provided token to look up the invitation. + let token_bytes = hex::decode(invite_token).map_err(|_| { + ( + StatusCode::BAD_REQUEST, + "Invalid invite token format".to_string(), + ) + })?; + let mut hasher = Sha256::new(); + hasher.update(&token_bytes); + let hash: [u8; 32] = hasher.finalize().into(); + + // Look up the invitation by hash. + let invitation = store + .get_invitation_by_hash(&hash) + .await + .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))? + .ok_or(( + StatusCode::NOT_FOUND, + "Invitation not found or already used".to_string(), + ))?; + + // Verify the invitation is still pending. + if invitation.status != "pending" { + return Err(( + StatusCode::CONFLICT, + format!("Invitation is already '{}'", invitation.status), + )); + } + + // Verify the invitation has not expired. + if invitation.expires_at < chrono::Utc::now() { + return Err((StatusCode::GONE, "Invitation has expired".to_string())); + } + + // Generate a user id from the display name. + let new_user_id = display_name + .to_ascii_lowercase() + .split_whitespace() + .collect::>() + .join("-"); + + let now = chrono::Utc::now(); + let user_record = UserRecord { + id: new_user_id.clone(), + email: invitation.email.clone(), + display_name: display_name.clone(), + status: "active".to_string(), + created_at: now, + updated_at: now, + last_login_at: None, + created_by: Some(invitation.invited_by.clone()), + metadata: serde_json::json!({}), + }; + + store + .create_user(&user_record) + .await + .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; + + // Create a first API token for the new user. + let mut api_token_bytes = [0u8; 32]; + OsRng.fill_bytes(&mut api_token_bytes); + let plaintext_api_token = hex::encode(api_token_bytes); + + let mut api_hasher = Sha256::new(); + api_hasher.update(api_token_bytes); + let api_hash: [u8; 32] = api_hasher.finalize().into(); + + let api_prefix = &plaintext_api_token[..8]; + + let api_token_record = store + .create_api_token(&new_user_id, "default", &api_hash, api_prefix, None) + .await + .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; + + // Mark the invitation as accepted. + store + .accept_invitation(invitation.id, &new_user_id) + .await + .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; + + Ok(Json(serde_json::json!({ + "user": { + "id": new_user_id, + "email": user_record.email, + "display_name": user_record.display_name, + "status": "active", + "created_at": now.to_rfc3339(), + }, + "api_token": { + "token": plaintext_api_token, + "id": api_token_record.id.to_string(), + "name": api_token_record.name, + "token_prefix": api_token_record.token_prefix, + "created_at": api_token_record.created_at.to_rfc3339(), + }, + "invitation_id": invitation.id.to_string(), + }))) +} diff --git a/src/channels/web/handlers/mod.rs b/src/channels/web/handlers/mod.rs index 50c7a0b9..e67403e8 100644 --- a/src/channels/web/handlers/mod.rs +++ b/src/channels/web/handlers/mod.rs @@ -2,10 +2,13 @@ //! //! Each module groups related endpoint handlers by domain. +pub mod invitations; pub mod jobs; pub mod memory; pub mod routines; pub mod skills; +pub mod tokens; +pub mod users; // Modules not yet wired into server.rs router -- suppress dead_code until // they replace their inline counterparts. diff --git a/src/channels/web/handlers/tokens.rs b/src/channels/web/handlers/tokens.rs new file mode 100644 index 00000000..097f87bf --- /dev/null +++ b/src/channels/web/handlers/tokens.rs @@ -0,0 +1,132 @@ +//! API token management handlers. + +use std::sync::Arc; + +use axum::{ + Json, + extract::{Path, State}, + http::StatusCode, +}; +use rand::RngCore; +use rand::rngs::OsRng; +use sha2::{Digest, Sha256}; +use uuid::Uuid; + +use crate::channels::web::auth::AuthenticatedUser; +use crate::channels::web::server::GatewayState; + +/// POST /api/tokens — create a new API token (returns plaintext ONCE). +pub async fn tokens_create_handler( + State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, + Json(body): Json, +) -> Result, (StatusCode, String)> { + let store = state.store.as_ref().ok_or(( + StatusCode::SERVICE_UNAVAILABLE, + "Database not available".to_string(), + ))?; + + let name = body + .get("name") + .and_then(|v| v.as_str()) + .ok_or(( + StatusCode::BAD_REQUEST, + "Missing required field 'name'".to_string(), + ))? + .to_string(); + + let expires_in_days = body.get("expires_in_days").and_then(|v| v.as_u64()); + + let expires_at = + expires_in_days.map(|days| chrono::Utc::now() + chrono::Duration::days(days as i64)); + + // Generate 32 random bytes for the token. + let mut token_bytes = [0u8; 32]; + OsRng.fill_bytes(&mut token_bytes); + let plaintext_token = hex::encode(token_bytes); + + // SHA-256 hash for storage — plaintext is never persisted. + let mut hasher = Sha256::new(); + hasher.update(token_bytes); + let hash: [u8; 32] = hasher.finalize().into(); + + // First 8 chars of the hex token as a prefix for identification. + let token_prefix = &plaintext_token[..8]; + + let record = store + .create_api_token(&user.user_id, &name, &hash, token_prefix, expires_at) + .await + .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; + + // Return the plaintext token — this is the ONLY time it is shown. + Ok(Json(serde_json::json!({ + "token": plaintext_token, + "id": record.id.to_string(), + "name": record.name, + "token_prefix": record.token_prefix, + "expires_at": record.expires_at.map(|dt| dt.to_rfc3339()), + "created_at": record.created_at.to_rfc3339(), + }))) +} + +/// GET /api/tokens — list the current user's tokens (no hashes). +pub async fn tokens_list_handler( + State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, +) -> Result, (StatusCode, String)> { + let store = state.store.as_ref().ok_or(( + StatusCode::SERVICE_UNAVAILABLE, + "Database not available".to_string(), + ))?; + + let tokens = store + .list_api_tokens(&user.user_id) + .await + .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; + + let tokens_json: Vec = tokens + .into_iter() + .map(|t| { + serde_json::json!({ + "id": t.id.to_string(), + "name": t.name, + "token_prefix": t.token_prefix, + "expires_at": t.expires_at.map(|dt| dt.to_rfc3339()), + "last_used_at": t.last_used_at.map(|dt| dt.to_rfc3339()), + "created_at": t.created_at.to_rfc3339(), + "revoked_at": t.revoked_at.map(|dt| dt.to_rfc3339()), + }) + }) + .collect(); + + Ok(Json(serde_json::json!({ "tokens": tokens_json }))) +} + +/// DELETE /api/tokens/{id} — revoke a token. +pub async fn tokens_revoke_handler( + State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, + Path(id): Path, +) -> Result, (StatusCode, String)> { + let store = state.store.as_ref().ok_or(( + StatusCode::SERVICE_UNAVAILABLE, + "Database not available".to_string(), + ))?; + + let token_id = Uuid::parse_str(&id) + .map_err(|_| (StatusCode::BAD_REQUEST, "Invalid token ID".to_string()))?; + + let revoked = store + .revoke_api_token(token_id, &user.user_id) + .await + .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; + + if !revoked { + return Err((StatusCode::NOT_FOUND, "Token not found".to_string())); + } + + Ok(Json(serde_json::json!({ + "status": "revoked", + "id": token_id.to_string(), + }))) +} diff --git a/src/channels/web/handlers/users.rs b/src/channels/web/handlers/users.rs new file mode 100644 index 00000000..4df5a432 --- /dev/null +++ b/src/channels/web/handlers/users.rs @@ -0,0 +1,245 @@ +//! User management API handlers (admin). + +use std::sync::Arc; + +use axum::{ + Json, + extract::{Path, State}, + http::StatusCode, +}; + +use crate::channels::web::auth::AuthenticatedUser; +use crate::channels::web::server::GatewayState; +use crate::db::UserRecord; + +/// POST /api/admin/users — create a new user. +pub async fn users_create_handler( + State(state): State>, + AuthenticatedUser(user): AuthenticatedUser, + Json(body): Json, +) -> Result, (StatusCode, String)> { + let store = state.store.as_ref().ok_or(( + StatusCode::SERVICE_UNAVAILABLE, + "Database not available".to_string(), + ))?; + + let display_name = body + .get("display_name") + .and_then(|v| v.as_str()) + .ok_or(( + StatusCode::BAD_REQUEST, + "Missing required field 'display_name'".to_string(), + ))? + .to_string(); + + let email = body.get("email").and_then(|v| v.as_str()).map(String::from); + + // Generate user id: prefer email if provided, otherwise derive from display_name. + let user_id = if let Some(ref e) = email { + e.clone() + } else { + display_name + .to_ascii_lowercase() + .split_whitespace() + .collect::>() + .join("-") + }; + + let now = chrono::Utc::now(); + let user_record = UserRecord { + id: user_id.clone(), + email, + display_name: display_name.clone(), + status: "active".to_string(), + created_at: now, + updated_at: now, + last_login_at: None, + created_by: Some(user.user_id.clone()), + metadata: serde_json::json!({}), + }; + + store + .create_user(&user_record) + .await + .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; + + Ok(Json(serde_json::json!({ + "id": user_record.id, + "email": user_record.email, + "display_name": user_record.display_name, + "status": user_record.status, + "created_at": user_record.created_at.to_rfc3339(), + "created_by": user_record.created_by, + }))) +} + +/// GET /api/admin/users — list all users. +pub async fn users_list_handler( + State(state): State>, + AuthenticatedUser(_user): AuthenticatedUser, +) -> Result, (StatusCode, String)> { + let store = state.store.as_ref().ok_or(( + StatusCode::SERVICE_UNAVAILABLE, + "Database not available".to_string(), + ))?; + + let users = store + .list_users(None) + .await + .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; + + let users_json: Vec = users + .into_iter() + .map(|u| { + serde_json::json!({ + "id": u.id, + "email": u.email, + "display_name": u.display_name, + "status": u.status, + "created_at": u.created_at.to_rfc3339(), + "updated_at": u.updated_at.to_rfc3339(), + "last_login_at": u.last_login_at.map(|dt| dt.to_rfc3339()), + "created_by": u.created_by, + }) + }) + .collect(); + + Ok(Json(serde_json::json!({ "users": users_json }))) +} + +/// GET /api/admin/users/{id} — get a single user. +pub async fn users_detail_handler( + State(state): State>, + AuthenticatedUser(_user): AuthenticatedUser, + Path(id): Path, +) -> Result, (StatusCode, String)> { + let store = state.store.as_ref().ok_or(( + StatusCode::SERVICE_UNAVAILABLE, + "Database not available".to_string(), + ))?; + + let user_record = store + .get_user(&id) + .await + .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))? + .ok_or((StatusCode::NOT_FOUND, "User not found".to_string()))?; + + Ok(Json(serde_json::json!({ + "id": user_record.id, + "email": user_record.email, + "display_name": user_record.display_name, + "status": user_record.status, + "created_at": user_record.created_at.to_rfc3339(), + "updated_at": user_record.updated_at.to_rfc3339(), + "last_login_at": user_record.last_login_at.map(|dt| dt.to_rfc3339()), + "created_by": user_record.created_by, + "metadata": user_record.metadata, + }))) +} + +/// PATCH /api/admin/users/{id} — update a user's profile. +pub async fn users_update_handler( + State(state): State>, + AuthenticatedUser(_user): AuthenticatedUser, + Path(id): Path, + Json(body): Json, +) -> Result, (StatusCode, String)> { + let store = state.store.as_ref().ok_or(( + StatusCode::SERVICE_UNAVAILABLE, + "Database not available".to_string(), + ))?; + + // Verify the user exists. + let existing = store + .get_user(&id) + .await + .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))? + .ok_or((StatusCode::NOT_FOUND, "User not found".to_string()))?; + + let display_name = body + .get("display_name") + .and_then(|v| v.as_str()) + .unwrap_or(&existing.display_name); + + let metadata = body.get("metadata").unwrap_or(&existing.metadata); + + store + .update_user_profile(&id, display_name, metadata) + .await + .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; + + // Re-fetch the updated record to return consistent data. + let updated = store + .get_user(&id) + .await + .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))? + .ok_or((StatusCode::NOT_FOUND, "User not found".to_string()))?; + + Ok(Json(serde_json::json!({ + "id": updated.id, + "email": updated.email, + "display_name": updated.display_name, + "status": updated.status, + "created_at": updated.created_at.to_rfc3339(), + "updated_at": updated.updated_at.to_rfc3339(), + "metadata": updated.metadata, + }))) +} + +/// POST /api/admin/users/{id}/suspend — suspend a user. +pub async fn users_suspend_handler( + State(state): State>, + AuthenticatedUser(_user): AuthenticatedUser, + Path(id): Path, +) -> Result, (StatusCode, String)> { + let store = state.store.as_ref().ok_or(( + StatusCode::SERVICE_UNAVAILABLE, + "Database not available".to_string(), + ))?; + + // Verify the user exists. + store + .get_user(&id) + .await + .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))? + .ok_or((StatusCode::NOT_FOUND, "User not found".to_string()))?; + + store + .update_user_status(&id, "suspended") + .await + .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; + + Ok(Json(serde_json::json!({ + "id": id, + "status": "suspended", + }))) +} + +/// POST /api/admin/users/{id}/activate — activate a user. +pub async fn users_activate_handler( + State(state): State>, + AuthenticatedUser(_user): AuthenticatedUser, + Path(id): Path, +) -> Result, (StatusCode, String)> { + let store = state.store.as_ref().ok_or(( + StatusCode::SERVICE_UNAVAILABLE, + "Database not available".to_string(), + ))?; + + // Verify the user exists. + store + .get_user(&id) + .await + .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))? + .ok_or((StatusCode::NOT_FOUND, "User not found".to_string()))?; + + store + .update_user_status(&id, "active") + .await + .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; + + Ok(Json(serde_json::json!({ + "id": id, + "status": "active", + }))) +} diff --git a/src/channels/web/mod.rs b/src/channels/web/mod.rs index b26a7829..4f1362c7 100644 --- a/src/channels/web/mod.rs +++ b/src/channels/web/mod.rs @@ -55,7 +55,7 @@ use crate::workspace::Workspace; use self::log_layer::{LogBroadcaster, LogLevelHandle}; -use self::auth::MultiAuthState; +use self::auth::{CombinedAuthState, DbAuthenticator, MultiAuthState}; use self::server::GatewayState; use self::sse::SseManager; use self::types::SseEvent; @@ -64,8 +64,8 @@ use self::types::SseEvent; pub struct GatewayChannel { config: GatewayConfig, state: Arc, - /// Multi-user auth state (replaces bare auth_token). - auth: MultiAuthState, + /// Combined auth state: env-var tokens + optional DB-backed tokens. + auth: CombinedAuthState, } impl GatewayChannel { @@ -82,7 +82,10 @@ impl GatewayChannel { bytes.iter().map(|b| format!("{b:02x}")).collect() }); - let auth = MultiAuthState::single(auth_token, config.user_id.clone()); + let auth = CombinedAuthState { + env_auth: MultiAuthState::single(auth_token, config.user_id.clone()), + db_auth: None, + }; let state = Arc::new(GatewayState { msg_tx: tokio::sync::RwLock::new(None), @@ -123,6 +126,10 @@ impl GatewayChannel { /// Create a gateway channel with a pre-built multi-user auth state. pub fn new_multi_auth(config: GatewayConfig, auth: MultiAuthState) -> Self { + let auth = CombinedAuthState { + env_auth: auth, + db_auth: None, + }; let state = Arc::new(GatewayState { msg_tx: tokio::sync::RwLock::new(None), sse: Arc::new(SseManager::new()), @@ -238,6 +245,12 @@ impl GatewayChannel { self } + /// Enable DB-backed token authentication alongside env-var tokens. + pub fn with_db_auth(mut self, store: Arc) -> Self { + self.auth.db_auth = Some(DbAuthenticator::new(store)); + self + } + /// Inject the container job manager for sandbox operations. pub fn with_job_manager(mut self, jm: Arc) -> Self { self.rebuild_state(|s| s.job_manager = Some(jm)); @@ -316,7 +329,7 @@ impl GatewayChannel { /// Get the first auth token (for printing to console on startup). pub fn auth_token(&self) -> &str { - self.auth.first_token().unwrap_or("") + self.auth.env_auth.first_token().unwrap_or("") } /// Get a reference to the shared gateway state (for the agent to push SSE events). diff --git a/src/channels/web/server.rs b/src/channels/web/server.rs index 63dd6eb2..e44e97b2 100644 --- a/src/channels/web/server.rs +++ b/src/channels/web/server.rs @@ -31,7 +31,7 @@ use crate::bootstrap::ironclaw_base_dir; use crate::channels::IncomingMessage; use crate::channels::relay::DEFAULT_RELAY_NAME; use crate::channels::web::auth::{ - AuthenticatedUser, MultiAuthState, UserIdentity, auth_middleware, + AuthenticatedUser, CombinedAuthState, UserIdentity, auth_middleware, }; use crate::channels::web::handlers::jobs::{ job_files_list_handler, job_files_read_handler, jobs_cancel_handler, jobs_detail_handler, @@ -384,7 +384,7 @@ pub struct GatewayState { pub async fn start_server( addr: SocketAddr, state: Arc, - auth: MultiAuthState, + auth: CombinedAuthState, ) -> Result { let listener = tokio::net::TcpListener::bind(addr).await.map_err(|e| { crate::error::ChannelError::StartupFailed { @@ -510,6 +510,45 @@ pub async fn start_server( "/api/settings/{key}", axum::routing::delete(settings_delete_handler), ) + // User management (admin) + .route( + "/api/admin/users", + get(super::handlers::users::users_list_handler) + .post(super::handlers::users::users_create_handler), + ) + .route( + "/api/admin/users/{id}", + get(super::handlers::users::users_detail_handler) + .patch(super::handlers::users::users_update_handler), + ) + .route( + "/api/admin/users/{id}/suspend", + post(super::handlers::users::users_suspend_handler), + ) + .route( + "/api/admin/users/{id}/activate", + post(super::handlers::users::users_activate_handler), + ) + // Token management + .route( + "/api/tokens", + get(super::handlers::tokens::tokens_list_handler) + .post(super::handlers::tokens::tokens_create_handler), + ) + .route( + "/api/tokens/{id}", + axum::routing::delete(super::handlers::tokens::tokens_revoke_handler), + ) + // Invitations + .route( + "/api/invitations", + get(super::handlers::invitations::invitations_list_handler) + .post(super::handlers::invitations::invitations_create_handler), + ) + .route( + "/api/invitations/accept", + post(super::handlers::invitations::invitations_accept_handler), + ) // Gateway control plane .route("/api/gateway/status", get(gateway_status_handler)) // OpenAI-compatible API @@ -3195,7 +3234,10 @@ mod tests { let state = test_gateway_state(None); let addr: SocketAddr = "127.0.0.1:0".parse().unwrap(); - let auth = MultiAuthState::single("test-token".to_string(), "test".to_string()); + let auth = CombinedAuthState::from(crate::channels::web::auth::MultiAuthState::single( + "test-token".to_string(), + "test".to_string(), + )); let bound = start_server(addr, state.clone(), auth) .await .expect("server should start"); diff --git a/src/channels/web/test_helpers.rs b/src/channels/web/test_helpers.rs index 802512a6..6156e694 100644 --- a/src/channels/web/test_helpers.rs +++ b/src/channels/web/test_helpers.rs @@ -105,7 +105,7 @@ impl TestGatewayBuilder { let addr: SocketAddr = "127.0.0.1:0" .parse() .expect("hard-coded address must parse"); // safety: constant literal - let bound = start_server(addr, state.clone(), auth).await?; + let bound = start_server(addr, state.clone(), auth.into()).await?; Ok((bound, state)) } @@ -119,7 +119,7 @@ impl TestGatewayBuilder { let addr: SocketAddr = "127.0.0.1:0" .parse() .expect("hard-coded address must parse"); // safety: constant literal - let bound = start_server(addr, state.clone(), auth).await?; + let bound = start_server(addr, state.clone(), auth.into()).await?; Ok((bound, state)) } } diff --git a/src/channels/web/tests/multi_tenant.rs b/src/channels/web/tests/multi_tenant.rs index 55010831..56ce03da 100644 --- a/src/channels/web/tests/multi_tenant.rs +++ b/src/channels/web/tests/multi_tenant.rs @@ -340,7 +340,10 @@ mod jobs_isolation { .route("/api/jobs/{id}/cancel", post(jobs_cancel_handler)) .route("/api/jobs/{id}/restart", post(jobs_restart_handler)) .route("/api/jobs/{id}/prompt", post(jobs_prompt_handler)) - .layer(middleware::from_fn_with_state(auth, auth_middleware)) + .layer(middleware::from_fn_with_state( + crate::channels::web::auth::CombinedAuthState::from(auth), + auth_middleware, + )) .with_state(state) } @@ -546,7 +549,10 @@ mod routines_isolation { .route("/api/routines/{id}", get(routines_detail_handler)) .route("/api/routines/{id}/toggle", post(routines_toggle_handler)) .route("/api/routines/{id}", delete(routines_delete_handler)) - .layer(middleware::from_fn_with_state(auth, auth_middleware)) + .layer(middleware::from_fn_with_state( + crate::channels::web::auth::CombinedAuthState::from(auth), + auth_middleware, + )) .with_state(state) } @@ -671,7 +677,10 @@ mod auth_enforcement { .route("/api/logs/level", get(authed_handler).put(authed_handler)) // Gateway status .route("/api/gateway/status", get(authed_handler)) - .layer(middleware::from_fn_with_state(auth, auth_middleware)) + .layer(middleware::from_fn_with_state( + crate::channels::web::auth::CombinedAuthState::from(auth), + auth_middleware, + )) .with_state(state) } diff --git a/src/main.rs b/src/main.rs index eab01264..cffbfceb 100644 --- a/src/main.rs +++ b/src/main.rs @@ -649,6 +649,7 @@ async fn async_main() -> anyhow::Result<()> { } if let Some(ref d) = components.db { gw = gw.with_store(Arc::clone(d)); + gw = gw.with_db_auth(Arc::clone(d)); } if let Some(ref jm) = container_job_manager { gw = gw.with_job_manager(Arc::clone(jm)); diff --git a/tests/multi_tenant_integration.rs b/tests/multi_tenant_integration.rs index 02eb60e8..22c38511 100644 --- a/tests/multi_tenant_integration.rs +++ b/tests/multi_tenant_integration.rs @@ -73,7 +73,10 @@ fn user_echo_app(auth: MultiAuthState) -> Router { .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)) + .layer(middleware::from_fn_with_state( + ironclaw::channels::web::auth::CombinedAuthState::from(auth), + auth_middleware, + )) } // =========================================================================== @@ -905,7 +908,7 @@ async fn start_multi_user_server_with_db() -> ( }); let addr: SocketAddr = "127.0.0.1:0".parse().unwrap(); - let bound = ironclaw::channels::web::server::start_server(addr, state.clone(), auth) + let bound = ironclaw::channels::web::server::start_server(addr, state.clone(), auth.into()) .await .expect("Failed to start server with DB"); diff --git a/tests/openai_compat_integration.rs b/tests/openai_compat_integration.rs index 16568246..4da8c099 100644 --- a/tests/openai_compat_integration.rs +++ b/tests/openai_compat_integration.rs @@ -224,7 +224,7 @@ async fn start_test_server_with_provider( "test-user".to_string(), ); let addr: SocketAddr = "127.0.0.1:0".parse().unwrap(); - let bound_addr = start_server(addr, state.clone(), auth) + let bound_addr = start_server(addr, state.clone(), auth.into()) .await .expect("Failed to start test server"); @@ -722,7 +722,7 @@ async fn test_no_llm_provider_returns_503() { "test-user".to_string(), ); let addr: SocketAddr = "127.0.0.1:0".parse().unwrap(); - let bound_addr = start_server(addr, state, auth).await.unwrap(); + let bound_addr = start_server(addr, state, auth.into()).await.unwrap(); let url = format!("http://{}/v1/chat/completions", bound_addr); let resp = client() @@ -760,7 +760,7 @@ async fn test_chat_completions_body_too_large() { post(ironclaw::channels::web::openai_compat::chat_completions_handler), ) .route_layer(middleware::from_fn_with_state( - auth_state, + ironclaw::channels::web::auth::CombinedAuthState::from(auth_state), ironclaw::channels::web::auth::auth_middleware, )) .layer(DefaultBodyLimit::max(10 * 1024 * 1024)) diff --git a/tests/support/gateway_workflow_harness.rs b/tests/support/gateway_workflow_harness.rs index e4620f70..70dc5452 100644 --- a/tests/support/gateway_workflow_harness.rs +++ b/tests/support/gateway_workflow_harness.rs @@ -297,7 +297,7 @@ impl GatewayWorkflowHarness { let addr = start_server( "127.0.0.1:0".parse().expect("valid localhost addr"), Arc::clone(&gateway_state), - auth, + auth.into(), ) .await .expect("failed to start gateway server"); diff --git a/tests/ws_gateway_integration.rs b/tests/ws_gateway_integration.rs index 43277389..4fd6d7b8 100644 --- a/tests/ws_gateway_integration.rs +++ b/tests/ws_gateway_integration.rs @@ -72,7 +72,7 @@ async fn start_test_server() -> ( "test-user".to_string(), ); let addr: SocketAddr = "127.0.0.1:0".parse().unwrap(); - let bound_addr = start_server(addr, state.clone(), auth) + let bound_addr = start_server(addr, state.clone(), auth.into()) .await .expect("Failed to start test server");