use std::time::{Duration, Instant};
use async_trait::async_trait;
use scylla::statement::SerialConsistency;
use super::{CanonicalStore, DurabilityToken};
use crate::runtime::executors::cassandra::CassandraClient;
const OUTBOX_SEQ_ID: &str = "outbox";
const OUTBOX_CAS_MAX_ATTEMPTS: u32 = 64;
pub struct CassandraCanonicalStore {
pub(super) client: CassandraClient,
pub(super) instance_name: String,
pub(super) keyspace: String,
pub(super) outbox_table: String,
}
impl CassandraCanonicalStore {
pub fn new(
client: CassandraClient,
instance_name: impl Into<String>,
keyspace: impl Into<String>,
outbox_table: impl Into<String>,
) -> Self {
Self {
client,
instance_name: instance_name.into(),
keyspace: keyspace.into(),
outbox_table: outbox_table.into(),
}
}
fn keyspace_ident(&self) -> Result<&str, String> {
safe_ident(&self.keyspace).map(|_| self.keyspace.as_str())
}
fn outbox_ident(&self) -> Result<&str, String> {
safe_ident(&self.outbox_table).map(|_| self.outbox_table.as_str())
}
pub(super) fn qualified(&self, table: &str) -> String {
format!("\"{}\".\"{}\"", self.keyspace, table)
}
pub(super) fn client(&self) -> &CassandraClient {
&self.client
}
pub(super) fn keyspace(&self) -> &str {
&self.keyspace
}
pub(super) async fn ensure_keyspace(&self) -> Result<(), String> {
let ks = self.keyspace_ident()?;
let ks_ddl = format!(
"CREATE KEYSPACE IF NOT EXISTS {ks} WITH replication = \
{{'class':'SimpleStrategy','replication_factor':1}}"
);
self.client.cql_execute(&ks_ddl, ()).await
}
fn seq_table(&self) -> String {
self.qualified("udb_seq")
}
fn lease_table(&self) -> String {
self.qualified("udb_advisory_leases")
}
async fn read_outbox_seq(&self) -> Result<i64, String> {
let cql = format!(
"SELECT seq FROM {seq} WHERE id = '{id}'",
seq = self.seq_table(),
id = OUTBOX_SEQ_ID,
);
Ok(self
.client
.cql_query_first_i64(&cql, ())
.await?
.unwrap_or(0))
}
}
pub(super) fn safe_ident(name: &str) -> Result<(), String> {
let mut chars = name.chars();
match chars.next() {
Some(c) if c.is_ascii_alphabetic() || c == '_' => {}
_ => return Err(format!("unsafe cassandra identifier '{name}'")),
}
if name.chars().all(|c| c.is_ascii_alphanumeric() || c == '_') {
Ok(())
} else {
Err(format!("unsafe cassandra identifier '{name}'"))
}
}
pub(super) fn now_unix_ms() -> i64 {
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_millis() as i64)
.unwrap_or(0)
}
#[async_trait]
impl CanonicalStore for CassandraCanonicalStore {
fn backend_label(&self) -> &'static str {
"cassandra"
}
fn instance_name(&self) -> &str {
&self.instance_name
}
async fn ensure_system_tables(&self) -> Result<(), String> {
let ks = self.keyspace_ident()?;
let _ = self.outbox_ident()?;
let ks_ddl = format!(
"CREATE KEYSPACE IF NOT EXISTS {ks} WITH replication = \
{{'class':'SimpleStrategy','replication_factor':1}}"
);
self.client.cql_execute(&ks_ddl, ()).await?;
let outbox_ddl = format!(
"CREATE TABLE IF NOT EXISTS {tbl} ( \
shard text, \
event_seq bigint, \
event_id text, \
topic text, \
partition_key text, \
payload text, \
created_at timestamp, \
PRIMARY KEY (shard, event_seq) \
) WITH CLUSTERING ORDER BY (event_seq DESC)",
tbl = self.qualified(&self.outbox_table),
);
self.client.cql_execute(&outbox_ddl, ()).await?;
let seq_ddl = format!(
"CREATE TABLE IF NOT EXISTS {seq} ( id text PRIMARY KEY, seq bigint )",
seq = self.seq_table(),
);
self.client.cql_execute(&seq_ddl, ()).await?;
Ok(())
}
async fn enqueue_outbox_event(
&self,
event_id: &str,
topic: &str,
partition_key: &str,
payload: &serde_json::Value,
) -> Result<i64, String> {
let seq_table = self.seq_table();
let mut attempt = 0u32;
let new_seq = loop {
attempt += 1;
if attempt > OUTBOX_CAS_MAX_ATTEMPTS {
return Err(format!(
"outbox seq CAS did not converge after {OUTBOX_CAS_MAX_ATTEMPTS} attempts \
(contention on {seq_table})"
));
}
match self
.client
.cql_query_first_i64(
&format!("SELECT seq FROM {seq_table} WHERE id = '{OUTBOX_SEQ_ID}'"),
(),
)
.await?
{
None => {
let seed = format!(
"INSERT INTO {seq_table} (id, seq) VALUES ('{OUTBOX_SEQ_ID}', ?) IF NOT EXISTS"
);
let applied = self
.client
.cql_lwt_applied(&seed, (1i64,), SerialConsistency::Serial)
.await?;
if applied {
break 1i64;
}
continue;
}
Some(current) => {
let next = current + 1;
let cas = format!(
"UPDATE {seq_table} SET seq = ? WHERE id = '{OUTBOX_SEQ_ID}' IF seq = ?"
);
let applied = self
.client
.cql_lwt_applied(&cas, (next, current), SerialConsistency::Serial)
.await?;
if applied {
break next;
}
continue;
}
}
};
let insert = format!(
"INSERT INTO {tbl} (shard, event_seq, event_id, topic, partition_key, payload, created_at) \
VALUES ('{OUTBOX_SEQ_ID}', ?, ?, ?, ?, ?, ?)",
tbl = self.qualified(&self.outbox_table),
);
let payload_text = serde_json::to_string(payload)
.map_err(|e| format!("outbox payload serialise failed: {e}"))?;
self.client
.cql_execute(
&insert,
(
new_seq,
event_id,
topic,
partition_key,
payload_text,
scylla::frame::value::CqlTimestamp(now_unix_ms()),
),
)
.await?;
Ok(new_seq)
}
async fn outbox_max_seq(&self) -> Result<i64, String> {
self.read_outbox_seq().await
}
async fn current_durability_token(&self) -> Result<DurabilityToken, String> {
let seq = self.read_outbox_seq().await?;
Ok(DurabilityToken::new("cassandra", seq.to_string()))
}
async fn wait_for_token(
&self,
token: &DurabilityToken,
timeout: Duration,
) -> Result<bool, String> {
if !token.is_for("cassandra") {
return Err(format!(
"CassandraCanonicalStore cannot wait on a '{}' token",
token.backend_label
));
}
let target: i64 = token.value.parse().map_err(|e| {
format!(
"malformed cassandra durability token '{}': {e}",
token.value
)
})?;
let started = Instant::now();
let poll = super::durability_poll_interval(timeout, super::CASSANDRA_DURABILITY_POLL_MS);
loop {
if self.read_outbox_seq().await? >= target {
return Ok(true);
}
if started.elapsed() >= timeout {
return Ok(false);
}
tokio::time::sleep(poll).await;
}
}
async fn ensure_advisory_lease_table(&self) -> Result<(), String> {
let ks = self.keyspace_ident()?;
let ks_ddl = format!(
"CREATE KEYSPACE IF NOT EXISTS {ks} WITH replication = \
{{'class':'SimpleStrategy','replication_factor':1}}"
);
self.client.cql_execute(&ks_ddl, ()).await?;
let ddl = format!(
"CREATE TABLE IF NOT EXISTS {tbl} ( \
lease_name text PRIMARY KEY, \
owner_id text, \
expires_at timestamp \
)",
tbl = self.lease_table(),
);
self.client.cql_execute(&ddl, ()).await?;
Ok(())
}
async fn try_acquire_advisory_lease(
&self,
lease_name: &str,
owner_id: &str,
ttl: Duration,
) -> Result<bool, String> {
let lease_tbl = self.lease_table();
let now = now_unix_ms();
let new_expires = now + (ttl.as_millis() as i64);
let insert = format!(
"INSERT INTO {lease_tbl} (lease_name, owner_id, expires_at) VALUES (?, ?, ?) IF NOT EXISTS"
);
let inserted = self
.client
.cql_lwt_applied(
&insert,
(
lease_name,
owner_id,
scylla::frame::value::CqlTimestamp(new_expires),
),
SerialConsistency::Serial,
)
.await?;
if inserted {
return Ok(true);
}
let refresh = format!(
"UPDATE {lease_tbl} SET owner_id = ?, expires_at = ? WHERE lease_name = ? IF owner_id = ?"
);
let refreshed = self
.client
.cql_lwt_applied(
&refresh,
(
owner_id,
scylla::frame::value::CqlTimestamp(new_expires),
lease_name,
owner_id,
),
SerialConsistency::Serial,
)
.await?;
if refreshed {
return Ok(true);
}
let takeover = format!(
"UPDATE {lease_tbl} SET owner_id = ?, expires_at = ? WHERE lease_name = ? IF expires_at <= ?"
);
let took = self
.client
.cql_lwt_applied(
&takeover,
(
owner_id,
scylla::frame::value::CqlTimestamp(new_expires),
lease_name,
scylla::frame::value::CqlTimestamp(now),
),
SerialConsistency::Serial,
)
.await?;
Ok(took)
}
async fn release_advisory_lease(&self, lease_name: &str, owner_id: &str) -> Result<(), String> {
let lease_tbl = self.lease_table();
let del = format!("DELETE FROM {lease_tbl} WHERE lease_name = ? IF owner_id = ?");
let _ = self
.client
.cql_lwt_applied(&del, (lease_name, owner_id), SerialConsistency::Serial)
.await?;
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn backend_label_is_pinned() {
assert!(safe_ident("udb_conf_abc123").is_ok());
assert!(safe_ident("_internal").is_ok());
}
#[test]
fn unsafe_idents_are_rejected() {
assert!(safe_ident("").is_err());
assert!(safe_ident("1leading_digit").is_err());
assert!(safe_ident("evil; DROP").is_err());
assert!(safe_ident("has-dash").is_err());
assert!(safe_ident("has.dot").is_err());
assert!(safe_ident("quote\"inject").is_err());
}
}