use std::sync::Arc;
use std::time::Duration;
use sqlx::{PgPool, Row};
use super::source::{MessageSource, ReceivedMessage};
use super::{
run_source, Bus, BusConsumer, BusTopologyConfig, MessageRouter, RunOptions, TransportError,
};
use super::{Message, MessageKind};
const DEFAULT_LEASE: Duration = Duration::from_secs(30);
const SCHEMA: &str = "\
CREATE TABLE IF NOT EXISTS bus_queue (
seq BIGSERIAL PRIMARY KEY,
name TEXT NOT NULL,
message_id TEXT,
kind TEXT NOT NULL,
payload BYTEA NOT NULL,
content_type TEXT NOT NULL DEFAULT 'application/json',
metadata TEXT NOT NULL DEFAULT '[]',
available_at TIMESTAMPTZ NOT NULL DEFAULT now(),
locked_until TIMESTAMPTZ,
attempts INTEGER NOT NULL DEFAULT 0
);
CREATE INDEX IF NOT EXISTS bus_queue_claim_idx ON bus_queue (name, available_at, seq);
CREATE TABLE IF NOT EXISTS bus_log (
seq BIGSERIAL PRIMARY KEY,
name TEXT NOT NULL,
message_id TEXT,
kind TEXT NOT NULL,
payload BYTEA NOT NULL,
content_type TEXT NOT NULL DEFAULT 'application/json',
metadata TEXT NOT NULL DEFAULT '[]',
appended_at TIMESTAMPTZ NOT NULL DEFAULT now()
);
CREATE INDEX IF NOT EXISTS bus_log_name_seq_idx ON bus_log (name, seq);
CREATE TABLE IF NOT EXISTS bus_offset (
consumer TEXT PRIMARY KEY,
last_seq BIGINT NOT NULL DEFAULT 0
)";
fn db_err(context: &str, err: sqlx::Error) -> TransportError {
TransportError::retryable(format!("postgres bus {context}: {err}"))
}
fn message_from_row(row: &sqlx::postgres::PgRow) -> Result<Message, TransportError> {
fn decode_err(column: &str, err: sqlx::Error) -> TransportError {
TransportError::permanent(format!(
"postgres bus corrupt row: required column '{column}' failed to decode: {err}"
))
}
let name: String = row.try_get("name").map_err(|err| decode_err("name", err))?;
let kind: String = row.try_get("kind").map_err(|err| decode_err("kind", err))?;
let payload: Vec<u8> = row
.try_get("payload")
.map_err(|err| decode_err("payload", err))?;
let metadata_json: String = row.try_get("metadata").unwrap_or_else(|_| "[]".to_string());
let metadata =
serde_json::from_str::<Vec<(String, String)>>(&metadata_json).unwrap_or_default();
Ok(Message {
id: row
.try_get::<Option<String>, _>("message_id")
.unwrap_or(None),
name,
kind: MessageKind::from_str_lossy(&kind),
payload,
content_type: row
.try_get("content_type")
.unwrap_or_else(|_| "application/json".to_string()),
metadata,
})
}
#[derive(Clone)]
pub struct PostgresBus {
pool: PgPool,
topology: BusTopologyConfig,
lease: Duration,
}
impl PostgresBus {
pub fn new(pool: PgPool) -> Self {
Self {
pool,
topology: BusTopologyConfig::default(),
lease: DEFAULT_LEASE,
}
}
pub fn new_with_group(pool: PgPool, group: impl Into<String>) -> Self {
Self::new(pool).group(group)
}
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 async fn ensure_tables(&self) -> Result<(), TransportError> {
for statement in SCHEMA.split(';') {
let statement = statement.trim();
if statement.is_empty() {
continue;
}
sqlx::query(statement)
.execute(&self.pool)
.await
.map_err(|err| db_err("ensure_tables", err))?;
}
Ok(())
}
async fn enqueue(&self, message: Message) -> Result<(), TransportError> {
let metadata = serde_json::to_string(&message.metadata).unwrap_or_else(|_| "[]".into());
sqlx::query(
"INSERT INTO bus_queue (name, message_id, kind, payload, content_type, metadata) \
VALUES ($1, $2, $3, $4, $5, $6)",
)
.bind(&message.name)
.bind(&message.id)
.bind(message.kind.as_str())
.bind(&message.payload)
.bind(&message.content_type)
.bind(metadata)
.execute(&self.pool)
.await
.map_err(|err| db_err("enqueue", err))?;
Ok(())
}
async fn append(&self, message: Message) -> Result<(), TransportError> {
let metadata = serde_json::to_string(&message.metadata).unwrap_or_else(|_| "[]".into());
sqlx::query(
"INSERT INTO bus_log (name, message_id, kind, payload, content_type, metadata) \
VALUES ($1, $2, $3, $4, $5, $6)",
)
.bind(&message.name)
.bind(&message.id)
.bind(message.kind.as_str())
.bind(&message.payload)
.bind(&message.content_type)
.bind(metadata)
.execute(&self.pool)
.await
.map_err(|err| db_err("append", err))?;
Ok(())
}
}
impl Bus for PostgresBus {
async fn send(&self, name: &str, payload: Vec<u8>) -> Result<(), TransportError> {
self.enqueue(Message::new(name, MessageKind::Command, payload))
.await
}
async fn publish(&self, name: &str, payload: Vec<u8>) -> Result<(), TransportError> {
self.append(Message::new(name, MessageKind::Event, payload))
.await
}
async fn send_message(&self, message: Message) -> Result<(), TransportError> {
self.enqueue(message).await
}
async fn publish_message(&self, message: Message) -> Result<(), TransportError> {
self.append(message).await
}
}
impl BusConsumer for PostgresBus {
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 = QueueSource {
pool: self.pool.clone(),
names,
lease_secs: self.lease.as_secs_f64(),
};
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(), "postgres")?;
let source = LogSource {
pool: self.pool.clone(),
names,
consumer: group,
};
run_source(router, source, options).await
}
}
struct QueueSource {
pool: PgPool,
names: Vec<String>,
lease_secs: f64,
}
impl MessageSource for QueueSource {
type Received = QueueReceived;
async fn recv(&mut self) -> Result<Option<Self::Received>, TransportError> {
let row = sqlx::query(
"UPDATE bus_queue SET locked_until = now() + ($1 * interval '1 second'), \
attempts = attempts + 1 \
WHERE seq = ( \
SELECT seq FROM bus_queue \
WHERE (name = ANY($2) OR name IS NULL) AND available_at <= now() \
AND (locked_until IS NULL OR locked_until < now()) \
ORDER BY seq FOR UPDATE SKIP LOCKED LIMIT 1 \
) \
RETURNING seq, name, message_id, kind, payload, content_type, metadata",
)
.bind(self.lease_secs)
.bind(&self.names)
.fetch_optional(&self.pool)
.await
.map_err(|err| db_err("claim", err))?;
Ok(row.map(|row| {
let seq: i64 = row.try_get("seq").unwrap_or_default();
let (message, decode_error) = decode_or_placeholder(&row);
QueueReceived {
pool: self.pool.clone(),
seq,
message,
decode_error,
}
}))
}
}
fn decode_or_placeholder(row: &sqlx::postgres::PgRow) -> (Message, Option<TransportError>) {
match message_from_row(row) {
Ok(message) => (message, None),
Err(error) => (
Message::new("", MessageKind::Event, Vec::new()),
Some(error),
),
}
}
pub struct QueueReceived {
pool: PgPool,
seq: i64,
message: Message,
decode_error: Option<TransportError>,
}
impl QueueReceived {
async fn delete(&self) -> Result<(), TransportError> {
sqlx::query("DELETE FROM bus_queue WHERE seq = $1")
.bind(self.seq)
.execute(&self.pool)
.await
.map_err(|err| db_err("delete", err))?;
Ok(())
}
}
impl ReceivedMessage for QueueReceived {
fn message(&self) -> &Message {
&self.message
}
fn decode_error(&self) -> Option<&TransportError> {
self.decode_error.as_ref()
}
async fn ack(self) -> Result<(), TransportError> {
self.delete().await
}
async fn nack(self, _reason: &str) -> Result<(), TransportError> {
sqlx::query("UPDATE bus_queue SET locked_until = NULL WHERE seq = $1")
.bind(self.seq)
.execute(&self.pool)
.await
.map_err(|err| db_err("nack", err))?;
Ok(())
}
async fn dead_letter(self, _reason: &str) -> Result<(), TransportError> {
self.delete().await
}
async fn park(self, _reason: &str) -> Result<(), TransportError> {
self.delete().await
}
}
struct LogSource {
pool: PgPool,
names: Vec<String>,
consumer: String,
}
impl MessageSource for LogSource {
type Received = LogReceived;
async fn recv(&mut self) -> Result<Option<Self::Received>, TransportError> {
let row = sqlx::query(
"SELECT seq, name, message_id, kind, payload, content_type, metadata FROM bus_log \
WHERE (name = ANY($1) OR name IS NULL) \
AND seq > COALESCE((SELECT last_seq FROM bus_offset WHERE consumer = $2), 0) \
ORDER BY seq LIMIT 1",
)
.bind(&self.names)
.bind(&self.consumer)
.fetch_optional(&self.pool)
.await
.map_err(|err| db_err("log read", err))?;
Ok(row.map(|row| {
let seq: i64 = row.try_get("seq").unwrap_or_default();
let (message, decode_error) = decode_or_placeholder(&row);
LogReceived {
pool: self.pool.clone(),
consumer: self.consumer.clone(),
seq,
message,
decode_error,
}
}))
}
}
pub struct LogReceived {
pool: PgPool,
consumer: String,
seq: i64,
message: Message,
decode_error: Option<TransportError>,
}
impl LogReceived {
async fn advance_offset(&self) -> Result<(), TransportError> {
sqlx::query(
"INSERT INTO bus_offset (consumer, last_seq) VALUES ($1, $2) \
ON CONFLICT (consumer) DO UPDATE SET last_seq = EXCLUDED.last_seq \
WHERE bus_offset.last_seq < EXCLUDED.last_seq",
)
.bind(&self.consumer)
.bind(self.seq)
.execute(&self.pool)
.await
.map_err(|err| db_err("advance offset", err))?;
Ok(())
}
}
impl ReceivedMessage for LogReceived {
fn message(&self) -> &Message {
&self.message
}
fn decode_error(&self) -> Option<&TransportError> {
self.decode_error.as_ref()
}
async fn ack(self) -> Result<(), TransportError> {
self.advance_offset().await
}
async fn nack(self, _reason: &str) -> Result<(), TransportError> {
Ok(())
}
async fn dead_letter(self, _reason: &str) -> Result<(), TransportError> {
self.advance_offset().await
}
async fn park(self, _reason: &str) -> Result<(), TransportError> {
self.advance_offset().await
}
}