use std::collections::{HashMap, HashSet};
#[derive(Debug, Clone)]
pub struct AccessRule {
pub table: String,
pub row_filter: Option<String>,
pub allowed_columns: Option<HashSet<String>>,
pub denied_columns: HashSet<String>,
}
#[derive(Debug, Clone, Default)]
pub struct AccessContext {
pub tenant_id: Option<String>,
pub user_id: Option<String>,
pub roles: Vec<String>,
rules: HashMap<String, AccessRule>,
}
impl AccessContext {
pub fn new() -> Self {
Self::default()
}
pub fn with_tenant(mut self, tenant_id: impl Into<String>) -> Self {
self.tenant_id = Some(tenant_id.into());
self
}
pub fn with_user(mut self, user_id: impl Into<String>) -> Self {
self.user_id = Some(user_id.into());
self
}
pub fn add_rule(&mut self, rule: AccessRule) {
self.rules.insert(rule.table.clone(), rule);
}
pub fn row_filter(&self, table: &str) -> Option<&str> {
self.rules.get(table).and_then(|r| r.row_filter.as_deref())
}
pub fn is_column_allowed(&self, table: &str, column: &str) -> bool {
if let Some(rule) = self.rules.get(table) {
if rule.denied_columns.contains(column) {
return false;
}
if let Some(ref allowed) = rule.allowed_columns {
return allowed.contains(column);
}
}
true
}
pub fn filter_columns(&self, table: &str, columns: &[String]) -> Vec<String> {
columns
.iter()
.filter(|col| self.is_column_allowed(table, col))
.cloned()
.collect()
}
}
pub struct RowLevelSecurity {
context: AccessContext,
}
impl RowLevelSecurity {
pub fn new(context: AccessContext) -> Self {
Self { context }
}
pub fn tenant_isolation(mut self, table: &str, tenant_column: &str) -> Self {
if let Some(ref tenant_id) = self.context.tenant_id {
if crate::sql_safety::validate_identifier(table, "table").is_err() {
return self;
}
if crate::sql_safety::validate_identifier(tenant_column, "tenant_column").is_err() {
return self;
}
let escaped_id = escape_sql_literal(tenant_id);
self.context.add_rule(AccessRule {
table: table.to_string(),
row_filter: Some(format!("{} = '{}'", tenant_column, escaped_id)),
allowed_columns: None,
denied_columns: HashSet::new(),
});
}
self
}
pub fn user_isolation(mut self, table: &str, user_column: &str) -> Self {
if let Some(ref user_id) = self.context.user_id {
if crate::sql_safety::validate_identifier(table, "table").is_err() {
return self;
}
if crate::sql_safety::validate_identifier(user_column, "user_column").is_err() {
return self;
}
let escaped_id = escape_sql_literal(user_id);
self.context.add_rule(AccessRule {
table: table.to_string(),
row_filter: Some(format!("{} = '{}'", user_column, escaped_id)),
allowed_columns: None,
denied_columns: HashSet::new(),
});
}
self
}
pub fn deny_columns(mut self, table: &str, columns: &[&str]) -> Self {
let rule = self
.context
.rules
.entry(table.to_string())
.or_insert(AccessRule {
table: table.to_string(),
row_filter: None,
allowed_columns: None,
denied_columns: HashSet::new(),
});
for col in columns {
rule.denied_columns.insert(col.to_string());
}
self
}
pub fn build(self) -> AccessContext {
self.context
}
}
fn escape_sql_literal(s: &str) -> String {
let mut out = String::with_capacity(s.len() + 8);
for ch in s.chars() {
match ch {
'\'' => out.push_str("''"),
'\\' => out.push_str("\\\\"),
'\0' => out.push_str("\\0"),
'\n' => out.push_str("\\n"),
'\r' => out.push_str("\\r"),
'\x1a' => out.push_str("\\Z"),
other => out.push(other),
}
}
out
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_access_context_default() {
let ctx = AccessContext::new();
assert!(ctx.tenant_id.is_none());
assert!(ctx.user_id.is_none());
assert!(ctx.roles.is_empty());
}
#[test]
fn test_with_tenant_and_user() {
let ctx = AccessContext::new()
.with_tenant("tenant-1")
.with_user("user-1");
assert_eq!(ctx.tenant_id.as_deref(), Some("tenant-1"));
assert_eq!(ctx.user_id.as_deref(), Some("user-1"));
}
#[test]
fn test_column_allowed_by_default() {
let ctx = AccessContext::new();
assert!(ctx.is_column_allowed("users", "id"));
assert!(ctx.is_column_allowed("users", "password"));
}
#[test]
fn test_deny_columns() {
let mut ctx = AccessContext::new();
ctx.add_rule(AccessRule {
table: "users".to_string(),
row_filter: None,
allowed_columns: None,
denied_columns: ["password".to_string()].into_iter().collect(),
});
assert!(!ctx.is_column_allowed("users", "password"));
assert!(ctx.is_column_allowed("users", "id"));
}
#[test]
fn test_allowed_columns_whitelist() {
let mut ctx = AccessContext::new();
let mut allowed: HashSet<String> = HashSet::new();
allowed.insert("id".to_string());
allowed.insert("name".to_string());
ctx.add_rule(AccessRule {
table: "users".to_string(),
row_filter: None,
allowed_columns: Some(allowed),
denied_columns: HashSet::new(),
});
assert!(ctx.is_column_allowed("users", "id"));
assert!(ctx.is_column_allowed("users", "name"));
assert!(!ctx.is_column_allowed("users", "secret"));
}
#[test]
fn test_filter_columns() {
let mut ctx = AccessContext::new();
ctx.add_rule(AccessRule {
table: "users".to_string(),
row_filter: None,
allowed_columns: None,
denied_columns: ["password".to_string()].into_iter().collect(),
});
let cols = vec!["id".to_string(), "name".to_string(), "password".to_string()];
let filtered = ctx.filter_columns("users", &cols);
assert_eq!(filtered, vec!["id".to_string(), "name".to_string()]);
}
#[test]
fn test_row_filter() {
let mut ctx = AccessContext::new();
ctx.add_rule(AccessRule {
table: "orders".to_string(),
row_filter: Some("tenant_id = 't1'".to_string()),
allowed_columns: None,
denied_columns: HashSet::new(),
});
assert_eq!(ctx.row_filter("orders"), Some("tenant_id = 't1'"));
assert_eq!(ctx.row_filter("users"), None);
}
#[test]
fn test_row_level_security_tenant_isolation() {
let ctx = AccessContext::new().with_tenant("tenant-42");
let built = RowLevelSecurity::new(ctx)
.tenant_isolation("orders", "tenant_id")
.build();
assert_eq!(built.row_filter("orders"), Some("tenant_id = 'tenant-42'"));
}
#[test]
fn test_row_level_security_user_isolation() {
let ctx = AccessContext::new().with_user("u-1");
let built = RowLevelSecurity::new(ctx)
.user_isolation("profiles", "user_id")
.build();
assert_eq!(built.row_filter("profiles"), Some("user_id = 'u-1'"));
}
#[test]
fn test_row_level_security_deny_columns() {
let ctx = AccessContext::new();
let built = RowLevelSecurity::new(ctx)
.deny_columns("users", &["password", "salt"])
.build();
assert!(!built.is_column_allowed("users", "password"));
assert!(!built.is_column_allowed("users", "salt"));
assert!(built.is_column_allowed("users", "id"));
}
#[test]
fn test_tenant_isolation_skipped_without_tenant() {
let ctx = AccessContext::new();
let built = RowLevelSecurity::new(ctx)
.tenant_isolation("orders", "tenant_id")
.build();
assert_eq!(built.row_filter("orders"), None);
}
#[test]
fn test_escape_sql_literal_single_quote() {
assert_eq!(escape_sql_literal("O'Brien"), "O''Brien");
}
#[test]
fn test_escape_sql_literal_backslash() {
assert_eq!(escape_sql_literal(r"a\b"), r"a\\b");
assert_eq!(escape_sql_literal(r"\"), r"\\");
}
#[test]
fn test_escape_sql_literal_classic_injection() {
let escaped = escape_sql_literal("' OR '1'='1");
let quote_count = escaped.matches('\'').count();
assert_eq!(quote_count % 2, 0, "escaped quotes must be paired");
assert_eq!(escaped, "'' OR ''1''=''1");
}
#[test]
fn test_escape_sql_literal_mysql_backslash_injection() {
let payload = r"\' OR 1=1--";
let escaped = escape_sql_literal(payload);
assert_eq!(escaped, r"\\'' OR 1=1--");
let quote_count = escaped.matches('\'').count();
assert_eq!(quote_count % 2, 0, "escaped quotes must be paired");
}
#[test]
fn test_escape_sql_literal_null_byte() {
assert_eq!(escape_sql_literal("a\0b"), "a\\0b");
}
#[test]
fn test_escape_sql_literal_newline_carriage_return() {
assert_eq!(escape_sql_literal("a\nb\rc"), r"a\nb\rc");
}
#[test]
fn test_escape_sql_literal_ctrl_z() {
assert_eq!(escape_sql_literal("a\x1ab"), "a\\Zb");
}
#[test]
fn test_tenant_isolation_rejects_invalid_table_name() {
let ctx = AccessContext::new().with_tenant("t1");
let built = RowLevelSecurity::new(ctx)
.tenant_isolation("orders; DROP TABLE users", "tenant_id")
.build();
assert_eq!(built.row_filter("orders; DROP TABLE users"), None);
}
#[test]
fn test_tenant_isolation_rejects_invalid_column_name() {
let ctx = AccessContext::new().with_tenant("t1");
let built = RowLevelSecurity::new(ctx)
.tenant_isolation("orders", "tenant_id; DROP TABLE users")
.build();
assert_eq!(built.row_filter("orders"), None);
}
#[test]
fn test_tenant_isolation_escapes_tenant_id_injection() {
let ctx = AccessContext::new().with_tenant("' OR '1'='1");
let built = RowLevelSecurity::new(ctx)
.tenant_isolation("orders", "tenant_id")
.build();
let filter = built.row_filter("orders").unwrap();
let quote_count = filter.matches('\'').count();
assert_eq!(
quote_count % 2,
0,
"tenant_id injection not escaped: {filter}"
);
assert_eq!(filter, "tenant_id = ''' OR ''1''=''1'");
}
#[test]
fn test_user_isolation_escapes_user_id_backslash_injection() {
let ctx = AccessContext::new().with_user(r"\' OR 1=1--");
let built = RowLevelSecurity::new(ctx)
.user_isolation("profiles", "user_id")
.build();
let filter = built.row_filter("profiles").unwrap();
let quote_count = filter.matches('\'').count();
assert_eq!(
quote_count % 2,
0,
"user_id backslash injection not escaped: {filter}"
);
}
}