Skip to main content

systemprompt_security/authz/repository/
entities.rs

1//! Entity-catalog persistence for access control.
2//!
3//! Copyright (c) systemprompt.io — Business Source License 1.1.
4//! See <https://systemprompt.io> for licensing details.
5
6use 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}