use postgres_types::ToSql;
use std::fmt::Write;
const ACCUMULATOR_OVERFLOW_MSG: &str = "djogi accumulator exceeded u32::MAX bind positions -- this is a framework-internal invariant break";
#[doc(hidden)]
pub struct SqlAccumulator {
sql: String,
binds: Vec<Box<dyn ToSql + Sync + Send>>,
next_param: u32,
}
impl SqlAccumulator {
pub fn new(initial_sql: &str) -> Self {
SqlAccumulator {
sql: initial_sql.to_owned(),
binds: Vec::new(),
next_param: 1,
}
}
pub fn push_sql(&mut self, s: &str) {
self.sql.push_str(s);
}
pub fn push_bind<T>(&mut self, v: T)
where
T: ToSql + Sync + Send + 'static,
{
let _ = write!(self.sql, "${}", self.next_param);
self.binds.push(Box::new(v));
self.next_param = self
.next_param
.checked_add(1)
.expect(ACCUMULATOR_OVERFLOW_MSG);
}
#[doc(hidden)]
pub fn push_boxed_bind(&mut self, v: Box<dyn ToSql + Sync + Send>) {
let _ = write!(self.sql, "${}", self.next_param);
self.binds.push(v);
self.next_param = self
.next_param
.checked_add(1)
.expect(ACCUMULATOR_OVERFLOW_MSG);
}
pub fn push_list_binds<T, I>(&mut self, iter: I)
where
T: ToSql + Sync + Send + 'static,
I: IntoIterator<Item = T>,
{
let mut first = true;
for v in iter {
if !first {
self.sql.push_str(", ");
}
first = false;
self.push_bind(v);
}
}
pub fn push_csv<'a, I: IntoIterator<Item = &'a str>>(&mut self, items: I) {
let mut first = true;
for s in items {
if !first {
self.sql.push_str(", ");
}
first = false;
self.sql.push_str(s);
}
}
pub fn push_null_literal(&mut self) {
self.sql.push_str("NULL");
}
pub fn extend_with(&mut self, other: SqlAccumulator) {
let SqlAccumulator {
sql: other_sql,
binds: other_binds,
next_param: _,
} = other;
let offset = self.next_param - 1;
if offset == 0 {
self.sql.push_str(&other_sql);
} else {
let reserve_extra = 9_usize
.checked_mul(other_binds.len())
.and_then(|extra| other_sql.len().checked_add(extra))
.expect(ACCUMULATOR_OVERFLOW_MSG);
self.sql.reserve(reserve_extra);
let bytes = other_sql.as_bytes();
let mut start = 0;
let mut i = 0;
while i < bytes.len() {
if bytes[i] == b'$' && i + 1 < bytes.len() && bytes[i + 1].is_ascii_digit() {
if start < i {
self.sql.push_str(&other_sql[start..i]);
}
let mut j = i + 1;
let mut n: u32 = 0;
while j < bytes.len() && bytes[j].is_ascii_digit() {
let digit = u32::from(bytes[j] - b'0');
n = n
.checked_mul(10)
.and_then(|n10| n10.checked_add(digit))
.expect(ACCUMULATOR_OVERFLOW_MSG);
j += 1;
}
let renumbered = n.checked_add(offset).expect(ACCUMULATOR_OVERFLOW_MSG);
let _ = write!(self.sql, "${}", renumbered);
i = j;
start = j;
} else {
i += 1;
}
}
if start < bytes.len() {
self.sql.push_str(&other_sql[start..]);
}
}
let other_count = u32::try_from(other_binds.len()).expect(ACCUMULATOR_OVERFLOW_MSG);
self.next_param = self
.next_param
.checked_add(other_count)
.expect(ACCUMULATOR_OVERFLOW_MSG);
self.binds.extend(other_binds);
}
pub fn into_parts(self) -> (String, Vec<Box<dyn ToSql + Sync + Send>>) {
(self.sql, self.binds)
}
pub fn sql(&self) -> &str {
&self.sql
}
pub fn bind_count(&self) -> u32 {
self.next_param - 1
}
pub fn pop_sql_suffix(&mut self, suffix: &str) -> bool {
if self.sql.ends_with(suffix) {
let new_len = self.sql.len() - suffix.len();
self.sql.truncate(new_len);
true
} else {
false
}
}
}
pub fn as_params(binds: &[Box<dyn ToSql + Sync + Send>]) -> Vec<&(dyn ToSql + Sync)> {
binds
.iter()
.map(|b| b.as_ref() as &(dyn ToSql + Sync))
.collect()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn new_accumulator_starts_empty() {
let acc = SqlAccumulator::new("SELECT 1");
assert_eq!(acc.sql(), "SELECT 1");
assert_eq!(acc.bind_count(), 0);
}
#[test]
fn push_sql_appends_raw_text() {
let mut acc = SqlAccumulator::new("SELECT * FROM t");
acc.push_sql(" WHERE active = ");
acc.push_sql("TRUE");
assert_eq!(acc.sql(), "SELECT * FROM t WHERE active = TRUE");
assert_eq!(acc.bind_count(), 0);
}
#[test]
fn push_bind_inserts_positional_placeholder() {
let mut acc = SqlAccumulator::new("SELECT * FROM t WHERE id = ");
acc.push_bind(42_i64);
assert_eq!(acc.sql(), "SELECT * FROM t WHERE id = $1");
assert_eq!(acc.bind_count(), 1);
acc.push_sql(" AND name = ");
acc.push_bind("alice".to_owned());
assert_eq!(acc.sql(), "SELECT * FROM t WHERE id = $1 AND name = $2");
assert_eq!(acc.bind_count(), 2);
}
#[test]
fn push_list_binds_produces_comma_separated_params() {
let mut acc = SqlAccumulator::new("SELECT * FROM t WHERE id IN (");
acc.push_list_binds([1_i32, 2, 3]);
acc.push_sql(")");
assert_eq!(acc.sql(), "SELECT * FROM t WHERE id IN ($1, $2, $3)");
assert_eq!(acc.bind_count(), 3);
}
#[test]
fn push_list_binds_empty_is_noop() {
let mut acc = SqlAccumulator::new("SELECT 1");
acc.push_list_binds(std::iter::empty::<i32>());
assert_eq!(acc.sql(), "SELECT 1");
assert_eq!(acc.bind_count(), 0);
}
#[test]
fn push_null_literal_does_not_allocate_slot() {
let mut acc = SqlAccumulator::new("SELECT * FROM t WHERE col IS ");
acc.push_null_literal();
assert_eq!(acc.sql(), "SELECT * FROM t WHERE col IS NULL");
assert_eq!(acc.bind_count(), 0);
}
#[test]
fn into_parts_returns_sql_and_binds() {
let mut acc = SqlAccumulator::new("SELECT * FROM t WHERE id = ");
acc.push_bind(99_i64);
let (sql, binds) = acc.into_parts();
assert_eq!(sql, "SELECT * FROM t WHERE id = $1");
assert_eq!(binds.len(), 1);
}
#[test]
fn extend_with_renumbers_inner_dollar_n_relative_to_outer_offset() {
let mut inner = SqlAccumulator::new("SELECT * FROM t WHERE id = ");
inner.push_bind(7_i64);
inner.push_sql(" AND name = ");
inner.push_bind("alice");
let mut outer = SqlAccumulator::new("SELECT * FROM (");
outer.push_sql("inline ");
outer.push_bind(99_i32);
outer.push_sql(", ");
outer.extend_with(inner);
outer.push_sql(") WHERE rank <= ");
outer.push_bind(3_i32);
let (sql, binds) = outer.into_parts();
assert!(
sql.starts_with("SELECT * FROM (inline $1, SELECT * FROM t WHERE id = $2"),
"got: {sql}"
);
assert!(sql.contains("AND name = $3"), "got: {sql}");
assert!(sql.ends_with("WHERE rank <= $4"), "got: {sql}");
assert_eq!(binds.len(), 4);
}
#[test]
fn extend_with_at_offset_zero_preserves_inner_dollar_numbers() {
let mut inner = SqlAccumulator::new("SELECT a = ");
inner.push_bind(1_i64);
inner.push_sql(", b = ");
inner.push_bind(2_i64);
let mut outer = SqlAccumulator::new("");
outer.extend_with(inner);
let (sql, binds) = outer.into_parts();
assert_eq!(sql, "SELECT a = $1, b = $2");
assert_eq!(binds.len(), 2);
}
#[test]
fn bind_count_tracks_push_bind_calls() {
let mut acc = SqlAccumulator::new("");
assert_eq!(acc.bind_count(), 0);
acc.push_bind(1_i32);
assert_eq!(acc.bind_count(), 1);
acc.push_bind(2_i32);
assert_eq!(acc.bind_count(), 2);
acc.push_null_literal(); assert_eq!(acc.bind_count(), 2);
}
#[test]
fn push_bind_never_leaks_into_sql_text() {
let mut acc = SqlAccumulator::new("SELECT * FROM t WHERE name = ");
acc.push_bind("'; DROP TABLE users; --".to_owned());
let sql = acc.sql();
assert!(
sql.contains("$1"),
"expected $1 placeholder in SQL, got: {sql}"
);
assert!(
!sql.contains("DROP"),
"user-supplied value leaked into SQL text: {sql}"
);
let (sql_out, binds) = acc.into_parts();
assert_eq!(sql_out, "SELECT * FROM t WHERE name = $1");
assert_eq!(binds.len(), 1);
}
#[test]
fn push_list_binds_emits_one_placeholder_per_element() {
let mut acc = SqlAccumulator::new("SELECT * FROM t WHERE id IN (");
acc.push_list_binds(["1 OR 1=1".to_owned(), "2".to_owned(), "3".to_owned()]);
acc.push_sql(")");
let sql = acc.sql();
assert_eq!(
sql, "SELECT * FROM t WHERE id IN ($1, $2, $3)",
"expected exactly three placeholders, got: {sql}"
);
assert!(
!sql.contains("OR"),
"user-supplied value leaked into SQL text: {sql}"
);
assert_eq!(acc.bind_count(), 3);
}
#[test]
fn push_null_literal_emits_sql_null_not_placeholder() {
let mut acc = SqlAccumulator::new("SELECT * FROM t WHERE col IS ");
acc.push_null_literal();
let sql = acc.sql();
assert_eq!(
sql, "SELECT * FROM t WHERE col IS NULL",
"expected literal NULL in SQL, got: {sql}"
);
assert_eq!(
acc.bind_count(),
0,
"push_null_literal must not allocate a bind slot"
);
let (_, binds) = acc.into_parts();
assert!(
binds.is_empty(),
"no bind values expected after push_null_literal"
);
}
#[test]
#[should_panic(expected = "djogi accumulator exceeded u32::MAX bind positions")]
fn push_bind_panics_on_counter_overflow_at_u32_max() {
let mut acc = SqlAccumulator::new("");
acc.next_param = u32::MAX;
acc.push_bind(1_i64);
}
#[test]
#[should_panic(expected = "djogi accumulator exceeded u32::MAX bind positions")]
fn extend_with_panics_on_placeholder_digit_parse_overflow() {
let mut inner = SqlAccumulator::new("");
inner.push_sql("$9999999999");
let mut outer = SqlAccumulator::new("");
outer.next_param = 2; outer.extend_with(inner);
}
#[test]
#[should_panic(expected = "djogi accumulator exceeded u32::MAX bind positions")]
fn extend_with_panics_on_renumber_offset_overflow() {
let mut inner = SqlAccumulator::new("");
inner.push_sql("$2");
let mut outer = SqlAccumulator::new("");
outer.next_param = u32::MAX;
outer.extend_with(inner);
}
#[test]
#[should_panic(expected = "djogi accumulator exceeded u32::MAX bind positions")]
fn extend_with_panics_on_post_splice_bind_count_overflow() {
let mut inner = SqlAccumulator::new("");
inner.push_bind(1_i64);
let mut outer = SqlAccumulator::new("");
outer.next_param = u32::MAX;
outer.extend_with(inner);
}
}