use deadpool_postgres::{Config, Pool, Runtime};
use serde_json::Value;
use tokio_postgres::{NoTls, Row};
use tokio_postgres::types::{FromSql, Type};
use uuid::Uuid;
use crate::models::{
Account, AccountRole, AppState, DatabaseAccess, FieldMetadata, Permission, QueryOptions,
QueryResult,
};
struct RawBytes(Vec<u8>);
impl<'a> FromSql<'a> for RawBytes {
fn from_sql(_ty: &Type, raw: &'a [u8]) -> Result<Self, Box<dyn std::error::Error + Sync + Send>> {
Ok(RawBytes(raw.to_vec()))
}
fn accepts(_ty: &Type) -> bool {
true
}
}
fn numeric_binary_to_string(raw: &[u8]) -> Option<String> {
if raw.len() < 8 {
return None;
}
let ndigits = i16::from_be_bytes([raw[0], raw[1]]) as usize;
let weight = i16::from_be_bytes([raw[2], raw[3]]);
let sign = u16::from_be_bytes([raw[4], raw[5]]);
let dscale = u16::from_be_bytes([raw[6], raw[7]]) as usize;
if raw.len() != 8 + ndigits * 2 {
return None;
}
if sign == 0xC000 {
return Some("NaN".to_string());
}
let is_negative = sign == 0x4000;
let mut digits = Vec::with_capacity(ndigits);
for i in 0..ndigits {
let offset = 8 + i * 2;
digits.push(i16::from_be_bytes([raw[offset], raw[offset + 1]]));
}
let mut result = String::new();
if is_negative {
result.push('-');
}
let int_groups = (weight + 1).max(0) as usize;
if ndigits == 0 {
result.push('0');
} else {
for i in 0..int_groups {
let d = if i < ndigits { digits[i] } else { 0 };
if i == 0 {
result.push_str(&d.to_string());
} else {
result.push_str(&format!("{:04}", d));
}
}
if int_groups == 0 {
result.push('0');
}
}
if dscale > 0 {
result.push('.');
let mut dec = String::new();
for i in int_groups..ndigits {
dec.push_str(&format!("{:04}", digits[i]));
}
while dec.len() < dscale {
dec.push('0');
}
if dec.len() > dscale {
dec.truncate(dscale);
}
result.push_str(&dec);
}
Some(result)
}
pub async fn get_or_create_pool(
state: &AppState,
pool_key: &str,
account: &Account,
db_access: &DatabaseAccess,
) -> Result<Pool, String> {
if let Some(pool) = state.connections.get(pool_key) {
return Ok(pool.clone());
}
let instances = state.instances.read().await;
let instance = instances.get(&account.instance_id)
.ok_or("Instance not found")?;
let mut config = Config::new();
config.host = Some(instance.host.clone());
config.port = Some(instance.port);
config.user = Some(db_access.username.clone());
config.password = Some(db_access.password.clone());
config.dbname = Some(db_access.database.clone());
let pool = config.create_pool(Some(Runtime::Tokio1), NoTls)
.map_err(|e| e.to_string())?;
state.connections.insert(pool_key.to_string(), pool.clone());
Ok(pool)
}
pub fn validate_query_permissions(
query: &str,
permissions: &[Permission],
role: &AccountRole,
) -> Result<(), String> {
if *role == AccountRole::Superuser {
return Ok(());
}
if query.len() > 1_000_000 {
return Err("STATEMENT_NOT_ALLOWED: query exceeds 1MB limit".to_string());
}
let code = strip_sql_comments(query);
let statements: Vec<&str> = split_statements(&code)
.into_iter()
.map(str::trim)
.filter(|s| !s.is_empty())
.collect();
if statements.len() != 1 {
return Err(
"MULTI_STATEMENT: only a single statement per request is allowed".to_string(),
);
}
let stmt = statements[0];
let blanked = blank_string_literals(stmt);
let bupper = blanked.to_uppercase();
for func in DENIED_FUNCTIONS {
if contains_word(&bupper, func) {
return Err(format!(
"STATEMENT_NOT_ALLOWED: function '{func}' is not permitted through this API"
));
}
}
let keyword = first_keyword(stmt);
if DENIED_FIRST_KEYWORDS.contains(&keyword.as_str()) {
return Err(format!(
"STATEMENT_NOT_ALLOWED: '{keyword}' is not permitted through this API"
));
}
let mut required: Vec<Permission> = Vec::new();
match keyword.as_str() {
"SELECT" | "VALUES" | "TABLE" | "SHOW" => {
required.push(Permission::Select);
}
"INSERT" => required.push(Permission::Insert),
"UPDATE" => required.push(Permission::Update),
"DELETE" => required.push(Permission::Delete),
"MERGE" => {
required.push(Permission::Insert);
required.push(Permission::Update);
required.push(Permission::Delete);
}
"CREATE" => required.push(Permission::Create),
"DROP" => required.push(Permission::Drop),
"TRUNCATE" => required.push(Permission::Truncate),
"EXPLAIN" => {
if contains_word(&bupper, "ANALYZE") {
return Err(
"STATEMENT_NOT_ALLOWED: use EXPLAIN without ANALYZE".to_string(),
);
}
required.push(Permission::Select);
}
"WITH" => {
if contains_word(&bupper, "INSERT") {
required.push(Permission::Insert);
}
if contains_word(&bupper, "UPDATE") {
required.push(Permission::Update);
}
if contains_word(&bupper, "DELETE") {
required.push(Permission::Delete);
}
if contains_word(&bupper, "MERGE") {
required.push(Permission::Insert);
required.push(Permission::Update);
required.push(Permission::Delete);
}
required.push(Permission::Select);
}
_ => {
return Err(format!(
"STATEMENT_NOT_ALLOWED: unsupported statement '{keyword}'"
));
}
}
if permissions.contains(&Permission::All) {
return Ok(());
}
let missing: Vec<String> = required
.iter()
.filter(|p| !permissions.contains(p))
.map(|p| format!("{p:?}"))
.collect();
if missing.is_empty() {
Ok(())
} else {
Err(format!("MISSING_PERMISSION: missing permissions: {}", missing.join(", ")))
}
}
const DENIED_FIRST_KEYWORDS: &[&str] = &[
"COPY", "GRANT", "REVOKE", "ALTER", "SET", "RESET", "LISTEN", "NOTIFY", "UNLISTEN", "VACUUM",
"ANALYZE", "CLUSTER", "REINDEX", "DISCARD", "CHECKPOINT", "CALL", "DO", "EXECUTE", "PREPARE",
"DEALLOCATE", "BEGIN", "START", "COMMIT", "ROLLBACK", "SAVEPOINT", "LOCK", "COMMENT",
"SECURITY", "LOAD", "IMPORT",
];
const DENIED_FUNCTIONS: &[&str] = &[
"PG_READ_FILE",
"PG_READ_BINARY_FILE",
"PG_LS_DIR",
"PG_STAT_FILE",
"DBLINK",
"DBLINK_EXEC",
"DBLINK_CONNECT",
"LO_IMPORT",
"LO_EXPORT",
"PG_EXECUTE_FROM_FILE",
];
fn strip_sql_comments(q: &str) -> String {
let bytes = q.as_bytes();
let mut out = String::with_capacity(q.len());
let mut i = 0;
let mut state = ScanState::Code;
let mut dollar_tag = String::new();
while i < bytes.len() {
let c = bytes[i] as char;
match state {
ScanState::Code => {
if c == '-' && peek(bytes, i + 1) == Some('-') {
state = ScanState::LineComment;
i += 2;
} else if c == '/' && peek(bytes, i + 1) == Some('*') {
state = ScanState::BlockComment(1);
i += 2;
} else if c == '\'' {
state = ScanState::SingleQuote;
out.push(c);
i += 1;
} else if c == '"' {
state = ScanState::DoubleQuote;
out.push(c);
i += 1;
} else if c == '$' {
if let Some(tag) = match_dollar_tag(bytes, i) {
dollar_tag = tag.clone();
state = ScanState::DollarQuote;
out.push_str(&tag);
i += tag.len();
} else {
out.push(c);
i += 1;
}
} else {
out.push(c);
i += 1;
}
}
ScanState::LineComment => {
if c == '\n' {
state = ScanState::Code;
out.push(c);
}
i += 1;
}
ScanState::BlockComment(depth) => {
if c == '/' && peek(bytes, i + 1) == Some('*') {
state = ScanState::BlockComment(depth + 1);
i += 2;
} else if c == '*' && peek(bytes, i + 1) == Some('/') {
if depth == 1 {
state = ScanState::Code;
} else {
state = ScanState::BlockComment(depth - 1);
}
i += 2;
} else {
i += 1;
}
}
ScanState::SingleQuote => {
out.push(c);
if c == '\'' {
if peek(bytes, i + 1) == Some('\'') {
out.push('\'');
i += 2;
} else {
state = ScanState::Code;
i += 1;
}
} else {
i += 1;
}
}
ScanState::DoubleQuote => {
out.push(c);
if c == '"' {
if peek(bytes, i + 1) == Some('"') {
out.push('"');
i += 2;
} else {
state = ScanState::Code;
i += 1;
}
} else {
i += 1;
}
}
ScanState::DollarQuote => {
if c == '$' {
if let Some(tag) = match_dollar_tag(bytes, i) {
if tag == dollar_tag {
out.push_str(&tag);
state = ScanState::Code;
i += tag.len();
continue;
}
}
}
out.push(c);
i += 1;
}
}
}
out
}
#[derive(Clone, Copy)]
enum ScanState {
Code,
LineComment,
BlockComment(u32),
SingleQuote,
DoubleQuote,
DollarQuote,
}
fn peek(bytes: &[u8], i: usize) -> Option<char> {
bytes.get(i).map(|b| *b as char)
}
fn match_dollar_tag(bytes: &[u8], i: usize) -> Option<String> {
let mut j = i + 1;
while j < bytes.len() {
let c = bytes[j] as char;
if c == '$' {
let tag: String = bytes[i..=j].iter().map(|b| *b as char).collect();
let inner = &tag[1..tag.len() - 1];
if inner.is_empty() {
return Some(tag);
}
if inner.bytes().all(|b| b.is_ascii_alphanumeric() || b == b'_')
&& !inner.bytes().next().is_some_and(|b| b.is_ascii_digit())
{
return Some(tag);
}
return None;
}
if !(c.is_ascii_alphanumeric() || c == '_') {
return None;
}
j += 1;
if j - i > 64 {
return None;
}
}
None
}
fn split_statements(code: &str) -> Vec<&str> {
let bytes = code.as_bytes();
let mut parts = Vec::new();
let mut start = 0;
let mut i = 0;
let mut state = ScanState::Code;
let mut dollar_tag = String::new();
while i < bytes.len() {
let c = bytes[i] as char;
match state {
ScanState::Code => {
if c == ';' {
parts.push(&code[start..i]);
start = i + 1;
i += 1;
} else if c == '\'' {
state = ScanState::SingleQuote;
i += 1;
} else if c == '"' {
state = ScanState::DoubleQuote;
i += 1;
} else if c == '$' {
if let Some(tag) = match_dollar_tag(bytes, i) {
dollar_tag = tag.clone();
state = ScanState::DollarQuote;
i += tag.len();
} else {
i += 1;
}
} else {
i += 1;
}
}
ScanState::SingleQuote => {
if c == '\'' {
if peek(bytes, i + 1) == Some('\'') {
i += 2;
} else {
state = ScanState::Code;
i += 1;
}
} else {
i += 1;
}
}
ScanState::DoubleQuote => {
if c == '"' {
if peek(bytes, i + 1) == Some('"') {
i += 2;
} else {
state = ScanState::Code;
i += 1;
}
} else {
i += 1;
}
}
ScanState::DollarQuote => {
if c == '$' {
if let Some(tag) = match_dollar_tag(bytes, i) {
if tag == dollar_tag {
state = ScanState::Code;
i += tag.len();
continue;
}
}
}
i += 1;
}
_ => {
i += 1;
}
}
}
parts.push(&code[start..]);
parts
}
fn blank_string_literals(stmt: &str) -> String {
let bytes = stmt.as_bytes();
let mut out = String::with_capacity(stmt.len());
let mut i = 0;
let mut state = ScanState::Code;
let mut dollar_tag = String::new();
while i < bytes.len() {
let c = bytes[i] as char;
match state {
ScanState::Code => {
if c == '\'' {
state = ScanState::SingleQuote;
out.push(' ');
i += 1;
} else if c == '"' {
state = ScanState::DoubleQuote;
out.push(' ');
i += 1;
} else if c == '$' {
if let Some(tag) = match_dollar_tag(bytes, i) {
dollar_tag = tag.clone();
state = ScanState::DollarQuote;
for _ in 0..tag.len() {
out.push(' ');
}
i += tag.len();
} else {
out.push(c);
i += 1;
}
} else {
out.push(c);
i += 1;
}
}
ScanState::SingleQuote | ScanState::DoubleQuote => {
let quote = if matches!(state, ScanState::SingleQuote) {
'\''
} else {
'"'
};
if c == quote {
if peek(bytes, i + 1) == Some(quote) {
out.push_str(" ");
i += 2;
} else {
out.push(' ');
state = ScanState::Code;
i += 1;
}
} else {
out.push(if c == '\n' { '\n' } else { ' ' });
i += 1;
}
}
ScanState::DollarQuote => {
if c == '$' {
if let Some(tag) = match_dollar_tag(bytes, i) {
if tag == dollar_tag {
for _ in 0..tag.len() {
out.push(' ');
}
state = ScanState::Code;
i += tag.len();
continue;
}
}
}
out.push(if c == '\n' { '\n' } else { ' ' });
i += 1;
}
_ => {
out.push(c);
i += 1;
}
}
}
out
}
fn first_keyword(stmt: &str) -> String {
let s = stmt
.trim_start()
.trim_start_matches(|c| c == '(')
.trim_start();
let end = s
.find(|c: char| !c.is_ascii_alphanumeric() && c != '_')
.unwrap_or(s.len());
s[..end].to_uppercase()
}
fn contains_word(haystack: &str, word: &str) -> bool {
let h = haystack.as_bytes();
let w = word.as_bytes();
if w.len() > h.len() {
return false;
}
let is_word = |c: u8| c.is_ascii_alphanumeric() || c == b'_' || c == b'$';
let mut i = 0;
while i + w.len() <= h.len() {
if &h[i..i + w.len()] == w {
let left_ok = i == 0 || !is_word(h[i - 1]);
let right_ok = i + w.len() == h.len() || !is_word(h[i + w.len()]);
if left_ok && right_ok {
return true;
}
}
i += 1;
}
false
}
pub fn rows_to_json(rows: &[Row]) -> Vec<Value> {
rows.iter().map(|row| {
let mut obj = serde_json::Map::new();
for (i, column) in row.columns().iter().enumerate() {
let value = match column.type_().name() {
"bool" => row.try_get::<_, bool>(i).ok().map(Value::Bool),
"int2" | "int4" => row.try_get::<_, i32>(i).ok().map(|v| Value::Number(v.into())),
"int8" => row.try_get::<_, i64>(i).ok().map(|v| Value::Number(v.into())),
"float4" => row.try_get::<_, f32>(i).ok().map(|v| Value::Number(serde_json::Number::from_f64(v as f64).unwrap_or_else(|| 0.into()))),
"float8" => row.try_get::<_, f64>(i).ok().map(|v| Value::Number(serde_json::Number::from_f64(v).unwrap_or_else(|| 0.into()))),
"text" | "varchar" | "char" | "bpchar" => row.try_get::<_, String>(i).ok().map(Value::String),
"json" | "jsonb" => row.try_get::<_, Value>(i).ok(),
"timestamp" => row.try_get::<_, chrono::NaiveDateTime>(i).ok().map(|dt| Value::String(dt.format("%Y-%m-%d %H:%M:%S").to_string())),
"timestamptz" => row.try_get::<_, chrono::DateTime<chrono::Utc>>(i).ok().map(|dt| Value::String(dt.to_rfc3339())),
"date" => row.try_get::<_, chrono::NaiveDate>(i).ok().map(|d| Value::String(d.format("%Y-%m-%d").to_string())),
"uuid" => row.try_get::<_, Uuid>(i).ok().map(|u| Value::String(u.to_string())),
"numeric" | "decimal" => {
row.try_get::<_, RawBytes>(i)
.ok()
.and_then(|raw| numeric_binary_to_string(&raw.0))
.map(Value::String)
},
"time" | "timetz" => row.try_get::<_, chrono::NaiveTime>(i).ok().map(|t| Value::String(t.format("%H:%M:%S").to_string())),
_ => row.try_get::<_, String>(i).ok().map(Value::String),
};
obj.insert(column.name().to_string(), value.unwrap_or(Value::Null));
}
Value::Object(obj)
}).collect()
}
fn json_to_sql_literal(value: &Value) -> Result<String, String> {
match value {
Value::Null => Ok("NULL".to_string()),
Value::Bool(b) => Ok(if *b { "TRUE".to_string() } else { "FALSE".to_string() }),
Value::Number(n) => {
if let Some(i) = n.as_i64() {
Ok(i.to_string())
} else if let Some(u) = n.as_u64() {
Ok(u.to_string())
} else if let Some(f) = n.as_f64() {
if !f.is_finite() {
return Err(
"STATEMENT_NOT_ALLOWED: non-finite number parameter".to_string(),
);
}
Ok(format!("{f:?}"))
} else {
Err("STATEMENT_NOT_ALLOWED: unsupported number parameter".to_string())
}
}
Value::String(s) => Ok(quote_string_literal(s)),
Value::Array(_) | Value::Object(_) => Ok(quote_string_literal(&value.to_string())),
}
}
fn quote_string_literal(s: &str) -> String {
if s.contains('\\') {
format!("E'{}'", s.replace('\\', "\\\\").replace('\'', "\\'"))
} else {
format!("'{}'", s.replace('\'', "''"))
}
}
pub fn interpolate_query(query: &str, params: &[Value]) -> Result<String, String> {
let bytes = query.as_bytes();
let mut out = String::with_capacity(query.len() + params.len() * 8);
let mut i = 0;
let mut state = ScanState::Code;
let mut dollar_tag = String::new();
while i < bytes.len() {
let c = bytes[i] as char;
match state {
ScanState::Code => {
if c == '-' && peek(bytes, i + 1) == Some('-') {
state = ScanState::LineComment;
out.push_str("--");
i += 2;
} else if c == '/' && peek(bytes, i + 1) == Some('*') {
state = ScanState::BlockComment(1);
out.push_str("/*");
i += 2;
} else if c == '\'' {
state = ScanState::SingleQuote;
out.push(c);
i += 1;
} else if c == '"' {
state = ScanState::DoubleQuote;
out.push(c);
i += 1;
} else if c == '$' {
if let Some(tag) = match_dollar_tag(bytes, i) {
dollar_tag = tag.clone();
state = ScanState::DollarQuote;
out.push_str(&tag);
i += tag.len();
} else if let Some((n, len)) = match_positional_param(bytes, i) {
if n == 0 || n > params.len() {
return Err(format!(
"STATEMENT_NOT_ALLOWED: missing value for parameter ${n}"
));
}
out.push_str(&json_to_sql_literal(¶ms[n - 1])?);
i += len;
} else {
out.push(c);
i += 1;
}
} else {
out.push(c);
i += 1;
}
}
ScanState::LineComment => {
out.push(c);
if c == '\n' {
state = ScanState::Code;
}
i += 1;
}
ScanState::BlockComment(depth) => {
if c == '/' && peek(bytes, i + 1) == Some('*') {
state = ScanState::BlockComment(depth + 1);
out.push_str("/*");
i += 2;
} else if c == '*' && peek(bytes, i + 1) == Some('/') {
if depth == 1 {
state = ScanState::Code;
} else {
state = ScanState::BlockComment(depth - 1);
}
out.push_str("*/");
i += 2;
} else {
out.push(c);
i += 1;
}
}
ScanState::SingleQuote => {
out.push(c);
if c == '\'' {
if peek(bytes, i + 1) == Some('\'') {
out.push('\'');
i += 2;
} else {
state = ScanState::Code;
i += 1;
}
} else {
i += 1;
}
}
ScanState::DoubleQuote => {
out.push(c);
if c == '"' {
if peek(bytes, i + 1) == Some('"') {
out.push('"');
i += 2;
} else {
state = ScanState::Code;
i += 1;
}
} else {
i += 1;
}
}
ScanState::DollarQuote => {
if c == '$' {
if let Some(tag) = match_dollar_tag(bytes, i) {
if tag == dollar_tag {
out.push_str(&tag);
state = ScanState::Code;
i += tag.len();
continue;
}
}
}
out.push(c);
i += 1;
}
}
}
Ok(out)
}
fn match_positional_param(bytes: &[u8], i: usize) -> Option<(usize, usize)> {
let mut j = i + 1;
while j < bytes.len() && (bytes[j] as char).is_ascii_digit() {
j += 1;
}
if j == i + 1 {
return None;
}
let n: usize = std::str::from_utf8(&bytes[i + 1..j]).ok()?.parse().ok()?;
Some((n, j - i))
}
const DEFAULT_STATEMENT_TIMEOUT_MS: u64 = 60_000;
const MAX_STATEMENT_TIMEOUT_MS: u64 = 300_000;
fn is_plain_write(query: &str) -> bool {
let blanked = blank_string_literals(query);
let upper = blanked.to_uppercase();
if !matches!(
first_keyword(&upper).as_str(),
"INSERT" | "UPDATE" | "DELETE" | "MERGE"
) {
return false;
}
!contains_word(&upper, "RETURNING")
}
pub async fn execute_query_with_pool(
pool: Pool,
query: String,
params: Vec<Value>,
options: &QueryOptions,
) -> Result<QueryResult, String> {
let client = pool.get().await.map_err(|e| e.to_string())?;
let timeout_ms = options
.timeout_ms
.unwrap_or(DEFAULT_STATEMENT_TIMEOUT_MS)
.clamp(1, MAX_STATEMENT_TIMEOUT_MS);
client
.batch_execute(if options.read_only {
"BEGIN TRANSACTION READ ONLY"
} else {
"BEGIN"
})
.await
.map_err(|e| e.to_string())?;
client
.batch_execute(&format!("SET LOCAL statement_timeout = '{timeout_ms}ms'"))
.await
.map_err(|e| e.to_string())?;
let final_query = interpolate_query(&query, ¶ms)?;
let (rows, affected) = if is_plain_write(&query) {
match client.execute(&final_query, &[]).await {
Ok(n) => (vec![], Some(n)),
Err(e) => {
let e = e.to_string();
let _ = client.batch_execute("ROLLBACK").await;
if e.contains("57014")
|| e.contains("canceling statement due to statement timeout")
{
return Err(
"STATEMENT_TIMEOUT: query exceeded its time budget".to_string(),
);
}
return Err(e);
}
}
} else {
match client.query(&final_query, &[]).await {
Ok(rows) => (rows, None),
Err(e) => {
let e = e.to_string();
let _ = client.batch_execute("ROLLBACK").await;
if e.contains("57014") || e.contains("canceling statement due to statement timeout") {
return Err("STATEMENT_TIMEOUT: query exceeded its time budget".to_string());
}
return Err(e);
}
}
};
client
.batch_execute("COMMIT")
.await
.map_err(|e| e.to_string())?;
let mut result = rows_to_result(rows);
result.rows_affected = affected;
Ok(result)
}
pub async fn execute_transaction_with_pool(
pool: Pool,
queries: Vec<(String, Vec<Value>, QueryOptions)>,
) -> Result<Vec<QueryResult>, String> {
if queries.is_empty() {
return Err("STATEMENT_NOT_ALLOWED: empty transaction".to_string());
}
if queries.len() > MAX_BATCH_ITEMS {
return Err(format!(
"STATEMENT_NOT_ALLOWED: transaction exceeds {MAX_BATCH_ITEMS} statements"
));
}
let mut client = pool.get().await.map_err(|e| e.to_string())?;
let txn = client
.transaction()
.await
.map_err(|e| e.to_string())?;
let mut results = Vec::with_capacity(queries.len());
for (query, params, options) in &queries {
let timeout_ms = options
.timeout_ms
.unwrap_or(DEFAULT_STATEMENT_TIMEOUT_MS)
.clamp(1, MAX_STATEMENT_TIMEOUT_MS);
txn.batch_execute(&format!("SET LOCAL statement_timeout = '{timeout_ms}ms'"))
.await
.map_err(|e| e.to_string())?;
let final_query = interpolate_query(query, params)?;
if is_plain_write(query) {
match txn.execute(&final_query, &[]).await {
Ok(n) => {
let mut r = rows_to_result(vec![]);
r.rows_affected = Some(n);
results.push(r);
}
Err(e) => {
let msg = e.to_string();
txn.rollback().await.map_err(|rb| rb.to_string())?;
if msg.contains("57014")
|| msg.contains("canceling statement due to statement timeout")
{
return Err(
"STATEMENT_TIMEOUT: query exceeded its time budget".to_string(),
);
}
return Err(msg);
}
}
continue;
}
let rows = txn.query(&final_query, &[]).await;
match rows {
Ok(rows) => results.push(rows_to_result(rows)),
Err(e) => {
let msg = e.to_string();
txn.rollback().await.map_err(|rb| rb.to_string())?;
if msg.contains("57014")
|| msg.contains("canceling statement due to statement timeout")
{
return Err("STATEMENT_TIMEOUT: query exceeded its time budget".to_string());
}
return Err(msg);
}
}
}
txn.commit().await.map_err(|e| e.to_string())?;
Ok(results)
}
pub const MAX_BATCH_ITEMS: usize = 50;
fn rows_to_result(rows: Vec<Row>) -> QueryResult {
let fields = if !rows.is_empty() {
rows[0].columns().iter().map(|col| {
FieldMetadata {
name: col.name().to_string(),
data_type: col.type_().name().to_string(),
nullable: true,
max_length: None,
}
}).collect()
} else {
vec![]
};
let json_rows = rows_to_json(&rows);
QueryResult {
rows: json_rows,
fields,
query_plan: None,
rows_affected: None,
}
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
#[test]
fn test_is_plain_write() {
assert!(is_plain_write("UPDATE t SET a = $1 WHERE id = $2"));
assert!(is_plain_write("INSERT INTO t (a) VALUES ($1)"));
assert!(is_plain_write("DELETE FROM t WHERE id = $1"));
assert!(is_plain_write(" update t set a=null where id=$1 "));
assert!(!is_plain_write(
"UPDATE t SET a = $1 WHERE id = $2 RETURNING id"
));
assert!(!is_plain_write("DELETE FROM t WHERE id = $1 RETURNING *"));
assert!(!is_plain_write("SELECT * FROM t WHERE a = $1"));
assert!(!is_plain_write(
"WITH x AS (SELECT 1) SELECT * FROM x"
));
assert!(is_plain_write(
"UPDATE t SET a = 'RETURNING' WHERE id = $1"
));
assert!(!is_plain_write("SELECT 'RETURNING'"));
}
#[test]
fn test_validate_query_permissions_select() {
let permissions = vec![Permission::Select];
let role = AccountRole::Developer;
let result = validate_query_permissions("SELECT * FROM users", &permissions, &role);
assert!(result.is_ok());
}
#[test]
fn test_validate_query_permissions_insert_denied() {
let permissions = vec![Permission::Select];
let role = AccountRole::Developer;
let result = validate_query_permissions("INSERT INTO users VALUES (1)", &permissions, &role);
assert!(result.is_err());
assert!(result.unwrap_err().contains("MISSING_PERMISSION"));
}
#[test]
fn test_validate_query_permissions_owner_bypass() {
let permissions = vec![];
let role = AccountRole::Superuser;
let result = validate_query_permissions("DROP TABLE users", &permissions, &role);
assert!(result.is_ok());
}
#[test]
fn test_validate_query_permissions_all_permission() {
let permissions = vec![Permission::All];
let role = AccountRole::Developer;
let result = validate_query_permissions("DELETE FROM users", &permissions, &role);
assert!(result.is_ok());
}
#[test]
fn test_validate_query_permissions_update() {
let permissions = vec![Permission::Update, Permission::Select];
let role = AccountRole::Developer;
let result = validate_query_permissions("UPDATE users SET name = 'test'", &permissions, &role);
assert!(result.is_ok());
}
#[test]
fn test_validate_query_permissions_delete() {
let permissions = vec![Permission::Delete];
let role = AccountRole::Developer;
let result = validate_query_permissions("DELETE FROM users WHERE id = 1", &permissions, &role);
assert!(result.is_ok());
}
#[test]
fn test_validate_query_permissions_create() {
let permissions = vec![Permission::Create];
let role = AccountRole::Developer;
let result = validate_query_permissions("CREATE TABLE test (id INT)", &permissions, &role);
assert!(result.is_ok());
}
#[test]
fn test_validate_query_permissions_drop() {
let permissions = vec![Permission::Drop];
let role = AccountRole::Developer;
let result = validate_query_permissions("DROP TABLE test", &permissions, &role);
assert!(result.is_ok());
}
#[test]
fn test_validate_query_permissions_truncate() {
let permissions = vec![Permission::Truncate];
let role = AccountRole::Developer;
let result = validate_query_permissions("TRUNCATE TABLE test", &permissions, &role);
assert!(result.is_ok());
}
#[test]
fn test_validate_query_permissions_unsupported() {
let permissions = vec![Permission::Select];
let role = AccountRole::Developer;
let result = validate_query_permissions("ALTER TABLE users ADD COLUMN x INT", &permissions, &role);
assert!(result.is_err());
assert!(result.unwrap_err().contains("STATEMENT_NOT_ALLOWED"));
}
#[test]
fn test_validate_query_permissions_case_insensitive() {
let permissions = vec![Permission::Select];
let role = AccountRole::Developer;
let result = validate_query_permissions("select * from users", &permissions, &role);
assert!(result.is_ok());
}
#[test]
fn test_validate_query_permissions_with_cte() {
let permissions = vec![Permission::Select];
let role = AccountRole::Developer;
let result = validate_query_permissions("WITH cte AS (SELECT * FROM users) SELECT * FROM cte", &permissions, &role);
assert!(result.is_ok());
}
#[test]
fn test_interpolate_basic_types() {
assert_eq!(
interpolate_query("SELECT $1, $2, $3, $4", &[json!(42), json!("o'clock"), json!(true), Value::Null]).unwrap(),
"SELECT 42, 'o''clock', TRUE, NULL"
);
}
#[test]
fn test_interpolate_backslash_uses_escape_string() {
assert_eq!(
interpolate_query("SELECT $1", &[json!("a\\b'c")]).unwrap(),
"SELECT E'a\\\\b\\'c'"
);
}
#[test]
fn test_interpolate_ignores_dollar_in_literals() {
assert_eq!(
interpolate_query("SELECT '$1', 'R$ 5', $1", &[json!(7)]).unwrap(),
"SELECT '$1', 'R$ 5', 7"
);
}
#[test]
fn test_interpolate_injection_attempt_stays_literal() {
let evil = "x'; DROP TABLE users; --";
let out = interpolate_query("SELECT * FROM t WHERE name = $1", &[json!(evil)]).unwrap();
assert_eq!(out, "SELECT * FROM t WHERE name = 'x''; DROP TABLE users; --'");
}
#[test]
fn test_interpolate_extra_params_ignored() {
assert_eq!(
interpolate_query("SELECT 1", &[json!(1), json!(2)]).unwrap(),
"SELECT 1"
);
}
#[test]
fn test_interpolate_missing_param_errors() {
let r = interpolate_query("SELECT $1, $3", &[json!(1), json!(2)]);
assert!(r.is_err());
assert!(r.unwrap_err().contains("missing value for parameter $3"));
}
#[test]
fn test_interpolate_repeated_and_multi_digit() {
assert_eq!(
interpolate_query("SELECT $1, $1, $10", &[
json!(1), json!(2), json!(3), json!(4), json!(5),
json!(6), json!(7), json!(8), json!(9), json!("dez"),
]).unwrap(),
"SELECT 1, 1, 'dez'"
);
}
#[test]
fn test_interpolate_json_values() {
assert_eq!(
interpolate_query("SELECT $1", &[json!({"a": 1})]).unwrap(),
"SELECT '{\"a\":1}'"
);
}
#[test]
fn test_interpolate_big_u64() {
let big = serde_json::Number::from(u64::MAX);
assert_eq!(
interpolate_query("SELECT $1", &[Value::Number(big)]).unwrap(),
format!("SELECT {}", u64::MAX)
);
}
#[test]
fn test_numeric_binary_to_string_zero() {
let raw = vec![0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00];
let result = numeric_binary_to_string(&raw);
assert_eq!(result, Some("0".to_string()));
}
#[test]
fn test_numeric_binary_to_string_nan() {
let raw = vec![0x00, 0x00, 0x00, 0x00, 0xC0, 0x00, 0x00, 0x00];
let result = numeric_binary_to_string(&raw);
assert_eq!(result, Some("NaN".to_string()));
}
#[test]
fn test_numeric_binary_to_string_invalid_short() {
let raw = vec![0x00, 0x00, 0x00];
let result = numeric_binary_to_string(&raw);
assert_eq!(result, None);
}
fn perms(p: &[Permission]) -> Vec<Permission> {
p.to_vec()
}
#[test]
fn test_rejects_stacked_statements() {
let role = AccountRole::Developer;
let all = perms(&[Permission::All]);
for q in [
"SELECT 1; DROP TABLE users",
"SELECT 1 ; SELECT 2",
"WITH a AS (SELECT 1) SELECT * FROM a; DELETE FROM users",
] {
let r = validate_query_permissions(q, &all, &role);
assert!(r.is_err(), "{q} should be rejected");
assert!(r.unwrap_err().contains("MULTI_STATEMENT"), "{q}");
}
}
#[test]
fn test_semicolon_inside_literal_is_single_statement() {
let role = AccountRole::Developer;
let sel = perms(&[Permission::Select]);
assert!(validate_query_permissions("SELECT ';' AS s", &sel, &role).is_ok());
assert!(validate_query_permissions("SELECT $1", &sel, &role).is_ok());
}
#[test]
fn test_comment_smuggling_blocked() {
let role = AccountRole::Developer;
let sel = perms(&[Permission::Select]);
assert!(validate_query_permissions("/* SELECT */ DELETE FROM users", &sel, &role).is_err());
assert!(validate_query_permissions("-- SELECT\nDELETE FROM users", &sel, &role).is_err());
assert!(validate_query_permissions("SELECT 1 /* trailing */", &sel, &role).is_ok());
}
#[test]
fn test_with_writes_require_write_permission() {
let role = AccountRole::Developer;
let sel = perms(&[Permission::Select]);
let del = perms(&[Permission::Select, Permission::Delete]);
let q = "WITH moved AS (DELETE FROM a RETURNING *) SELECT * FROM moved";
assert!(validate_query_permissions(q, &sel, &role).is_err());
assert!(validate_query_permissions(q, &del, &role).is_ok());
assert!(validate_query_permissions(
"WITH cte AS (SELECT * FROM users) SELECT * FROM cte",
&sel, &role
).is_ok());
}
#[test]
fn test_row_lock_query_is_allowed_for_queue_claims() {
let role = AccountRole::Developer;
let permissions = perms(&[Permission::Select]);
let query = "SELECT id FROM queue WHERE status = 'pending' ORDER BY id FOR UPDATE SKIP LOCKED LIMIT 1";
assert!(validate_query_permissions(query, &permissions, &role).is_ok());
}
#[test]
fn test_denied_statements_and_functions() {
let role = AccountRole::Developer;
let all = perms(&[Permission::All]);
for q in [
"COPY users TO STDOUT",
"GRANT SELECT ON t TO u",
"VACUUM users",
"CALL my_proc()",
"DO $$ BEGIN RAISE NOTICE 'hi'; END $$",
"EXPLAIN ANALYZE SELECT * FROM users",
"BEGIN",
"SET ROLE postgres",
"SELECT pg_read_file('pg_hba.conf')",
"SELECT dblink('c', 'SELECT 1')",
"SELECT lo_import('/etc/passwd')",
] {
let r = validate_query_permissions(q, &all, &role);
assert!(r.is_err(), "{q} should be denied even with ALL");
assert!(r.unwrap_err().contains("STATEMENT_NOT_ALLOWED"), "{q}");
}
}
#[test]
fn test_dangerous_words_as_values_are_allowed() {
let role = AccountRole::Developer;
let w = perms(&[Permission::Select, Permission::Insert, Permission::Update]);
assert!(validate_query_permissions(
"SELECT * FROM posts WHERE comment = 'DROP TABLE x'", &w, &role).is_ok());
assert!(validate_query_permissions("UPDATE users SET name = 'x'", &w, &role).is_ok());
assert!(validate_query_permissions("SELECT set_config('a','b',false)", &w, &role).is_ok());
}
#[test]
fn test_explain_without_analyze_allowed() {
let role = AccountRole::Developer;
let sel = perms(&[Permission::Select]);
assert!(validate_query_permissions("EXPLAIN SELECT * FROM users", &sel, &role).is_ok());
}
#[test]
fn test_superuser_bypass_kept() {
let role = AccountRole::Superuser;
assert!(validate_query_permissions("DROP TABLE users", &[], &role).is_ok());
}
}