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:
1 parent
4c697e9adf
commit
dc5417334b
26 files changed
+2477
-183
No files matched your search
+97
-17
@@ -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 })))
|
||||
}
|
||||
@@ -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);
|
||||
@@ -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()
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
@@ -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
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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];
|
||||
|
||||
@@ -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
@@ -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
@@ -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 {
|
||||
|
||||
Reference in new issue
Block a user