use std::sync::Arc;
use async_trait::async_trait;
#[derive(Debug, Clone, PartialEq)]
pub enum SqlValue {
Null,
Boolean(bool),
Integer(i64),
Real(f64),
Text(String),
Blob(Vec<u8>),
Json(String),
}
#[derive(Debug, Clone, Default, PartialEq)]
pub struct SqlRows {
pub columns: Vec<String>,
pub rows: Vec<Vec<SqlValue>>,
}
#[derive(Debug, Clone, thiserror::Error)]
pub enum SqlError {
#[error("sql syntax error: {0}")]
Syntax(String),
#[error("sql constraint error: {0}")]
Constraint(String),
#[error("sql error: {0}")]
Other(String),
}
impl SqlError {
pub fn other<E: std::fmt::Display>(err: E) -> Self {
Self::Other(err.to_string())
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum Dialect {
#[default]
Sqlite,
Postgres,
Mysql,
}
#[async_trait]
pub trait SqlBackend: Send + Sync {
fn dialect(&self) -> Dialect {
Dialect::Sqlite
}
async fn begin(&self) -> Result<Box<dyn SqlTransaction>, SqlError>;
async fn begin_read_only(&self) -> Result<Box<dyn SqlTransaction>, SqlError> {
self.begin().await
}
async fn run_script(&self, _sql: &str) -> Result<(), SqlError> {
Err(SqlError::Other(
"this database does not support running a raw SQL script".into(),
))
}
async fn run_query(&self, sql: &str) -> Result<SqlRows, SqlError> {
let mut tx = self.begin_read_only().await?;
let result = tx.query(sql, &[]).await;
let _ = tx.rollback().await;
result
}
fn injects_session_context(&self) -> bool {
false
}
}
pub fn reject_reserved_session_writes(sql: &str) -> Result<(), SqlError> {
use sqlparser::dialect::GenericDialect;
use sqlparser::tokenizer::{Token, Tokenizer, Word};
const GUC_NAMESPACE: &str = "boatramp";
const MYSQL_VAR_PREFIX: &str = "@boatramp_";
let refused = || {
Err(SqlError::Other(
"setting a boatramp-reserved session key (boatramp.* / @boatramp_*) is not \
permitted from a handler: it is managed by rls_session and reserved for \
per-request tenant isolation"
.to_string(),
))
};
let dialect = GenericDialect {};
let Ok(raw) = Tokenizer::new(&dialect, sql).tokenize() else {
return refused();
};
let toks: Vec<&Token> = raw
.iter()
.filter(|t| !matches!(t, Token::Whitespace(_)))
.collect();
fn word_lc(tok: &Token) -> Option<String> {
match tok {
Token::Word(Word { value, .. }) => Some(value.to_ascii_lowercase()),
_ => None,
}
}
let is_reserved_var = |w: &str| w.starts_with(MYSQL_VAR_PREFIX);
if toks
.iter()
.any(|t| matches!(t, Token::DollarQuotedString(_)))
{
return refused();
}
if toks
.iter()
.any(|t| word_lc(t).is_some_and(|w| is_reserved_var(&w)))
{
return refused();
}
{
let leading = toks.first().and_then(|t| word_lc(t));
let has_word = |w: &str| toks.iter().any(|t| word_lc(t).as_deref() == Some(w));
let names_reserved = || {
toks.iter()
.any(|t| word_lc(t).is_some_and(|w| w == GUC_NAMESPACE || is_reserved_var(&w)))
};
match leading.as_deref() {
Some("do") | Some("call") | Some("prepare") | Some("execute") => return refused(),
Some("create") | Some("alter") if has_word("function") || has_word("procedure") => {
return refused()
}
Some("alter")
if matches!(
toks.get(1).and_then(|t| word_lc(t)).as_deref(),
Some("role") | Some("database") | Some("user") | Some("system")
) && names_reserved() =>
{
return refused()
}
_ => {}
}
}
if let Some(first) = toks.first().and_then(|t| word_lc(t)) {
match first.as_str() {
"discard" => return refused(),
"reset" => {
if let Some(target) = toks.get(1).and_then(|t| word_lc(t)) {
if target == "all" || target == GUC_NAMESPACE || is_reserved_var(&target) {
return refused();
}
}
}
"set" => {
let mut idx = 1;
if matches!(
toks.get(idx).and_then(|t| word_lc(t)).as_deref(),
Some("session") | Some("local")
) {
idx += 1;
}
if let Some(target) = toks.get(idx).and_then(|t| word_lc(t)) {
if target == GUC_NAMESPACE || is_reserved_var(&target) {
return refused();
}
}
}
_ => {}
}
}
for (i, tok) in toks.iter().enumerate() {
if word_lc(tok).as_deref() != Some("set_config") {
continue;
}
if !matches!(toks.get(i + 1), Some(Token::LParen)) {
continue;
}
let arg0 = toks.get(i + 2);
let after = toks.get(i + 3);
match (arg0, after) {
(Some(Token::SingleQuotedString(s)), Some(Token::Comma | Token::RParen)) => {
let name = s.to_ascii_lowercase();
if name == GUC_NAMESPACE || name.starts_with(&format!("{GUC_NAMESPACE}.")) {
return refused();
}
}
_ => return refused(),
}
}
Ok(())
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum PreviewSqlMode {
#[default]
Empty,
Branch,
Shared,
}
#[async_trait]
pub trait SqlBackends: Send + Sync {
async fn database(
&self,
project: &str,
site: &str,
name: &str,
) -> Result<Arc<dyn SqlBackend>, SqlError>;
async fn preview_database(
&self,
project: &str,
site: &str,
name: &str,
preview: &str,
) -> Result<Arc<dyn SqlBackend>, SqlError> {
let qualified = crate::project::ProjectRef::new(project).qualified(site);
self.database(
crate::project::DEFAULT_PROJECT,
&format!("{qualified}/_preview/{preview}"),
name,
)
.await
}
}
#[async_trait]
pub trait OperatorSql: Send + Sync {
async fn exec_script(&self, project: &str, db: &str, script: &str) -> Result<(), SqlError>;
async fn query(&self, project: &str, db: &str, sql: &str) -> Result<SqlRows, SqlError>;
async fn ping(&self, project: &str, db: &str) -> Result<Vec<SqlPingReplica>, SqlError>;
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct SqlPingReplica {
pub endpoint: String,
pub healthy: bool,
pub phase: String,
pub tcp_reachable: bool,
}
#[async_trait]
pub trait TenantDeprovisioner: Send + Sync {
async fn deprovision_project(&self, project: &str);
async fn deprovision_site(&self, project: &str, site: &str);
}
#[async_trait]
pub trait SqlTransaction: Send {
async fn query(&mut self, sql: &str, params: &[SqlValue]) -> Result<SqlRows, SqlError>;
async fn execute(&mut self, sql: &str, params: &[SqlValue]) -> Result<u64, SqlError>;
async fn commit(self: Box<Self>) -> Result<(), SqlError>;
async fn rollback(self: Box<Self>) -> Result<(), SqlError>;
}
#[cfg(test)]
mod reserved_session_writes_tests {
use super::reject_reserved_session_writes as check;
fn rejected(sql: &str) -> bool {
check(sql).is_err()
}
#[test]
fn set_config_on_reserved_guc_is_rejected() {
assert!(rejected(
"SELECT set_config('boatramp.project','victim',false)"
));
assert!(rejected("select set_config('boatramp.site', 'x', true)"));
assert!(rejected(
"SELECT set_config ( 'boatramp.project' , 'v', false )"
));
assert!(rejected(
"SELECT set_config(\"boatramp.project\", 'v', false)"
));
assert!(rejected(
"SELECT set_config('search_path','app',false), \
set_config('boatramp.project','v',false)"
));
}
#[test]
fn set_reserved_guc_is_rejected() {
assert!(rejected("SET boatramp.project = 'victim'"));
assert!(rejected("set boatramp.project='victim'")); assert!(rejected("SET SESSION boatramp.site = 'x'"));
assert!(rejected("SET LOCAL boatramp.project TO 'x'"));
}
#[test]
fn set_reserved_mysql_var_is_rejected() {
assert!(rejected("SET @boatramp_project = 'victim'"));
assert!(rejected("set @boatramp_site='x'"));
assert!(rejected("SET @boatramp_project := 'x'")); assert!(rejected("SET SESSION @boatramp_project = 'x'"));
}
#[test]
fn reset_and_discard_of_reserved_state_is_rejected() {
assert!(rejected("RESET boatramp.project"));
assert!(rejected("RESET ALL")); assert!(rejected("DISCARD ALL"));
assert!(rejected("discard all"));
}
#[test]
fn unrelated_set_statements_are_allowed() {
assert!(!rejected("SET statement_timeout = 5000"));
assert!(!rejected("SET search_path TO app, public"));
assert!(!rejected("SET SESSION time_zone = '+00:00'"));
assert!(!rejected("SET @my_var = 1")); assert!(!rejected("RESET statement_timeout"));
}
#[test]
fn a_select_mentioning_set_in_an_identifier_is_allowed() {
assert!(!rejected("SELECT settings FROM boatramp_projects"));
assert!(!rejected(
"SELECT * FROM offset_table WHERE reset_at > now()"
));
assert!(!rejected("SELECT * FROM t WHERE name = 'boatramp.project'"));
}
#[test]
fn set_config_on_a_non_reserved_guc_is_allowed() {
assert!(!rejected("SELECT set_config('search_path','app',false)"));
assert!(!rejected(
"SELECT set_config('statement_timeout', '5000', true)"
));
}
#[test]
fn inline_comment_splitting_the_keyword_is_rejected() {
assert!(rejected("SET/*x*/ boatramp.project='x'"));
assert!(rejected("set_config/*c*/('boatramp.project','x')"));
}
#[test]
fn leading_comment_before_set_is_rejected() {
assert!(rejected("/*c*/SET boatramp.project='x'"));
assert!(rejected("/* hi */ set_config('boatramp.site','x')"));
}
#[test]
fn set_config_with_concatenated_name_is_rejected() {
assert!(rejected(
"SELECT set_config('boat'||'ramp.project','x',false)"
));
assert!(rejected(
"SELECT set_config('boatramp.'||'project','x',false)"
));
}
#[test]
fn mysql_quoted_reserved_var_is_rejected() {
assert!(rejected("SET `@boatramp_project`=1"));
assert!(rejected("SET @`boatramp_project`=1"));
}
#[test]
fn case_variants_are_rejected() {
assert!(rejected("sEt boatramp.project=1"));
assert!(rejected("SeT_config('boatramp.project','x')"));
}
#[test]
fn set_config_edge_forms_are_rejected() {
assert!(rejected(
"SELECT set_config ( 'boatramp.project' , 'v', false )"
));
assert!(rejected(
"SELECT set_config('search_path','app',false), \
set_config('boatramp.project','v',false)"
));
}
#[test]
fn dollar_quoted_do_block_reserved_write_is_rejected() {
assert!(rejected(
"DO $$ BEGIN PERFORM set_config('boatramp.project','victim',false); END $$;"
));
assert!(rejected(
"DO $$ BEGIN SET boatramp.project = 'victim'; END $$;"
));
assert!(rejected(
"DO $tag$ PERFORM set_config('boatramp.project','v',false); $tag$;"
));
assert!(rejected(
"SELECT set_config($$boatramp.project$$, 'v', false)"
));
}
#[test]
fn procedural_and_persistent_constructs_are_rejected() {
assert!(rejected(
"DO 'BEGIN PERFORM set_config(''boatramp.project'',''v'',false); END'"
));
assert!(rejected("CALL do_evil()"));
assert!(rejected(
"CREATE FUNCTION e() RETURNS void AS $$ SELECT set_config('boatramp.project','v',false) $$ LANGUAGE sql"
));
assert!(rejected(
"CREATE FUNCTION e() RETURNS void AS 'BEGIN PERFORM set_config(''boatramp.project'',''v'',false); END' LANGUAGE plpgsql"
));
assert!(rejected(
"CREATE OR REPLACE PROCEDURE p() LANGUAGE sql AS $$ SELECT 1 $$"
));
assert!(rejected(
"ALTER ROLE tenant_role SET boatramp.project = 'victim'"
));
assert!(rejected("ALTER DATABASE app SET boatramp.site = 'victim'"));
}
#[test]
fn mysql_reserved_var_anywhere_is_rejected() {
assert!(rejected("SET @x=1, @boatramp_project='victim'"));
assert!(rejected("SET @a=1, @b=2, @boatramp_project='victim'"));
assert!(rejected("SET @x:=1, @boatramp_project:='victim'"));
assert!(rejected("SELECT 'victim' INTO @boatramp_project"));
assert!(rejected("SELECT 'victim' AS v INTO @boatramp_project"));
assert!(rejected("SELECT 1,'victim' INTO @junk, @boatramp_project"));
assert!(rejected("select 'victim' into @boatramp_project"));
assert!(rejected("SELECT 'v' INTO @boatramp_site"));
}
#[test]
fn prepared_statement_indirection_is_rejected() {
assert!(rejected(
"PREPARE s FROM 'SET @boatramp_project=''victim'''"
));
assert!(rejected("EXECUTE s"));
assert!(rejected(
"prepare s from 'SELECT ''v'' INTO @boatramp_site'"
));
}
#[test]
fn legit_set_and_set_config_forms_still_pass() {
assert!(!rejected("SET statement_timeout = '5s'"));
assert!(!rejected("SET search_path TO myschema"));
assert!(!rejected("SET SESSION time_zone = '+00:00'"));
assert!(!rejected("SET @my_var = 1"));
assert!(!rejected("RESET statement_timeout"));
assert!(!rejected("set_config('search_path','x',false)"));
assert!(!rejected("set_config('statement_timeout','5s',true)"));
assert!(!rejected(
"SELECT settings FROM t WHERE k = 'boatramp.project'"
));
assert!(!rejected("SELECT * FROM orders WHERE id = $1"));
assert!(!rejected("INSERT INTO orders (id, total) VALUES ($1, $2)"));
assert!(!rejected("UPDATE orders SET total = $1 WHERE id = $2"));
assert!(!rejected(
"CREATE TABLE orders (id bigint primary key, total numeric)"
));
assert!(!rejected("ALTER TABLE orders ADD COLUMN note text"));
assert!(!rejected("SET @x = 1, @y = 2"));
assert!(!rejected("SELECT 42 INTO @myvar"));
assert!(!rejected("SELECT total INTO @t FROM orders WHERE id = $1"));
}
}