use sqlx::pool::PoolConnection;
use sqlx::postgres::{PgArguments, PgRow};
use sqlx::query::{Query, QueryAs, QueryScalar};
use sqlx::{FromRow, PgPool, Postgres};
use std::future::Future;
use std::sync::Arc;
use tokio::sync::Mutex;
use uuid::Uuid;
tokio::task_local! {
static COMPANY: Option<Uuid>;
static REQUEST_CONN: Arc<Mutex<PoolConnection<Postgres>>>;
}
pub async fn with_request_scope<F, R>(pool: &PgPool, company: Uuid, f: F) -> Result<R, sqlx::Error>
where
F: Future<Output = R>,
{
let mut conn = pool.acquire().await?;
sqlx::query("SELECT set_config('app.company_id', $1, false)")
.bind(company.to_string())
.execute(&mut *conn)
.await?;
let holder = Arc::new(Mutex::new(conn));
let result = COMPANY
.scope(Some(company), REQUEST_CONN.scope(holder.clone(), f))
.await;
{
let mut guard = holder.lock().await;
if let Err(e) = sqlx::query("SELECT set_config('app.company_id', '', false)")
.execute(&mut **guard)
.await
{
tracing::error!(
target: "backbone_orm::company_scope",
error = %e,
"failed to reset app.company_id on request connection; the pool connection may carry \
the previous tenant's company_id — treat as a fence-hygiene incident",
);
}
}
Ok(result)
}
fn request_conn() -> Option<Arc<Mutex<PoolConnection<Postgres>>>> {
REQUEST_CONN.try_with(|c| c.clone()).ok()
}
pub(crate) fn current_request_conn() -> Option<Arc<Mutex<PoolConnection<Postgres>>>> {
request_conn()
}
pub async fn with_company_scope<F, R>(company: Option<Uuid>, f: F) -> R
where
F: Future<Output = R>,
{
COMPANY.scope(company, f).await
}
pub(crate) async fn with_company_scope_internal<F, R>(company: Option<Uuid>, f: F) -> R
where
F: Future<Output = R>,
{
COMPANY.scope(company, f).await
}
pub(crate) async fn with_request_conn_internal<F, R>(
holder: Arc<Mutex<PoolConnection<Postgres>>>,
f: F,
) -> R
where
F: Future<Output = R>,
{
REQUEST_CONN.scope(holder, f).await
}
pub fn current_company() -> Option<Uuid> {
COMPANY.try_with(|c| *c).ok().flatten()
}
async fn bind_company(conn: &mut sqlx::PgConnection, company: Uuid) -> Result<(), sqlx::Error> {
sqlx::query("SELECT set_config('app.company_id', $1, true)")
.bind(company.to_string())
.execute(conn)
.await?;
Ok(())
}
pub async fn bind_company_on(
conn: &mut sqlx::PgConnection,
company: Uuid,
) -> Result<(), sqlx::Error> {
bind_company(conn, company).await
}
pub async fn bind_current_company(conn: &mut sqlx::PgConnection) -> Result<(), sqlx::Error> {
if let Some(company) = current_company() {
bind_company(conn, company).await?;
}
Ok(())
}
pub async fn fetch_all_scoped<'q, T>(
pool: &PgPool,
query: QueryAs<'q, Postgres, T, PgArguments>,
) -> Result<Vec<T>, sqlx::Error>
where
T: for<'r> FromRow<'r, PgRow> + Send + Unpin,
{
if let Some(conn) = request_conn() {
let mut g = conn.lock().await;
return query.fetch_all(&mut **g).await;
}
match current_company() {
None => query.fetch_all(pool).await,
Some(company) => {
let mut tx = pool.begin().await?;
bind_company(&mut tx, company).await?;
let rows = query.fetch_all(&mut *tx).await?;
tx.commit().await?;
Ok(rows)
}
}
}
pub async fn fetch_one_scoped<'q, T>(
pool: &PgPool,
query: QueryAs<'q, Postgres, T, PgArguments>,
) -> Result<T, sqlx::Error>
where
T: for<'r> FromRow<'r, PgRow> + Send + Unpin,
{
if let Some(conn) = request_conn() {
let mut g = conn.lock().await;
return query.fetch_one(&mut **g).await;
}
match current_company() {
None => query.fetch_one(pool).await,
Some(company) => {
let mut tx = pool.begin().await?;
bind_company(&mut tx, company).await?;
let row = query.fetch_one(&mut *tx).await?;
tx.commit().await?;
Ok(row)
}
}
}
pub async fn fetch_optional_scoped<'q, T>(
pool: &PgPool,
query: QueryAs<'q, Postgres, T, PgArguments>,
) -> Result<Option<T>, sqlx::Error>
where
T: for<'r> FromRow<'r, PgRow> + Send + Unpin,
{
if let Some(conn) = request_conn() {
let mut g = conn.lock().await;
return query.fetch_optional(&mut **g).await;
}
match current_company() {
None => query.fetch_optional(pool).await,
Some(company) => {
let mut tx = pool.begin().await?;
bind_company(&mut tx, company).await?;
let row = query.fetch_optional(&mut *tx).await?;
tx.commit().await?;
Ok(row)
}
}
}
pub async fn fetch_one_scalar_scoped<'q, S>(
pool: &PgPool,
query: QueryScalar<'q, Postgres, S, PgArguments>,
) -> Result<S, sqlx::Error>
where
S: Send + Unpin,
(S,): for<'r> FromRow<'r, PgRow>,
{
if let Some(conn) = request_conn() {
let mut g = conn.lock().await;
return query.fetch_one(&mut **g).await;
}
match current_company() {
None => query.fetch_one(pool).await,
Some(company) => {
let mut tx = pool.begin().await?;
bind_company(&mut tx, company).await?;
let val = query.fetch_one(&mut *tx).await?;
tx.commit().await?;
Ok(val)
}
}
}
pub async fn fetch_optional_scalar_scoped<'q, S>(
pool: &PgPool,
query: QueryScalar<'q, Postgres, S, PgArguments>,
) -> Result<Option<S>, sqlx::Error>
where
S: Send + Unpin,
(S,): for<'r> FromRow<'r, PgRow>,
{
if let Some(conn) = request_conn() {
let mut g = conn.lock().await;
return query.fetch_optional(&mut **g).await;
}
match current_company() {
None => query.fetch_optional(pool).await,
Some(company) => {
let mut tx = pool.begin().await?;
bind_company(&mut tx, company).await?;
let val = query.fetch_optional(&mut *tx).await?;
tx.commit().await?;
Ok(val)
}
}
}
pub async fn fetch_optional_row_scoped<'q>(
pool: &PgPool,
query: Query<'q, Postgres, PgArguments>,
) -> Result<Option<PgRow>, sqlx::Error> {
if let Some(conn) = request_conn() {
let mut g = conn.lock().await;
return query.fetch_optional(&mut **g).await;
}
match current_company() {
None => query.fetch_optional(pool).await,
Some(company) => {
let mut tx = pool.begin().await?;
bind_company(&mut tx, company).await?;
let row = query.fetch_optional(&mut *tx).await?;
tx.commit().await?;
Ok(row)
}
}
}
pub async fn fetch_one_row_scoped<'q>(
pool: &PgPool,
query: Query<'q, Postgres, PgArguments>,
) -> Result<PgRow, sqlx::Error> {
if let Some(conn) = request_conn() {
let mut g = conn.lock().await;
return query.fetch_one(&mut **g).await;
}
match current_company() {
None => query.fetch_one(pool).await,
Some(company) => {
let mut tx = pool.begin().await?;
bind_company(&mut tx, company).await?;
let row = query.fetch_one(&mut *tx).await?;
tx.commit().await?;
Ok(row)
}
}
}
pub async fn fetch_all_rows_scoped<'q>(
pool: &PgPool,
query: Query<'q, Postgres, PgArguments>,
) -> Result<Vec<PgRow>, sqlx::Error> {
if let Some(conn) = request_conn() {
let mut g = conn.lock().await;
return query.fetch_all(&mut **g).await;
}
match current_company() {
None => query.fetch_all(pool).await,
Some(company) => {
let mut tx = pool.begin().await?;
bind_company(&mut tx, company).await?;
let rows = query.fetch_all(&mut *tx).await?;
tx.commit().await?;
Ok(rows)
}
}
}
pub async fn execute_scoped<'q>(
pool: &PgPool,
query: Query<'q, Postgres, PgArguments>,
) -> Result<sqlx::postgres::PgQueryResult, sqlx::Error> {
if let Some(conn) = request_conn() {
let mut g = conn.lock().await;
return query.execute(&mut **g).await;
}
match current_company() {
None => query.execute(pool).await,
Some(company) => {
let mut tx = pool.begin().await?;
bind_company(&mut tx, company).await?;
let res = query.execute(&mut *tx).await?;
tx.commit().await?;
Ok(res)
}
}
}
#[cfg(test)]
mod tests {
use super::{request_conn, with_request_scope};
use sqlx::postgres::PgPoolOptions;
use sqlx::PgPool;
use uuid::Uuid;
fn dsn() -> Option<String> {
std::env::var("BACKBONE_ORM_RLS_DSN").ok()
}
async fn admin_pool(dsn: &str) -> PgPool {
PgPoolOptions::new().max_connections(4).connect(dsn).await.unwrap()
}
async fn app_pool(dsn: &str, role: &str) -> PgPool {
let after_at = dsn.rsplit('@').next().unwrap();
let url = format!("postgresql://{role}:rlspw@{after_at}");
PgPoolOptions::new().max_connections(1).connect(&url).await.unwrap()
}
fn role_name() -> String {
format!("rls_reset_app_{}", &Uuid::new_v4().simple().to_string()[..8])
}
async fn setup(admin: &PgPool, role: &str) {
sqlx::raw_sql(&format!(
"DROP SCHEMA IF EXISTS rls_reset_test CASCADE; \
CREATE SCHEMA rls_reset_test; \
CREATE ROLE {role} LOGIN PASSWORD 'rlspw'; \
GRANT USAGE ON SCHEMA rls_reset_test TO {role}; \
CREATE TABLE rls_reset_test.t (id uuid PRIMARY KEY, company_id uuid NOT NULL); \
GRANT SELECT, INSERT, UPDATE, DELETE ON rls_reset_test.t TO {role};",
))
.execute(admin).await.unwrap();
}
#[tokio::test]
async fn lingering_request_conn_clone_does_not_dirty_the_pooled_connection() {
let Some(dsn) = dsn() else { eprintln!("skipping: set BACKBONE_ORM_RLS_DSN"); return; };
let role = role_name();
let admin = admin_pool(&dsn).await;
setup(&admin, &role).await;
let pool = app_pool(&dsn, &role).await;
let company_a = Uuid::new_v4();
let (tx, rx) = tokio::sync::oneshot::channel();
with_request_scope(&pool, company_a, async {
if let Some(conn) = request_conn() {
let _ = tx.send(conn);
}
})
.await
.unwrap();
let held = rx.await.unwrap();
drop(held);
let mut conn = pool.acquire().await.unwrap();
let setting: String =
sqlx::query_scalar("SELECT current_setting('app.company_id', true)")
.fetch_one(&mut *conn)
.await
.unwrap();
assert_eq!(
setting, "",
"app.company_id leaked onto the pooled connection after scope exit — a lingering \
REQUEST_CONN clone must not bypass the session-var reset (cross-tenant leak)"
);
}
}