feat: harden runtime lifecycle and application credentials

- enforce deterministic runtime lifecycle state transitions
- add live graceful-to-forced shutdown escalation
- align HTTP draining and worker shutdown with global deadline
- guarantee deterministic shutdown hook ordering
- add secure application client IDs and one-time client secrets
- hash application secrets with BLAKE3 and constant-time verification
- make credential creation and rotation transactionally auditable
- enforce strict client_id authentication and redirect URI validation
- add SQLite and PostgreSQL credential migrations
- add application credential and runtime lifecycle acceptance tests
- update Dioxus application management workflows
- update security and architecture documentation
This commit is contained in:
thakares committed 2026-07-23 15:17:30 +05:30
1 parent 4c697e9adf
commit dc5417334b
26 files changed
+2477 -183

No files matched your search

+97 -17
View File
@@ -1,6 +1,7 @@
use axum::{
Json,
extract::{Path, State},
http::{HeaderMap, HeaderValue, header},
};
use serde::{Deserialize, Serialize};
use serde_json::{Value, json};
@@ -13,14 +14,17 @@ use crate::{
state::AppState,
};
pub const MANAGE_PERM: &str = "applications:manage";
#[derive(Serialize)]
pub struct ApplicationResponse {
pub id: String,
pub name: String,
pub slug: String,
/// Client ID — currently the application slug (OAuth2-ready).
pub client_id: String,
pub description: Option<String>,
pub enabled: bool,
pub credentials_configured: bool,
pub redirect_urls: Vec<String>,
pub scopes: Vec<String>,
pub created_at: String,
@@ -29,30 +33,51 @@ pub struct ApplicationResponse {
impl From<Application> for ApplicationResponse {
fn from(a: Application) -> Self {
let client_id = a.get_client_id().to_string();
let redirect_urls = a.redirect_urls();
let scopes = a.scopes();
let credentials_configured = a.has_credentials();
Self {
id: a.id,
name: a.name,
client_id: a.slug.clone().unwrap_or_default(),
slug: a.slug.unwrap_or_default(),
client_id,
description: a.description,
enabled: a.enabled,
// Placeholder until OAuth2 tables land
redirect_urls: Vec::new(),
scopes: Vec::new(),
credentials_configured,
redirect_urls,
scopes,
created_at: a.created_at,
updated_at: a.updated_at,
}
}
}
#[derive(Serialize)]
pub struct CreateApplicationResponse {
pub application: ApplicationResponse,
pub client_secret: String,
}
#[derive(Serialize)]
pub struct RotateSecretResponse {
pub client_secret: String,
}
fn no_store_headers() -> HeaderMap {
let mut headers = HeaderMap::new();
headers.insert(header::CACHE_CONTROL, HeaderValue::from_static("no-store"));
headers
}
/// GET /api/v1/applications
pub async fn list_applications(
State(state): State<AppState>,
auth: AuthUser,
) -> Result<Json<Value>> {
// Any authenticated user can see registered apps; mutations need roles:manage
let _ = auth;
let apps = identity::list(&state.provider, Tenant::DEFAULT_ID).await?;
let views: Vec<ApplicationResponse> = apps.into_iter().map(ApplicationResponse::from).collect();
let _ = auth;
Ok(Json(json!({ "applications": views })))
}
@@ -60,6 +85,9 @@ pub async fn list_applications(
pub struct CreateApplicationRequest {
pub name: String,
pub slug: String,
pub description: Option<String>,
pub redirect_urls: Option<Vec<String>>,
pub scopes: Option<Vec<String>>,
}
/// POST /api/v1/applications
@@ -67,13 +95,29 @@ pub async fn create_application(
State(state): State<AppState>,
auth: AuthUser,
Json(body): Json<CreateApplicationRequest>,
) -> Result<Json<Value>> {
require(&state.provider, &auth.user.id, "roles:manage").await?;
) -> Result<(HeaderMap, Json<CreateApplicationResponse>)> {
require(&state.provider, &auth.user.id, MANAGE_PERM).await?;
let app = identity::create(&state.provider, Tenant::DEFAULT_ID, &body.name, &body.slug).await?;
Ok(Json(
json!({ "application": ApplicationResponse::from(app) }),
))
let (app, raw_secret) = identity::create(
&state.provider,
Tenant::DEFAULT_ID,
&body.name,
&body.slug,
body.description.as_deref(),
body.redirect_urls,
body.scopes,
Some(&auth.user.id),
None,
None,
)
.await?;
let resp = CreateApplicationResponse {
application: ApplicationResponse::from(app),
client_secret: raw_secret,
};
Ok((no_store_headers(), Json(resp)))
}
/// GET /api/v1/applications/:id
@@ -90,9 +134,13 @@ pub async fn get_application(
}
#[derive(Debug, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct UpdateApplicationRequest {
pub name: String,
pub slug: String,
pub description: Option<String>,
pub redirect_urls: Option<Vec<String>>,
pub scopes: Option<Vec<String>>,
pub enabled: bool,
}
@@ -103,21 +151,53 @@ pub async fn update_application(
Path(id): Path<String>,
Json(body): Json<UpdateApplicationRequest>,
) -> Result<Json<Value>> {
require(&state.provider, &auth.user.id, "roles:manage").await?;
require(&state.provider, &auth.user.id, MANAGE_PERM).await?;
let app = identity::update(
&state.provider,
&id,
&body.name,
&body.slug,
body.description.as_deref(),
body.redirect_urls,
body.scopes,
body.enabled,
Some(&auth.user.id),
None,
None,
)
.await?;
let app = identity::update(&state.provider, &id, &body.name, &body.slug, body.enabled).await?;
Ok(Json(
json!({ "application": ApplicationResponse::from(app) }),
))
}
/// POST /api/v1/applications/:id/secret
pub async fn rotate_application_secret(
State(state): State<AppState>,
auth: AuthUser,
Path(id): Path<String>,
) -> Result<(HeaderMap, Json<RotateSecretResponse>)> {
require(&state.provider, &auth.user.id, MANAGE_PERM).await?;
let raw_secret =
identity::rotate_secret(&state.provider, &id, Some(&auth.user.id), None, None).await?;
let resp = RotateSecretResponse {
client_secret: raw_secret,
};
Ok((no_store_headers(), Json(resp)))
}
/// DELETE /api/v1/applications/:id
pub async fn delete_application(
State(state): State<AppState>,
auth: AuthUser,
Path(id): Path<String>,
) -> Result<Json<Value>> {
require(&state.provider, &auth.user.id, "roles:manage").await?;
identity::delete(&state.provider, &id).await?;
require(&state.provider, &auth.user.id, MANAGE_PERM).await?;
identity::delete(&state.provider, &id, Some(&auth.user.id), None, None).await?;
Ok(Json(json!({ "success": true })))
}
+4
View File
@@ -82,6 +82,10 @@ pub fn build(state: AppState) -> Router {
.patch(applications::update_application)
.delete(applications::delete_application),
)
.route(
"/applications/{id}/secret",
post(applications::rotate_application_secret),
)
// Service accounts
.route(
"/service-accounts",
@@ -0,0 +1,23 @@
-- ── Add Application Credentials Columns & Permissions (PostgreSQL) ───────────
ALTER TABLE applications ADD COLUMN IF NOT EXISTS client_id TEXT;
ALTER TABLE applications ADD COLUMN IF NOT EXISTS description TEXT;
ALTER TABLE applications ADD COLUMN IF NOT EXISTS client_secret_hash TEXT;
ALTER TABLE applications ADD COLUMN IF NOT EXISTS redirect_uris TEXT;
ALTER TABLE applications ADD COLUMN IF NOT EXISTS scopes TEXT;
-- Backfill client_id for existing applications
UPDATE applications SET client_id = 'nx9_app_' || replace(id, '-', '') WHERE client_id IS NULL;
-- Create unique index on client_id
CREATE UNIQUE INDEX IF NOT EXISTS idx_applications_client_id ON applications(client_id);
-- Seed applications:manage permission
INSERT INTO permissions (id, name, description) VALUES
('20000000-0000-0000-0000-000000000008', 'applications:manage', 'Manage registered application credentials')
ON CONFLICT (name) DO NOTHING;
-- Grant permission to admin role
INSERT INTO role_permissions (role_id, permission_id) VALUES
('10000000-0000-0000-0000-000000000001', '20000000-0000-0000-0000-000000000008')
ON CONFLICT DO NOTHING;
@@ -0,0 +1,10 @@
-- ── Application Credentials Production Hardening (PostgreSQL) ───────────────
-- Backfill any remaining applications with client_id if missing
UPDATE applications SET client_id = 'nx9_app_' || replace(id, '-', '') WHERE client_id IS NULL;
-- Enforce NOT NULL constraint on client_id
ALTER TABLE applications ALTER COLUMN client_id SET NOT NULL;
-- Ensure unique index on client_id exists
CREATE UNIQUE INDEX IF NOT EXISTS idx_applications_client_id ON applications(client_id);
@@ -0,0 +1,21 @@
-- ── Add Application Credentials Columns & Permissions (SQLite) ────────────────
ALTER TABLE applications ADD COLUMN client_id TEXT;
ALTER TABLE applications ADD COLUMN description TEXT;
ALTER TABLE applications ADD COLUMN client_secret_hash TEXT;
ALTER TABLE applications ADD COLUMN redirect_uris TEXT;
ALTER TABLE applications ADD COLUMN scopes TEXT;
-- Backfill client_id for existing applications
UPDATE applications SET client_id = 'nx9_app_' || replace(id, '-', '') WHERE client_id IS NULL;
-- Create unique index on client_id
CREATE UNIQUE INDEX IF NOT EXISTS idx_applications_client_id ON applications(client_id);
-- Seed applications:manage permission
INSERT OR IGNORE INTO permissions (id, name, description) VALUES
('20000000-0000-0000-0000-000000000008', 'applications:manage', 'Manage registered application credentials');
-- Grant permission to admin role
INSERT OR IGNORE INTO role_permissions (role_id, permission_id) VALUES
('10000000-0000-0000-0000-000000000001', '20000000-0000-0000-0000-000000000008');
@@ -0,0 +1,7 @@
-- ── Application Credentials Production Hardening (SQLite) ───────────────────
-- Backfill any remaining applications with client_id if missing
UPDATE applications SET client_id = 'nx9_app_' || replace(id, '-', '') WHERE client_id IS NULL;
-- Ensure unique index on client_id exists
CREATE UNIQUE INDEX IF NOT EXISTS idx_applications_client_id ON applications(client_id);
+42
View File
@@ -8,9 +8,51 @@ pub struct Application {
pub name: String,
pub description: Option<String>,
pub slug: Option<String>,
pub client_id: String,
pub enabled: bool,
pub client_secret_hash: Option<String>,
pub redirect_uris: Option<String>,
pub scopes: Option<String>,
pub created_at: String,
pub updated_at: String,
}
impl Application {
/// Return effective client ID string.
pub fn get_client_id(&self) -> &str {
&self.client_id
}
/// Parse configured redirect URLs.
pub fn redirect_urls(&self) -> Vec<String> {
let Some(raw) = &self.redirect_uris else {
return Vec::new();
};
if let Ok(vec) = serde_json::from_str::<Vec<String>>(raw) {
return vec;
}
raw.split([',', '\n', ' '])
.map(|s| s.trim().to_string())
.filter(|s| !s.is_empty())
.collect()
}
/// Parse configured scopes.
pub fn scopes(&self) -> Vec<String> {
let Some(raw) = &self.scopes else {
return Vec::new();
};
if let Ok(vec) = serde_json::from_str::<Vec<String>>(raw) {
return vec;
}
raw.split([',', ' '])
.map(|s| s.trim().to_string())
.filter(|s| !s.is_empty())
.collect()
}
/// Return true if application has configured client secret credentials.
pub fn has_credentials(&self) -> bool {
self.client_secret_hash.is_some()
}
}
+130 -12
View File
@@ -10,17 +10,25 @@ pub struct PostgresApplicationsRepository {
#[async_trait]
impl ApplicationsRepository for PostgresApplicationsRepository {
async fn create(
async fn create_with_audit(
&self,
id: &str,
tenant_id: &str,
name: &str,
slug: &str,
client_id: &str,
client_secret_hash: Option<&str>,
description: Option<&str>,
redirect_uris: Option<&str>,
scopes: Option<&str>,
audit_event: Option<crate::audit::AuditEvent<'_>>,
) -> Result<Application, sqlx::Error> {
sqlx::query_as::<_, Application>(
let mut tx = self.pool.begin().await?;
let app = sqlx::query_as::<_, Application>(
r#"
INSERT INTO applications (id, tenant_id, name, slug)
VALUES ($1, $2, $3, $4)
INSERT INTO applications (id, tenant_id, name, slug, client_id, client_secret_hash, description, redirect_uris, scopes)
VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9)
RETURNING *
"#,
)
@@ -28,8 +36,43 @@ impl ApplicationsRepository for PostgresApplicationsRepository {
.bind(tenant_id)
.bind(name)
.bind(slug)
.fetch_one(&self.pool)
.await
.bind(client_id)
.bind(client_secret_hash)
.bind(description)
.bind(redirect_uris)
.bind(scopes)
.fetch_one(&mut *tx)
.await?;
if let Some(event) = audit_event {
let audit_id = uuid::Uuid::new_v4().to_string();
let severity_str = event.severity.to_string();
sqlx::query(
r#"
INSERT INTO audit_logs (
id, actor_user_id, target_user_id,
action, resource_type, resource_id,
severity, ip_address, user_agent, metadata_json
)
VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10)
"#,
)
.bind(&audit_id)
.bind(event.actor_id)
.bind(event.target_id)
.bind(event.action)
.bind(event.resource_type)
.bind(event.resource_id)
.bind(&severity_str)
.bind(event.ip)
.bind(event.ua)
.bind(event.metadata)
.execute(&mut *tx)
.await?;
}
tx.commit().await?;
Ok(app)
}
async fn find_by_slug(&self, slug: &str) -> Result<Option<Application>, sqlx::Error> {
@@ -39,6 +82,13 @@ impl ApplicationsRepository for PostgresApplicationsRepository {
.await
}
async fn find_by_client_id(&self, client_id: &str) -> Result<Option<Application>, sqlx::Error> {
sqlx::query_as::<_, Application>("SELECT * FROM applications WHERE client_id = $1")
.bind(client_id)
.fetch_optional(&self.pool)
.await
}
async fn find_by_id(&self, id: &str) -> Result<Option<Application>, sqlx::Error> {
sqlx::query_as::<_, Application>("SELECT * FROM applications WHERE id = $1")
.bind(id)
@@ -59,10 +109,72 @@ impl ApplicationsRepository for PostgresApplicationsRepository {
sqlx::query(
"UPDATE applications SET enabled = $1, updated_at = to_char(clock_timestamp() AT TIME ZONE 'UTC', 'YYYY-MM-DD\"T\"HH24:MI:SS\"Z\"') WHERE id = $2",
)
.bind(enabled)
.bind(id)
.execute(&self.pool)
.await?;
.bind(enabled)
.bind(id)
.execute(&self.pool)
.await?;
Ok(())
}
async fn update_secret_hash(&self, id: &str, secret_hash: &str) -> Result<(), sqlx::Error> {
sqlx::query(
"UPDATE applications SET client_secret_hash = $1, updated_at = to_char(clock_timestamp() AT TIME ZONE 'UTC', 'YYYY-MM-DD\"T\"HH24:MI:SS\"Z\"') WHERE id = $2",
)
.bind(secret_hash)
.bind(id)
.execute(&self.pool)
.await?;
Ok(())
}
async fn rotate_secret_with_audit(
&self,
id: &str,
secret_hash: &str,
audit_event: Option<crate::audit::AuditEvent<'_>>,
) -> Result<(), sqlx::Error> {
let mut tx = self.pool.begin().await?;
let res = sqlx::query(
"UPDATE applications SET client_secret_hash = $1, updated_at = to_char(clock_timestamp() AT TIME ZONE 'UTC', 'YYYY-MM-DD\"T\"HH24:MI:SS\"Z\"') WHERE id = $2",
)
.bind(secret_hash)
.bind(id)
.execute(&mut *tx)
.await?;
if res.rows_affected() == 0 {
return Err(sqlx::Error::RowNotFound);
}
if let Some(event) = audit_event {
let audit_id = uuid::Uuid::new_v4().to_string();
let severity_str = event.severity.to_string();
sqlx::query(
r#"
INSERT INTO audit_logs (
id, actor_user_id, target_user_id,
action, resource_type, resource_id,
severity, ip_address, user_agent, metadata_json
)
VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10)
"#,
)
.bind(&audit_id)
.bind(event.actor_id)
.bind(event.target_id)
.bind(event.action)
.bind(event.resource_type)
.bind(event.resource_id)
.bind(&severity_str)
.bind(event.ip)
.bind(event.ua)
.bind(event.metadata)
.execute(&mut *tx)
.await?;
}
tx.commit().await?;
Ok(())
}
@@ -71,18 +183,24 @@ impl ApplicationsRepository for PostgresApplicationsRepository {
id: &str,
name: &str,
slug: &str,
description: Option<&str>,
redirect_uris: Option<&str>,
scopes: Option<&str>,
enabled: bool,
) -> Result<(), sqlx::Error> {
sqlx::query(
r#"
UPDATE applications
SET name = $1, slug = $2, enabled = $3,
SET name = $1, slug = $2, description = $3, redirect_uris = $4, scopes = $5, enabled = $6,
updated_at = to_char(clock_timestamp() AT TIME ZONE 'UTC', 'YYYY-MM-DD"T"HH24:MI:SS"Z"')
WHERE id = $4
WHERE id = $7
"#,
)
.bind(name)
.bind(slug)
.bind(description)
.bind(redirect_uris)
.bind(scopes)
.bind(enabled)
.bind(id)
.execute(&self.pool)
+135 -17
View File
@@ -10,37 +10,87 @@ pub struct SqliteApplicationsRepository {
#[async_trait]
impl ApplicationsRepository for SqliteApplicationsRepository {
async fn create(
async fn create_with_audit(
&self,
id: &str,
tenant_id: &str,
name: &str,
slug: &str,
client_id: &str,
client_secret_hash: Option<&str>,
description: Option<&str>,
redirect_uris: Option<&str>,
scopes: Option<&str>,
audit_event: Option<crate::audit::AuditEvent<'_>>,
) -> Result<Application, sqlx::Error> {
sqlx::query_as::<_, Application>(
let mut tx = self.pool.begin().await?;
let app = sqlx::query_as::<_, Application>(
r#"
INSERT INTO applications (id, tenant_id, name, slug)
VALUES (?, ?, ?, ?)
RETURNING id, tenant_id, name, slug, enabled, created_at, updated_at, NULL as description, NULL as client_secret_hash, NULL as redirect_uris
INSERT INTO applications (id, tenant_id, name, slug, client_id, client_secret_hash, description, redirect_uris, scopes)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)
RETURNING id, tenant_id, name, slug, client_id, description, enabled, client_secret_hash, redirect_uris, scopes, created_at, updated_at
"#,
)
.bind(id)
.bind(tenant_id)
.bind(name)
.bind(slug)
.fetch_one(&self.pool)
.await
.bind(client_id)
.bind(client_secret_hash)
.bind(description)
.bind(redirect_uris)
.bind(scopes)
.fetch_one(&mut *tx)
.await?;
if let Some(event) = audit_event {
let audit_id = uuid::Uuid::new_v4().to_string();
let severity_str = event.severity.to_string();
sqlx::query(
r#"
INSERT INTO audit_logs (
id, actor_user_id, target_user_id,
action, resource_type, resource_id,
severity, ip_address, user_agent, metadata_json
)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
"#,
)
.bind(&audit_id)
.bind(event.actor_id)
.bind(event.target_id)
.bind(event.action)
.bind(event.resource_type)
.bind(event.resource_id)
.bind(&severity_str)
.bind(event.ip)
.bind(event.ua)
.bind(event.metadata)
.execute(&mut *tx)
.await?;
}
tx.commit().await?;
Ok(app)
}
async fn find_by_slug(&self, slug: &str) -> Result<Option<Application>, sqlx::Error> {
sqlx::query_as::<_, Application>("SELECT id, tenant_id, name, slug, enabled, created_at, updated_at, NULL as description, NULL as client_secret_hash, NULL as redirect_uris FROM applications WHERE slug = ?")
sqlx::query_as::<_, Application>("SELECT id, tenant_id, name, slug, client_id, description, enabled, client_secret_hash, redirect_uris, scopes, created_at, updated_at FROM applications WHERE slug = ?")
.bind(slug)
.fetch_optional(&self.pool)
.await
}
async fn find_by_client_id(&self, client_id: &str) -> Result<Option<Application>, sqlx::Error> {
sqlx::query_as::<_, Application>("SELECT id, tenant_id, name, slug, client_id, description, enabled, client_secret_hash, redirect_uris, scopes, created_at, updated_at FROM applications WHERE client_id = ?")
.bind(client_id)
.fetch_optional(&self.pool)
.await
}
async fn find_by_id(&self, id: &str) -> Result<Option<Application>, sqlx::Error> {
sqlx::query_as::<_, Application>("SELECT id, tenant_id, name, slug, enabled, created_at, updated_at, NULL as description, NULL as client_secret_hash, NULL as redirect_uris FROM applications WHERE id = ?")
sqlx::query_as::<_, Application>("SELECT id, tenant_id, name, slug, client_id, description, enabled, client_secret_hash, redirect_uris, scopes, created_at, updated_at FROM applications WHERE id = ?")
.bind(id)
.fetch_optional(&self.pool)
.await
@@ -48,7 +98,7 @@ impl ApplicationsRepository for SqliteApplicationsRepository {
async fn list(&self, tenant_id: &str) -> Result<Vec<Application>, sqlx::Error> {
sqlx::query_as::<_, Application>(
"SELECT id, tenant_id, name, slug, enabled, created_at, updated_at, NULL as description, NULL as client_secret_hash, NULL as redirect_uris FROM applications WHERE tenant_id = ? ORDER BY name",
"SELECT id, tenant_id, name, slug, client_id, description, enabled, client_secret_hash, redirect_uris, scopes, created_at, updated_at FROM applications WHERE tenant_id = ? ORDER BY name",
)
.bind(tenant_id)
.fetch_all(&self.pool)
@@ -57,12 +107,74 @@ impl ApplicationsRepository for SqliteApplicationsRepository {
async fn set_enabled(&self, id: &str, enabled: bool) -> Result<(), sqlx::Error> {
sqlx::query(
"UPDATE applications SET enabled = ?, updated_at = strftime('%Y-%m-%dT%H:%M:%SZ', 'now') WHERE id = ?",
)
.bind(enabled)
.bind(id)
.execute(&self.pool)
.await?;
"UPDATE applications SET enabled = ?, updated_at = strftime('%Y-%m-%dT%H:%M:%SZ', 'now') WHERE id = ?",
)
.bind(enabled)
.bind(id)
.execute(&self.pool)
.await?;
Ok(())
}
async fn update_secret_hash(&self, id: &str, secret_hash: &str) -> Result<(), sqlx::Error> {
sqlx::query(
"UPDATE applications SET client_secret_hash = ?, updated_at = strftime('%Y-%m-%dT%H:%M:%SZ', 'now') WHERE id = ?",
)
.bind(secret_hash)
.bind(id)
.execute(&self.pool)
.await?;
Ok(())
}
async fn rotate_secret_with_audit(
&self,
id: &str,
secret_hash: &str,
audit_event: Option<crate::audit::AuditEvent<'_>>,
) -> Result<(), sqlx::Error> {
let mut tx = self.pool.begin().await?;
let res = sqlx::query(
"UPDATE applications SET client_secret_hash = ?, updated_at = strftime('%Y-%m-%dT%H:%M:%SZ', 'now') WHERE id = ?",
)
.bind(secret_hash)
.bind(id)
.execute(&mut *tx)
.await?;
if res.rows_affected() == 0 {
return Err(sqlx::Error::RowNotFound);
}
if let Some(event) = audit_event {
let audit_id = uuid::Uuid::new_v4().to_string();
let severity_str = event.severity.to_string();
sqlx::query(
r#"
INSERT INTO audit_logs (
id, actor_user_id, target_user_id,
action, resource_type, resource_id,
severity, ip_address, user_agent, metadata_json
)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
"#,
)
.bind(&audit_id)
.bind(event.actor_id)
.bind(event.target_id)
.bind(event.action)
.bind(event.resource_type)
.bind(event.resource_id)
.bind(&severity_str)
.bind(event.ip)
.bind(event.ua)
.bind(event.metadata)
.execute(&mut *tx)
.await?;
}
tx.commit().await?;
Ok(())
}
@@ -71,18 +183,24 @@ impl ApplicationsRepository for SqliteApplicationsRepository {
id: &str,
name: &str,
slug: &str,
description: Option<&str>,
redirect_uris: Option<&str>,
scopes: Option<&str>,
enabled: bool,
) -> Result<(), sqlx::Error> {
sqlx::query(
r#"
UPDATE applications
SET name = ?, slug = ?, enabled = ?,
SET name = ?, slug = ?, description = ?, redirect_uris = ?, scopes = ?, enabled = ?,
updated_at = strftime('%Y-%m-%dT%H:%M:%SZ', 'now')
WHERE id = ?
"#,
)
.bind(name)
.bind(slug)
.bind(description)
.bind(redirect_uris)
.bind(scopes)
.bind(enabled)
.bind(id)
.execute(&self.pool)
+20 -1
View File
@@ -96,22 +96,41 @@ pub trait SessionsRepository: Send + Sync {
#[async_trait::async_trait]
pub trait ApplicationsRepository: Send + Sync {
async fn create(
#[allow(clippy::too_many_arguments)]
async fn create_with_audit(
&self,
id: &str,
tenant_id: &str,
name: &str,
slug: &str,
client_id: &str,
client_secret_hash: Option<&str>,
description: Option<&str>,
redirect_uris: Option<&str>,
scopes: Option<&str>,
audit_event: Option<crate::audit::AuditEvent<'_>>,
) -> Result<Application, sqlx::Error>;
async fn find_by_slug(&self, slug: &str) -> Result<Option<Application>, sqlx::Error>;
async fn find_by_client_id(&self, client_id: &str) -> Result<Option<Application>, sqlx::Error>;
async fn find_by_id(&self, id: &str) -> Result<Option<Application>, sqlx::Error>;
async fn list(&self, tenant_id: &str) -> Result<Vec<Application>, sqlx::Error>;
async fn set_enabled(&self, id: &str, enabled: bool) -> Result<(), sqlx::Error>;
async fn update_secret_hash(&self, id: &str, secret_hash: &str) -> Result<(), sqlx::Error>;
async fn rotate_secret_with_audit(
&self,
id: &str,
secret_hash: &str,
audit_event: Option<crate::audit::AuditEvent<'_>>,
) -> Result<(), sqlx::Error>;
#[allow(clippy::too_many_arguments)]
async fn update(
&self,
id: &str,
name: &str,
slug: &str,
description: Option<&str>,
redirect_uris: Option<&str>,
scopes: Option<&str>,
enabled: bool,
) -> Result<(), sqlx::Error>;
async fn delete(&self, id: &str) -> Result<(), sqlx::Error>;
+369 -13
View File
@@ -1,11 +1,106 @@
use crate::db::repository::traits::AuditRepositoryExt;
use crate::{db::models::Application, error::AppError};
use subtle::ConstantTimeEq;
pub const CLIENT_ID_PREFIX: &str = "nx9_app_";
pub const CLIENT_SECRET_PREFIX: &str = "nx9_secret_";
/// Generate a new unique server-side Client ID.
pub fn generate_client_id() -> String {
let mut bytes = [0u8; 16];
rand::RngCore::fill_bytes(&mut rand::thread_rng(), &mut bytes);
format!("{}{}", CLIENT_ID_PREFIX, hex::encode(bytes))
}
/// Generate a new CSPRNG Client Secret.
pub fn generate_client_secret() -> String {
let mut bytes = [0u8; 32];
rand::RngCore::fill_bytes(&mut rand::thread_rng(), &mut bytes);
format!("{}{}", CLIENT_SECRET_PREFIX, hex::encode(bytes))
}
/// Hash a raw client secret string into a hex-encoded BLAKE3 digest.
pub fn hash_client_secret(raw: &str) -> String {
hex::encode(blake3::hash(raw.as_bytes()).as_bytes())
}
/// Hash a raw client secret string into a 32-byte BLAKE3 digest.
pub fn hash_secret_bytes(raw: &str) -> [u8; 32] {
*blake3::hash(raw.as_bytes()).as_bytes()
}
/// Constant-time byte array comparison using subtle::ConstantTimeEq.
pub fn constant_time_compare(a: &[u8], b: &[u8]) -> bool {
if a.len() != b.len() {
return false;
}
a.ct_eq(b).into()
}
pub fn validate_redirect_uris(uris: &[String]) -> Result<(), AppError> {
if uris.len() > 10 {
return Err(AppError::InvalidInput(
"maximum 10 redirect URIs allowed".into(),
));
}
for uri in uris {
let trimmed = uri.trim();
if trimmed.is_empty() {
return Err(AppError::InvalidInput(
"redirect URI cannot be empty".into(),
));
}
if trimmed.len() > 512 {
return Err(AppError::InvalidInput(
"redirect URI exceeds maximum length of 512 characters".into(),
));
}
let parsed = url::Url::parse(trimmed).map_err(|e| {
AppError::InvalidInput(format!("invalid redirect URI '{trimmed}': {e}"))
})?;
if parsed.fragment().is_some() {
return Err(AppError::InvalidInput(format!(
"redirect URI '{trimmed}' must not contain a fragment"
)));
}
if !parsed.username().is_empty() || parsed.password().is_some() {
return Err(AppError::InvalidInput(format!(
"redirect URI '{trimmed}' must not contain user credentials"
)));
}
match parsed.scheme() {
"https" => {}
"http" => {
let host = parsed.host_str().unwrap_or("");
if host != "localhost" && host != "127.0.0.1" && host != "[::1]" && host != "::1" {
return Err(AppError::InvalidInput(format!(
"redirect URI '{trimmed}' with http scheme is only allowed for localhost/loopback development"
)));
}
}
other => {
return Err(AppError::InvalidInput(format!(
"redirect URI '{trimmed}' has unsupported scheme '{other}'; only https (or http for localhost) is allowed"
)));
}
}
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub async fn create(
provider: &std::sync::Arc<dyn crate::db::provider::DatabaseProvider>,
tenant_id: &str,
name: &str,
slug: &str,
) -> Result<Application, AppError> {
description: Option<&str>,
redirect_uris: Option<Vec<String>>,
scopes: Option<Vec<String>>,
audit_actor_id: Option<&str>,
audit_ip: Option<&str>,
audit_ua: Option<&str>,
) -> Result<(Application, String), AppError> {
let name = name.trim();
let slug = slug.trim();
if name.is_empty() || slug.is_empty() {
@@ -13,6 +108,9 @@ pub async fn create(
"name and slug cannot be empty".into(),
));
}
if let Some(ref uris) = redirect_uris {
validate_redirect_uris(uris)?;
}
if provider
.applications()
.find_by_slug(slug)
@@ -22,12 +120,52 @@ pub async fn create(
{
return Err(AppError::Conflict(format!("slug '{slug}' already exists")));
}
let id = uuid::Uuid::new_v4().to_string();
provider
let client_id = generate_client_id();
let raw_secret = generate_client_secret();
let secret_hash = hash_client_secret(&raw_secret);
let redirect_json = redirect_uris.map(|v| serde_json::to_string(&v).unwrap_or_default());
let scopes_json = scopes.map(|v| serde_json::to_string(&v).unwrap_or_default());
let metadata = serde_json::json!({
"application_id": id,
"name": name,
"client_id": client_id,
})
.to_string();
let audit_event = crate::audit::AuditEvent {
actor_id: audit_actor_id,
target_id: None,
action: "application.created",
resource_type: "application",
resource_id: Some(&id),
severity: crate::db::models::AuditSeverity::Info,
ip: audit_ip,
ua: audit_ua,
metadata: Some(&metadata),
};
let app = provider
.applications()
.create(&id, tenant_id, name, slug)
.create_with_audit(
&id,
tenant_id,
name,
slug,
&client_id,
Some(&secret_hash),
description,
redirect_json.as_deref(),
scopes_json.as_deref(),
Some(audit_event),
)
.await
.map_err(AppError::Database)
.map_err(AppError::Database)?;
Ok((app, raw_secret))
}
pub async fn list(
@@ -65,14 +203,108 @@ pub async fn find_by_slug(
.ok_or(AppError::NotFound)
}
pub async fn rotate_secret(
provider: &std::sync::Arc<dyn crate::db::provider::DatabaseProvider>,
id: &str,
audit_actor_id: Option<&str>,
audit_ip: Option<&str>,
audit_ua: Option<&str>,
) -> Result<String, AppError> {
let app = provider
.applications()
.find_by_id(id)
.await
.map_err(AppError::Database)?
.ok_or(AppError::NotFound)?;
let raw_secret = generate_client_secret();
let secret_hash = hash_client_secret(&raw_secret);
let metadata = serde_json::json!({
"application_id": id,
"name": app.name,
"client_id": app.get_client_id(),
})
.to_string();
let audit_event = crate::audit::AuditEvent {
actor_id: audit_actor_id,
target_id: None,
action: "application.secret_rotated",
resource_type: "application",
resource_id: Some(id),
severity: crate::db::models::AuditSeverity::Warning,
ip: audit_ip,
ua: audit_ua,
metadata: Some(&metadata),
};
provider
.applications()
.rotate_secret_with_audit(id, &secret_hash, Some(audit_event))
.await
.map_err(AppError::Database)?;
Ok(raw_secret)
}
pub async fn validate_client_credentials(
provider: &std::sync::Arc<dyn crate::db::provider::DatabaseProvider>,
client_id: &str,
client_secret: &str,
) -> Result<Application, AppError> {
let supplied_digest = hash_secret_bytes(client_secret);
let app = provider
.applications()
.find_by_client_id(client_id)
.await
.map_err(AppError::Database)?;
let dummy_digest = [0u8; 32];
let (valid_app, stored_digest_opt) = match app {
Some(ref a) if a.enabled => {
let digest_opt = a
.client_secret_hash
.as_ref()
.and_then(|h| hex::decode(h).ok())
.and_then(|vec| <[u8; 32]>::try_from(vec).ok());
(digest_opt.is_some(), digest_opt)
}
_ => (false, None),
};
let target_digest = stored_digest_opt.as_ref().unwrap_or(&dummy_digest);
let matches = constant_time_compare(&supplied_digest, target_digest);
if valid_app && matches {
Ok(app.unwrap())
} else {
Err(AppError::Unauthorized)
}
}
#[allow(clippy::too_many_arguments)]
pub async fn update(
provider: &std::sync::Arc<dyn crate::db::provider::DatabaseProvider>,
id: &str,
name: &str,
slug: &str,
description: Option<&str>,
redirect_uris: Option<Vec<String>>,
scopes: Option<Vec<String>>,
enabled: bool,
audit_actor_id: Option<&str>,
audit_ip: Option<&str>,
audit_ua: Option<&str>,
) -> Result<Application, AppError> {
let _ = provider.applications().find_by_id(id).await?;
let existing = provider
.applications()
.find_by_id(id)
.await
.map_err(AppError::Database)?
.ok_or(AppError::NotFound)?;
let name = name.trim();
let slug = slug.trim();
if name.is_empty() || slug.is_empty() {
@@ -80,6 +312,9 @@ pub async fn update(
"name and slug cannot be empty".into(),
));
}
if let Some(ref uris) = redirect_uris {
validate_redirect_uris(uris)?;
}
if let Some(other) = provider
.applications()
.find_by_slug(slug)
@@ -90,40 +325,161 @@ pub async fn update(
return Err(AppError::Conflict(format!("slug '{slug}' already exists")));
}
}
let redirect_json = redirect_uris.map(|v| serde_json::to_string(&v).unwrap_or_default());
let scopes_json = scopes.map(|v| serde_json::to_string(&v).unwrap_or_default());
provider
.applications()
.update(id, name, slug, enabled)
.update(
id,
name,
slug,
description,
redirect_json.as_deref(),
scopes_json.as_deref(),
enabled,
)
.await
.map_err(AppError::Database)?;
provider
let updated = provider
.applications()
.find_by_id(id)
.await
.map_err(crate::error::AppError::Database)?
.ok_or_else(|| crate::error::AppError::NotFound)
.map_err(AppError::Database)?
.ok_or(AppError::NotFound)?;
let action = if existing.enabled != enabled {
if enabled {
"application.enabled"
} else {
"application.disabled"
}
} else {
"application.updated"
};
let metadata = serde_json::json!({
"application_id": id,
"name": updated.name,
"client_id": updated.get_client_id(),
"enabled": enabled,
})
.to_string();
provider
.audit()
.log(crate::audit::AuditEvent {
actor_id: audit_actor_id,
target_id: None,
action,
resource_type: "application",
resource_id: Some(id),
severity: crate::db::models::AuditSeverity::Info,
ip: audit_ip,
ua: audit_ua,
metadata: Some(&metadata),
})
.await?;
Ok(updated)
}
pub async fn set_enabled(
provider: &std::sync::Arc<dyn crate::db::provider::DatabaseProvider>,
id: &str,
enabled: bool,
audit_actor_id: Option<&str>,
audit_ip: Option<&str>,
audit_ua: Option<&str>,
) -> Result<(), AppError> {
let _ = provider.applications().find_by_id(id).await?;
let app = provider
.applications()
.find_by_id(id)
.await
.map_err(AppError::Database)?
.ok_or(AppError::NotFound)?;
provider
.applications()
.set_enabled(id, enabled)
.await
.map_err(AppError::Database)
.map_err(AppError::Database)?;
let action = if enabled {
"application.enabled"
} else {
"application.disabled"
};
let metadata = serde_json::json!({
"application_id": id,
"name": app.name,
"client_id": app.get_client_id(),
"enabled": enabled,
})
.to_string();
provider
.audit()
.log(crate::audit::AuditEvent {
actor_id: audit_actor_id,
target_id: None,
action,
resource_type: "application",
resource_id: Some(id),
severity: crate::db::models::AuditSeverity::Info,
ip: audit_ip,
ua: audit_ua,
metadata: Some(&metadata),
})
.await?;
Ok(())
}
pub async fn delete(
provider: &std::sync::Arc<dyn crate::db::provider::DatabaseProvider>,
id: &str,
audit_actor_id: Option<&str>,
audit_ip: Option<&str>,
audit_ua: Option<&str>,
) -> Result<(), AppError> {
let _ = provider.applications().find_by_id(id).await?;
let app = provider
.applications()
.find_by_id(id)
.await
.map_err(AppError::Database)?
.ok_or(AppError::NotFound)?;
provider
.applications()
.delete(id)
.await
.map_err(AppError::Database)
.map_err(AppError::Database)?;
let metadata = serde_json::json!({
"application_id": id,
"name": app.name,
"client_id": app.get_client_id(),
})
.to_string();
provider
.audit()
.log(crate::audit::AuditEvent {
actor_id: audit_actor_id,
target_id: None,
action: "application.deleted",
resource_type: "application",
resource_id: Some(id),
severity: crate::db::models::AuditSeverity::Warning,
ip: audit_ip,
ua: audit_ua,
metadata: Some(&metadata),
})
.await?;
Ok(())
}
+78 -55
View File
@@ -29,6 +29,8 @@ pub struct Application {
pub signals: SignalManager,
pub shutdown: ShutdownCoordinator,
pub metrics: RuntimeMetrics,
pub local_addr: Option<std::net::SocketAddr>,
pub bound_port: Arc<std::sync::atomic::AtomicU16>,
}
impl Application {
@@ -82,52 +84,57 @@ impl Application {
&self.metrics
}
/// Force a runtime state update.
/// Force advance runtime state (monotonic, forward-only).
pub fn set_state(&self, state: RuntimeState) {
self.state.force_set(state);
self.state.force_advance(state);
}
/// Perform graceful shutdown flow explicitly.
pub async fn perform_shutdown(&mut self) -> Result<()> {
if !self.state.initiate_shutdown() {
if self.state.load().is_shutting_down() {
return Ok(());
let current_state = self.state.load();
if current_state == RuntimeState::Running {
let _ = self.state.initiate_shutdown();
} else if !current_state.is_shutting_down() {
self.state.force_advance(RuntimeState::Draining);
}
if self.state.load() == RuntimeState::Draining {
tracing::info!("draining active connections");
let _ = self
.state
.transition(RuntimeState::Draining, RuntimeState::StoppingWorkers);
}
if self.state.load() == RuntimeState::StoppingWorkers {
tracing::info!("stopping background workers");
self.workers
.shutdown_all_with_coordinator(Duration::from_secs(10), Some(&self.shutdown))
.await;
let _ = self
.state
.transition(RuntimeState::StoppingWorkers, RuntimeState::ExecutingHooks);
}
if self.state.load() == RuntimeState::ExecutingHooks {
tracing::info!("executing shutdown hooks");
self.hooks.execute_all().await;
let _ = self
.state
.transition(RuntimeState::ExecutingHooks, RuntimeState::ClosingResources);
}
if self.state.load() == RuntimeState::ClosingResources {
tracing::info!("closing database connection pool and resources");
if let Some(pool) = self.pool_handle.take() {
pool.close().await;
}
self.state.force_set(RuntimeState::Draining);
let _ = self
.state
.transition(RuntimeState::ClosingResources, RuntimeState::Stopped);
}
println!("Draining");
tracing::info!("draining active connections");
let _ = self
.state
.transition(RuntimeState::Draining, RuntimeState::StoppingWorkers);
println!("StoppingWorkers");
tracing::info!("stopping background workers");
self.workers.shutdown_all(Duration::from_secs(10)).await;
let _ = self
.state
.transition(RuntimeState::StoppingWorkers, RuntimeState::ExecutingHooks);
println!("ExecutingHooks");
tracing::info!("executing shutdown hooks");
self.hooks.execute_all().await;
let _ = self
.state
.transition(RuntimeState::ExecutingHooks, RuntimeState::ClosingResources);
println!("ClosingResources");
tracing::info!("closing database connection pool and resources");
if let Some(pool) = self.pool_handle.take() {
pool.close().await;
}
let _ = self
.state
.transition(RuntimeState::ClosingResources, RuntimeState::Stopped);
println!("Stopped");
tracing::info!("application stopped cleanly");
Ok(())
}
}
@@ -135,12 +142,10 @@ impl Application {
#[async_trait::async_trait]
impl Lifecycle for Application {
async fn initialize(&mut self) -> Result<()> {
println!("Initializing");
let _ = self
.state
.transition(RuntimeState::Initializing, RuntimeState::Starting);
println!("Starting");
let config = match &self.config {
Some(cfg) => cfg.clone(),
None => {
@@ -179,8 +184,6 @@ impl Lifecycle for Application {
.transition(RuntimeState::Starting, RuntimeState::Running);
}
println!("Running");
let config = self.config.as_ref().cloned().unwrap_or_default();
let addr_str = format!("{}:{}", config.server.host, config.server.port);
let listener = tokio::net::TcpListener::bind(&addr_str)
@@ -188,7 +191,9 @@ impl Lifecycle for Application {
.with_context(|| format!("failed to bind TCP listener to {addr_str}"))?;
let local_addr = listener.local_addr()?;
println!("Listening on {}", local_addr);
self.local_addr = Some(local_addr);
self.bound_port
.store(local_addr.port(), std::sync::atomic::Ordering::Release);
tracing::info!(address = %local_addr, "Listening on {}", local_addr);
let router = match self.router.take() {
@@ -205,24 +210,42 @@ impl Lifecycle for Application {
let signal_mgr = self.signals.clone();
let shutdown_coord = self.shutdown.clone();
let state = self.state.clone();
let server = axum::serve(listener, router).with_graceful_shutdown(async move {
tokio::select! {
sig = signals::wait_for_shutdown_signal() => {
tracing::info!(signal = sig, "received shutdown signal");
signal_mgr.record_signal();
shutdown_coord.cancel();
}
_ = shutdown_coord.cancelled() => {
tracing::info!("shutdown coordinator cancelled");
}
}
let signal_task = tokio::spawn(signals::listen_for_signals(
signal_mgr,
shutdown_coord.clone(),
self.state.clone(),
));
let graceful_token = shutdown_coord.token().clone();
let forced_token = shutdown_coord.forced_token().clone();
let state_for_shutdown = state.clone();
let server_fut = axum::serve(listener, router).with_graceful_shutdown(async move {
graceful_token.cancelled().await;
tracing::info!("graceful shutdown triggered; initiating HTTP connection draining");
let _ = state_for_shutdown.initiate_shutdown();
});
if let Err(err) = server.await {
tracing::error!(error = %err, "HTTP server error");
let mut server_task = tokio::spawn(async move { server_fut.await });
tokio::select! {
res = &mut server_task => {
match res {
Ok(Ok(())) => tracing::info!("HTTP server stopped gracefully"),
Ok(Err(err)) => tracing::error!(error = %err, "HTTP server error"),
Err(join_err) => tracing::debug!(error = %join_err, "HTTP server task finished"),
}
}
_ = forced_token.cancelled() => {
tracing::warn!("live forced shutdown escalation received during HTTP drain; aborting server task immediately");
server_task.abort();
let _ = server_task.await;
}
}
signal_task.abort();
self.perform_shutdown().await
}
+47 -7
View File
@@ -2,36 +2,76 @@
use tokio_util::sync::CancellationToken;
/// Dual-token shutdown coordinator that supports graceful termination
/// (1st signal) and live forced escalation (2nd signal).
#[derive(Clone)]
pub struct ShutdownCoordinator {
root: CancellationToken,
graceful: CancellationToken,
forced: CancellationToken,
}
impl ShutdownCoordinator {
pub fn new() -> Self {
Self {
root: CancellationToken::new(),
graceful: CancellationToken::new(),
forced: CancellationToken::new(),
}
}
/// Access the primary graceful cancellation token.
pub fn token(&self) -> &CancellationToken {
&self.root
&self.graceful
}
/// Access the forced cancellation token.
pub fn forced_token(&self) -> &CancellationToken {
&self.forced
}
/// Create a child token linked to graceful cancellation.
pub fn child_token(&self) -> CancellationToken {
self.root.child_token()
self.graceful.child_token()
}
/// Trigger graceful shutdown.
pub fn cancel(&self) {
self.root.cancel();
self.graceful.cancel();
}
/// Trigger graceful shutdown explicitly.
pub fn cancel_graceful(&self) {
self.graceful.cancel();
}
/// Trigger forced shutdown escalation live.
pub fn cancel_forced(&self) {
self.graceful.cancel();
self.forced.cancel();
}
/// Check if graceful shutdown has been initiated.
pub fn is_cancelled(&self) -> bool {
self.root.is_cancelled()
self.graceful.is_cancelled()
}
/// Check if forced escalation has been triggered.
pub fn is_forced(&self) -> bool {
self.forced.is_cancelled()
}
/// Await graceful cancellation.
pub async fn cancelled(&self) {
self.root.cancelled().await;
self.graceful.cancelled().await;
}
/// Await graceful cancellation explicitly.
pub async fn graceful_cancelled(&self) {
self.graceful.cancelled().await;
}
/// Await forced escalation live.
pub async fn forced_cancelled(&self) {
self.forced.cancelled().await;
}
}
+6 -1
View File
@@ -40,7 +40,12 @@ impl HookRegistry {
}
let mut indices: Vec<usize> = (0..self.hooks.len()).collect();
indices.sort_by_key(|&i| self.hooks[i].priority());
indices.sort_by(
|&a, &b| match self.hooks[a].priority().cmp(&self.hooks[b].priority()) {
std::cmp::Ordering::Equal => a.cmp(&b),
ord => ord,
},
);
for i in indices {
let hook = &self.hooks[i];
+34
View File
@@ -3,6 +3,8 @@
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
use super::{AtomicRuntimeState, ShutdownCoordinator};
#[derive(Clone)]
pub struct SignalManager {
signal_count: Arc<AtomicUsize>,
@@ -32,6 +34,17 @@ impl SignalManager {
}
count
}
/// Record a signal and trigger live escalation on the coordinator.
pub fn handle_signal(&self, coordinator: &ShutdownCoordinator) -> usize {
let count = self.record_signal();
if count == 1 {
coordinator.cancel_graceful();
} else if count >= 2 {
coordinator.cancel_forced();
}
count
}
}
impl Default for SignalManager {
@@ -65,3 +78,24 @@ pub async fn wait_for_shutdown_signal() -> &'static str {
name = sigterm => name,
}
}
/// Continuous signal monitor that remains active during graceful shutdown
/// to observe and trigger forced escalation live.
pub async fn listen_for_signals(
signal_mgr: SignalManager,
coordinator: ShutdownCoordinator,
state: AtomicRuntimeState,
) {
loop {
let sig = wait_for_shutdown_signal().await;
let count = signal_mgr.handle_signal(&coordinator);
tracing::info!(signal = sig, count, "received OS signal");
if count == 1 {
let _ = state.initiate_shutdown();
} else {
// 2nd signal received: forced escalation
tracing::warn!("second signal received; escalating to forced shutdown");
break;
}
}
}
+75 -13
View File
@@ -70,19 +70,22 @@ impl fmt::Display for RuntimeState {
}
}
use std::sync::Arc;
/// Lock-free atomic runtime state container.
///
/// Uses `AtomicU8` with `compare_exchange` to ensure deterministic,
/// race-free state transitions without mutex contention.
#[derive(Clone)]
pub struct AtomicRuntimeState {
state: AtomicU8,
state: Arc<AtomicU8>,
}
impl AtomicRuntimeState {
/// Create a new state machine in the `Initializing` state.
pub fn new() -> Self {
Self {
state: AtomicU8::new(RuntimeState::Initializing as u8),
state: Arc::new(AtomicU8::new(RuntimeState::Initializing as u8)),
}
}
@@ -91,15 +94,32 @@ impl AtomicRuntimeState {
RuntimeState::from_u8(self.state.load(Ordering::Acquire)).unwrap_or(RuntimeState::Stopped)
}
/// Check if a state transition follows the valid lifecycle graph.
pub fn is_valid_transition(expected: RuntimeState, new: RuntimeState) -> bool {
(new as u8) == (expected as u8) + 1
}
/// Attempt an atomic state transition from `expected` to `new`.
///
/// Returns `Ok(new)` if the transition succeeded, or `Err(actual)` if the
/// current state did not match `expected`.
/// The transition is validated against the lifecycle graph. Returns `Ok(new)`
/// if the transition succeeded, or `Err(actual)` if the transition was invalid
/// or the current state did not match `expected`.
pub fn transition(
&self,
expected: RuntimeState,
new: RuntimeState,
) -> Result<RuntimeState, RuntimeState> {
if !Self::is_valid_transition(expected, new) {
let actual = self.load();
tracing::warn!(
expected = %expected,
actual = %actual,
target = %new,
"illegal lifecycle graph transition rejected"
);
return Err(actual);
}
match self.state.compare_exchange(
expected as u8,
new as u8,
@@ -123,13 +143,39 @@ impl AtomicRuntimeState {
}
}
/// Unconditionally advance the state. Used during forced shutdown when
/// intermediate states may have been skipped.
pub fn force_set(&self, new: RuntimeState) {
let prev = self.state.swap(new as u8, Ordering::AcqRel);
let prev_state = RuntimeState::from_u8(prev).unwrap_or(RuntimeState::Stopped);
if prev_state != new {
tracing::info!(from = %prev_state, to = %new, "runtime state forced");
/// Unconditionally advance the state forward. Used during emergency recovery
/// when intermediate states are skipped.
///
/// Restricted to `pub(crate)` visibility to preserve lifecycle graph invariants.
/// Guarantees monotonic forward movement (`new >= current_state`) and rejects
/// backward state regressions.
pub(crate) fn force_advance(&self, new: RuntimeState) {
loop {
let current = self.load();
if (new as u8) < (current as u8) {
tracing::warn!(
current = %current,
target = %new,
"rejected state regression in force_advance"
);
break;
}
if current == new {
break;
}
if self
.state
.compare_exchange(
current as u8,
new as u8,
Ordering::AcqRel,
Ordering::Acquire,
)
.is_ok()
{
tracing::info!(from = %current, to = %new, "runtime state force advanced");
break;
}
}
}
@@ -180,9 +226,25 @@ mod tests {
}
#[test]
fn test_force_set() {
fn test_invalid_graph_transition_rejected() {
let state = AtomicRuntimeState::new();
state.force_set(RuntimeState::ClosingResources);
// Initializing -> ClosingResources is invalid in the normal graph
assert!(
state
.transition(RuntimeState::Initializing, RuntimeState::ClosingResources)
.is_err()
);
assert_eq!(state.load(), RuntimeState::Initializing);
}
#[test]
fn test_force_advance() {
let state = AtomicRuntimeState::new();
state.force_advance(RuntimeState::ClosingResources);
assert_eq!(state.load(), RuntimeState::ClosingResources);
// State regression must be rejected
state.force_advance(RuntimeState::Initializing);
assert_eq!(state.load(), RuntimeState::ClosingResources);
}
+97 -2
View File
@@ -6,6 +6,8 @@ use std::time::Duration;
use tokio::task::JoinSet;
use super::ShutdownCoordinator;
pub struct TaskGroup {
name: String,
tasks: JoinSet<()>,
@@ -100,11 +102,104 @@ impl WorkerManager {
}
}
pub async fn shutdown_all(&mut self, timeout: Duration) {
pub async fn drain_all(&mut self) {
for group in self.groups.values_mut() {
group.shutdown(timeout).await;
group.abort_all();
while group.tasks.join_next().await.is_some() {}
}
}
/// Shut down all worker groups concurrently under a single global deadline,
/// while observing live forced shutdown escalation.
pub async fn shutdown_all_with_coordinator(
&mut self,
timeout: Duration,
coordinator: Option<&ShutdownCoordinator>,
) {
let active = self.active_tasks();
if active == 0 {
return;
}
tracing::info!(
active_tasks = active,
timeout_secs = timeout.as_secs(),
"shutting down background worker task groups under global deadline"
);
let is_already_forced = coordinator.map(|c| c.is_forced()).unwrap_or(false);
if is_already_forced {
tracing::warn!("forced shutdown active; aborting all worker tasks immediately");
self.drain_all().await;
return;
}
let groups = std::mem::take(&mut self.groups);
let mut group_joiners = JoinSet::new();
let mut group_map = HashMap::new();
for (name, mut group) in groups {
group_joiners.spawn(async move {
while group.tasks.join_next().await.is_some() {}
(name, group)
});
}
let join_all_fut = async {
while let Some(res) = group_joiners.join_next().await {
if let Ok((name, group)) = res {
group_map.insert(name, group);
}
}
};
let forced_fut = async {
if let Some(coord) = coordinator {
coord.forced_cancelled().await;
} else {
std::future::pending::<()>().await;
}
};
tokio::select! {
_ = join_all_fut => {
tracing::info!("all worker task groups shut down cleanly");
}
_ = tokio::time::sleep(timeout) => {
tracing::warn!("global worker shutdown timeout expired; aborting remaining tasks");
group_joiners.abort_all();
while let Some(res) = group_joiners.join_next().await {
if let Ok((name, mut group)) = res {
group.abort_all();
group_map.insert(name, group);
}
}
}
_ = forced_fut => {
tracing::warn!("live forced shutdown escalation received during worker wait; aborting remaining tasks immediately");
group_joiners.abort_all();
while let Some(res) = group_joiners.join_next().await {
if let Ok((name, mut group)) = res {
group.abort_all();
group_map.insert(name, group);
}
}
}
}
for group in group_map.values_mut() {
if !group.is_empty() {
group.abort_all();
while group.tasks.join_next().await.is_some() {}
}
}
self.groups = group_map;
}
pub async fn shutdown_all(&mut self, timeout: Duration) {
self.shutdown_all_with_coordinator(timeout, None).await;
}
}
impl Default for WorkerManager {