uqa_sql/catalog/roles/
inquiry.rs1use super::{
10 guards::RoleCatalogGuards,
11 memberships::{
12 parse_pg_has_role_privileges, pg_has_role_privilege, resolve_pg_has_role_identifier,
13 role_privilege_text,
14 },
15 RoleReferenceNames,
16};
17use crate::catalog::roles::RoleReference;
18use crate::SQLError;
19use uqa_core::Value;
20
21pub fn pg_get_userbyid_value(
23 catalog: &dyn RoleCatalogGuards,
24 arguments: &[Value],
25) -> Result<Value, SQLError> {
26 let [argument] = arguments else {
27 return Err(SQLError::BadArity {
28 name: "pg_get_userbyid".into(),
29 expected: "1".into(),
30 actual: arguments.len(),
31 });
32 };
33 let oid = match argument {
34 Value::Null => return Ok(Value::Null),
35 Value::Int(oid) if u32::try_from(*oid).is_ok() => *oid,
36 _ => {
37 return Err(SQLError::TypeMismatch(
38 "pg_get_userbyid argument must be oid".into(),
39 ))
40 }
41 };
42 let roles = catalog.role_definitions();
43 Ok(Value::Str(
44 roles
45 .values()
46 .find(|role| role.oid == oid)
47 .map_or_else(|| format!("unknown (OID={oid})"), |role| role.name.clone()),
48 ))
49}
50
51pub fn pg_has_role_value(
52 names: &dyn RoleReferenceNames,
53 catalog: &dyn RoleCatalogGuards,
54 arguments: &[Value],
55) -> Result<Value, SQLError> {
56 if arguments.iter().any(|argument| argument == &Value::Null) {
57 return Ok(Value::Null);
58 }
59 let (subject_value, target_value, privilege_value) = match arguments {
60 [target, privilege] => (None, target, privilege),
61 [subject, target, privilege] => (Some(subject), target, privilege),
62 _ => {
63 return Err(SQLError::BadArity {
64 name: "pg_has_role".into(),
65 expected: "2 or 3".into(),
66 actual: arguments.len(),
67 });
68 }
69 };
70 let current_user = subject_value.is_none().then(|| names.current_role());
71 let roles = catalog.role_definitions();
72 let subject = subject_value.map_or_else(
73 || Ok(current_user),
74 |value| {
75 resolve_pg_has_role_identifier(value, &roles).map(|role| role.map(RoleReference::from))
76 },
77 )?;
78 let target = resolve_pg_has_role_identifier(target_value, &roles)?;
79 let privileges = parse_pg_has_role_privileges(role_privilege_text(privilege_value)?)?;
80 let memberships = catalog.role_memberships();
81 let allowed = privileges.into_iter().any(|privilege| {
82 pg_has_role_privilege(
83 &roles,
84 &memberships,
85 subject.as_ref(),
86 target.as_deref(),
87 privilege,
88 )
89 });
90 Ok(Value::Bool(allowed))
91}
92
93#[cfg(test)]
94mod tests {
95 use super::*;
96 use crate::catalog::roles::{
97 guards::{RoleDefinitionRead, RoleMembershipRead},
98 RoleDefinition,
99 };
100 use std::{cell::Cell, collections::BTreeMap};
101
102 struct Catalog {
103 roles: BTreeMap<String, RoleDefinition>,
104 reads: Cell<usize>,
105 }
106 impl RoleCatalogGuards for Catalog {
107 fn role_definitions(&self) -> RoleDefinitionRead<'_> {
108 self.reads.set(self.reads.get() + 1);
109 Box::new(&self.roles)
110 }
111 fn role_memberships(&self) -> RoleMembershipRead<'_> {
112 panic!("role names do not depend on membership privileges")
113 }
114 }
115
116 #[test]
117 fn pg_get_userbyid_uses_current_identity_without_privilege_or_null_reads() {
118 let mut role = RoleDefinition::bootstrap();
119 role.name = "selected_owner".into();
120 let mut catalog = Catalog {
121 roles: BTreeMap::from([(role.name.clone(), role)]),
122 reads: Cell::new(0),
123 };
124 assert_eq!(
125 pg_get_userbyid_value(&catalog, &[Value::Null]).unwrap(),
126 Value::Null
127 );
128 assert_eq!(catalog.reads.get(), 0);
129 for (oid, expected) in [
130 (10, "selected_owner"),
131 (0, "unknown (OID=0)"),
132 (4_294_967_295, "unknown (OID=4294967295)"),
133 ] {
134 assert_eq!(
135 pg_get_userbyid_value(&catalog, &[Value::Int(oid)]).unwrap(),
136 Value::Str(expected.into())
137 );
138 }
139 let mut renamed = catalog.roles.remove("selected_owner").unwrap();
140 renamed.name = "renamed_owner".into();
141 catalog.roles.insert(renamed.name.clone(), renamed);
142 assert_eq!(
143 pg_get_userbyid_value(&catalog, &[Value::Int(10)]).unwrap(),
144 Value::Str("renamed_owner".into())
145 );
146 catalog.roles.clear();
147 assert_eq!(
148 pg_get_userbyid_value(&catalog, &[Value::Int(10)]).unwrap(),
149 Value::Str("unknown (OID=10)".into())
150 );
151 }
152}