use super::poll::PollBackoff;
use crate::canonical_message::tracing_support::LazyMessageIds;
use crate::models::SqlxConfig;
use crate::traits::{
BoxFuture, ConsumerError, EndpointStatus, MessageConsumer, MessageDisposition,
MessagePublisher, PublisherError, ReceivedBatch, Sent, SentBatch,
};
use crate::CanonicalMessage;
use anyhow::{anyhow, Context};
use async_trait::async_trait;
use sqlx::any::AnyPoolOptions;
use sqlx::postgres::{PgPool, PgPoolCopyExt, PgPoolOptions};
use sqlx::{AnyPool, AssertSqlSafe, Column, Row};
use std::sync::{Arc, Mutex};
use std::time::Duration;
use tracing::{info, trace, warn};
#[cfg(feature = "dedup")]
mod dedup;
#[cfg(feature = "dedup")]
pub(crate) use dedup::build_sql_dedup_store;
fn is_deadlock_error(e: &sqlx::Error) -> bool {
if let Some(db_err) = e.as_database_error() {
match db_err.code() {
Some(code) => {
let c = code.as_ref();
c == "1213" || c == "40001" || c == "40P01" || c == "1205"
}
None => false,
}
} else {
false
}
}
fn is_valid_table_name(name: &str) -> bool {
if name.is_empty() || name.starts_with('.') || name.ends_with('.') || name.contains("..") {
return false;
}
name.split('.')
.all(|part| !part.is_empty() && part.chars().all(|c| c.is_ascii_alphanumeric() || c == '_'))
}
fn find_top_level_values(sql: &str) -> Option<usize> {
fn is_word_byte(b: Option<u8>) -> bool {
matches!(b, Some(c) if c.is_ascii_alphanumeric() || c == b'_')
}
let bytes = sql.as_bytes();
let mut depth: i32 = 0;
for i in 0..bytes.len() {
match bytes[i] {
b'(' => depth += 1,
b')' => depth = depth.saturating_sub(1),
_ => {
if depth == 0
&& bytes.len() - i >= 6
&& bytes[i..i + 6].eq_ignore_ascii_case(b"VALUES")
&& !is_word_byte(i.checked_sub(1).map(|p| bytes[p]))
&& !is_word_byte(bytes.get(i + 6).copied())
{
return Some(i);
}
}
}
}
None
}
fn contains_payload_clause(query: &str) -> bool {
let lower_query = query.to_lowercase();
let mut search_start = 0;
while let Some(open_paren_idx) = lower_query[search_start..].find('(') {
let absolute_open_idx = search_start + open_paren_idx;
if let Some(close_paren_idx) = lower_query[absolute_open_idx..].find(')') {
let absolute_close_idx = absolute_open_idx + close_paren_idx;
let content = &lower_query[absolute_open_idx + 1..absolute_close_idx];
if content.trim() == "payload" {
return true;
}
search_start = absolute_close_idx + 1;
} else {
break;
}
}
false
}
fn audited_sql(sql: &str) -> AssertSqlSafe<&str> {
AssertSqlSafe(sql)
}
#[derive(Debug, Clone, PartialEq)]
enum ColumnSource {
Metadata(String),
Payload(String),
}
#[derive(Debug, Clone, PartialEq)]
enum BindValue {
Null,
Int(i64),
Float(f64),
Bool(bool),
Text(String),
}
fn positional_placeholder(driver_name: &str, index: usize) -> String {
match driver_name {
"PostgreSQL" => format!("${}", index),
"Microsoft SQL Server" => format!("@p{}", index),
_ => "?".to_string(),
}
}
fn parse_insert_template(
query: &str,
driver_name: &str,
) -> anyhow::Result<(String, Vec<ColumnSource>)> {
let mut out = String::with_capacity(query.len());
let mut sources: Vec<ColumnSource> = Vec::new();
let bytes = query.as_bytes();
let mut i = 0;
while i < bytes.len() {
if bytes[i] == b'$' && i + 1 < bytes.len() && bytes[i + 1] == b'{' {
let close = query[i + 2..].find('}').map(|off| i + 2 + off);
let close = close.ok_or_else(|| {
anyhow!(
"Malformed token in insert_query: unclosed '${{' near '{}'",
&query[i..]
)
})?;
let inner = &query[i + 2..close];
let (prefix, name) = inner.split_once(':').ok_or_else(|| {
anyhow!("Malformed token in insert_query: '${{{}}}' is missing a ':' separator (expected ${{metadata:key}} or ${{payload:field}})", inner)
})?;
let name = name.trim();
if name.is_empty() {
return Err(anyhow!(
"Malformed token in insert_query: '${{{}}}' has an empty key/field name",
inner
));
}
let source = match prefix.trim() {
"metadata" => ColumnSource::Metadata(name.to_string()),
"payload" => ColumnSource::Payload(name.to_string()),
other => {
return Err(anyhow!(
"Malformed token in insert_query: unknown prefix '{}' in '${{{}}}' (expected 'metadata' or 'payload')",
other,
inner
))
}
};
sources.push(source);
out.push_str(&positional_placeholder(driver_name, sources.len()));
i = close + 1;
} else {
let ch = query[i..].chars().next().unwrap();
out.push(ch);
i += ch.len_utf8();
}
}
Ok((out, sources))
}
fn resolve_source(
msg: &CanonicalMessage,
source: &ColumnSource,
payload_json: &Option<serde_json::Value>,
) -> BindValue {
match source {
ColumnSource::Metadata(key) => match msg.metadata.get(key) {
Some(v) => BindValue::Text(v.clone()),
None => BindValue::Null,
},
ColumnSource::Payload(field) => match payload_json.as_ref().and_then(|v| v.get(field)) {
Some(serde_json::Value::String(s)) => BindValue::Text(s.clone()),
Some(serde_json::Value::Bool(b)) => BindValue::Bool(*b),
Some(serde_json::Value::Number(n)) => {
if let Some(i) = n.as_i64() {
BindValue::Int(i)
} else if let Some(f) = n.as_f64() {
BindValue::Float(f)
} else {
BindValue::Null
}
}
_ => BindValue::Null,
},
}
}
type AnyQuery<'q> = sqlx::query::Query<'q, sqlx::Any, sqlx::any::AnyArguments>;
fn bind_value(query: AnyQuery<'_>, value: BindValue) -> AnyQuery<'_> {
match value {
BindValue::Null => query.bind(None::<String>),
BindValue::Int(i) => query.bind(i),
BindValue::Float(f) => query.bind(f),
BindValue::Bool(b) => query.bind(b),
BindValue::Text(s) => query.bind(s),
}
}
fn bind_message_sources<'q>(
mut query: AnyQuery<'q>,
msg: &CanonicalMessage,
sources: &[ColumnSource],
) -> AnyQuery<'q> {
let payload_json: Option<serde_json::Value> = serde_json::from_slice(&msg.payload).ok();
for source in sources {
query = bind_value(query, resolve_source(msg, source, &payload_json));
}
query
}
fn build_sqlx_url_with_tls(config: &SqlxConfig) -> anyhow::Result<String> {
let mut url = url::Url::parse(&config.url)?;
if let Some(username) = &config.username {
url.set_username(username)
.map_err(|_| anyhow!("Cannot set username on sqlx URL"))?;
}
if let Some(password) = &config.password {
url.set_password(Some(password))
.map_err(|_| anyhow!("Cannot set password on sqlx URL"))?;
}
if config.tls.required {
let scheme = url.scheme().to_string();
match scheme.as_str() {
"postgres" | "postgresql" => {
let mut query_pairs = url.query_pairs_mut();
if config.tls.accept_invalid_certs {
query_pairs.append_pair("sslmode", "require");
} else {
query_pairs.append_pair("sslmode", "verify-full");
}
if let Some(ca) = &config.tls.ca_file {
query_pairs.append_pair("sslrootcert", ca);
}
if let Some(cert) = &config.tls.cert_file {
query_pairs.append_pair("sslcert", cert);
}
if let Some(key) = &config.tls.key_file {
query_pairs.append_pair("sslkey", key);
}
if let Some(pass) = &config.tls.cert_password {
query_pairs.append_pair("sslpassword", pass);
}
}
"mysql" | "mariadb" => {
warn!("For complex MySQL/MariaDB TLS setups, using a client configuration file (my.cnf) is recommended over URL parameters.");
let mut query_pairs = url.query_pairs_mut();
if config.tls.accept_invalid_certs {
query_pairs.append_pair("ssl-mode", "REQUIRED");
} else if config.tls.ca_file.is_some() {
query_pairs.append_pair("ssl-mode", "VERIFY_CA");
} else {
query_pairs.append_pair("ssl-mode", "VERIFY_IDENTITY");
}
if let Some(ca) = &config.tls.ca_file {
query_pairs.append_pair("ssl-ca", ca);
}
}
"mssql" | "sqlserver" => {
let mut query_pairs = url.query_pairs_mut();
if config.tls.accept_invalid_certs {
query_pairs.append_pair("encrypt", "true");
query_pairs.append_pair("trust-server-certificate", "true");
} else {
query_pairs.append_pair("encrypt", "strict");
}
}
_ => {}
}
}
Ok(url.to_string())
}
async fn create_sqlx_pool(config: &SqlxConfig) -> anyhow::Result<AnyPool> {
let url = build_sqlx_url_with_tls(config)?;
let mut pool_options = AnyPoolOptions::new();
if let Some(max_conn) = config.max_connections {
pool_options = pool_options.max_connections(max_conn);
}
if let Some(min_conn) = config.min_connections {
pool_options = pool_options.min_connections(min_conn);
}
if let Some(timeout) = config.acquire_timeout_ms {
pool_options = pool_options.acquire_timeout(Duration::from_millis(timeout));
}
if let Some(timeout) = config.idle_timeout_ms {
pool_options = pool_options.idle_timeout(Duration::from_millis(timeout));
}
if let Some(lifetime) = config.max_lifetime_ms {
pool_options = pool_options.max_lifetime(Duration::from_millis(lifetime));
}
Ok(pool_options.connect(&url).await?)
}
async fn create_shared_sqlx_pool(config: &SqlxConfig) -> anyhow::Result<std::sync::Arc<AnyPool>> {
let identity = crate::support::connection_registry::connection_identity((
&config.url,
&config.username,
&config.password,
config.tls.required,
&config.tls.ca_file,
&config.tls.cert_file,
&config.tls.key_file,
&config.tls.cert_password,
config.tls.accept_invalid_certs,
(
config.max_connections,
config.min_connections,
config.acquire_timeout_ms,
config.idle_timeout_ms,
config.max_lifetime_ms,
),
));
let config_clone = config.clone();
crate::support::connection_registry::get_or_create(
"sqlx-pool",
identity,
config.shared.unwrap_or(true),
move || async move { create_sqlx_pool(&config_clone).await },
)
.await
}
pub struct SqlxPublisher {
pool: AnyPool,
_shared_pool: std::sync::Arc<AnyPool>,
insert_query: String,
column_sources: Vec<ColumnSource>,
driver_name: String,
table: String,
copy: Option<PgCopySink>,
}
struct PgCopySink {
pool: PgPool,
table: String,
columns: Vec<String>,
sources: Vec<ColumnSource>,
}
fn copy_escape_text(s: &str) -> String {
let mut out = String::with_capacity(s.len());
for ch in s.chars() {
match ch {
'\\' => out.push_str("\\\\"),
'\t' => out.push_str("\\t"),
'\n' => out.push_str("\\n"),
'\r' => out.push_str("\\r"),
_ => out.push(ch),
}
}
out
}
fn is_deterministic_sqlstate(code: &str) -> bool {
matches!(&code.get(..2), Some("42") | Some("22") | Some("23"))
}
fn sql_error_is_deterministic(e: &sqlx::Error) -> bool {
use sqlx::error::ErrorKind;
let Some(db_err) = e.as_database_error() else {
return false;
};
if matches!(
db_err.kind(),
ErrorKind::UniqueViolation
| ErrorKind::ForeignKeyViolation
| ErrorKind::NotNullViolation
| ErrorKind::CheckViolation
| ErrorKind::ExclusionViolation
) {
return true;
}
if db_err.code().is_some_and(|c| is_deterministic_sqlstate(&c)) {
return true;
}
let msg = db_err.message().to_ascii_lowercase();
msg.contains("no such column") || msg.contains("no such table")
}
fn classify_sql_error(e: sqlx::Error) -> PublisherError {
if sql_error_is_deterministic(&e) {
return PublisherError::NonRetryable(anyhow!(e));
}
PublisherError::Retryable(anyhow!(e))
}
fn classify_sql_consumer_error(e: sqlx::Error) -> ConsumerError {
if sql_error_is_deterministic(&e) {
ConsumerError::Permanent(anyhow!(e))
} else {
ConsumerError::Connection(anyhow!(e))
}
}
fn extract_copy_columns(raw_query: &str, token_count: usize) -> anyhow::Result<Vec<String>> {
let upper = raw_query.to_uppercase();
if upper.contains("ON CONFLICT")
|| upper.contains("RETURNING")
|| upper.contains("ON DUPLICATE")
{
return Err(anyhow!(
"bulk_copy cannot be used with ON CONFLICT/RETURNING/ON DUPLICATE clauses (COPY does not support them)."
));
}
let values_pos = find_top_level_values(raw_query).ok_or_else(|| {
anyhow!("bulk_copy requires an INSERT ... VALUES query with a column list.")
})?;
let prefix = &raw_query[..values_pos];
let open = prefix.find('(').ok_or_else(|| {
anyhow!(
"bulk_copy requires an explicit column list, e.g. INSERT INTO t (a, b) VALUES (...)."
)
})?;
let close = prefix[open..]
.find(')')
.map(|off| open + off)
.ok_or_else(|| anyhow!("bulk_copy: unbalanced parentheses in the column list."))?;
let columns: Vec<String> = prefix[open + 1..close]
.split(',')
.map(|c| c.trim().to_string())
.collect();
if columns.iter().any(|c| c.is_empty()) || columns.len() != token_count {
return Err(anyhow!(
"bulk_copy: the column list ({} columns) must match the {} `${{...}}` value token(s), one token per column.",
columns.len(),
token_count
));
}
let after = &raw_query[values_pos + "VALUES".len()..];
let vopen = after
.find('(')
.ok_or_else(|| anyhow!("bulk_copy: could not find the VALUES tuple."))?;
let vclose = after[vopen..]
.find(')')
.map(|off| vopen + off)
.ok_or_else(|| anyhow!("bulk_copy: unbalanced parentheses in the VALUES tuple."))?;
let mut residue = after[vopen + 1..vclose].to_string();
while let Some(s) = residue.find("${") {
match residue[s..].find('}') {
Some(e) => residue.replace_range(s..s + e + 1, ""),
None => break,
}
}
if residue.chars().any(|c| c != ',' && !c.is_whitespace()) {
return Err(anyhow!(
"bulk_copy requires every VALUES entry to be a single `${{...}}` token (no literals, expressions, or functions)."
));
}
Ok(columns)
}
impl SqlxPublisher {
pub async fn new(config: &SqlxConfig) -> anyhow::Result<Self> {
sqlx::any::install_default_drivers();
if !is_valid_table_name(&config.table) {
return Err(anyhow!(
"Invalid table name: '{}'. Only alphanumeric characters and underscores are allowed.",
config.table
));
}
let shared_pool = create_shared_sqlx_pool(config).await?;
let pool = (*shared_pool).clone();
let table = config.table.clone();
let conn = pool.acquire().await?;
let driver_name = conn.backend_name().to_string();
drop(conn);
info!(table = %config.table, driver = %driver_name, "SQLx publisher connected");
let raw_insert_query =
config
.insert_query
.clone()
.unwrap_or_else(|| match driver_name.as_str() {
"PostgreSQL" => format!("INSERT INTO {} (payload) VALUES ($1)", config.table),
"Microsoft SQL Server" => {
format!("INSERT INTO {} (payload) VALUES (@p1)", config.table)
}
_ => format!("INSERT INTO {} (payload) VALUES (?)", config.table),
});
let (insert_query, column_sources) =
parse_insert_template(&raw_insert_query, &driver_name)?;
if config.auto_create_table && !column_sources.is_empty() {
return Err(anyhow!(
"auto_create_table is not supported with a multi-column insert_query; create the table manually."
));
}
if config.auto_create_table {
let create_table_query = match driver_name.as_str() {
"PostgreSQL" => format!(
"CREATE TABLE IF NOT EXISTS {} (id BIGSERIAL PRIMARY KEY, payload BYTEA NOT NULL, locked_until TIMESTAMPTZ, created_at TIMESTAMPTZ DEFAULT NOW())",
config.table
),
"MySQL" | "MariaDB" => format!(
"CREATE TABLE IF NOT EXISTS {} (id BIGINT AUTO_INCREMENT PRIMARY KEY, payload BLOB NOT NULL, locked_until DATETIME, created_at DATETIME DEFAULT CURRENT_TIMESTAMP)",
config.table
),
"SQLite" => format!(
"CREATE TABLE IF NOT EXISTS {} (id INTEGER PRIMARY KEY AUTOINCREMENT, payload BLOB NOT NULL, locked_until DATETIME, created_at DATETIME DEFAULT CURRENT_TIMESTAMP)",
config.table
),
"Microsoft SQL Server" => format!(
"IF NOT EXISTS (SELECT * FROM sys.objects WHERE object_id = OBJECT_ID(N'{0}') AND type in (N'U'))
CREATE TABLE {0} (id BIGINT IDENTITY(1,1) PRIMARY KEY, payload VARBINARY(MAX) NOT NULL, locked_until DATETIME2, created_at DATETIME2 DEFAULT GETUTCDATE())",
config.table
),
_ => "".to_string(), };
if !create_table_query.is_empty() {
if let Err(e) = sqlx::query(audited_sql(&create_table_query))
.execute(&pool)
.await
{
warn!(
"Failed to auto-create table '{}': {}. Please ensure it exists.",
config.table, e
);
} else {
let table_name_for_index =
config.table.split('.').next_back().unwrap_or(&config.table);
let index_name = format!("idx_{}_locked_until", table_name_for_index);
let create_index_query = match driver_name.as_str() {
"PostgreSQL" | "SQLite" | "MariaDB" => {
format!(
"CREATE INDEX IF NOT EXISTS {} ON {} (locked_until)",
index_name, config.table
)
}
"MySQL" => {
format!(
"CREATE INDEX {} ON {} (locked_until)",
index_name, config.table
)
}
"Microsoft SQL Server" => {
format!(
"IF NOT EXISTS (SELECT * FROM sys.indexes WHERE name = N'{}' AND object_id = OBJECT_ID(N'{}'))
CREATE INDEX {} ON {} (locked_until)",
index_name, config.table, index_name, config.table
)
}
_ => "".to_string(),
};
if !create_index_query.is_empty() {
if let Err(e) = sqlx::query(audited_sql(&create_index_query))
.execute(&pool)
.await
{
let driver_lc = driver_name.to_lowercase();
if (driver_lc.contains("mysql") || driver_lc.contains("mariadb"))
&& e.as_database_error()
.is_some_and(|db_err| db_err.code().as_deref() == Some("1061"))
{
trace!("Index {} on {} already exists.", index_name, config.table);
} else {
warn!("Failed to create index on '{}': {}", config.table, e);
}
}
}
}
}
}
let copy = if config.bulk_copy {
if driver_name != "PostgreSQL" {
return Err(anyhow!(
"bulk_copy is only supported for PostgreSQL (driver: {}).",
driver_name
));
}
if column_sources.is_empty() {
return Err(anyhow!(
"bulk_copy requires a token-based insert_query (e.g. INSERT INTO t (a, b) VALUES (${{payload:a}}, ${{payload:b}})); single-payload COPY is not supported."
));
}
let columns = extract_copy_columns(&raw_insert_query, column_sources.len())?;
let url = build_sqlx_url_with_tls(config)?;
let pg_pool = PgPoolOptions::new()
.max_connections(config.max_connections.unwrap_or(5))
.connect(&url)
.await
.context("bulk_copy: failed to open native PostgreSQL pool")?;
Some(PgCopySink {
pool: pg_pool,
table: table.clone(),
columns,
sources: column_sources.clone(),
})
} else {
None
};
Ok(Self {
pool,
_shared_pool: shared_pool,
insert_query,
column_sources,
driver_name,
table,
copy,
})
}
}
#[async_trait]
impl MessagePublisher for SqlxPublisher {
async fn send(&self, message: CanonicalMessage) -> Result<Sent, PublisherError> {
trace!(message_id = %format!("{:032x}", message.message_id), table = %self.table, "Publishing to SQL");
let query = sqlx::query(audited_sql(&self.insert_query));
let query = if self.column_sources.is_empty() {
query.bind(message.payload.to_vec())
} else {
bind_message_sources(query, &message, &self.column_sources)
};
query
.execute(&self.pool)
.await
.map_err(classify_sql_error)?;
Ok(Sent::Ack)
}
async fn send_batch(
&self,
messages: Vec<CanonicalMessage>,
) -> Result<SentBatch, PublisherError> {
if messages.is_empty() {
return Ok(SentBatch::Ack);
}
if let Some(sink) = &self.copy {
return self.send_batch_copy(sink, messages).await;
}
trace!(count = messages.len(), message_ids = ?LazyMessageIds(&messages), "Publishing batch to SQLx");
let values_pos = match find_top_level_values(&self.insert_query) {
Some(pos) => pos,
None => {
warn!("Could not optimize batch insert due to custom query format. Falling back to iterative inserts.");
return self.send_batch_iterative(messages).await;
}
};
let base_query = &self.insert_query[..values_pos];
let after_values = values_pos + "VALUES".len();
let values_suffix = match self.insert_query[after_values..].find('(') {
Some(rel_open) => {
let open = after_values + rel_open;
let mut depth = 0usize;
let mut end = None;
for (idx, ch) in self.insert_query[open..].char_indices() {
match ch {
'(' => depth += 1,
')' => {
depth -= 1;
if depth == 0 {
end = Some(open + idx + ch.len_utf8());
break;
}
}
_ => {}
}
}
end.map(|e| &self.insert_query[e..]).unwrap_or("")
}
None => "",
};
if self.column_sources.is_empty() && !contains_payload_clause(base_query) {
warn!("Could not optimize batch insert due to custom query format. Falling back to iterative inserts.");
return self.send_batch_iterative(messages).await;
}
let per_row = self.column_sources.len().max(1);
let mut placeholders = String::new();
let mut param_idx = 1;
for i in 0..messages.len() {
if i > 0 {
placeholders.push_str(", ");
}
placeholders.push('(');
for j in 0..per_row {
if j > 0 {
placeholders.push_str(", ");
}
placeholders.push_str(&positional_placeholder(&self.driver_name, param_idx));
param_idx += 1;
}
placeholders.push(')');
}
let sql = format!("{} VALUES {}{}", base_query, placeholders, values_suffix);
let mut query = sqlx::query(audited_sql(&sql));
for msg in &messages {
if self.column_sources.is_empty() {
query = query.bind(msg.payload.to_vec());
} else {
query = bind_message_sources(query, msg, &self.column_sources);
}
}
query
.execute(&self.pool)
.await
.map_err(classify_sql_error)?;
Ok(SentBatch::Ack)
}
async fn status(&self) -> EndpointStatus {
let (healthy, error) = match self.pool.acquire().await {
Ok(_) => (true, None),
Err(e) => (false, Some(e.to_string())),
};
EndpointStatus {
healthy,
target: self.table.clone(),
error,
details: serde_json::json!({ "driver": self.driver_name, "pool_size": self.pool.size(), "pool_idle": self.pool.num_idle() }),
..Default::default()
}
}
fn as_any(&self) -> &dyn std::any::Any {
self
}
}
impl SqlxPublisher {
async fn send_batch_copy(
&self,
sink: &PgCopySink,
messages: Vec<CanonicalMessage>,
) -> Result<SentBatch, PublisherError> {
let stmt = format!(
"COPY {} ({}) FROM STDIN WITH (FORMAT text)",
sink.table,
sink.columns.join(", ")
);
let mut buf = String::new();
for msg in &messages {
let payload_json: Option<serde_json::Value> = serde_json::from_slice(&msg.payload).ok();
for (i, source) in sink.sources.iter().enumerate() {
if i > 0 {
buf.push('\t');
}
match resolve_source(msg, source, &payload_json) {
BindValue::Null => buf.push_str("\\N"),
BindValue::Int(n) => buf.push_str(&n.to_string()),
BindValue::Float(f) => buf.push_str(&f.to_string()),
BindValue::Bool(b) => buf.push_str(if b { "t" } else { "f" }),
BindValue::Text(s) => buf.push_str(©_escape_text(&s)),
}
}
buf.push('\n');
}
let mut copier = sink
.pool
.copy_in_raw(&stmt)
.await
.map_err(classify_sql_error)?;
copier
.send(buf.as_bytes())
.await
.map_err(classify_sql_error)?;
copier.finish().await.map_err(classify_sql_error)?;
trace!(count = messages.len(), table = %sink.table, "Bulk-copied batch to PostgreSQL");
Ok(SentBatch::Ack)
}
async fn send_batch_iterative(
&self,
messages: Vec<CanonicalMessage>,
) -> Result<SentBatch, PublisherError> {
let mut tx = self
.pool
.begin()
.await
.map_err(|e| PublisherError::Retryable(anyhow!(e)))?;
for msg in &messages {
let query = sqlx::query(audited_sql(&self.insert_query));
let query = if self.column_sources.is_empty() {
query.bind(msg.payload.to_vec())
} else {
bind_message_sources(query, msg, &self.column_sources)
};
query.execute(&mut *tx).await.map_err(classify_sql_error)?;
}
tx.commit()
.await
.map_err(|e| PublisherError::Retryable(anyhow!(e)))?;
Ok(SentBatch::Ack)
}
}
pub struct SqlxConsumer {
pool: AnyPool,
select_query: String,
delete_after_read: bool,
table: String,
backoff: PollBackoff,
driver_name: String,
}
impl SqlxConsumer {
pub async fn new(config: &SqlxConfig) -> anyhow::Result<Self> {
sqlx::any::install_default_drivers();
if !is_valid_table_name(&config.table) {
return Err(anyhow!(
"Invalid table name: '{}'. Only alphanumeric characters and underscores are allowed.",
config.table
));
}
let pool = create_sqlx_pool(config).await?;
let conn = pool.acquire().await?;
let driver_name = conn.backend_name().to_string();
drop(conn);
info!(table = %config.table, driver = %driver_name, "SQLx consumer connected");
let select_query = if let Some(query) = &config.select_query {
match driver_name.as_str() {
"PostgreSQL" => {
if !query.contains("$1") {
return Err(anyhow!("Custom select_query for PostgreSQL must contain a '$1' placeholder for the batch size limit."));
}
query.clone()
}
"Microsoft SQL Server" => {
if !query.contains("@p1") {
return Err(anyhow!("Custom select_query for SQL Server must contain a '@p1' placeholder for the batch size limit."));
}
query.clone()
}
_ => {
return Err(anyhow!("Custom select_query is not supported for the '{}' driver. It is only supported for PostgreSQL and Microsoft SQL Server.", driver_name));
}
}
} else {
match driver_name.as_str() {
"PostgreSQL" => {
format!(
r#"
WITH available AS (
SELECT id FROM {0}
WHERE locked_until IS NULL OR locked_until < NOW()
ORDER BY id
LIMIT $1
FOR UPDATE SKIP LOCKED
),
updated AS (
UPDATE {0}
SET locked_until = NOW() + interval '60 seconds'
WHERE id IN (SELECT id FROM available)
RETURNING id, payload
)
SELECT id, payload FROM updated"#,
config.table,
)
}
"Microsoft SQL Server" => {
format!(
r#"
UPDATE {0}
SET locked_until = DATEADD(second, 60, GETUTCDATE())
OUTPUT INSERTED.id, INSERTED.payload
WHERE id IN (SELECT TOP (@p1) id FROM {0} WITH (UPDLOCK, READPAST) WHERE locked_until IS NULL OR locked_until < GETUTCDATE() ORDER BY id)"#,
config.table
)
}
_ => format!("SELECT id, payload FROM {}", config.table),
}
};
Ok(Self {
pool,
select_query,
delete_after_read: config.delete_after_read,
table: config.table.clone(),
backoff: PollBackoff::new(
Duration::from_millis(config.polling_interval_ms.unwrap_or(100)),
config.max_polling_interval_ms.map(Duration::from_millis),
),
driver_name,
})
}
}
impl SqlxConsumer {
async fn fetch_and_lock_mysql(
&self,
limit: usize,
) -> Result<Vec<sqlx::any::AnyRow>, ConsumerError> {
let mut tx = self
.pool
.begin()
.await
.map_err(classify_sql_consumer_error)?;
let lock_query = format!(
"SELECT id FROM {} WHERE locked_until IS NULL OR locked_until < NOW() ORDER BY id LIMIT ? FOR UPDATE SKIP LOCKED",
self.table
);
let locked_ids: Vec<i64> = sqlx::query(audited_sql(&lock_query))
.bind(limit as i64)
.fetch_all(&mut *tx)
.await
.map_err(classify_sql_consumer_error)?
.into_iter()
.map(|row| row.get("id"))
.collect();
if locked_ids.is_empty() {
tx.commit().await.ok(); return Ok(vec![]);
}
let placeholders = locked_ids
.iter()
.map(|_| "?")
.collect::<Vec<_>>()
.join(", ");
let update_query = format!(
"UPDATE {} SET locked_until = NOW() + INTERVAL 60 SECOND WHERE id IN ({})",
self.table, placeholders
);
let mut query = sqlx::query(audited_sql(&update_query));
for id in &locked_ids {
query = query.bind(*id);
}
query
.execute(&mut *tx)
.await
.map_err(classify_sql_consumer_error)?;
let select_query = format!(
"SELECT id, payload FROM {} WHERE id IN ({})",
self.table, placeholders
);
let mut query = sqlx::query(audited_sql(&select_query));
for id in &locked_ids {
query = query.bind(*id);
}
let rows = query
.fetch_all(&mut *tx)
.await
.map_err(classify_sql_consumer_error)?;
tx.commit().await.map_err(classify_sql_consumer_error)?;
Ok(rows)
}
async fn fetch_and_lock_sqlite(
&self,
limit: usize,
) -> Result<Vec<sqlx::any::AnyRow>, ConsumerError> {
let mut tx = self
.pool
.begin_with("BEGIN IMMEDIATE")
.await
.map_err(classify_sql_consumer_error)?;
let select_query = format!(
"SELECT id FROM {} WHERE locked_until IS NULL OR locked_until < datetime('now') ORDER BY id LIMIT ?",
self.table
);
let locked_ids: Vec<i64> = sqlx::query(audited_sql(&select_query))
.bind(limit as i64)
.fetch_all(&mut *tx)
.await
.map_err(classify_sql_consumer_error)?
.into_iter()
.map(|row| row.get("id"))
.collect();
if locked_ids.is_empty() {
tx.commit().await.ok();
return Ok(vec![]);
}
let placeholders = locked_ids
.iter()
.map(|_| "?")
.collect::<Vec<_>>()
.join(", ");
let update_query = format!(
"UPDATE {} SET locked_until = datetime('now', '+60 seconds') WHERE id IN ({})",
self.table, placeholders
);
let mut query = sqlx::query(audited_sql(&update_query));
for id in &locked_ids {
query = query.bind(*id);
}
query
.execute(&mut *tx)
.await
.map_err(classify_sql_consumer_error)?;
let select_payload_query = format!(
"SELECT id, payload FROM {} WHERE id IN ({})",
self.table, placeholders
);
let mut query = sqlx::query(audited_sql(&select_payload_query));
for id in &locked_ids {
query = query.bind(*id);
}
let rows = query
.fetch_all(&mut *tx)
.await
.map_err(classify_sql_consumer_error)?;
tx.commit().await.map_err(classify_sql_consumer_error)?;
Ok(rows)
}
async fn get_pending_count(&self) -> anyhow::Result<usize> {
let query = match self.driver_name.as_str() {
"PostgreSQL" | "MySQL" | "MariaDB" => format!(
"SELECT COUNT(*) FROM {} WHERE locked_until IS NULL OR locked_until < NOW()",
self.table
),
"SQLite" => format!(
"SELECT COUNT(*) FROM {} WHERE locked_until IS NULL OR locked_until < datetime('now')",
self.table
),
"Microsoft SQL Server" => format!(
"SELECT COUNT(*) FROM {} WHERE locked_until IS NULL OR locked_until < GETUTCDATE()",
self.table
),
_ => anyhow::bail!("Unsupported driver for pending count: {}", self.driver_name),
};
let row: sqlx::any::AnyRow = sqlx::query(audited_sql(&query))
.fetch_one(&self.pool)
.await?;
if let Ok(c) = row.try_get::<i64, _>(0) {
usize::try_from(c).map_err(|e| anyhow!("i64 to usize conversion failed: {}", e))
} else {
let c: i32 = row.try_get(0)?;
usize::try_from(c).map_err(|e| anyhow!("i32 to usize conversion failed: {}", e))
}
}
}
#[async_trait]
impl MessageConsumer for SqlxConsumer {
fn commit_requires_order(&self) -> bool {
false
}
async fn receive_batch(&mut self, max_messages: usize) -> Result<ReceivedBatch, ConsumerError> {
if max_messages == 0 {
return Ok(ReceivedBatch {
messages: Vec::new(),
commit: Box::new(|_| Box::pin(async { Ok(()) })),
});
}
let rows = match self.driver_name.as_str() {
"PostgreSQL" | "Microsoft SQL Server" => sqlx::query(audited_sql(&self.select_query))
.bind(max_messages as i64)
.fetch_all(&self.pool)
.await
.map_err(classify_sql_consumer_error)?,
"MySQL" | "MariaDB" => self.fetch_and_lock_mysql(max_messages).await?,
"SQLite" => self.fetch_and_lock_sqlite(max_messages).await?,
_ => {
warn!("SQLx consumer for driver '{}' is using a non-locking read strategy. This is not safe for concurrent consumers.", self.driver_name);
let final_query = format!("{} LIMIT ?", self.select_query);
sqlx::query(audited_sql(&final_query))
.bind(max_messages as i64)
.fetch_all(&self.pool)
.await
.map_err(classify_sql_consumer_error)?
}
};
if rows.is_empty() {
tokio::time::sleep(self.backoff.idle_delay()).await;
return Ok(ReceivedBatch {
messages: Vec::new(),
commit: Box::new(|_| Box::pin(async { Ok(()) })),
});
}
self.backoff.reset();
let mut messages = Vec::new();
let mut ids_to_delete = Vec::new();
for row in rows.into_iter().take(max_messages) {
let payload: Vec<u8> = row
.try_get("payload")
.context("Failed to get 'payload' column")?;
let id: i64 = row.try_get("id").context("Failed to get 'id' column")?;
messages.push(CanonicalMessage::new(payload, None));
ids_to_delete.push(id);
}
trace!(count = messages.len(), "Received batch of SQLx messages");
let pool = self.pool.clone();
let table = self.table.clone();
let delete = self.delete_after_read;
let driver_name = self.driver_name.clone();
let commit = Box::new(move |dispositions: Vec<MessageDisposition>| {
let pool = pool.clone();
let table = table.clone();
let ids = ids_to_delete.clone();
let driver_name = driver_name.clone();
Box::pin(async move {
if !delete {
return Ok(());
}
let mut ids_to_ack = Vec::new();
for (i, disp) in dispositions.iter().enumerate() {
let should_ack = match disp {
MessageDisposition::Ack => true,
MessageDisposition::Reply(_) => {
tracing::warn!("SQLx consumer received a Reply/StreamReply, but replying is not supported by this endpoint. The reply payload is dropped, and the original message is acknowledged.");
true
}
MessageDisposition::Nack => false,
};
if should_ack {
if let Some(id) = ids.get(i) {
ids_to_ack.push(*id);
}
}
}
if !ids_to_ack.is_empty() {
let mut placeholders = String::new();
for i in 0..ids_to_ack.len() {
if i > 0 {
placeholders.push_str(", ");
}
match driver_name.as_str() {
"PostgreSQL" => placeholders.push_str(&format!("${}", i + 1)),
"Microsoft SQL Server" => {
placeholders.push_str(&format!("@p{}", i + 1))
}
_ => placeholders.push('?'),
}
}
let sql = format!("DELETE FROM {} WHERE id IN ({})", table, placeholders);
let mut attempts = 0;
loop {
let mut query = sqlx::query(audited_sql(&sql));
for id in &ids_to_ack {
query = query.bind(*id);
}
match query.execute(&pool).await {
Ok(_) => break,
Err(e) => {
if is_deadlock_error(&e) && attempts < 5 {
attempts += 1;
warn!(
attempts,
error = %e,
"Deadlock detected during SQLx commit, retrying..."
);
tokio::time::sleep(Duration::from_millis(attempts * 50)).await;
continue;
}
return Err(anyhow!("Failed to delete acked messages: {}", e));
}
}
}
}
Ok(())
}) as BoxFuture<'static, anyhow::Result<()>>
});
Ok(ReceivedBatch { messages, commit })
}
async fn status(&self) -> EndpointStatus {
let (mut healthy, mut error) = match self.pool.acquire().await {
Ok(_) => (true, None),
Err(e) => (false, Some(e.to_string())),
};
let mut pending = None;
if healthy {
match self.get_pending_count().await {
Ok(c) => pending = Some(c),
Err(e) => {
healthy = false;
error = Some(e.to_string());
}
}
};
EndpointStatus {
healthy,
target: self.table.clone(),
pending,
error,
details: serde_json::json!({ "driver": self.driver_name, "pool_size": self.pool.size(), "pool_idle": self.pool.num_idle() }),
..Default::default()
}
}
fn as_any(&self) -> &dyn std::any::Any {
self
}
}
#[derive(Clone, Debug, PartialEq)]
enum SqlCursor {
Int(i64),
Text(String),
}
impl SqlCursor {
fn encode(&self) -> String {
match self {
SqlCursor::Int(n) => format!("int:{}", n),
SqlCursor::Text(s) => format!("str:{}", s),
}
}
fn decode(s: &str) -> Option<SqlCursor> {
let (tag, val) = s.split_once(':')?;
match tag {
"int" => val.parse::<i64>().ok().map(SqlCursor::Int),
"str" => Some(SqlCursor::Text(val.to_string())),
_ => None,
}
}
}
fn row_to_json(row: &sqlx::any::AnyRow) -> serde_json::Value {
let mut map = serde_json::Map::with_capacity(row.columns().len());
for col in row.columns() {
map.insert(col.name().to_string(), extract_json_value(row, col));
}
serde_json::Value::Object(map)
}
fn extract_json_value(
row: &sqlx::any::AnyRow,
col: &<sqlx::Any as sqlx::Database>::Column,
) -> serde_json::Value {
use serde_json::Value;
use sqlx::any::AnyTypeInfoKind;
let idx = col.ordinal();
match col.type_info().kind() {
AnyTypeInfoKind::Null => Value::Null,
AnyTypeInfoKind::Bool => row
.try_get::<Option<bool>, _>(idx)
.ok()
.flatten()
.map(Value::from)
.unwrap_or(Value::Null),
AnyTypeInfoKind::SmallInt | AnyTypeInfoKind::Integer | AnyTypeInfoKind::BigInt => row
.try_get::<Option<i64>, _>(idx)
.ok()
.flatten()
.map(Value::from)
.unwrap_or(Value::Null),
AnyTypeInfoKind::Real => row
.try_get::<Option<f32>, _>(idx)
.ok()
.flatten()
.map(|v| Value::from(v as f64))
.unwrap_or(Value::Null),
AnyTypeInfoKind::Double => row
.try_get::<Option<f64>, _>(idx)
.ok()
.flatten()
.map(Value::from)
.unwrap_or(Value::Null),
AnyTypeInfoKind::Text => row
.try_get::<Option<String>, _>(idx)
.ok()
.flatten()
.map(Value::from)
.unwrap_or(Value::Null),
AnyTypeInfoKind::Blob => row
.try_get::<Option<Vec<u8>>, _>(idx)
.ok()
.flatten()
.map(|b| {
const HEX: &[u8; 16] = b"0123456789abcdef";
let mut s = String::with_capacity(b.len() * 2);
for x in &b {
s.push(HEX[(x >> 4) as usize] as char);
s.push(HEX[(x & 0x0f) as usize] as char);
}
Value::from(s)
})
.unwrap_or(Value::Null),
}
}
fn resolve_cursor_column(
row: &sqlx::any::AnyRow,
column: &str,
) -> Option<(usize, sqlx::any::AnyTypeInfoKind)> {
let col = row.try_column(column).ok()?;
Some((col.ordinal(), col.type_info().kind()))
}
fn extract_cursor_at(
row: &sqlx::any::AnyRow,
idx: usize,
kind: sqlx::any::AnyTypeInfoKind,
) -> Option<SqlCursor> {
use sqlx::any::AnyTypeInfoKind;
match kind {
AnyTypeInfoKind::SmallInt | AnyTypeInfoKind::Integer | AnyTypeInfoKind::BigInt => row
.try_get::<Option<i64>, _>(idx)
.ok()
.flatten()
.map(SqlCursor::Int),
AnyTypeInfoKind::Text => row
.try_get::<Option<String>, _>(idx)
.ok()
.flatten()
.map(SqlCursor::Text),
_ => None,
}
}
fn pg_typname_is_any_safe(typname: &str) -> bool {
matches!(
typname,
"bool"
| "int2"
| "int4"
| "int8"
| "float4"
| "float8"
| "bytea"
| "text"
| "varchar"
| "citext"
)
}
fn quote_ident(name: &str) -> String {
format!("\"{}\"", name.replace('"', "\"\""))
}
async fn build_cursor_projection(pool: &AnyPool, driver_name: &str, table: &str) -> String {
if driver_name != "PostgreSQL" {
return "*".to_string();
}
let sql = "SELECT a.attname::text AS name, t.typname::text AS typname \
FROM pg_attribute a JOIN pg_type t ON t.oid = a.atttypid \
WHERE a.attrelid = $1::regclass AND a.attnum > 0 AND NOT a.attisdropped \
ORDER BY a.attnum";
let rows = match sqlx::query(sql).bind(table).fetch_all(pool).await {
Ok(rows) if !rows.is_empty() => rows,
Ok(_) => return "*".to_string(),
Err(e) => {
warn!(table = %table, error = %e, "Could not introspect PostgreSQL columns; falling back to SELECT * (timestamptz/uuid/etc. columns may fail to decode)");
return "*".to_string();
}
};
let mut parts = Vec::with_capacity(rows.len());
for row in &rows {
let name: String = match row.try_get("name") {
Ok(n) => n,
Err(_) => return "*".to_string(),
};
let typname: String = row.try_get("typname").unwrap_or_default();
let ident = quote_ident(&name);
if pg_typname_is_any_safe(&typname) {
parts.push(ident);
} else {
parts.push(format!("{ident}::text AS {ident}"));
}
}
parts.join(", ")
}
fn is_permanent_decode_error(e: &sqlx::Error) -> bool {
if matches!(e, sqlx::Error::ColumnDecode { .. }) {
return true;
}
let msg = e.to_string();
msg.contains("Any driver does not support") || msg.contains("Any driver mapping")
}
struct SqlTableCheckpointStore {
pool: AnyPool,
driver_name: String,
meta_table: String,
cursor_id: String,
}
impl SqlTableCheckpointStore {
async fn ensure_table(&self) -> anyhow::Result<()> {
let sql = format!(
"CREATE TABLE IF NOT EXISTS {} (cursor_id VARCHAR(255) PRIMARY KEY, last_value TEXT)",
self.meta_table
);
sqlx::query(audited_sql(&sql))
.execute(&self.pool)
.await
.with_context(|| format!("Failed to create meta table '{}'", self.meta_table))?;
Ok(())
}
}
#[async_trait]
impl crate::checkpoint::CheckpointStore for SqlTableCheckpointStore {
async fn load(&self) -> anyhow::Result<Option<String>> {
let sql = format!(
"SELECT last_value FROM {} WHERE cursor_id = {}",
self.meta_table,
positional_placeholder(&self.driver_name, 1)
);
let row = sqlx::query(audited_sql(&sql))
.bind(self.cursor_id.clone())
.fetch_optional(&self.pool)
.await?;
Ok(row.and_then(|r| r.try_get::<Option<String>, _>("last_value").ok().flatten()))
}
async fn save(&self, value: &str) -> anyhow::Result<()> {
let p1 = positional_placeholder(&self.driver_name, 1);
let p2 = positional_placeholder(&self.driver_name, 2);
let sql = match self.driver_name.as_str() {
"MySQL" | "MariaDB" => format!(
"INSERT INTO {0} (cursor_id, last_value) VALUES ({1}, {2}) \
ON DUPLICATE KEY UPDATE last_value = VALUES(last_value)",
self.meta_table, p1, p2
),
_ => format!(
"INSERT INTO {0} (cursor_id, last_value) VALUES ({1}, {2}) \
ON CONFLICT (cursor_id) DO UPDATE SET last_value = excluded.last_value",
self.meta_table, p1, p2
),
};
sqlx::query(audited_sql(&sql))
.bind(self.cursor_id.clone())
.bind(value.to_string())
.execute(&self.pool)
.await
.with_context(|| format!("Failed to persist cursor to '{}'", self.meta_table))?;
Ok(())
}
}
pub(crate) async fn build_sql_checkpoint_store(
url: &str,
table: Option<String>,
source_name: &str,
cursor_id: &str,
) -> anyhow::Result<Arc<dyn crate::checkpoint::CheckpointStore>> {
sqlx::any::install_default_drivers();
let pool = AnyPool::connect(url)
.await
.with_context(|| format!("Failed to connect checkpoint store at '{}'", url))?;
let driver_name = {
let conn = pool.acquire().await?;
let name = conn.backend_name().to_string();
drop(conn);
name
};
let meta_table = table.unwrap_or_else(|| crate::checkpoint::default_meta_name(source_name));
source_sql_checkpoint_store(pool, driver_name, meta_table, source_name, cursor_id).await
}
async fn source_sql_checkpoint_store(
pool: AnyPool,
driver_name: String,
meta_table: String,
source_name: &str,
cursor_id: &str,
) -> anyhow::Result<Arc<dyn crate::checkpoint::CheckpointStore>> {
if !is_valid_table_name(&meta_table) {
return Err(anyhow!("Invalid checkpoint table name: '{}'.", meta_table));
}
let store = SqlTableCheckpointStore {
pool,
driver_name,
meta_table,
cursor_id: crate::checkpoint::checkpoint_key(source_name, cursor_id),
};
store.ensure_table().await?;
Ok(Arc::new(store))
}
pub struct SqlxCursorReader {
pool: AnyPool,
table: String,
cursor_column: String,
driver_name: String,
projection: String,
backoff: PollBackoff,
checkpoint: Option<Arc<dyn crate::checkpoint::CheckpointStore>>,
last_value: Arc<Mutex<Option<SqlCursor>>>,
}
impl SqlxCursorReader {
pub async fn new(config: &SqlxConfig) -> anyhow::Result<Self> {
sqlx::any::install_default_drivers();
if config.delete_after_read {
return Err(anyhow!(
"SQLx `cursor_column` (non-destructive) and `delete_after_read` are mutually exclusive"
));
}
if !is_valid_table_name(&config.table) {
return Err(anyhow!("Invalid table name: '{}'.", config.table));
}
let cursor_column = config
.cursor_column
.clone()
.ok_or_else(|| anyhow!("cursor_column is required for the SQLx cursor reader"))?;
if !is_valid_table_name(&cursor_column) {
return Err(anyhow!("Invalid cursor_column name: '{}'.", cursor_column));
}
let pool = create_sqlx_pool(config).await?;
let conn = pool.acquire().await?;
let driver_name = conn.backend_name().to_string();
drop(conn);
if driver_name == "Microsoft SQL Server" {
return Err(anyhow!(
"cursor_column mode is not supported for Microsoft SQL Server"
));
}
info!(table = %config.table, column = %cursor_column, driver = %driver_name, "SQLx cursor reader connected");
let checkpoint: Option<Arc<dyn crate::checkpoint::CheckpointStore>> = if let Some(cid) =
&config.cursor_id
{
use crate::checkpoint::CheckpointBackend;
let backend = match &config.checkpoint_store {
None => CheckpointBackend::Source {
name: crate::checkpoint::default_meta_name(&config.table),
},
Some(spec) => crate::checkpoint::parse_checkpoint_store(spec)?,
};
let store = match backend {
CheckpointBackend::Source { name } => {
source_sql_checkpoint_store(
pool.clone(),
driver_name.clone(),
name,
&config.table,
cid,
)
.await?
}
external => {
crate::checkpoint::build_external_store(external, &config.table, cid).await?
}
};
Some(store)
} else {
warn!(
table = %config.table,
"SQLx cursor reader has no cursor_id; resume is disabled and every restart re-copies from the beginning. Set cursor_id to persist progress."
);
None
};
let last_value = match &checkpoint {
Some(cp) => cp.load().await?.and_then(|s| {
let decoded = SqlCursor::decode(&s);
if decoded.is_none() {
warn!(value = %s, "Ignoring unparseable sql cursor; starting from beginning");
}
decoded
}),
None => None,
};
info!(table = %config.table, cursor_id = ?config.cursor_id, has_checkpoint = %last_value.is_some(), "SQLx cursor reader initialized");
let projection = build_cursor_projection(&pool, &driver_name, &config.table).await;
Ok(Self {
pool,
table: config.table.clone(),
cursor_column,
driver_name,
projection,
backoff: PollBackoff::new(
Duration::from_millis(config.polling_interval_ms.unwrap_or(100)),
config.max_polling_interval_ms.map(Duration::from_millis),
),
checkpoint,
last_value: Arc::new(Mutex::new(last_value)),
})
}
}
#[async_trait]
impl MessageConsumer for SqlxCursorReader {
async fn receive_batch(&mut self, max_messages: usize) -> Result<ReceivedBatch, ConsumerError> {
if max_messages == 0 {
return Ok(ReceivedBatch {
messages: Vec::new(),
commit: Box::new(|_| Box::pin(async { Ok(()) })),
});
}
let last = self.last_value.lock().unwrap().clone();
let sql = match &last {
Some(_) => format!(
"SELECT {0} FROM {1} WHERE {2} > {3} ORDER BY {2} ASC LIMIT {4}",
self.projection,
self.table,
self.cursor_column,
positional_placeholder(&self.driver_name, 1),
positional_placeholder(&self.driver_name, 2),
),
None => format!(
"SELECT {0} FROM {1} ORDER BY {2} ASC LIMIT {3}",
self.projection,
self.table,
self.cursor_column,
positional_placeholder(&self.driver_name, 1),
),
};
let mut query = sqlx::query(audited_sql(&sql));
if let Some(c) = &last {
query = match c {
SqlCursor::Int(n) => query.bind(*n),
SqlCursor::Text(s) => query.bind(s.clone()),
};
}
let fetch_limit = (max_messages as i64).saturating_add(1);
query = query.bind(fetch_limit);
let rows = match query.fetch_all(&self.pool).await {
Ok(rows) => rows,
Err(e) if is_permanent_decode_error(&e) => {
return Err(ConsumerError::Connection(anyhow::Error::new(
crate::errors::ProcessingError::NonRetryable(anyhow!(
"SQLx cursor reader on table '{}' hit a column of a type the SQL `Any` \
driver cannot decode: {e}. This is permanent; expose the column as \
TEXT/BIGINT (e.g. via a view that CASTs it) and point the reader there.",
self.table
)),
)));
}
Err(e) => return Err(classify_sql_consumer_error(e)),
};
if rows.is_empty() {
tokio::time::sleep(self.backoff.idle_delay()).await;
return Ok(ReceivedBatch {
messages: Vec::new(),
commit: Box::new(|_| Box::pin(async { Ok(()) })),
});
}
self.backoff.reset();
let cursor_col = resolve_cursor_column(&rows[0], &self.cursor_column);
let mut fetched: Vec<(SqlCursor, CanonicalMessage)> = Vec::with_capacity(rows.len());
for row in &rows {
let cursor = cursor_col
.and_then(|(idx, kind)| extract_cursor_at(row, idx, kind))
.ok_or_else(|| {
ConsumerError::Connection(anyhow!(
"cursor_column '{}' is missing or of a type the SQL `Any` driver cannot decode \
(only integer and text cursors are supported). CAST it to BIGINT/TEXT in a view, \
or point cursor_column at an integer or text column.",
self.cursor_column
))
})?;
let payload = serde_json::to_vec(&row_to_json(row)).unwrap_or_default();
fetched.push((cursor, CanonicalMessage::new(payload, None)));
}
let had_more = fetched.len() > max_messages;
let mut emit_len = fetched.len().min(max_messages);
if had_more {
let peek_val = fetched[max_messages].0.clone();
while emit_len > 0 && fetched[emit_len - 1].0 == peek_val {
emit_len -= 1;
}
if emit_len == 0 {
return Err(ConsumerError::Connection(anyhow!(
"cursor_column '{}' has a group of equal values larger than batch_size ({}); \
cannot page without skipping rows. Increase batch_size above the size of the \
largest equal-value group.",
self.cursor_column,
max_messages
)));
}
}
fetched.truncate(emit_len);
let mut messages = Vec::with_capacity(fetched.len());
let mut cursors: Vec<SqlCursor> = Vec::with_capacity(fetched.len());
for (cursor, msg) in fetched {
cursors.push(cursor.clone());
messages.push(msg);
*self.last_value.lock().unwrap() = Some(cursor);
}
trace!(count = messages.len(), "Received batch of SQLx cursor rows");
let checkpoint = self.checkpoint.clone();
let last_value = self.last_value.clone();
let resume_from = last; let commit = Box::new(move |dispositions: Vec<MessageDisposition>| {
Box::pin(async move {
let mut acked = 0usize;
for disp in dispositions.iter().take(cursors.len()) {
if matches!(disp, MessageDisposition::Ack | MessageDisposition::Reply(_)) {
acked += 1;
} else {
break;
}
}
let boundary = if acked == 0 {
resume_from
} else {
Some(cursors[acked - 1].clone())
};
if acked < cursors.len() {
*last_value.lock().unwrap() = boundary.clone();
}
if let (Some(cur), Some(cp)) = (boundary, checkpoint) {
if let Err(e) = cp.save(&cur.encode()).await {
tracing::warn!(error = %e, "Failed to persist sql cursor. Rows may be reprocessed on restart.");
}
}
Ok(())
}) as BoxFuture<'static, anyhow::Result<()>>
});
Ok(ReceivedBatch { messages, commit })
}
async fn status(&self) -> EndpointStatus {
let (healthy, error) = match self.pool.acquire().await {
Ok(_) => (true, None),
Err(e) => (false, Some(e.to_string())),
};
EndpointStatus {
healthy,
target: self.table.clone(),
error,
details: serde_json::json!({ "driver": self.driver_name, "mode": "cursor_column", "cursor_column": self.cursor_column }),
..Default::default()
}
}
fn as_any(&self) -> &dyn std::any::Any {
self
}
}
#[cfg(test)]
mod tests;