use axum::Extension;
use axum::Json;
use axum::Router;
use axum::http::StatusCode;
use axum::response::{IntoResponse, Response};
use axum::routing::get;
use serde::Serialize;
use sqlx::{Pool, Sqlite};
use crate::writer::{LEASE_STALE_MS, WriterHandle, epoch_ms_now};
pub fn health_router(pool: Pool<Sqlite>, writer: WriterHandle) -> Router {
let state = HealthState { pool, writer };
Router::new()
.route("/health", get(handle_health))
.route("/ready", get(handle_ready))
.layer(Extension(state))
}
#[derive(Clone)]
struct HealthState {
pool: Pool<Sqlite>,
writer: WriterHandle,
}
#[derive(Debug, Clone, Copy, Serialize)]
#[serde(rename_all = "lowercase")]
enum OkOrDegraded {
Ok,
Degraded,
}
#[derive(Debug, Clone, Copy, Serialize)]
#[serde(rename_all = "lowercase")]
enum OkOrFailed {
Ok,
Failed,
}
#[derive(Debug, Clone, Copy, Serialize)]
#[serde(rename_all = "lowercase")]
enum HealthStatus {
Ok,
}
#[derive(Debug, Serialize)]
struct HealthResponse {
status: HealthStatus,
version: &'static str,
}
#[derive(Debug, Serialize)]
struct ReadyResponse {
status: OkOrDegraded,
version: &'static str,
checks: ReadyChecks,
}
#[derive(Debug, Serialize)]
struct ReadyChecks {
database: OkOrFailed,
signing_key: OkOrFailed,
label_stream: OkOrDegraded,
}
async fn handle_health() -> Response {
let body = HealthResponse {
status: HealthStatus::Ok,
version: env!("CARGO_PKG_VERSION"),
};
(StatusCode::OK, Json(body)).into_response()
}
async fn handle_ready(Extension(state): Extension<HealthState>) -> Response {
let shutdown_rx = state.writer.shutdown_signal();
let (database, signing_key, label_stream) = tokio::join!(
check_database(&state.pool),
check_signing_key(&state.pool),
check_label_stream(&state.pool, &shutdown_rx),
);
let overall = match (database, signing_key, label_stream) {
(OkOrFailed::Ok, OkOrFailed::Ok, OkOrDegraded::Ok) => OkOrDegraded::Ok,
_ => OkOrDegraded::Degraded,
};
let code = match overall {
OkOrDegraded::Ok => StatusCode::OK,
OkOrDegraded::Degraded => StatusCode::SERVICE_UNAVAILABLE,
};
let body = ReadyResponse {
status: overall,
version: env!("CARGO_PKG_VERSION"),
checks: ReadyChecks {
database,
signing_key,
label_stream,
},
};
(code, Json(body)).into_response()
}
async fn check_database(pool: &Pool<Sqlite>) -> OkOrFailed {
match sqlx::query("SELECT 1").execute(pool).await {
Ok(_) => OkOrFailed::Ok,
Err(_) => OkOrFailed::Failed,
}
}
async fn check_signing_key(pool: &Pool<Sqlite>) -> OkOrFailed {
match sqlx::query_scalar!("SELECT COUNT(*) FROM signing_keys WHERE valid_to IS NULL")
.fetch_one(pool)
.await
{
Ok(n) if n > 0 => OkOrFailed::Ok,
_ => OkOrFailed::Failed,
}
}
async fn check_label_stream(
pool: &Pool<Sqlite>,
shutdown_rx: &tokio::sync::watch::Receiver<bool>,
) -> OkOrDegraded {
if *shutdown_rx.borrow() {
return OkOrDegraded::Degraded;
}
let row = sqlx::query_scalar!("SELECT last_heartbeat FROM server_instance_lease WHERE id = 1")
.fetch_optional(pool)
.await;
match row {
Ok(Some(hb)) => {
let age_ms = (epoch_ms_now() - hb).max(0);
if age_ms < LEASE_STALE_MS {
OkOrDegraded::Ok
} else {
OkOrDegraded::Degraded
}
}
_ => OkOrDegraded::Degraded,
}
}