use std::time::Duration;
use axum::Json;
use axum::extract::State;
use axum::http::StatusCode;
use serde::Serialize;
use crate::state::AppState;
const DB_CHECK_TIMEOUT: Duration = Duration::from_secs(2);
#[derive(Serialize)]
pub struct HealthResponse {
pub status: &'static str,
}
pub async fn up_check() -> (StatusCode, Json<HealthResponse>) {
(StatusCode::OK, Json(HealthResponse { status: "ok" }))
}
pub async fn health_check(State(state): State<AppState>) -> (StatusCode, Json<HealthResponse>) {
match check_db(&state).await {
Ok(()) => (StatusCode::OK, Json(HealthResponse { status: "ok" })),
Err(err) => {
tracing::warn!(error = %err, "health check: database unreachable");
(
StatusCode::SERVICE_UNAVAILABLE,
Json(HealthResponse {
status: "unavailable",
}),
)
}
}
}
async fn check_db(state: &AppState) -> Result<(), sqlx::Error> {
let pool = state.tenant_db.pool().clone();
let probe = async move {
let mut conn = pool.acquire().await?;
sqlx::query("SELECT 1").execute(conn.as_mut()).await?;
Ok::<(), sqlx::Error>(())
};
match tokio::time::timeout(DB_CHECK_TIMEOUT, probe).await {
Ok(result) => result,
Err(_) => Err(sqlx::Error::PoolTimedOut),
}
}
#[cfg(test)]
#[path = "../../../tests/http/controllers/health.rs"]
mod tests;