use std::collections::HashMap;
use std::future::Future;
use std::marker::PhantomData;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, Mutex};
use std::time::{Duration, Instant, SystemTime, UNIX_EPOCH};
use sqlx::{Database, Encode, FromRow, IntoArguments, Pool, Type};
use super::{InMemoryLock, Lock, LockError, LockManager};
#[derive(Debug, Clone)]
pub(crate) struct LeaseConfig {
pub lease_ttl: Duration,
pub retry_interval: Duration,
pub max_wait: Option<Duration>,
}
impl Default for LeaseConfig {
fn default() -> Self {
Self {
lease_ttl: Duration::from_secs(30),
retry_interval: Duration::from_millis(50),
max_wait: None,
}
}
}
pub(crate) struct LockShared {
pub gate: InMemoryLock,
pub token: Mutex<Option<String>>,
}
impl LockShared {
pub fn new() -> Self {
Self {
gate: InMemoryLock::new(),
token: Mutex::new(None),
}
}
}
pub(crate) trait LeaseBackend: Send + Sync {
fn shared(&self) -> &LockShared;
fn config(&self) -> &LeaseConfig;
fn mint_token(&self) -> String;
fn db_acquire(&self, token: &str) -> impl Future<Output = Result<bool, LockError>> + Send;
fn db_release(&self, token: &str) -> impl Future<Output = Result<(), LockError>> + Send;
}
struct GateGuard<'a> {
gate: &'a InMemoryLock,
armed: bool,
}
impl<'a> GateGuard<'a> {
fn new(gate: &'a InMemoryLock) -> Self {
Self { gate, armed: true }
}
fn disarm(&mut self) {
self.armed = false;
}
}
impl Drop for GateGuard<'_> {
fn drop(&mut self) {
if self.armed {
let _ = self.gate.unlock_core();
}
}
}
pub(crate) async fn lease_lock(backend: &impl LeaseBackend) -> Result<(), LockError> {
let started = Instant::now();
let max_wait = backend.config().max_wait;
match max_wait {
Some(max) => match tokio::time::timeout(max, backend.shared().gate.lock()).await {
Ok(result) => result?,
Err(_elapsed) => {
return Err(LockError::AcquireFailed(format!(
"lease acquire timed out after {max:?} waiting for the in-process gate"
)))
}
},
None => backend.shared().gate.lock().await?,
}
let mut guard = GateGuard::new(&backend.shared().gate);
let token = backend.mint_token();
loop {
match backend.db_acquire(&token).await {
Ok(true) => {
store_token(backend.shared(), Some(token));
guard.disarm(); return Ok(());
}
Ok(false) => {}
Err(err) => return Err(err), }
if let Some(max) = max_wait {
if started.elapsed() >= max {
return Err(LockError::AcquireFailed(format!(
"lease acquire timed out after {max:?}"
)));
}
}
tokio::time::sleep(jittered(backend.config().retry_interval, &token)).await;
}
}
pub(crate) async fn lease_try_lock(backend: &impl LeaseBackend) -> Result<bool, LockError> {
if !backend.shared().gate.try_lock().await? {
return Ok(false);
}
let mut guard = GateGuard::new(&backend.shared().gate);
let token = backend.mint_token();
match backend.db_acquire(&token).await {
Ok(true) => {
store_token(backend.shared(), Some(token));
guard.disarm();
Ok(true)
}
Ok(false) => Ok(false), Err(err) => Err(err), }
}
pub(crate) async fn lease_unlock(backend: &impl LeaseBackend) -> Result<(), LockError> {
let _guard = GateGuard::new(&backend.shared().gate);
let token = store_token(backend.shared(), None);
if let Some(token) = token {
let _ = backend.db_release(&token).await;
}
Ok(())
}
fn store_token(shared: &LockShared, next: Option<String>) -> Option<String> {
let mut slot = shared
.token
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
std::mem::replace(&mut slot, next)
}
pub(crate) fn default_owner_id(prefix: &str) -> String {
static MANAGER_SEQ: AtomicU64 = AtomicU64::new(0);
let pid = std::process::id();
let nanos = SystemTime::now()
.duration_since(UNIX_EPOCH)
.map(|d| d.as_nanos())
.unwrap_or(0);
let seq = MANAGER_SEQ.fetch_add(1, Ordering::Relaxed);
format!("{prefix}-{pid}-{nanos}-{seq}")
}
pub(crate) fn mint_token(owner_id: &str, seq: &AtomicU64) -> String {
format!("{owner_id}:{}", seq.fetch_add(1, Ordering::Relaxed))
}
fn jittered(base: Duration, token: &str) -> Duration {
use std::hash::{Hash, Hasher};
let mut hasher = std::collections::hash_map::DefaultHasher::new();
token.hash(&mut hasher);
let frac = (hasher.finish() % 1000) as f64 / 1000.0;
base + base.mul_f64(0.25 * frac)
}
pub(crate) fn lease_acquire_error(err: sqlx::Error) -> LockError {
LockError::AcquireFailed(format!("sqlx lease acquire failed: {err}"))
}
pub(crate) fn lease_release_error(err: sqlx::Error) -> LockError {
LockError::ReleaseFailed(format!("sqlx lease release failed: {err}"))
}
pub trait LockDialect: Send + Sync + 'static {
type Db: Database;
const OWNER_PREFIX: &'static str;
const DDL: &'static str;
const ACQUIRE_SQL: &'static str;
const RELEASE_SQL: &'static str;
const SWEEP_SQL: &'static str;
fn busy_is_contention(_err: &sqlx::Error) -> bool {
false
}
fn rows_affected(result: <Self::Db as Database>::QueryResult) -> u64;
}
pub trait LeaseQueries: LockDialect {
fn acquire(
pool: &Pool<Self::Db>,
key: &str,
token: &str,
ttl_secs: f64,
) -> impl Future<Output = Result<bool, LockError>> + Send;
fn release(
pool: &Pool<Self::Db>,
key: &str,
token: &str,
) -> impl Future<Output = Result<(), LockError>> + Send;
fn sweep(pool: &Pool<Self::Db>) -> impl Future<Output = Result<u64, LockError>> + Send;
fn migrate(pool: &Pool<Self::Db>) -> impl Future<Output = Result<(), LockError>> + Send;
}
impl<D> LeaseQueries for D
where
D: LockDialect,
for<'c> &'c mut <D::Db as Database>::Connection: sqlx::Executor<'c, Database = D::Db>,
<D::Db as Database>::Arguments: IntoArguments<D::Db>,
str: Type<D::Db>,
for<'q> &'q str: Encode<'q, D::Db>,
f64: Type<D::Db> + for<'q> Encode<'q, D::Db>,
for<'r> (String,): FromRow<'r, <D::Db as Database>::Row>,
{
async fn acquire(
pool: &Pool<Self::Db>,
key: &str,
token: &str,
ttl_secs: f64,
) -> Result<bool, LockError> {
let outcome = sqlx::query_as::<_, (String,)>(D::ACQUIRE_SQL)
.bind(key)
.bind(token)
.bind(ttl_secs)
.fetch_optional(pool)
.await;
match outcome {
Ok(row) => Ok(matches!(row, Some((owner,)) if owner == token)),
Err(err) if D::busy_is_contention(&err) => Ok(false),
Err(err) => Err(lease_acquire_error(err)),
}
}
async fn release(pool: &Pool<Self::Db>, key: &str, token: &str) -> Result<(), LockError> {
sqlx::query(D::RELEASE_SQL)
.bind(key)
.bind(token)
.execute(pool)
.await
.map_err(lease_release_error)?;
Ok(())
}
async fn sweep(pool: &Pool<Self::Db>) -> Result<u64, LockError> {
let result = sqlx::query(D::SWEEP_SQL)
.execute(pool)
.await
.map_err(lease_release_error)?;
Ok(D::rows_affected(result))
}
async fn migrate(pool: &Pool<Self::Db>) -> Result<(), LockError> {
for statement in D::DDL.split(';') {
let statement = statement.trim();
if statement.is_empty() {
continue;
}
sqlx::query(statement).execute(pool).await.map_err(|err| {
LockError::Other(format!("migrate aggregate_locks failed: {err}"))
})?;
}
Ok(())
}
}
pub struct SqlxLockManager<D: LockDialect> {
pool: Pool<D::Db>,
owner_id: String,
config: LeaseConfig,
token_seq: Arc<AtomicU64>,
locks: Arc<Mutex<HashMap<String, Arc<SqlxLock<D>>>>>,
}
impl<D: LockDialect> Clone for SqlxLockManager<D> {
fn clone(&self) -> Self {
Self {
pool: self.pool.clone(),
owner_id: self.owner_id.clone(),
config: self.config.clone(),
token_seq: Arc::clone(&self.token_seq),
locks: Arc::clone(&self.locks),
}
}
}
impl<D: LockDialect> SqlxLockManager<D> {
pub fn new(pool: Pool<D::Db>) -> Self {
Self {
pool,
owner_id: default_owner_id(D::OWNER_PREFIX),
config: LeaseConfig::default(),
token_seq: Arc::new(AtomicU64::new(0)),
locks: Arc::new(Mutex::new(HashMap::new())),
}
}
pub fn with_lease_ttl(mut self, ttl: Duration) -> Self {
self.config.lease_ttl = ttl;
self
}
pub fn with_retry_interval(mut self, interval: Duration) -> Self {
self.config.retry_interval = interval;
self
}
pub fn with_max_wait(mut self, max_wait: Option<Duration>) -> Self {
self.config.max_wait = max_wait;
self
}
pub fn with_owner_id(mut self, owner_id: impl Into<String>) -> Self {
self.owner_id = owner_id.into();
self
}
}
impl<D: LeaseQueries> SqlxLockManager<D> {
pub async fn migrate(pool: &Pool<D::Db>) -> Result<(), LockError> {
D::migrate(pool).await
}
pub async fn sweep_expired(&self) -> Result<u64, LockError> {
D::sweep(&self.pool).await
}
}
impl<D: LeaseQueries> LockManager for SqlxLockManager<D> {
type Lock = SqlxLock<D>;
fn get_lock(&self, id: &str) -> Result<Arc<SqlxLock<D>>, LockError> {
let mut locks = self
.locks
.lock()
.map_err(|_| LockError::Poisoned("sqlx lock manager map poisoned".into()))?;
Ok(locks
.entry(id.to_string())
.or_insert_with(|| {
Arc::new(SqlxLock {
pool: self.pool.clone(),
owner_id: self.owner_id.clone(),
config: self.config.clone(),
token_seq: Arc::clone(&self.token_seq),
key: id.to_string(),
shared: LockShared::new(),
dialect: PhantomData,
})
})
.clone())
}
}
pub struct SqlxLock<D: LockDialect> {
pool: Pool<D::Db>,
owner_id: String,
config: LeaseConfig,
token_seq: Arc<AtomicU64>,
key: String,
shared: LockShared,
dialect: PhantomData<D>,
}
impl<D: LeaseQueries> LeaseBackend for SqlxLock<D> {
fn shared(&self) -> &LockShared {
&self.shared
}
fn config(&self) -> &LeaseConfig {
&self.config
}
fn mint_token(&self) -> String {
mint_token(&self.owner_id, &self.token_seq)
}
async fn db_acquire(&self, token: &str) -> Result<bool, LockError> {
D::acquire(
&self.pool,
&self.key,
token,
self.config.lease_ttl.as_secs_f64(),
)
.await
}
async fn db_release(&self, token: &str) -> Result<(), LockError> {
D::release(&self.pool, &self.key, token).await
}
}
impl<D: LeaseQueries> Lock for SqlxLock<D> {
fn lock(&self) -> impl Future<Output = Result<(), LockError>> + Send + '_ {
lease_lock(self)
}
fn try_lock(&self) -> impl Future<Output = Result<bool, LockError>> + Send + '_ {
lease_try_lock(self)
}
fn unlock(&self) -> impl Future<Output = Result<(), LockError>> + Send + '_ {
lease_unlock(self)
}
}