feat(web): DB-backed auth, user/token/invitation API handlers

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<MultiAuthState> 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) <[email protected]>
This commit is contained in:
2026-03-24 14:31:02 -07:00
co-authored by Claude Opus 4.6
parent 80cee7d742
commit 27c43e185f
14 changed files with 836 additions and 39 deletions
+139 -19
View File
@@ -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<dyn Database>,
/// LRU cache: token_hash → (identity, inserted_at).
cache: Arc<RwLock<HashMap<[u8; 32], (UserIdentity, Instant)>>>,
}
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<dyn Database>) -> 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<UserIdentity> {
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<DbAuthenticator>,
}
impl From<MultiAuthState> 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<String> {
/// Auth middleware that validates bearer token from header or query param.
///
/// SSE connections can't set headers from `EventSource`, so we also accept
/// `?token=xxx` as a query parameter, but only on SSE/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<MultiAuthState>,
State(auth): State<CombinedAuthState>,
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<String> {
// 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<String, UserIdentity>) -> 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));
+229
View File
@@ -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<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Json(body): Json<serde_json::Value>,
) -> Result<Json<serde_json::Value>, (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<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
) -> Result<Json<serde_json::Value>, (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<serde_json::Value> = 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<Arc<GatewayState>>,
Json(body): Json<serde_json::Value>,
) -> Result<Json<serde_json::Value>, (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::<Vec<_>>()
.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(),
})))
}
+3
View File
@@ -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.
+132
View File
@@ -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<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Json(body): Json<serde_json::Value>,
) -> Result<Json<serde_json::Value>, (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<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
) -> Result<Json<serde_json::Value>, (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<serde_json::Value> = 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<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Path(id): Path<String>,
) -> Result<Json<serde_json::Value>, (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(),
})))
}
+245
View File
@@ -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<Arc<GatewayState>>,
AuthenticatedUser(user): AuthenticatedUser,
Json(body): Json<serde_json::Value>,
) -> Result<Json<serde_json::Value>, (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::<Vec<_>>()
.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<Arc<GatewayState>>,
AuthenticatedUser(_user): AuthenticatedUser,
) -> Result<Json<serde_json::Value>, (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<serde_json::Value> = 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<Arc<GatewayState>>,
AuthenticatedUser(_user): AuthenticatedUser,
Path(id): Path<String>,
) -> Result<Json<serde_json::Value>, (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<Arc<GatewayState>>,
AuthenticatedUser(_user): AuthenticatedUser,
Path(id): Path<String>,
Json(body): Json<serde_json::Value>,
) -> Result<Json<serde_json::Value>, (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<Arc<GatewayState>>,
AuthenticatedUser(_user): AuthenticatedUser,
Path(id): Path<String>,
) -> Result<Json<serde_json::Value>, (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<Arc<GatewayState>>,
AuthenticatedUser(_user): AuthenticatedUser,
Path(id): Path<String>,
) -> Result<Json<serde_json::Value>, (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",
})))
}
+18 -5
View File
@@ -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<GatewayState>,
/// 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<dyn Database>) -> 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<ContainerJobManager>) -> 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).
+45 -3
View File
@@ -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<GatewayState>,
auth: MultiAuthState,
auth: CombinedAuthState,
) -> Result<SocketAddr, crate::error::ChannelError> {
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");
+2 -2
View File
@@ -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))
}
}
+12 -3
View File
@@ -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)
}
+1
View File
@@ -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));
+5 -2
View File
@@ -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");
+3 -3
View File
@@ -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))
+1 -1
View File
@@ -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");
+1 -1
View File
@@ -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");