use std::collections::HashMap;
use std::time::{Duration, Instant, SystemTime, UNIX_EPOCH};
use crate::storage::s3::{LockClient, LockItem, StorageError};
use maplit::hashmap;
use rusoto_core::RusotoError;
use rusoto_dynamodb::*;
use uuid::Uuid;
mod options {
pub const PARTITION_KEY_VALUE: &str = "DYNAMO_LOCK_PARTITION_KEY_VALUE";
pub const TABLE_NAME: &str = "DYNAMO_LOCK_TABLE_NAME";
pub const OWNER_NAME: &str = "DYNAMO_LOCK_OWNER_NAME";
pub const LEASE_DURATION: &str = "DYNAMO_LOCK_LEASE_DURATION";
pub const REFRESH_PERIOD_MILLIS: &str = "DYNAMO_LOCK_REFRESH_PERIOD_MILLIS";
pub const ADDITIONAL_TIME_TO_WAIT_MILLIS: &str = "DYNAMO_LOCK_ADDITIONAL_TIME_TO_WAIT_MILLIS";
}
#[derive(Clone, Debug)]
pub struct Options {
pub partition_key_value: String,
pub table_name: String,
pub owner_name: String,
pub lease_duration: u64,
pub refresh_period: Duration,
pub additional_time_to_wait_for_lock: Duration,
}
impl Default for Options {
fn default() -> Self {
fn str_env(key: &str, default: String) -> String {
std::env::var(key).unwrap_or(default)
}
fn u64_env(key: &str, default: u64) -> u64 {
std::env::var(key)
.ok()
.and_then(|e| e.parse::<u64>().ok())
.unwrap_or(default)
}
let refresh_period = Duration::from_millis(u64_env(options::REFRESH_PERIOD_MILLIS, 1000));
let additional_time_to_wait_for_lock =
Duration::from_millis(u64_env(options::ADDITIONAL_TIME_TO_WAIT_MILLIS, 1000));
Self {
partition_key_value: str_env(options::PARTITION_KEY_VALUE, "delta-rs".to_string()),
table_name: str_env(options::TABLE_NAME, "delta_rs_lock_table".to_string()),
owner_name: str_env(options::OWNER_NAME, Uuid::new_v4().to_string()),
lease_duration: u64_env(options::LEASE_DURATION, 20),
refresh_period,
additional_time_to_wait_for_lock,
}
}
}
impl LockItem {
fn is_expired(&self) -> bool {
if self.is_released {
return true;
}
now_millis() - self.lookup_time > (self.lease_duration as u128) * 1000
}
}
#[derive(thiserror::Error, Debug)]
pub enum DynamoError {
#[error("Dynamo table not found")]
TableNotFound,
#[error("Conditional check failed")]
ConditionalCheckFailed,
#[error("DynamoDB item has invalid schema")]
InvalidItemSchema,
#[error("Could not acquire lock for {0} sec")]
TimedOut(u64),
#[error("Maximum allowed provisioned throughput for the table exceeded")]
ProvisionedThroughputExceeded,
#[error("Put item error: {0}")]
PutItemError(RusotoError<PutItemError>),
#[error("Update item error: {0}")]
UpdateItemError(#[from] RusotoError<UpdateItemError>),
#[error("Get item error: {0}")]
GetItemError(RusotoError<GetItemError>),
}
impl From<RusotoError<PutItemError>> for DynamoError {
fn from(error: RusotoError<PutItemError>) -> Self {
match error {
RusotoError::Service(PutItemError::ConditionalCheckFailed(_)) => {
DynamoError::ConditionalCheckFailed
}
RusotoError::Service(PutItemError::ProvisionedThroughputExceeded(_)) => {
DynamoError::ProvisionedThroughputExceeded
}
_ => DynamoError::PutItemError(error),
}
}
}
impl From<RusotoError<GetItemError>> for DynamoError {
fn from(error: RusotoError<GetItemError>) -> Self {
match error {
RusotoError::Service(GetItemError::ResourceNotFound(_)) => DynamoError::TableNotFound,
RusotoError::Service(GetItemError::ProvisionedThroughputExceeded(_)) => {
DynamoError::ProvisionedThroughputExceeded
}
_ => DynamoError::GetItemError(error),
}
}
}
pub const PARTITION_KEY_NAME: &str = "key";
pub const OWNER_NAME: &str = "ownerName";
pub const RECORD_VERSION_NUMBER: &str = "recordVersionNumber";
pub const IS_RELEASED: &str = "isReleased";
pub const LEASE_DURATION: &str = "leaseDuration";
pub const DATA: &str = "data";
mod expressions {
pub const ACQUIRE_LOCK_THAT_DOESNT_EXIST: &str = "attribute_not_exists(#pk)";
pub const PK_EXISTS_AND_IS_RELEASED: &str = "attribute_exists(#pk) AND #ir = :ir";
pub const PK_EXISTS_AND_RVN_MATCHES: &str = "attribute_exists(#pk) AND #rvn = :rvn";
pub const PK_EXISTS_AND_OWNER_RVN_MATCHES: &str =
"attribute_exists(#pk) AND #rvn = :rvn AND #on = :on";
pub const UPDATE_IS_RELEASED_AND_DATA: &str = "SET #ir = :ir, #d = :d";
pub const UPDATE_IS_RELEASED: &str = "SET #ir = :ir";
}
mod vars {
pub const PK_PATH: &str = "#pk";
pub const RVN_PATH: &str = "#rvn";
pub const RVN_VALUE: &str = ":rvn";
pub const IS_RELEASED_PATH: &str = "#ir";
pub const IS_RELEASED_VALUE: &str = ":ir";
pub const OWNER_NAME_PATH: &str = "#on";
pub const OWNER_NAME_VALUE: &str = ":on";
pub const DATA_PATH: &str = "#d";
pub const DATA_VALUE: &str = ":d";
}
pub struct DynamoDbLockClient {
client: DynamoDbClient,
opts: Options,
}
impl std::fmt::Debug for DynamoDbLockClient {
fn fmt(&self, fmt: &mut std::fmt::Formatter<'_>) -> Result<(), std::fmt::Error> {
write!(fmt, "DynamoDbLockClient")
}
}
#[async_trait::async_trait]
impl LockClient for DynamoDbLockClient {
async fn try_acquire_lock(&self) -> Result<Option<LockItem>, StorageError> {
Ok(self.try_acquire_lock().await?)
}
async fn get_lock(&self) -> Result<Option<LockItem>, StorageError> {
Ok(self.get_lock().await?)
}
async fn release_lock(&self, lock: &LockItem) -> Result<bool, StorageError> {
Ok(self.release_lock(lock).await?)
}
}
impl DynamoDbLockClient {
pub fn new(client: DynamoDbClient, opts: Options) -> Self {
Self { client, opts }
}
pub async fn try_acquire_lock(&self) -> Result<Option<LockItem>, DynamoError> {
match self.acquire_lock().await {
Ok(lock) => Ok(Some(lock)),
Err(DynamoError::TimedOut(_)) => Ok(None),
Err(DynamoError::ProvisionedThroughputExceeded) => Ok(None),
Err(e) => Err(e),
}
}
pub async fn acquire_lock(&self) -> Result<LockItem, DynamoError> {
let mut state = AcquireLockState {
client: self,
cached_lock: None,
started: Instant::now(),
timeout_in: self.opts.additional_time_to_wait_for_lock,
};
loop {
match state.try_acquire_lock().await {
Ok(lock) => return Ok(lock),
Err(DynamoError::ConditionalCheckFailed) => {
if state.has_timed_out() {
return Err(DynamoError::TimedOut(state.started.elapsed().as_secs()));
}
tokio::time::sleep(self.opts.refresh_period).await;
}
Err(e) => return Err(e),
}
}
}
pub async fn get_lock(&self) -> Result<Option<LockItem>, DynamoError> {
let output = self
.client
.get_item(GetItemInput {
consistent_read: Some(true),
table_name: self.opts.table_name.clone(),
key: hashmap! {
PARTITION_KEY_NAME.to_string() => attr(self.opts.partition_key_value.clone())
},
..Default::default()
})
.await?;
if let Some(item) = output.item {
let get_value = |key| -> Result<String, DynamoError> {
Ok(item
.get(key)
.and_then(|r| r.s.as_ref())
.ok_or(DynamoError::InvalidItemSchema)?
.clone())
};
let lease_duration = get_value(LEASE_DURATION)?
.parse::<u64>()
.map_err(|_| DynamoError::InvalidItemSchema)?;
return Ok(Some(LockItem {
owner_name: get_value(OWNER_NAME)?,
record_version_number: get_value(RECORD_VERSION_NUMBER)?,
lease_duration,
is_released: item.contains_key(IS_RELEASED),
data: get_value(DATA).ok(),
lookup_time: now_millis(),
}));
}
Ok(None)
}
pub async fn release_lock(&self, lock: &LockItem) -> Result<bool, DynamoError> {
let mut names = hashmap! {
vars::PK_PATH.to_string() => PARTITION_KEY_NAME.to_string(),
vars::RVN_PATH.to_string() => RECORD_VERSION_NUMBER.to_string(),
vars::OWNER_NAME_PATH.to_string() => OWNER_NAME.to_string(),
vars::IS_RELEASED_PATH.to_string() => IS_RELEASED.to_string(),
};
let mut values = hashmap! {
vars::IS_RELEASED_VALUE.to_string() => attr("1"),
vars::RVN_VALUE.to_string() => attr(&lock.record_version_number),
vars::OWNER_NAME_VALUE.to_string() => attr(&lock.owner_name),
};
let update: &str;
if let Some(ref data) = lock.data {
update = expressions::UPDATE_IS_RELEASED_AND_DATA;
names.insert(vars::DATA_PATH.to_string(), DATA.to_string());
values.insert(vars::DATA_VALUE.to_string(), attr(data));
} else {
update = expressions::UPDATE_IS_RELEASED;
}
let result = self.client.update_item(UpdateItemInput {
table_name: self.opts.table_name.clone(),
key: hashmap! {
PARTITION_KEY_NAME.to_string() => attr(self.opts.partition_key_value.clone())
},
condition_expression: Some(expressions::PK_EXISTS_AND_OWNER_RVN_MATCHES.to_string()),
update_expression: Some(update.to_string()),
expression_attribute_names: Some(names),
expression_attribute_values: Some(values),
..Default::default()
});
match result.await {
Ok(_) => Ok(true),
Err(RusotoError::Service(UpdateItemError::ConditionalCheckFailed(_))) => Ok(false),
Err(e) => Err(DynamoError::UpdateItemError(e)),
}
}
async fn upsert_item(
&self,
data: Option<String>,
condition_expression: Option<String>,
expression_attribute_names: Option<HashMap<String, String>>,
expression_attribute_values: Option<HashMap<String, AttributeValue>>,
) -> Result<LockItem, DynamoError> {
let rvn = Uuid::new_v4().to_string();
let mut item = hashmap! {
PARTITION_KEY_NAME.to_string() => attr(self.opts.partition_key_value.clone()),
OWNER_NAME.to_string() => attr(&self.opts.owner_name),
RECORD_VERSION_NUMBER.to_string() => attr(&rvn),
LEASE_DURATION.to_string() => attr(&self.opts.lease_duration),
};
if let Some(ref d) = data {
item.insert(DATA.to_string(), attr(d));
}
self.client
.put_item(PutItemInput {
table_name: self.opts.table_name.clone(),
item,
condition_expression,
expression_attribute_names,
expression_attribute_values,
..Default::default()
})
.await?;
Ok(LockItem {
owner_name: self.opts.owner_name.clone(),
record_version_number: rvn,
lease_duration: self.opts.lease_duration,
is_released: false,
data,
lookup_time: now_millis(),
})
}
}
fn now_millis() -> u128 {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.as_millis()
}
pub fn attr<T: ToString>(s: T) -> AttributeValue {
AttributeValue {
s: Some(s.to_string()),
..Default::default()
}
}
struct AcquireLockState<'a> {
client: &'a DynamoDbLockClient,
cached_lock: Option<LockItem>,
started: Instant,
timeout_in: Duration,
}
impl<'a> AcquireLockState<'a> {
fn has_timed_out(&self) -> bool {
self.started.elapsed() > self.timeout_in
}
async fn try_acquire_lock(&mut self) -> Result<LockItem, DynamoError> {
match self.client.get_lock().await? {
None => {
Ok(self.upsert_new_lock().await?)
}
Some(existing) if existing.is_released => {
Ok(self.upsert_released_lock(existing.data).await?)
}
Some(existing) => {
let cached = match self.cached_lock.as_ref() {
None => {
self.timeout_in = Duration::from_secs(
self.timeout_in.as_secs() + existing.lease_duration,
);
self.cached_lock = Some(existing);
return Err(DynamoError::ConditionalCheckFailed);
}
Some(cached) => cached,
};
let cached_rvn = &cached.record_version_number;
if cached_rvn == &existing.record_version_number {
if cached.is_expired() {
self.upsert_expired_lock(cached_rvn, existing.data).await
} else {
Err(DynamoError::ConditionalCheckFailed)
}
} else {
self.cached_lock = Some(existing);
return Err(DynamoError::ConditionalCheckFailed);
}
}
}
}
async fn upsert_new_lock(&self) -> Result<LockItem, DynamoError> {
self.client
.upsert_item(
None,
Some(expressions::ACQUIRE_LOCK_THAT_DOESNT_EXIST.to_string()),
Some(hashmap! {
vars::PK_PATH.to_string() => PARTITION_KEY_NAME.to_string(),
}),
None,
)
.await
}
async fn upsert_released_lock(&self, data: Option<String>) -> Result<LockItem, DynamoError> {
self.client
.upsert_item(
data,
Some(expressions::PK_EXISTS_AND_IS_RELEASED.to_string()),
Some(hashmap! {
vars::PK_PATH.to_string() => PARTITION_KEY_NAME.to_string(),
vars::IS_RELEASED_PATH.to_string() => IS_RELEASED.to_string(),
}),
Some(hashmap! {
vars::IS_RELEASED_VALUE.to_string() => attr("1")
}),
)
.await
}
async fn upsert_expired_lock(
&self,
existing_rvn: &str,
data: Option<String>,
) -> Result<LockItem, DynamoError> {
self.client
.upsert_item(
data,
Some(expressions::PK_EXISTS_AND_RVN_MATCHES.to_string()),
Some(hashmap! {
vars::PK_PATH.to_string() => PARTITION_KEY_NAME.to_string(),
vars::RVN_PATH.to_string() => RECORD_VERSION_NUMBER.to_string(),
}),
Some(hashmap! {
vars::RVN_VALUE.to_string() => attr(existing_rvn)
}),
)
.await
}
}