systemprompt_security/authz/repository/
entities.rs1use std::collections::HashMap;
7
8use sqlx::PgConnection;
9
10use super::AccessControlRepository;
11use crate::authz::error::AuthzResult;
12use crate::authz::types::{EntityKind, EntityRow};
13
14impl AccessControlRepository {
15 pub async fn get_entity(
16 &self,
17 entity_type: EntityKind,
18 entity_id: &str,
19 ) -> AuthzResult<Option<EntityRow>> {
20 let row = sqlx::query_as!(
21 EntityRow,
22 r#"
23 SELECT entity_type AS "kind: EntityKind", entity_id AS id, default_included, source
24 FROM access_control_entities
25 WHERE entity_type = $1 AND entity_id = $2
26 "#,
27 entity_type.as_str(),
28 entity_id,
29 )
30 .fetch_optional(&*self.pool)
31 .await?;
32 Ok(row)
33 }
34
35 pub async fn list_entities_bulk(
36 &self,
37 entity_type: EntityKind,
38 entity_ids: &[String],
39 ) -> AuthzResult<HashMap<String, EntityRow>> {
40 if entity_ids.is_empty() {
41 return Ok(HashMap::new());
42 }
43 let rows = sqlx::query_as!(
44 EntityRow,
45 r#"
46 SELECT entity_type AS "kind: EntityKind", entity_id AS id, default_included, source
47 FROM access_control_entities
48 WHERE entity_type = $1 AND entity_id = ANY($2)
49 "#,
50 entity_type.as_str(),
51 entity_ids,
52 )
53 .fetch_all(&*self.pool)
54 .await?;
55 Ok(rows.into_iter().map(|row| (row.id.clone(), row)).collect())
56 }
57
58 pub async fn upsert_entity(
59 &self,
60 entity_type: EntityKind,
61 entity_id: &str,
62 default_included: bool,
63 source: &str,
64 ) -> AuthzResult<()> {
65 sqlx::query!(
66 r#"
67 INSERT INTO access_control_entities (entity_type, entity_id, default_included, source)
68 VALUES ($1, $2, $3, $4)
69 ON CONFLICT (entity_type, entity_id) DO UPDATE
70 SET default_included = EXCLUDED.default_included,
71 source = EXCLUDED.source,
72 updated_at = NOW()
73 "#,
74 entity_type.as_str(),
75 entity_id,
76 default_included,
77 source,
78 )
79 .execute(&*self.write_pool)
80 .await?;
81 Ok(())
82 }
83
84 pub async fn ensure_entity(
85 &self,
86 entity_type: EntityKind,
87 entity_id: &str,
88 source: &str,
89 ) -> AuthzResult<()> {
90 sqlx::query!(
91 r#"
92 INSERT INTO access_control_entities (entity_type, entity_id, default_included, source)
93 VALUES ($1, $2, false, $3)
94 ON CONFLICT (entity_type, entity_id) DO NOTHING
95 "#,
96 entity_type.as_str(),
97 entity_id,
98 source,
99 )
100 .execute(&*self.write_pool)
101 .await?;
102 Ok(())
103 }
104
105 pub async fn reconcile_entities(
106 &self,
107 entity_type: EntityKind,
108 keep: &[&str],
109 default_included: bool,
110 source: &str,
111 ) -> AuthzResult<u64> {
112 let mut tx = self.write_pool.begin().await?;
113 upsert_entities_on(&mut tx, entity_type, keep, default_included, source).await?;
114 let keep_owned: Vec<String> = keep.iter().map(|id| (*id).to_owned()).collect();
115 let res = sqlx::query!(
116 r#"
117 DELETE FROM access_control_entities
118 WHERE entity_type = $1
119 AND entity_id <> ALL($2::text[])
120 "#,
121 entity_type.as_str(),
122 &keep_owned,
123 )
124 .execute(&mut *tx)
125 .await?;
126 tx.commit().await?;
127 Ok(res.rows_affected())
128 }
129
130 pub async fn list_entities(&self, entity_type: EntityKind) -> AuthzResult<Vec<EntityRow>> {
131 let rows = sqlx::query_as!(
132 EntityRow,
133 r#"
134 SELECT entity_type AS "kind: EntityKind", entity_id AS id, default_included, source
135 FROM access_control_entities
136 WHERE entity_type = $1
137 ORDER BY entity_id
138 "#,
139 entity_type.as_str(),
140 )
141 .fetch_all(&*self.pool)
142 .await?;
143 Ok(rows)
144 }
145}
146
147async fn upsert_entities_on(
148 conn: &mut PgConnection,
149 entity_type: EntityKind,
150 ids: &[&str],
151 default_included: bool,
152 source: &str,
153) -> AuthzResult<()> {
154 if ids.is_empty() {
155 return Ok(());
156 }
157 let ids_owned: Vec<String> = ids.iter().map(|id| (*id).to_owned()).collect();
158 sqlx::query!(
159 r#"
160 INSERT INTO access_control_entities (entity_type, entity_id, default_included, source)
161 SELECT $1, id, $3, $4
162 FROM UNNEST($2::text[]) AS id
163 ON CONFLICT (entity_type, entity_id) DO UPDATE
164 SET default_included = EXCLUDED.default_included,
165 source = EXCLUDED.source,
166 updated_at = NOW()
167 "#,
168 entity_type.as_str(),
169 &ids_owned,
170 default_included,
171 source,
172 )
173 .execute(conn)
174 .await?;
175 Ok(())
176}