use sqlx::PgConnection;
pub const AUDIT_CONTEXT_VARS: [&str; 6] = [
"app.actor",
"app.correlation_id",
"app.client_ip",
"app.user_agent",
"app.http_method",
"app.resource_path",
];
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub struct RequestAuditContext {
pub actor: String,
pub correlation_id: String,
pub client_ip: String,
pub user_agent: String,
pub http_method: String,
pub resource_path: String,
}
impl RequestAuditContext {
pub fn new(actor: impl Into<String>) -> Self {
Self {
actor: actor.into(),
..Self::default()
}
}
pub fn pairs(&self) -> [(&'static str, &str); 6] {
[
("app.actor", self.actor.as_str()),
("app.correlation_id", self.correlation_id.as_str()),
("app.client_ip", self.client_ip.as_str()),
("app.user_agent", self.user_agent.as_str()),
("app.http_method", self.http_method.as_str()),
("app.resource_path", self.resource_path.as_str()),
]
}
pub async fn bind_on(&self, conn: &mut PgConnection, local: bool) -> Result<(), sqlx::Error> {
for (var, value) in self.pairs() {
sqlx::query("SELECT set_config($1, $2, $3)")
.bind(var)
.bind(value)
.bind(local)
.execute(&mut *conn)
.await?;
}
Ok(())
}
}
tokio::task_local! {
static REQUEST_AUDIT: RequestAuditContext;
}
pub(crate) async fn with_request_audit<F>(audit: RequestAuditContext, f: F) -> F::Output
where
F: std::future::Future,
{
REQUEST_AUDIT.scope(audit, f).await
}
pub fn current_request_audit() -> Option<RequestAuditContext> {
REQUEST_AUDIT.try_with(|audit| audit.clone()).ok()
}
pub async fn relay_ambient_audit_on(conn: &mut PgConnection) -> Result<(), sqlx::Error> {
if let Some(audit) = current_request_audit() {
audit.bind_on(conn, true).await?;
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn the_audit_context_travels_the_task_local() {
assert!(current_request_audit().is_none());
let audit = RequestAuditContext::new("user-1");
let seen = with_request_audit(audit.clone(), async { current_request_audit() }).await;
assert_eq!(seen, Some(audit));
assert!(current_request_audit().is_none());
}
#[tokio::test]
async fn relay_binds_the_actor_onto_a_transaction_local_channel() {
let url = std::env::var("DATABASE_URL").unwrap_or_default();
if url.is_empty() {
eprintln!("SKIP| audit relay: DATABASE_URL unset");
return;
}
let pool = match sqlx::PgPool::connect(&url).await {
Ok(p) => p,
Err(e) => {
eprintln!("SKIP| audit relay: {e}");
return;
}
};
let audit = RequestAuditContext {
actor: "relay-probe".to_string(),
correlation_id: "corr-1".to_string(),
..RequestAuditContext::default()
};
with_request_audit(audit, async {
let mut tx = pool.begin().await.expect("begin");
relay_ambient_audit_on(&mut tx).await.expect("relay");
let actor: String = sqlx::query_scalar("SELECT current_setting('app.actor', true)")
.fetch_one(&mut *tx)
.await
.expect("read actor");
let correlation: String =
sqlx::query_scalar("SELECT current_setting('app.correlation_id', true)")
.fetch_one(&mut *tx)
.await
.expect("read correlation");
assert_eq!(actor, "relay-probe");
assert_eq!(correlation, "corr-1");
tx.rollback().await.expect("rollback");
})
.await;
}
}