use crate::relational::catalog::{PolicyCmd, QualifiedName, Table};
use crate::sql::error::{Result, SqlError};
use crate::sql::exec::{Exec, Frame};
use crate::sql::store::RowValues;
use sqlparser::ast::{Expr, Statement, TableFactor};
use std::collections::{BTreeSet, HashMap};
const BYPASS_ROLES: &[&str] = &["service_role"];
const OWNER_ROLES: &[&str] = &["postgres", "guardian"];
pub fn role_bypasses_rls(role: &str) -> bool {
BYPASS_ROLES.contains(&role) || OWNER_ROLES.contains(&role)
}
pub fn role_bypasses_rls_on(role: &str, table: &Table) -> bool {
BYPASS_ROLES.contains(&role) || (OWNER_ROLES.contains(&role) && !table.rls_forced)
}
#[derive(Clone, Copy)]
enum Phase {
Using,
Check,
}
struct CompiledPolicy {
cmd: PolicyCmd,
permissive: bool,
using_expr: Option<Expr>,
check_expr: Option<Expr>,
}
impl CompiledPolicy {
fn expr_for(&self, phase: Phase) -> Option<&Expr> {
match phase {
Phase::Using => self.using_expr.as_ref(),
Phase::Check => self.check_expr.as_ref().or(self.using_expr.as_ref()),
}
}
}
#[derive(Default)]
pub struct RlsContext {
policies: HashMap<QualifiedName, Vec<CompiledPolicy>>,
select_hidden: HashMap<QualifiedName, BTreeSet<String>>,
dml_hidden: HashMap<QualifiedName, BTreeSet<String>>,
}
impl Exec {
pub fn init_rls(&mut self, stmt: &Statement) -> Result<()> {
if BYPASS_ROLES.contains(&self.username.as_str()) {
return Ok(());
}
let mut keys: Vec<QualifiedName> = self
.tables
.iter()
.filter(|(_, l)| l.meta.rls_enabled && !role_bypasses_rls_on(&self.username, &l.meta))
.map(|(q, _)| q.clone())
.collect();
keys.sort();
if keys.is_empty() {
return Ok(());
}
for q in &keys {
let compiled = compile_policies(&self.tables[q].meta, &self.username)?;
self.rls.policies.insert(q.clone(), compiled);
}
let dml_target: Option<(QualifiedName, PolicyCmd)> = match stmt {
Statement::Update(u) => self
.rls_target(&u.table.relation)
.map(|q| (q, PolicyCmd::Update)),
Statement::Delete(d) => {
let items = match &d.from {
sqlparser::ast::FromTable::WithFromKeyword(items)
| sqlparser::ast::FromTable::WithoutKeyword(items) => items,
};
items
.first()
.and_then(|twj| self.rls_target(&twj.relation))
.map(|q| (q, PolicyCmd::Delete))
}
_ => None,
};
for q in &keys {
let (meta, rows): (Table, Vec<(String, RowValues)>) = {
let loaded = &self.tables[q];
(
loaded.meta.clone(),
loaded
.rows
.iter()
.map(|(rid, v)| (rid.clone(), v.clone()))
.collect(),
)
};
let schema = crate::sql::dml::table_schema(&meta, &meta.name);
let dml_cmd = dml_target
.as_ref()
.filter(|(tq, _)| tq == q)
.map(|(_, cmd)| *cmd);
let mut select_hidden = BTreeSet::new();
let mut dml_hidden = BTreeSet::new();
for (rid, values) in &rows {
let tuple = crate::sql::dml::row_tuple(&meta, values);
if !self.rls_row_passes(q, PolicyCmd::Select, Phase::Using, &schema, &tuple)? {
select_hidden.insert(rid.clone());
}
if let Some(cmd) = dml_cmd
&& !self.rls_row_passes(q, cmd, Phase::Using, &schema, &tuple)?
{
dml_hidden.insert(rid.clone());
}
}
self.rls.select_hidden.insert(q.clone(), select_hidden);
if dml_cmd.is_some() {
self.rls.dml_hidden.insert(q.clone(), dml_hidden);
}
}
Ok(())
}
pub fn rls_select_hidden(&self, q: &QualifiedName) -> Option<&BTreeSet<String>> {
self.rls.select_hidden.get(q).filter(|h| !h.is_empty())
}
pub fn rls_dml_hidden(&self, q: &QualifiedName) -> Option<&BTreeSet<String>> {
self.rls.dml_hidden.get(q).filter(|h| !h.is_empty())
}
pub fn rls_check_new_row(
&self,
q: &QualifiedName,
table: &Table,
values: &RowValues,
cmd: PolicyCmd,
) -> Result<()> {
if !self.rls.policies.contains_key(q) {
return Ok(()); }
let schema = crate::sql::dml::table_schema(table, &table.name);
let tuple = crate::sql::dml::row_tuple(table, values);
if self.rls_row_passes(q, cmd, Phase::Check, &schema, &tuple)? {
Ok(())
} else {
Err(SqlError::InsufficientPrivilege(format!(
"new row violates row-level security policy for table \"{}\"",
table.name
)))
}
}
pub fn rls_old_row_visible(
&self,
q: &QualifiedName,
table: &Table,
values: &RowValues,
cmd: PolicyCmd,
) -> Result<bool> {
if !self.rls.policies.contains_key(q) {
return Ok(true);
}
let schema = crate::sql::dml::table_schema(table, &table.name);
let tuple = crate::sql::dml::row_tuple(table, values);
self.rls_row_passes(q, cmd, Phase::Using, &schema, &tuple)
}
fn rls_row_passes(
&self,
q: &QualifiedName,
cmd: PolicyCmd,
phase: Phase,
schema: &crate::sql::row::RowSchema,
tuple: &crate::sql::row::Tuple,
) -> Result<bool> {
let Some(policies) = self.rls.policies.get(q) else {
return Ok(true);
};
let mut permissive_ok = false;
for p in policies {
if p.cmd != PolicyCmd::All && p.cmd != cmd {
continue;
}
let Some(expr) = p.expr_for(phase) else {
continue;
};
let pass = if p.permissive && permissive_ok {
true } else {
self.eval(expr, &[Frame { schema, row: tuple }])?.truthy() == Some(true)
};
if p.permissive {
permissive_ok |= pass;
} else if !pass {
return Ok(false);
}
}
Ok(permissive_ok)
}
fn rls_target(&self, relation: &TableFactor) -> Option<QualifiedName> {
if let TableFactor::Table { name, .. } = relation {
let (schema, n) = crate::sql::names::split_schema_table(name);
self.catalog.resolve_table_name(schema.as_deref(), &n)
} else {
None
}
}
}
fn compile_policies(table: &Table, role: &str) -> Result<Vec<CompiledPolicy>> {
let mut out = Vec::new();
for p in &table.policies {
if !p.roles.is_empty() && !p.roles.iter().any(|r| r == role) {
continue;
}
out.push(CompiledPolicy {
cmd: p.cmd,
permissive: p.permissive,
using_expr: parse_stored(p.using_expr.as_deref(), &p.name)?,
check_expr: parse_stored(p.check_expr.as_deref(), &p.name)?,
});
}
Ok(out)
}
fn parse_stored(text: Option<&str>, policy: &str) -> Result<Option<Expr>> {
match text {
None => Ok(None),
Some(t) => crate::sql::parser::parse_expr(t).map(Some).map_err(|e| {
SqlError::Internal(format!(
"stored expression of policy \"{policy}\" failed to parse: {e}"
))
}),
}
}