use super::config::TenantConfig;
use super::context::TenantContext;
use super::strategy::ColumnType;
use super::task_local;
use crate::error::{QueryError, QueryResult};
use crate::middleware::{BoxFuture, Middleware, Next, QueryContext, QueryResponse, QueryType};
use std::sync::{Arc, RwLock};
pub struct TenantMiddleware {
config: TenantConfig,
current_tenant: Arc<RwLock<Option<TenantContext>>>,
}
impl TenantMiddleware {
pub fn new(config: TenantConfig) -> Self {
Self {
config,
current_tenant: Arc::new(RwLock::new(None)),
}
}
pub fn set_tenant(&self, ctx: TenantContext) {
*self.current_tenant.write().expect("lock poisoned") = Some(ctx);
}
pub fn clear_tenant(&self) {
*self.current_tenant.write().expect("lock poisoned") = None;
}
pub fn current_tenant(&self) -> Option<TenantContext> {
self.current_tenant.read().expect("lock poisoned").clone()
}
pub fn scoped(&self, ctx: TenantContext) -> TenantScope {
self.set_tenant(ctx);
TenantScope {
middleware: Arc::new(self.clone()),
}
}
fn apply_row_level_filter(&self, sql: &str, tenant_id: &str) -> QueryResult<String> {
let config = match self.config.row_level_config() {
Some(c) => c,
None => return Ok(sql.to_string()),
};
let tenant_value = validated_tenant_value(&config.column, config.column_type, tenant_id)?;
self.inject_tenant_filter(sql, &config.column, &tenant_value)
}
fn inject_tenant_filter(&self, sql: &str, column: &str, value: &str) -> QueryResult<String> {
let filter = format!("{} = {}", column, value);
let body = sql.trim();
if body.is_empty() {
return Ok(sql.to_string());
}
match classify_statement(body) {
shape @ (StatementShape::Select | StatementShape::Write) => {
let (body, semi) = match body.strip_suffix(';') {
Some(b) => (b.trim_end(), ";"),
None => (body, ""),
};
let rewritten = if shape == StatementShape::Select {
inject_where_filter(
body,
&filter,
&[
"GROUP BY", "HAVING", "WINDOW", "ORDER BY", "LIMIT", "OFFSET", "FETCH",
],
&["UNION", "INTERSECT", "EXCEPT", "FOR"],
)?
} else {
inject_where_filter(body, &filter, &["RETURNING"], &[])?
};
Ok(format!("{}{}", rewritten, semi))
}
StatementShape::Insert => {
if self
.config
.row_level_config()
.is_some_and(|c| c.auto_insert)
{
inject_insert_column(body, column, value)
} else {
Ok(sql.to_string())
}
}
StatementShape::WithSelect => Err(QueryError::invalid_input(
"sql",
"cannot safely apply a tenant filter to a CTE-wrapped SELECT statement",
)),
StatementShape::Other => Err(QueryError::invalid_input(
"sql",
"cannot safely apply a tenant filter: unrecognized statement shape",
)),
}
}
fn apply_schema_isolation(&self, tenant_id: &str) -> Option<String> {
self.config
.schema_config()
.map(|c| c.search_path(tenant_id))
}
}
impl Clone for TenantMiddleware {
fn clone(&self) -> Self {
Self {
config: self.config.clone(),
current_tenant: Arc::clone(&self.current_tenant),
}
}
}
impl std::fmt::Debug for TenantMiddleware {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("TenantMiddleware")
.field("config", &self.config)
.field("has_tenant", &self.current_tenant().is_some())
.finish()
}
}
impl Middleware for TenantMiddleware {
fn handle<'a>(
&'a self,
mut ctx: QueryContext,
next: Next<'a>,
) -> BoxFuture<'a, QueryResult<QueryResponse>> {
Box::pin(async move {
let tenant_ctx = match task_local::current_tenant().or_else(|| self.current_tenant()) {
Some(ctx) => ctx,
None => {
if self.config.require_tenant {
if let Some(default) = &self.config.default_tenant {
TenantContext::new(default.clone())
} else {
return Err(QueryError::internal(
"Tenant context required but not provided",
));
}
} else {
return next.run(ctx).await;
}
}
};
if self.config.allow_bypass && tenant_ctx.should_bypass() {
if self.config.log_tenant_context {
tracing::debug!(
tenant_id = %tenant_ctx.id,
bypass = true,
"Tenant filter bypassed"
);
}
return next.run(ctx).await;
}
if self.config.strategy.is_row_level() {
let query_type = ctx.query_type();
let modified_sql =
self.apply_row_level_filter(ctx.sql(), tenant_ctx.id.as_str())?;
ctx = ctx.with_sql(modified_sql);
let is_write = matches!(query_type, QueryType::Update | QueryType::Delete)
|| classify_statement(ctx.sql()) == StatementShape::Write;
if self.config.enforce_on_writes
&& is_write
&& let Some(row_level) = self.config.row_level_config()
&& row_level.validate_writes
&& !references_identifier(ctx.sql(), &row_level.column)
{
return Err(QueryError::invalid_input(
&row_level.column,
"UPDATE/DELETE does not reference the tenant column",
));
}
}
if self.config.strategy.is_schema_based()
&& let Some(search_path) = self.apply_schema_isolation(tenant_ctx.id.as_str())
{
ctx.metadata_mut().set_schema_override(Some(
self.config
.schema_config()
.unwrap()
.schema_name(tenant_ctx.id.as_str()),
));
if self.config.log_tenant_context {
tracing::debug!(
tenant_id = %tenant_ctx.id,
search_path = %search_path,
"Setting schema for tenant"
);
}
}
if self.config.log_tenant_context {
tracing::debug!(
tenant_id = %tenant_ctx.id,
strategy = ?self.config.strategy,
sql = %ctx.sql(),
"Executing query with tenant context"
);
}
ctx.metadata_mut().tenant_id = Some(tenant_ctx.id.to_string());
next.run(ctx).await
})
}
fn name(&self) -> &'static str {
"TenantMiddleware"
}
}
fn validated_tenant_value(
column: &str,
column_type: ColumnType,
tenant_id: &str,
) -> QueryResult<String> {
column_type.try_format_value(tenant_id).map_err(|_| {
let expectation = match column_type {
ColumnType::String => "a string tenant id matching ^[A-Za-z0-9_:-.@]+$",
ColumnType::Uuid => "a valid UUID",
ColumnType::Integer | ColumnType::BigInt => "a valid integer",
};
QueryError::invalid_input(
column,
format!("tenant id is not {expectation}: {tenant_id:?}"),
)
})
}
fn inject_where_filter(
sql: &str,
filter: &str,
terminators: &[&str],
forbidden: &[&str],
) -> QueryResult<String> {
if find_top_level_keyword(sql, forbidden).is_some() {
return Err(QueryError::invalid_input(
"sql",
"cannot safely apply a tenant filter to a compound or locking statement",
));
}
let where_hit = find_top_level_keyword(sql, &["WHERE"]);
let terminator_hit = find_top_level_keyword(sql, terminators);
match (where_hit, terminator_hit) {
(Some((where_start, _)), Some((term_start, _))) if term_start < where_start => {
Err(QueryError::invalid_input(
"sql",
"cannot safely apply a tenant filter: unrecognized clause ordering",
))
}
(Some((_, where_end)), terminator) => {
let close = terminator.map_or(sql.len(), |(start, _)| start);
let existing = sql[where_end..close].trim();
if existing.is_empty() {
return Err(QueryError::invalid_input(
"sql",
"cannot safely apply a tenant filter: empty WHERE clause",
));
}
Ok(format!(
"{} {} AND ({}) {}",
sql[..where_end].trim_end(),
filter,
existing,
sql[close..].trim_start()
))
}
(None, Some((term_start, _))) => Ok(format!(
"{} WHERE {} {}",
sql[..term_start].trim_end(),
filter,
sql[term_start..].trim_start()
)),
(None, None) => Ok(format!("{} WHERE {}", sql.trim_end(), filter)),
}
}
fn inject_insert_column(sql: &str, column: &str, value: &str) -> QueryResult<String> {
let reject = || {
QueryError::invalid_input(
"sql",
"cannot safely auto-insert the tenant column into this INSERT statement",
)
};
let Some((_, into_end)) = find_top_level_keyword(sql, &["INTO"]) else {
return Err(reject());
};
let name_start = skip_ws(sql, into_end);
let Some(name_end) = parse_table_name(sql, name_start) else {
return Err(reject());
};
let cols_open = skip_ws(sql, name_end);
if sql.as_bytes().get(cols_open) != Some(&b'(') {
return Err(reject());
}
let Some(cols_close) = find_matching_paren(sql, cols_open) else {
return Err(reject());
};
let columns = split_top_level_commas(&sql[cols_open + 1..cols_close]);
if columns.iter().all(|c| c.trim().is_empty()) {
return Err(reject());
}
let tenant_column_index = columns
.iter()
.position(|c| unquote_ident(c.trim()).eq_ignore_ascii_case(column));
let after_cols = skip_ws(sql, cols_close + 1);
let Some((_, values_end)) = match_keyword_at(sql, after_cols, &["VALUES"]) else {
return Err(reject());
};
let vals_open = skip_ws(sql, values_end);
if sql.as_bytes().get(vals_open) != Some(&b'(') {
return Err(reject());
}
let Some(vals_close) = find_matching_paren(sql, vals_open) else {
return Err(reject());
};
let values = split_top_level_commas(&sql[vals_open + 1..vals_close]);
if values.len() != columns.len() {
return Err(reject());
}
let after_vals = skip_ws(sql, vals_close + 1);
if sql.as_bytes().get(after_vals) == Some(&b',') {
return Err(reject());
}
if let Some(index) = tenant_column_index {
if !sql_literal_eq(values[index], value) {
return Err(QueryError::invalid_input(
column,
"INSERT supplies a tenant column value that does not match the current tenant",
));
}
return Ok(sql.to_string());
}
Ok(format!(
"{}, {}{}, {}{}",
&sql[..cols_close],
column,
&sql[cols_close..vals_close],
value,
&sql[vals_close..]
))
}
fn sql_literal_eq(a: &str, b: &str) -> bool {
a.split_ascii_whitespace().collect::<Vec<_>>().join(" ")
== b.split_ascii_whitespace().collect::<Vec<_>>().join(" ")
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum StatementShape {
Select,
Write,
Insert,
WithSelect,
Other,
}
fn classify_statement(sql: &str) -> StatementShape {
let pos = skip_ws(sql, 0);
if match_keyword_at(sql, pos, &["SELECT"]).is_some() {
StatementShape::Select
} else if match_keyword_at(sql, pos, &["UPDATE", "DELETE"]).is_some() {
StatementShape::Write
} else if match_keyword_at(sql, pos, &["INSERT"]).is_some() {
StatementShape::Insert
} else if let Some((_, with_end)) = match_keyword_at(sql, pos, &["WITH"]) {
match cte_body_offset(sql, with_end) {
Some(body) if match_keyword_at(sql, body, &["SELECT"]).is_some() => {
StatementShape::WithSelect
}
Some(body) if match_keyword_at(sql, body, &["UPDATE", "DELETE"]).is_some() => {
StatementShape::Write
}
_ => StatementShape::Other,
}
} else {
StatementShape::Other
}
}
fn cte_body_offset(sql: &str, mut pos: usize) -> Option<usize> {
let bytes = sql.as_bytes();
pos = skip_ws(sql, pos);
if let Some((_, end)) = match_keyword_at(sql, pos, &["RECURSIVE"]) {
pos = end;
}
loop {
pos = skip_ws(sql, pos);
match bytes.get(pos) {
Some(b'"') | Some(b'`') | Some(b'[') => pos = skip_quoted(sql, pos),
Some(b) if b.is_ascii_alphabetic() || *b == b'_' => {
pos += 1;
while pos < bytes.len() && is_ident_byte(bytes[pos]) {
pos += 1;
}
}
_ => return None,
}
pos = skip_ws(sql, pos);
if bytes.get(pos) == Some(&b'(') {
pos = find_matching_paren(sql, pos)? + 1;
pos = skip_ws(sql, pos);
}
let (_, as_end) = match_keyword_at(sql, pos, &["AS"])?;
pos = skip_ws(sql, as_end);
if bytes.get(pos) != Some(&b'(') {
return None;
}
pos = find_matching_paren(sql, pos)? + 1;
pos = skip_ws(sql, pos);
match bytes.get(pos) {
Some(&b',') => pos += 1,
_ => return Some(pos),
}
}
}
fn is_ident_byte(b: u8) -> bool {
b.is_ascii_alphanumeric() || b == b'_' || b == b'$'
}
fn skip_comment(sql: &str, pos: usize) -> Option<usize> {
let bytes = sql.as_bytes();
if pos + 1 < bytes.len() && bytes[pos] == b'-' && bytes[pos + 1] == b'-' {
let mut i = pos + 2;
while i < bytes.len() && bytes[i] != b'\n' {
i += 1;
}
Some(i)
} else if pos + 1 < bytes.len() && bytes[pos] == b'/' && bytes[pos + 1] == b'*' {
let mut depth = 1usize;
let mut i = pos + 2;
while i + 1 < bytes.len() && depth > 0 {
if bytes[i] == b'/' && bytes[i + 1] == b'*' {
depth += 1;
i += 2;
} else if bytes[i] == b'*' && bytes[i + 1] == b'/' {
depth -= 1;
i += 2;
} else {
i += 1;
}
}
Some(i.min(bytes.len()))
} else {
None
}
}
fn skip_ws(sql: &str, mut pos: usize) -> usize {
let bytes = sql.as_bytes();
loop {
while pos < bytes.len() && bytes[pos].is_ascii_whitespace() {
pos += 1;
}
match skip_comment(sql, pos) {
Some(next) => pos = next,
None => return pos,
}
}
}
fn skip_quoted(sql: &str, pos: usize) -> usize {
let bytes = sql.as_bytes();
let close = match bytes[pos] {
b'[' => b']',
quote => quote,
};
let mut i = pos + 1;
while i < bytes.len() {
if bytes[i] == close {
if i + 1 < bytes.len() && bytes[i + 1] == close {
i += 2; continue;
}
return i + 1;
}
i += 1;
}
bytes.len()
}
fn skip_dollar_quoted(sql: &str, pos: usize) -> Option<usize> {
let bytes = sql.as_bytes();
if bytes.get(pos) != Some(&b'$') {
return None;
}
if bytes.get(pos + 1).is_some_and(u8::is_ascii_digit) {
return None;
}
let mut i = pos + 1;
while i < bytes.len() && (bytes[i].is_ascii_alphanumeric() || bytes[i] == b'_') {
i += 1;
}
if bytes.get(i) != Some(&b'$') {
return None;
}
let tag = &sql[pos..=i];
match sql[i + 1..].find(tag) {
Some(off) => Some(i + 1 + off + tag.len()),
None => Some(bytes.len()),
}
}
fn match_keyword_at<'k>(sql: &str, pos: usize, keywords: &[&'k str]) -> Option<(&'k str, usize)> {
let bytes = sql.as_bytes();
'keywords: for keyword in keywords {
let mut p = pos;
for (index, word) in keyword.split_ascii_whitespace().enumerate() {
if index > 0 {
p = skip_ws(sql, p);
}
let end = p + word.len();
if end > bytes.len() || !bytes[p..end].eq_ignore_ascii_case(word.as_bytes()) {
continue 'keywords;
}
if end < bytes.len() && is_ident_byte(bytes[end]) {
continue 'keywords;
}
p = end;
}
return Some((keyword, p));
}
None
}
fn find_top_level_keyword(sql: &str, keywords: &[&str]) -> Option<(usize, usize)> {
if keywords.is_empty() {
return None;
}
let bytes = sql.as_bytes();
let mut depth = 0usize;
let mut i = 0usize;
while i < bytes.len() {
let b = bytes[i];
if matches!(b, b'\'' | b'"' | b'`' | b'[') {
i = skip_quoted(sql, i);
} else if b == b'$' {
i = skip_dollar_quoted(sql, i).unwrap_or(i + 1);
} else if let Some(next) = skip_comment(sql, i) {
i = next;
} else if b == b'(' {
depth += 1;
i += 1;
} else if b == b')' {
depth = depth.saturating_sub(1);
i += 1;
} else if depth == 0 && b.is_ascii_alphabetic() && (i == 0 || !is_ident_byte(bytes[i - 1]))
{
if let Some((_, end)) = match_keyword_at(sql, i, keywords) {
return Some((i, end));
}
i += 1;
} else {
i += 1;
}
}
None
}
fn find_matching_paren(sql: &str, open: usize) -> Option<usize> {
let bytes = sql.as_bytes();
let mut depth = 0usize;
let mut i = open;
while i < bytes.len() {
let b = bytes[i];
if matches!(b, b'\'' | b'"' | b'`' | b'[') {
i = skip_quoted(sql, i);
} else if b == b'$' {
i = skip_dollar_quoted(sql, i).unwrap_or(i + 1);
} else if let Some(next) = skip_comment(sql, i) {
i = next;
} else if b == b'(' {
depth += 1;
i += 1;
} else if b == b')' {
depth -= 1;
if depth == 0 {
return Some(i);
}
i += 1;
} else {
i += 1;
}
}
None
}
fn split_top_level_commas(list: &str) -> Vec<&str> {
let bytes = list.as_bytes();
let mut parts = Vec::new();
let mut depth = 0usize;
let mut start = 0usize;
let mut i = 0usize;
while i < bytes.len() {
let b = bytes[i];
if matches!(b, b'\'' | b'"' | b'`' | b'[') {
i = skip_quoted(list, i);
} else if b == b'$' {
i = skip_dollar_quoted(list, i).unwrap_or(i + 1);
} else if let Some(next) = skip_comment(list, i) {
i = next;
} else if b == b'(' {
depth += 1;
i += 1;
} else if b == b')' {
depth = depth.saturating_sub(1);
i += 1;
} else if b == b',' && depth == 0 {
parts.push(&list[start..i]);
i += 1;
start = i;
} else {
i += 1;
}
}
parts.push(&list[start..]);
parts
}
fn parse_table_name(sql: &str, pos: usize) -> Option<usize> {
let bytes = sql.as_bytes();
let mut i = pos;
loop {
i = skip_ws(sql, i);
match bytes.get(i) {
Some(b'"') | Some(b'`') | Some(b'[') => i = skip_quoted(sql, i),
Some(b) if b.is_ascii_alphabetic() || *b == b'_' => {
while i < bytes.len() && is_ident_byte(bytes[i]) {
i += 1;
}
}
_ => return None,
}
let next = skip_ws(sql, i);
if bytes.get(next) == Some(&b'.') {
i = next + 1;
} else {
return Some(i);
}
}
}
fn unquote_ident(ident: &str) -> &str {
let bytes = ident.as_bytes();
if ident.len() >= 2 {
let (first, last) = (bytes[0], bytes[ident.len() - 1]);
if (first == b'"' && last == b'"')
|| (first == b'`' && last == b'`')
|| (first == b'[' && last == b']')
{
return &ident[1..ident.len() - 1];
}
}
ident
}
fn references_identifier(sql: &str, ident: &str) -> bool {
let bytes = sql.as_bytes();
let mut i = 0usize;
while i < bytes.len() {
let b = bytes[i];
if b == b'\'' {
i = skip_quoted(sql, i);
} else if b == b'"' || b == b'`' || b == b'[' {
let end = skip_quoted(sql, i);
if end > i + 2 && bytes[i + 1..end - 1].eq_ignore_ascii_case(ident.as_bytes()) {
return true;
}
i = end;
} else if b == b'$' {
i = skip_dollar_quoted(sql, i).unwrap_or(i + 1);
} else if let Some(next) = skip_comment(sql, i) {
i = next;
} else if b.is_ascii_alphabetic() || b == b'_' {
let start = i;
while i < bytes.len() && is_ident_byte(bytes[i]) {
i += 1;
}
if bytes[start..i].eq_ignore_ascii_case(ident.as_bytes()) {
return true;
}
} else {
i += 1;
}
}
false
}
pub struct TenantScope {
middleware: Arc<TenantMiddleware>,
}
impl Drop for TenantScope {
fn drop(&mut self) {
self.middleware.clear_tenant();
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_row_level_filter_select() {
let config = TenantConfig::row_level("tenant_id");
let middleware = TenantMiddleware::new(config);
let sql = middleware
.apply_row_level_filter("SELECT * FROM users", "tenant-123")
.unwrap();
assert!(sql.contains("WHERE tenant_id = 'tenant-123'"));
let sql = middleware
.apply_row_level_filter("SELECT * FROM users WHERE active = true", "tenant-123")
.unwrap();
assert!(sql.contains("tenant_id = 'tenant-123' AND (active = true)"));
}
#[test]
fn test_row_level_filter_update() {
let config = TenantConfig::row_level("tenant_id");
let middleware = TenantMiddleware::new(config);
let sql = middleware
.apply_row_level_filter("UPDATE users SET name = 'Bob'", "tenant-123")
.unwrap();
assert!(sql.contains("WHERE tenant_id = 'tenant-123'"));
let sql = middleware
.apply_row_level_filter("UPDATE users SET name = 'Bob' WHERE id = 1", "tenant-123")
.unwrap();
assert!(sql.contains("tenant_id = 'tenant-123' AND (id = 1)"));
}
#[test]
fn test_row_level_filter_delete() {
let config = TenantConfig::row_level("tenant_id");
let middleware = TenantMiddleware::new(config);
let sql = middleware
.apply_row_level_filter("DELETE FROM users", "tenant-123")
.unwrap();
assert!(sql.contains("WHERE tenant_id = 'tenant-123'"));
}
#[test]
fn test_tenant_scope() {
let config = TenantConfig::row_level("tenant_id");
let middleware = TenantMiddleware::new(config);
{
let _scope = middleware.scoped(TenantContext::new("tenant-123"));
assert!(middleware.current_tenant().is_some());
assert_eq!(
middleware.current_tenant().unwrap().id.as_str(),
"tenant-123"
);
}
assert!(middleware.current_tenant().is_none());
}
#[test]
fn test_integer_tenant_injection_rejected() {
use super::super::strategy::{IsolationStrategy, RowLevelConfig};
let mut config = TenantConfig::row_level("tenant_id");
config.strategy = IsolationStrategy::RowLevel(
RowLevelConfig::new("tenant_id").with_column_type(ColumnType::Integer),
);
let middleware = TenantMiddleware::new(config);
assert!(
middleware
.apply_row_level_filter("SELECT * FROM users", "1 OR true--")
.is_err()
);
let sql = middleware
.apply_row_level_filter("SELECT * FROM users", "42")
.unwrap();
assert!(sql.contains("WHERE tenant_id = 42"));
}
#[test]
fn test_or_predicate_is_parenthesized() {
let config = TenantConfig::row_level("tenant_id");
let middleware = TenantMiddleware::new(config);
let sql = middleware
.apply_row_level_filter(
"SELECT * FROM users WHERE active = true OR admin = true",
"tenant-123",
)
.unwrap();
assert!(sql.contains("tenant_id = 'tenant-123' AND (active = true OR admin = true)"));
}
#[test]
fn test_group_by_where_placement() {
let config = TenantConfig::row_level("tenant_id");
let middleware = TenantMiddleware::new(config);
let sql = middleware
.apply_row_level_filter(
"SELECT role, COUNT(*) FROM users GROUP BY role",
"tenant-123",
)
.unwrap();
let where_pos = sql.find("WHERE").expect("WHERE injected");
let group_pos = sql.find("GROUP BY").expect("GROUP BY preserved");
assert!(where_pos < group_pos, "unexpected SQL: {sql}");
let sql = middleware
.apply_row_level_filter(
"SELECT role, COUNT(*) FROM users WHERE active = true GROUP BY role",
"tenant-123",
)
.unwrap();
assert!(sql.contains("tenant_id = 'tenant-123' AND (active = true) GROUP BY role"));
}
#[test]
fn test_string_tenant_whitelist_enforced() {
let config = TenantConfig::row_level("tenant_id");
let middleware = TenantMiddleware::new(config);
for bad in [
"' OR 1=1-- ",
"\\' OR 1=1-- ",
"tenant'; DROP TABLE users--",
"a b",
"",
] {
assert!(
middleware
.apply_row_level_filter("SELECT * FROM users", bad)
.is_err(),
"tenant id {bad:?} must be rejected"
);
}
for good in ["tenant-123", "a_b-c:d.e@f", "Tenant_01"] {
let sql = middleware
.apply_row_level_filter("SELECT * FROM users", good)
.unwrap();
assert!(
sql.contains(&format!("tenant_id = '{good}'")),
"unexpected SQL: {sql}"
);
}
}
#[test]
fn test_cte_and_comment_prefixed_statements() {
let config = TenantConfig::row_level("tenant_id");
let middleware = TenantMiddleware::new(config);
let err = middleware
.apply_row_level_filter(
"WITH active_users AS (SELECT * FROM users) SELECT * FROM active_users",
"tenant-123",
)
.unwrap_err();
assert!(
err.to_string().contains("CTE-wrapped SELECT"),
"unexpected error: {err}"
);
for sql in [
"-- list users\nSELECT * FROM users",
"/* audit */ SELECT * FROM users",
] {
let rewritten = middleware
.apply_row_level_filter(sql, "tenant-123")
.unwrap();
assert!(
rewritten.contains("WHERE tenant_id = 'tenant-123'"),
"unexpected SQL: {rewritten}"
);
}
let sql = middleware
.apply_row_level_filter(
"WITH t AS (SELECT 1) UPDATE users SET name = 'Bob'",
"tenant-123",
)
.unwrap();
assert!(
sql.contains("WHERE tenant_id = 'tenant-123'"),
"unexpected SQL: {sql}"
);
}
#[test]
fn test_unrecognized_statements_fail_closed() {
let config = TenantConfig::row_level("tenant_id");
let middleware = TenantMiddleware::new(config);
for sql in [
"REPLACE INTO users (id) VALUES (1)",
"MERGE INTO users USING src ON users.id = src.id WHEN MATCHED THEN UPDATE SET name = src.name",
"INSERT INTO users SELECT * FROM staging",
"VACUUM users",
] {
assert!(
middleware
.apply_row_level_filter(sql, "tenant-123")
.is_err(),
"statement must be rejected: {sql}"
);
}
}
#[test]
fn test_insert_existing_tenant_column_must_match() {
let config = TenantConfig::row_level("tenant_id");
let middleware = TenantMiddleware::new(config);
for sql in [
"INSERT INTO users (name, tenant_id) VALUES ('Bob', 'tenant-123')",
"INSERT INTO users (name, tenant_id) VALUES ('Bob', 'tenant-123' )",
] {
let rewritten = middleware
.apply_row_level_filter(sql, "tenant-123")
.unwrap();
assert_eq!(rewritten, sql);
}
assert!(
middleware
.apply_row_level_filter(
"INSERT INTO users (name, tenant_id) VALUES ('Bob', 'B')",
"tenant-123",
)
.is_err()
);
assert!(
middleware
.apply_row_level_filter(
"INSERT INTO users (name, tenant_id) VALUES ('Bob', 'tenant-123' OR '1'='1')",
"tenant-123",
)
.is_err()
);
}
#[test]
fn test_insert_auto_injects_tenant_column() {
let config = TenantConfig::row_level("tenant_id");
let middleware = TenantMiddleware::new(config);
let sql = middleware
.apply_row_level_filter("INSERT INTO users (name) VALUES ('Bob')", "tenant-123")
.unwrap();
assert!(sql.contains("tenant_id"), "unexpected SQL: {sql}");
assert!(sql.contains("'tenant-123'"), "unexpected SQL: {sql}");
}
#[test]
fn test_bracket_doubled_escape_scans_as_one_identifier() {
assert_eq!(skip_quoted("[a]]b]", 0), "[a]]b]".len());
let config = TenantConfig::row_level("tenant_id");
let middleware = TenantMiddleware::new(config);
let sql = middleware
.apply_row_level_filter("SELECT * FROM [a]] WHERE b]", "tenant-123")
.unwrap();
assert!(
sql.ends_with("WHERE tenant_id = 'tenant-123'"),
"unexpected SQL: {sql}"
);
}
#[test]
fn test_dollar_quoted_strings_are_not_scanned() {
let config = TenantConfig::row_level("tenant_id");
let middleware = TenantMiddleware::new(config);
let sql = middleware
.apply_row_level_filter(
"SELECT * FROM users WHERE note = $tag$x GROUP BY y$tag$",
"tenant-123",
)
.unwrap();
assert!(
sql.contains("tenant_id = 'tenant-123' AND (note = $tag$x GROUP BY y$tag$)"),
"unexpected SQL: {sql}"
);
let sql = middleware
.apply_row_level_filter("SELECT $$a LIMIT b$$ FROM users LIMIT 5", "tenant-123")
.unwrap();
assert_eq!(
sql,
"SELECT $$a LIMIT b$$ FROM users WHERE tenant_id = 'tenant-123' LIMIT 5"
);
let sql = middleware
.apply_row_level_filter("SELECT $$x WHERE y$$ AS v FROM users", "tenant-123")
.unwrap();
assert_eq!(
sql,
"SELECT $$x WHERE y$$ AS v FROM users WHERE tenant_id = 'tenant-123'"
);
let sql = middleware
.apply_row_level_filter("SELECT * FROM users WHERE id = $1", "tenant-123")
.unwrap();
assert!(
sql.contains("tenant_id = 'tenant-123' AND (id = $1)"),
"unexpected SQL: {sql}"
);
}
#[test]
fn test_nested_block_comments_are_skipped() {
let config = TenantConfig::row_level("tenant_id");
let middleware = TenantMiddleware::new(config);
let sql = middleware
.apply_row_level_filter(
"SELECT * FROM users /* outer /* inner */ WHERE x */",
"tenant-123",
)
.unwrap();
assert_eq!(
sql,
"SELECT * FROM users /* outer /* inner */ WHERE x */ WHERE tenant_id = 'tenant-123'"
);
}
#[test]
fn test_window_and_fetch_terminate_where() {
let config = TenantConfig::row_level("tenant_id");
let middleware = TenantMiddleware::new(config);
let sql = middleware
.apply_row_level_filter(
"SELECT row_number() OVER w FROM users WINDOW w AS (ORDER BY id)",
"tenant-123",
)
.unwrap();
assert_eq!(
sql,
"SELECT row_number() OVER w FROM users WHERE tenant_id = 'tenant-123' \
WINDOW w AS (ORDER BY id)"
);
let sql = middleware
.apply_row_level_filter("SELECT * FROM users FETCH FIRST 10 ROWS ONLY", "tenant-123")
.unwrap();
assert_eq!(
sql,
"SELECT * FROM users WHERE tenant_id = 'tenant-123' FETCH FIRST 10 ROWS ONLY"
);
}
#[tokio::test]
async fn test_task_local_tenant_resolution() {
let config = TenantConfig::row_level("tenant_id");
let middleware = TenantMiddleware::new(config);
middleware.set_tenant(TenantContext::new("slot-tenant"));
task_local::with_tenant("task-tenant", async {
let response = middleware
.handle(
QueryContext::new("SELECT * FROM users", vec![]),
echo_next(),
)
.await
.unwrap();
let sql = response.data["sql"].as_str().unwrap();
assert!(
sql.contains("tenant_id = 'task-tenant'"),
"unexpected SQL: {sql}"
);
})
.await;
let response = middleware
.handle(
QueryContext::new("SELECT * FROM users", vec![]),
echo_next(),
)
.await
.unwrap();
let sql = response.data["sql"].as_str().unwrap();
assert!(
sql.contains("tenant_id = 'slot-tenant'"),
"unexpected SQL: {sql}"
);
}
fn echo_next<'a>() -> Next<'a> {
Next {
inner: Box::new(|ctx: QueryContext| {
let sql = ctx.sql().to_string();
Box::pin(async move {
Ok::<QueryResponse, QueryError>(QueryResponse::new(
serde_json::json!({ "sql": sql }),
))
})
}),
}
}
}