1use alloc::boxed::Box;
17use alloc::string::String;
18use alloc::vec::Vec;
19
20use spg_sql::ast::{BinOp, Expr, FromClause, Literal, SelectItem, SelectStatement, TableRef};
21use spg_storage::{Catalog, ColumnSchema, PolicyCmd, Row, TableSchema, Value};
22
23use crate::eval;
24use crate::{Engine, EngineError};
25
26#[derive(Clone, Copy, PartialEq, Eq)]
28enum QualKind {
29 Using,
31 WithCheck,
34}
35
36impl Engine {
37 pub(crate) fn rls_select_predicate(
41 &self,
42 stmt: &SelectStatement,
43 ) -> Result<Option<Expr>, EngineError> {
44 if self.is_superuser() {
45 return Ok(None);
46 }
47 let Some(from) = &stmt.from else {
48 return Ok(None);
49 };
50 let cat = self.active_catalog();
51 if !from.joins.is_empty() {
55 return Ok(None);
56 }
57 if from.primary.lateral_subquery.is_some() {
58 return Ok(None);
59 }
60 let Some(table) = cat.get(&from.primary.name) else {
61 return Ok(None);
62 };
63 if !table.schema().row_security {
64 return Ok(None);
65 }
66 Ok(Some(build_policy_predicate(
67 table.schema(),
68 self.current_role(),
69 &self.users.memberships_of_transitive(self.current_role()),
70 PolicyCmd::Select,
71 QualKind::Using,
72 )))
73 }
74
75 pub(crate) fn rls_rewrite_joins(&self, stmt: &SelectStatement) -> Option<SelectStatement> {
84 if self.is_superuser() {
85 return None;
86 }
87 let from = stmt.from.as_ref()?;
88 if from.joins.is_empty() {
89 return None;
90 }
91 let cat = self.active_catalog();
92 let needs = is_rls_base(&from.primary, cat)
93 || from.joins.iter().any(|j| is_rls_base(&j.table, cat));
94 if !needs {
95 return None;
96 }
97 let mut s = stmt.clone();
98 let from = s.from.as_mut().expect("checked above");
99 wrap_rls_table(&mut from.primary, cat);
100 for j in &mut from.joins {
101 wrap_rls_table(&mut j.table, cat);
102 }
103 Some(s)
104 }
105
106 pub(crate) fn rls_write_using_predicate(&self, table: &str, cmd: PolicyCmd) -> Option<Expr> {
110 if self.is_superuser() {
111 return None;
112 }
113 let t = self.active_catalog().get(table)?;
114 if !t.schema().row_security {
115 return None;
116 }
117 Some(build_policy_predicate(
118 t.schema(),
119 self.current_role(),
120 &self.users.memberships_of_transitive(self.current_role()),
121 cmd,
122 QualKind::Using,
123 ))
124 }
125
126 pub(crate) fn rls_check_new_rows(
131 &self,
132 table: &str,
133 cmd: PolicyCmd,
134 columns: &[ColumnSchema],
135 rows: &[Vec<Value<'static>>],
136 ) -> Result<(), EngineError> {
137 if self.is_superuser() {
138 return Ok(());
139 }
140 let Some(t) = self.active_catalog().get(table) else {
141 return Ok(());
142 };
143 if !t.schema().row_security {
144 return Ok(());
145 }
146 let pred = build_policy_predicate(
147 t.schema(),
148 self.current_role(),
149 &self.users.memberships_of_transitive(self.current_role()),
150 cmd,
151 QualKind::WithCheck,
152 );
153 let ctx = eval::EvalContext::new(columns, None);
154 for values in rows {
155 let tmp = Row {
156 values: values.clone(),
157 };
158 let v = eval::eval_expr(&pred, &tmp, &ctx).map_err(EngineError::Eval)?;
159 if !matches!(v, Value::Bool(true)) {
162 return Err(EngineError::Unsupported(alloc::format!(
163 "new row violates row-level security policy for table {table:?}"
164 )));
165 }
166 }
167 Ok(())
168 }
169}
170
171fn build_policy_predicate(
177 schema: &TableSchema,
178 role: &str,
179 member_of: &alloc::collections::BTreeSet<alloc::string::String>,
180 target_cmd: PolicyCmd,
181 kind: QualKind,
182) -> Expr {
183 let mut permissive: Vec<Expr> = Vec::new();
184 let mut restrictive: Vec<Expr> = Vec::new();
185 for p in &schema.policies {
186 if !(p.cmd == target_cmd || p.cmd == PolicyCmd::All) {
187 continue;
188 }
189 if !(p.roles.is_empty()
195 || p.roles.iter().any(|r| {
196 r.eq_ignore_ascii_case(role) || member_of.contains(&r.to_ascii_lowercase())
197 }))
198 {
199 continue;
200 }
201 let src = match kind {
202 QualKind::Using => p.using_expr.as_ref(),
203 QualKind::WithCheck => p.with_check_expr.as_ref().or(p.using_expr.as_ref()),
204 };
205 let Some(src) = src else {
206 if p.permissive {
209 permissive.push(bool_lit(true));
210 }
211 continue;
212 };
213 let term = match spg_sql::parser::parse_expression(src) {
214 Ok(mut e) => {
215 fold_session_identity(&mut e, role);
216 e
217 }
218 Err(_) => bool_lit(false), };
220 if p.permissive {
221 permissive.push(term);
222 } else {
223 restrictive.push(term);
224 }
225 }
226 if permissive.is_empty() {
227 return bool_lit(false); }
229 let mut pred = or_fold(permissive);
230 for r in restrictive {
231 pred = and(pred, r);
232 }
233 pred
234}
235
236fn fold_session_identity(e: &mut Expr, role: &str) {
241 match e {
242 Expr::FunctionCall { name, args } if args.is_empty() => {
243 match name.to_ascii_lowercase().as_str() {
244 "current_user" | "current_role" | "user" => {
245 *e = Expr::Literal(Literal::String(String::from(role)));
246 }
247 "session_user" => {
248 *e = Expr::Literal(Literal::String(String::from("admin")));
249 }
250 _ => {}
251 }
252 }
253 Expr::Binary { lhs, rhs, .. } => {
254 fold_session_identity(lhs, role);
255 fold_session_identity(rhs, role);
256 }
257 Expr::Unary { expr, .. }
258 | Expr::Cast { expr, .. }
259 | Expr::IsNull { expr, .. }
260 | Expr::FieldAccess { base: expr, .. } => fold_session_identity(expr, role),
261 Expr::FunctionCall { args, .. } => {
262 for a in args {
263 fold_session_identity(a, role);
264 }
265 }
266 Expr::Like { expr, pattern, .. } => {
267 fold_session_identity(expr, role);
268 fold_session_identity(pattern, role);
269 }
270 Expr::InList { expr, list, .. } => {
271 fold_session_identity(expr, role);
272 for it in list {
273 fold_session_identity(it, role);
274 }
275 }
276 _ => {}
277 }
278}
279
280impl Engine {
283 pub(crate) fn select_reads_policy_subject_table(&self, stmt: &SelectStatement) -> bool {
296 if self.is_superuser() {
297 return false;
298 }
299 let Some(from) = &stmt.from else {
300 return false;
301 };
302 let cat = self.active_catalog();
303 is_rls_base(&from.primary, cat) || from.joins.iter().any(|j| is_rls_base(&j.table, cat))
304 }
305}
306
307fn is_rls_base(tref: &TableRef, cat: &Catalog) -> bool {
308 tref.lateral_subquery.is_none()
309 && tref.unnest_expr.is_none()
310 && tref.generate_series_args.is_none()
311 && cat.get(&tref.name).is_some_and(|t| t.schema().row_security)
312}
313
314fn wrap_rls_table(tref: &mut TableRef, cat: &Catalog) {
317 if !is_rls_base(tref, cat) {
318 return;
319 }
320 let base = tref.name.clone();
321 let alias = tref.alias.clone().unwrap_or_else(|| base.clone());
322 let inner = SelectStatement {
323 items: alloc::vec![SelectItem::Wildcard],
324 from: Some(FromClause {
325 primary: bare_table_ref(base),
326 joins: Vec::new(),
327 }),
328 ..SelectStatement::default()
329 };
330 tref.name = alias.clone();
331 tref.alias = Some(alias);
332 tref.lateral_subquery = Some(Box::new(inner));
333}
334
335fn bare_table_ref(name: String) -> TableRef {
337 TableRef {
338 name,
339 alias: None,
340 only: false,
341 as_of_segment: None,
342 unnest_expr: None,
343 unnest_column_aliases: Vec::new(),
344 with_ordinality: false,
345 generate_series_args: None,
346 lateral_subquery: None,
347 jsonb_each_text_arg: None,
348 table_fn_call: None,
349 scalar_fn_item: false,
350 rows_from: None,
351 json_table: None,
352 }
353}
354
355fn bool_lit(b: bool) -> Expr {
356 Expr::Literal(Literal::Bool(b))
357}
358
359fn and(a: Expr, b: Expr) -> Expr {
360 Expr::Binary {
361 lhs: Box::new(a),
362 op: BinOp::And,
363 rhs: Box::new(b),
364 }
365}
366
367fn or_fold(mut terms: Vec<Expr>) -> Expr {
368 let mut acc = terms.remove(0);
369 for t in terms {
370 acc = Expr::Binary {
371 lhs: Box::new(acc),
372 op: BinOp::Or,
373 rhs: Box::new(t),
374 };
375 }
376 acc
377}