use std::collections::HashMap;
#[derive(Debug, Clone)]
pub struct RowLevelPolicy {
pub name: String,
pub table: String,
pub filter_sql: String,
pub params: Vec<String>,
}
impl RowLevelPolicy {
pub fn new(name: &str, table: &str, filter_sql: &str) -> Self {
Self {
name: name.to_string(),
table: table.to_string(),
filter_sql: filter_sql.to_string(),
params: Vec::new(),
}
}
pub fn with_param(mut self, param: &str) -> Self {
self.params.push(param.to_string());
self
}
pub fn where_clause(&self) -> String {
if self.filter_sql.is_empty() {
String::new()
} else {
format!("WHERE {}", self.filter_sql)
}
}
}
pub struct RowLevelPolicyManager {
policies: HashMap<String, Vec<RowLevelPolicy>>,
}
impl RowLevelPolicyManager {
pub fn new() -> Self {
Self {
policies: HashMap::new(),
}
}
pub fn add_policy(&mut self, policy: RowLevelPolicy) {
self.policies
.entry(policy.table.clone())
.or_default()
.push(policy);
}
pub fn get_policies(&self, table: &str) -> &[RowLevelPolicy] {
match self.policies.get(table) {
Some(policies) => policies,
None => &[],
}
}
pub fn build_where_clause(&self, table: &str) -> (String, Vec<String>) {
let policies = self.get_policies(table);
if policies.is_empty() {
return (String::new(), Vec::new());
}
let conditions: Vec<&str> = policies.iter().map(|p| p.filter_sql.as_str()).collect();
let mut params = Vec::new();
for p in policies {
params.extend(p.params.iter().cloned());
}
(format!("WHERE {}", conditions.join(" AND ")), params)
}
pub fn policy_count(&self) -> usize {
self.policies.values().map(|v| v.len()).sum()
}
}
impl Default for RowLevelPolicyManager {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_policy_where_clause() {
let policy = RowLevelPolicy::new("dept_isolation", "orders", "department_id = ?");
assert_eq!(policy.where_clause(), "WHERE department_id = ?");
}
#[test]
fn test_manager_build_where() {
let mut mgr = RowLevelPolicyManager::new();
mgr.add_policy(RowLevelPolicy::new("p1", "orders", "department_id = ?").with_param("eng"));
mgr.add_policy(RowLevelPolicy::new("p2", "orders", "is_deleted = ?").with_param("false"));
let (sql, params) = mgr.build_where_clause("orders");
assert!(sql.contains("department_id = ?"));
assert!(sql.contains("is_deleted = ?"));
assert!(sql.contains("AND"));
assert_eq!(params, vec!["eng", "false"]);
}
#[test]
fn test_manager_no_policies() {
let mgr = RowLevelPolicyManager::new();
let (sql, params) = mgr.build_where_clause("orders");
assert!(sql.is_empty());
assert!(params.is_empty());
}
#[test]
fn test_manager_multiple_tables() {
let mut mgr = RowLevelPolicyManager::new();
mgr.add_policy(RowLevelPolicy::new("p1", "orders", "dept = ?"));
mgr.add_policy(RowLevelPolicy::new("p2", "users", "tenant = ?"));
assert_eq!(mgr.get_policies("orders").len(), 1);
assert_eq!(mgr.get_policies("users").len(), 1);
assert_eq!(mgr.policy_count(), 2);
}
#[test]
fn test_empty_filter() {
let policy = RowLevelPolicy::new("empty", "orders", "");
assert!(policy.where_clause().is_empty());
}
}