use std::sync::Arc;
use std::time::Duration;
use sqlx::{PgPool, Row};
use super::source::{AsyncMessageSource, ReceivedMessage};
use super::{run_source, Bus, BusConsumer, 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 kind_str(kind: MessageKind) -> &'static str {
match kind {
MessageKind::Command => "command",
MessageKind::Event => "event",
}
}
fn kind_from_str(value: &str) -> MessageKind {
match value {
"command" => MessageKind::Command,
_ => MessageKind::Event,
}
}
fn message_from_row(row: &sqlx::postgres::PgRow) -> Message {
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();
Message {
id: row
.try_get::<Option<String>, _>("message_id")
.unwrap_or(None),
name: row.try_get("name").unwrap_or_default(),
kind: kind_from_str(&row.try_get::<String, _>("kind").unwrap_or_default()),
payload: row.try_get("payload").unwrap_or_default(),
content_type: row
.try_get("content_type")
.unwrap_or_else(|_| "application/json".to_string()),
metadata,
}
}
#[derive(Clone)]
pub struct PostgresBus {
pool: PgPool,
group: String,
lease: Duration,
}
impl PostgresBus {
pub fn new(pool: PgPool, group: impl Into<String>) -> Self {
Self {
pool,
group: group.into(),
lease: DEFAULT_LEASE,
}
}
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(kind_str(message.kind))
.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(kind_str(message.kind))
.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 source = LogSource {
pool: self.pool.clone(),
names,
consumer: self.group.clone(),
};
run_source(router, source, options).await
}
}
struct QueueSource {
pool: PgPool,
names: Vec<String>,
lease_secs: f64,
}
impl AsyncMessageSource 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) 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();
QueueReceived {
pool: self.pool.clone(),
seq,
message: message_from_row(&row),
}
}))
}
}
pub struct QueueReceived {
pool: PgPool,
seq: i64,
message: Message,
}
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
}
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 AsyncMessageSource 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) \
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();
LogReceived {
pool: self.pool.clone(),
consumer: self.consumer.clone(),
seq,
message: message_from_row(&row),
}
}))
}
}
pub struct LogReceived {
pool: PgPool,
consumer: String,
seq: i64,
message: Message,
}
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
}
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
}
}