use std::collections::VecDeque;
use std::future::Future;
use std::sync::atomic::{AtomicI64, Ordering};
use std::sync::Arc;
use std::time::Duration;
use sqlx::{ColumnIndex, Decode, Row, Type};
use crate::projection_protocol::{ProjectionEpoch, ProjectionSource};
use crate::sqlx_repo::is_sqlx_transient;
use super::source::{MessageSource, ReceivedMessage};
use super::{
run_source, Bus, BusConsumer, BusTopologyConfig, MessageRouter, OrderedDelivery, RunOptions,
TransportError, TransportErrorKind,
};
use super::{Message, MessageKind};
pub(crate) const DEFAULT_LEASE: Duration = Duration::from_secs(30);
const SOURCE_BATCH: i64 = 16;
pub(crate) fn db_err(backend: &str, context: &str, err: sqlx::Error) -> TransportError {
let kind = if is_sqlx_transient(&err) {
TransportErrorKind::Retryable
} else {
TransportErrorKind::Permanent
};
TransportError::new(kind, format!("{backend} bus {context}: {err}")).with_source(err)
}
pub(crate) fn metadata_json(message: &Message) -> String {
serde_json::to_string(&message.metadata).unwrap_or_else(|_| "[]".into())
}
fn corrupt_row(backend: &str, message: impl Into<String>) -> TransportError {
TransportError::permanent(format!("{backend} bus corrupt row: {}", message.into()))
}
fn decode_err(backend: &str, column: &str, err: sqlx::Error) -> TransportError {
corrupt_row(
backend,
format!("required column '{column}' failed to decode: {err}"),
)
}
fn parse_message_kind(backend: &str, value: &str) -> Result<MessageKind, TransportError> {
match value {
"command" => Ok(MessageKind::Command),
"event" => Ok(MessageKind::Event),
_ => Err(corrupt_row(
backend,
format!("required column 'kind' has unsupported value {value:?}"),
)),
}
}
fn parse_metadata(backend: &str, value: &str) -> Result<Vec<(String, String)>, TransportError> {
serde_json::from_str(value).map_err(|err| {
corrupt_row(
backend,
format!("required column 'metadata' failed to parse as JSON metadata: {err}"),
)
})
}
pub(crate) fn message_from_row<R>(backend: &str, row: &R) -> Result<Message, TransportError>
where
R: Row,
for<'a> &'a str: ColumnIndex<R>,
for<'r> String: Decode<'r, R::Database> + Type<R::Database>,
for<'r> Vec<u8>: Decode<'r, R::Database> + Type<R::Database>,
{
let name: String = row
.try_get("name")
.map_err(|err| decode_err(backend, "name", err))?;
let kind: String = row
.try_get("kind")
.map_err(|err| decode_err(backend, "kind", err))?;
let payload: Vec<u8> = row
.try_get("payload")
.map_err(|err| decode_err(backend, "payload", err))?;
let content_type: String = row
.try_get("content_type")
.map_err(|err| decode_err(backend, "content_type", err))?;
if content_type.is_empty() {
return Err(corrupt_row(
backend,
"required column 'content_type' is empty",
));
}
let metadata_json: String = row
.try_get("metadata")
.map_err(|err| decode_err(backend, "metadata", err))?;
let metadata = parse_metadata(backend, &metadata_json)?;
Ok(Message {
id: row
.try_get::<Option<String>, _>("message_id")
.unwrap_or(None),
name,
kind: parse_message_kind(backend, &kind)?,
payload,
content_type,
metadata,
})
}
pub(crate) fn validate_log_retry(
backend: &str,
existing: &Message,
retry: &Message,
) -> Result<(), TransportError> {
let matches = existing.name == retry.name
&& existing.kind == retry.kind
&& existing.payload == retry.payload
&& existing.content_type == retry.content_type
&& existing.causation_id() == retry.causation_id();
if matches {
return Ok(());
}
Err(TransportError::permanent(format!(
"{backend} bus ordered-log message ID {:?} was reused with a different \
name, kind, payload, content type, or causation ID",
retry.id()
)))
}
fn fresh_log_epoch() -> ProjectionEpoch {
ProjectionEpoch::new(format!("sql-log-{}", uuid::Uuid::now_v7()))
.expect("a UUID-backed SQL log epoch is valid")
}
pub struct ReceivedRow {
seq: i64,
message: Message,
decode_error: Option<TransportError>,
}
impl ReceivedRow {
pub(crate) fn from_row<R>(backend: &str, row: &R) -> Self
where
R: Row,
for<'a> &'a str: ColumnIndex<R>,
for<'r> String: Decode<'r, R::Database> + Type<R::Database>,
for<'r> Vec<u8>: Decode<'r, R::Database> + Type<R::Database>,
for<'r> i64: Decode<'r, R::Database> + Type<R::Database>,
{
let seq = row.try_get("seq").unwrap_or_default();
let (message, decode_error) = match message_from_row(backend, row) {
Ok(message) => (message, None),
Err(error) => (
Message::new("", MessageKind::Event, Vec::new()),
Some(error),
),
};
Self {
seq,
message,
decode_error,
}
}
}
pub struct ClaimedRow {
pub(crate) row: ReceivedRow,
pub(crate) claim_token: String,
}
pub trait SqlBusDialect: Clone + Send + Sync + 'static {
const BACKEND: &'static str;
const SCHEMA: &'static str;
fn execute_ddl(
&self,
statement: &'static str,
) -> impl Future<Output = Result<(), TransportError>> + Send;
fn ensure_ordered_log_schema(&self) -> impl Future<Output = Result<(), TransportError>> + Send;
fn insert_queue(
&self,
message: &Message,
) -> impl Future<Output = Result<(), TransportError>> + Send;
fn insert_log(
&self,
message: &Message,
epoch_candidate: &ProjectionEpoch,
expected_epoch: Option<&ProjectionEpoch>,
) -> impl Future<Output = Result<(), TransportError>> + Send;
fn prepare_log_epoch(
&self,
epoch_candidate: &ProjectionEpoch,
expected_epoch: Option<&ProjectionEpoch>,
) -> impl Future<Output = Result<ProjectionEpoch, TransportError>> + Send;
fn reset_log(
&self,
expected_epoch: &ProjectionEpoch,
next_epoch: &ProjectionEpoch,
) -> impl Future<Output = Result<(), TransportError>> + Send;
fn claim(
&self,
names: &[String],
lease_secs: f64,
limit: i64,
) -> impl Future<Output = Result<Vec<ClaimedRow>, TransportError>> + Send;
fn log_read(
&self,
names: &[String],
consumer: &str,
limit: i64,
expected_epoch: &ProjectionEpoch,
) -> impl Future<Output = Result<Vec<ReceivedRow>, TransportError>> + Send;
fn verify_log_epoch(
&self,
expected_epoch: &ProjectionEpoch,
) -> impl Future<Output = Result<(), TransportError>> + Send;
fn delete_claimed(
&self,
seq: i64,
claim_token: &str,
) -> impl Future<Output = Result<(), TransportError>> + Send;
fn release_claim(
&self,
seq: i64,
claim_token: &str,
) -> impl Future<Output = Result<(), TransportError>> + Send;
fn advance_offset(
&self,
consumer: &str,
seq: i64,
expected_epoch: &ProjectionEpoch,
) -> impl Future<Output = Result<(), TransportError>> + Send;
}
#[derive(Clone)]
pub struct SqlBus<B> {
dialect: B,
topology: BusTopologyConfig,
lease: Duration,
source_epoch: Option<ProjectionEpoch>,
}
impl<B: SqlBusDialect> SqlBus<B> {
pub(crate) fn from_dialect(dialect: B) -> Self {
Self {
dialect,
topology: BusTopologyConfig::default(),
lease: DEFAULT_LEASE,
source_epoch: None,
}
}
pub fn group(mut self, group: impl Into<String>) -> Self {
self.topology = self.topology.group(group);
self
}
pub fn with_lease(mut self, lease: Duration) -> Self {
self.lease = lease;
self
}
pub fn with_source_epoch(mut self, epoch: ProjectionEpoch) -> Self {
self.source_epoch = Some(epoch);
self
}
pub async fn ensure_tables(&self) -> Result<(), TransportError> {
for statement in B::SCHEMA.split(';') {
let statement = statement.trim();
if statement.is_empty() {
continue;
}
self.dialect.execute_ddl(statement).await?;
}
self.dialect.ensure_ordered_log_schema().await?;
let epoch_candidate = self.source_epoch.clone().unwrap_or_else(fresh_log_epoch);
self.dialect
.prepare_log_epoch(&epoch_candidate, self.source_epoch.as_ref())
.await?;
Ok(())
}
pub async fn reset_ordered_log(
&self,
expected_epoch: &ProjectionEpoch,
next_epoch: &ProjectionEpoch,
) -> Result<(), TransportError> {
if expected_epoch == next_epoch {
return Err(TransportError::permanent(
"ordered-log reset requires a distinct next epoch",
));
}
self.dialect.reset_log(expected_epoch, next_epoch).await
}
}
impl<B: SqlBusDialect> Bus for SqlBus<B> {
async fn send_message(&self, message: Message) -> Result<(), TransportError> {
self.dialect.insert_queue(&message).await
}
async fn publish_message(&self, message: Message) -> Result<(), TransportError> {
let epoch_candidate = self.source_epoch.clone().unwrap_or_else(fresh_log_epoch);
self.dialect
.insert_log(&message, &epoch_candidate, self.source_epoch.as_ref())
.await
}
}
impl<B: SqlBusDialect> BusConsumer for SqlBus<B> {
async fn listen<R: MessageRouter>(
&self,
router: Arc<R>,
options: RunOptions,
) -> Result<(), TransportError> {
self.ensure_tables().await?;
let names = router.subscription_plan().commands;
if names.is_empty() {
return Ok(());
}
let source = SqlQueueSource {
dialect: self.dialect.clone(),
names,
lease_secs: self.lease.as_secs_f64(),
buffer: VecDeque::new(),
};
run_source(router, source, options).await
}
async fn subscribe<R: MessageRouter>(
&self,
router: Arc<R>,
options: RunOptions,
) -> Result<(), TransportError> {
self.ensure_tables().await?;
let names = router.subscription_plan().events;
if names.is_empty() {
return Ok(());
}
let group = self
.topology
.resolve_consumer_group(router.as_ref(), B::BACKEND)?;
let epoch_candidate = self.source_epoch.clone().unwrap_or_else(fresh_log_epoch);
let source_epoch = self
.dialect
.prepare_log_epoch(&epoch_candidate, self.source_epoch.as_ref())
.await?;
let source = SqlLogSource {
dialect: self.dialect.clone(),
names,
consumer: group,
buffer: VecDeque::new(),
last_delivered: None,
settled_seq: Arc::new(AtomicI64::new(0)),
source_epoch,
};
run_source(router, source, options).await
}
}
struct SqlQueueSource<B> {
dialect: B,
names: Vec<String>,
lease_secs: f64,
buffer: VecDeque<ClaimedRow>,
}
impl<B: SqlBusDialect> MessageSource for SqlQueueSource<B> {
type Received = SqlQueueReceived<B>;
fn transport_name(&self) -> &'static str {
B::BACKEND
}
async fn recv(&mut self) -> Result<Option<Self::Received>, TransportError> {
if self.buffer.is_empty() {
let mut claimed = self
.dialect
.claim(&self.names, self.lease_secs, SOURCE_BATCH)
.await?;
claimed.sort_by_key(|claim| claim.row.seq);
self.buffer.extend(claimed);
}
Ok(self.buffer.pop_front().map(|claimed| SqlQueueReceived {
dialect: self.dialect.clone(),
row: claimed.row,
claim_token: claimed.claim_token,
}))
}
}
pub struct SqlQueueReceived<B> {
dialect: B,
row: ReceivedRow,
claim_token: String,
}
impl<B: SqlBusDialect> ReceivedMessage for SqlQueueReceived<B> {
fn message(&self) -> &Message {
&self.row.message
}
fn decode_error(&self) -> Option<&TransportError> {
self.row.decode_error.as_ref()
}
async fn ack(self) -> Result<(), TransportError> {
self.dialect
.delete_claimed(self.row.seq, &self.claim_token)
.await
}
async fn nack(self, _reason: &str) -> Result<(), TransportError> {
self.dialect
.release_claim(self.row.seq, &self.claim_token)
.await
}
async fn dead_letter(self, _reason: &str) -> Result<(), TransportError> {
self.dialect
.delete_claimed(self.row.seq, &self.claim_token)
.await
}
async fn park(self, _reason: &str) -> Result<(), TransportError> {
self.dialect
.delete_claimed(self.row.seq, &self.claim_token)
.await
}
}
struct SqlLogSource<B> {
dialect: B,
names: Vec<String>,
consumer: String,
buffer: VecDeque<ReceivedRow>,
last_delivered: Option<i64>,
settled_seq: Arc<AtomicI64>,
source_epoch: ProjectionEpoch,
}
impl<B: SqlBusDialect> MessageSource for SqlLogSource<B> {
type Received = SqlLogReceived<B>;
fn transport_name(&self) -> &'static str {
B::BACKEND
}
async fn recv(&mut self) -> Result<Option<Self::Received>, TransportError> {
self.dialect.verify_log_epoch(&self.source_epoch).await?;
if let Some(last) = self.last_delivered {
if self.settled_seq.load(Ordering::Acquire) < last {
self.buffer.clear();
}
}
if self.buffer.is_empty() {
let rows = self
.dialect
.log_read(
&self.names,
&self.consumer,
SOURCE_BATCH,
&self.source_epoch,
)
.await?;
self.buffer.extend(rows);
}
let Some(row) = self.buffer.pop_front() else {
return Ok(None);
};
let position = u64::try_from(row.seq).map_err(|_| {
corrupt_row(
B::BACKEND,
format!(
"bus_log seq {} is outside the projection cursor domain",
row.seq
),
)
})?;
let source = ProjectionSource::new(format!("{}.bus_log", B::BACKEND), b"global".to_vec())
.map_err(|error| corrupt_row(B::BACKEND, error.to_string()))?;
let ordered = OrderedDelivery::new(source, self.source_epoch.clone(), position, false)
.map_err(|error| corrupt_row(B::BACKEND, error.to_string()))?;
self.last_delivered = Some(row.seq);
Ok(Some(SqlLogReceived {
dialect: self.dialect.clone(),
consumer: self.consumer.clone(),
settled_seq: self.settled_seq.clone(),
row,
ordered,
}))
}
}
pub struct SqlLogReceived<B> {
dialect: B,
consumer: String,
settled_seq: Arc<AtomicI64>,
row: ReceivedRow,
ordered: OrderedDelivery,
}
impl<B: SqlBusDialect> SqlLogReceived<B> {
async fn settle_forward(self) -> Result<(), TransportError> {
self.dialect
.advance_offset(&self.consumer, self.row.seq, self.ordered.epoch())
.await?;
self.settled_seq.store(self.row.seq, Ordering::Release);
Ok(())
}
}
impl<B: SqlBusDialect> ReceivedMessage for SqlLogReceived<B> {
fn message(&self) -> &Message {
&self.row.message
}
fn ordered_delivery(&self) -> Option<&OrderedDelivery> {
Some(&self.ordered)
}
fn decode_error(&self) -> Option<&TransportError> {
self.row.decode_error.as_ref()
}
async fn ack(self) -> Result<(), TransportError> {
self.settle_forward().await
}
async fn nack(self, _reason: &str) -> Result<(), TransportError> {
Ok(())
}
async fn dead_letter(self, _reason: &str) -> Result<(), TransportError> {
self.settle_forward().await
}
async fn park(self, _reason: &str) -> Result<(), TransportError> {
self.settle_forward().await
}
}