Skip to main content

haste_repository/pg/models/
tenant.rs

1use crate::{
2    admin::TenantModelAdmin,
3    pg::{PGConnection, StoreError},
4    types::tenant::{CreateTenant, Tenant, TenantSearchClaims},
5    utilities::{generate_id, validate_id},
6};
7use haste_fhir_operation_error::OperationOutcomeError;
8use haste_jwt::TenantId;
9use sqlx::{PgExecutor, QueryBuilder};
10
11fn validate_tenant_customization(
12    subscription_tier: &str,
13    display_name: Option<&String>,
14    logo_data: Option<&Vec<u8>>,
15    logo_content_type: Option<&String>,
16) -> Result<(), OperationOutcomeError> {
17    if logo_data.is_some() != logo_content_type.is_some() {
18        return Err(OperationOutcomeError::error(
19            haste_fhir_model::r4::generated::terminology::IssueType::invalid(),
20            "Tenant logo data and content type must be provided together".to_string(),
21        ));
22    }
23
24    if let Some(content_type) = logo_content_type
25        && !content_type.starts_with("image/")
26    {
27        return Err(OperationOutcomeError::error(
28            haste_fhir_model::r4::generated::terminology::IssueType::invalid(),
29            "Tenant logo content type must be an image MIME type".to_string(),
30        ));
31    }
32
33    if subscription_tier == "free"
34        && (display_name.is_some() || logo_data.is_some() || logo_content_type.is_some())
35    {
36        return Err(OperationOutcomeError::error(
37            haste_fhir_model::r4::generated::terminology::IssueType::forbidden(),
38            "Tenant customization requires a paid subscription tier".to_string(),
39        ));
40    }
41
42    Ok(())
43}
44
45async fn create_tenant<'a, 'e, E>(
46    executor: E,
47    tenant: CreateTenant,
48) -> Result<Tenant, OperationOutcomeError>
49where
50    E: PgExecutor<'e>,
51{
52    let id = tenant
53        .id
54        .unwrap_or_else(|| TenantId::new(generate_id(None)));
55    validate_id(id.as_ref())?;
56
57    validate_tenant_customization(
58        tenant.subscription_tier.as_deref().unwrap_or("free"),
59        tenant.display_name.as_ref(),
60        tenant.logo_data.as_ref(),
61        tenant.logo_content_type.as_ref(),
62    )?;
63
64    let result = sqlx::query_as::<_, Tenant>(
65        r"
66            INSERT INTO tenants (id, subscription_tier, display_name, logo_data, logo_content_type)
67            VALUES ($1, $2, $3, $4, $5)
68            RETURNING id, subscription_tier, display_name, logo_data, logo_content_type
69        ",
70    )
71    .bind(id)
72    .bind(
73        tenant
74            .subscription_tier
75            .unwrap_or_else(|| "free".to_string()),
76    )
77    .bind(tenant.display_name)
78    .bind(tenant.logo_data)
79    .bind(tenant.logo_content_type)
80    .fetch_one(executor)
81    .await;
82
83    match result {
84        Ok(tenant) => Ok(tenant),
85        Err(e) => {
86            if let sqlx::Error::Database(db_error) = &e
87                && db_error.code().as_deref() == Some("23505")
88            {
89                println!("Duplicate tenant ID detected");
90                Err(StoreError::Duplicate.into())
91            } else {
92                Err(StoreError::SQLXError(e).into())
93            }
94        }
95    }
96}
97
98async fn read_tenant<'a, 'e, E>(
99    executor: E,
100    id: &'a str,
101) -> Result<Option<Tenant>, OperationOutcomeError>
102where
103    E: PgExecutor<'e>,
104{
105    let tenant = sqlx::query_as::<_, Tenant>(
106        r"
107            SELECT id, subscription_tier, display_name, logo_data, logo_content_type
108            FROM tenants
109            WHERE id = $1
110        ",
111    )
112    .bind(id)
113    .fetch_optional(executor)
114    .await
115    .map_err(StoreError::SQLXError)?;
116
117    Ok(tenant)
118}
119
120async fn update_tenant<'a, 'e, E>(
121    executor: E,
122    tenant: Tenant,
123) -> Result<Tenant, OperationOutcomeError>
124where
125    E: PgExecutor<'e>,
126{
127    validate_tenant_customization(
128        &tenant.subscription_tier,
129        tenant.display_name.as_ref(),
130        tenant.logo_data.as_ref(),
131        tenant.logo_content_type.as_ref(),
132    )?;
133
134    let updated_tenant = sqlx::query_as::<_, Tenant>(
135        r"
136            UPDATE tenants
137            SET subscription_tier = $1, display_name = $2, logo_data = $3, logo_content_type = $4
138            WHERE id = $5
139            RETURNING id, subscription_tier, display_name, logo_data, logo_content_type
140        ",
141    )
142    .bind(tenant.subscription_tier)
143    .bind(tenant.display_name)
144    .bind(tenant.logo_data)
145    .bind(tenant.logo_content_type)
146    .bind(tenant.id)
147    .fetch_one(executor)
148    .await
149    .map_err(StoreError::SQLXError)?;
150
151    Ok(updated_tenant)
152}
153
154async fn delete_tenant<'a, 'e, E>(executor: E, id: &'a str) -> Result<(), OperationOutcomeError>
155where
156    E: PgExecutor<'e>,
157{
158    sqlx::query(
159        r"
160            DELETE FROM tenants
161            WHERE id = $1
162        ",
163    )
164    .bind(id)
165    .execute(executor)
166    .await
167    .map_err(StoreError::SQLXError)?;
168
169    Ok(())
170}
171
172async fn search_tenant<'a, 'e, E>(
173    executor: E,
174    clauses: &'a TenantSearchClaims,
175) -> Result<Vec<Tenant>, OperationOutcomeError>
176where
177    E: PgExecutor<'e>,
178{
179    let mut query_builder: QueryBuilder<sqlx::Postgres> = QueryBuilder::new(
180        r"SELECT id, subscription_tier, display_name, logo_data, logo_content_type FROM tenants WHERE ",
181    );
182
183    if let Some(subscription_tier) = clauses.subscription_tier.as_ref() {
184        query_builder
185            .push(" subscription_tier = ")
186            .push_bind(subscription_tier);
187    }
188
189    let query = query_builder.build_query_as::<Tenant>();
190
191    let tenants: Vec<Tenant> = query.fetch_all(executor).await.map_err(StoreError::from)?;
192
193    Ok(tenants)
194}
195
196impl<Key: AsRef<str> + Send + Sync>
197    TenantModelAdmin<CreateTenant, Tenant, TenantSearchClaims, Tenant, Key> for PGConnection
198{
199    async fn create(
200        &self,
201        _tenant: &TenantId,
202        new_tenant: CreateTenant,
203    ) -> Result<Tenant, OperationOutcomeError> {
204        match self {
205            PGConnection::Pool(pool, _) => create_tenant(pool, new_tenant).await,
206            PGConnection::Transaction(tx, _, _) => {
207                let mut tx = tx.lock().await;
208                create_tenant(&mut **tx, new_tenant).await
209            }
210        }
211    }
212
213    async fn read(
214        &self,
215        _tenant: &TenantId,
216        id: &Key,
217    ) -> Result<Option<Tenant>, haste_fhir_operation_error::OperationOutcomeError> {
218        match self {
219            PGConnection::Pool(pool, _) => read_tenant(pool, id.as_ref()).await,
220            PGConnection::Transaction(tx, _, _) => {
221                let mut tx = tx.lock().await;
222                read_tenant(&mut **tx, id.as_ref()).await
223            }
224        }
225    }
226
227    async fn update(
228        &self,
229        _tenant: &TenantId,
230        model: Tenant,
231    ) -> Result<Tenant, haste_fhir_operation_error::OperationOutcomeError> {
232        match self {
233            PGConnection::Pool(pool, _) => update_tenant(pool, model).await,
234            PGConnection::Transaction(tx, _, _) => {
235                let mut tx = tx.lock().await;
236                update_tenant(&mut **tx, model).await
237            }
238        }
239    }
240
241    async fn delete(
242        &self,
243        _tenant: &TenantId,
244        id: &Key,
245    ) -> Result<(), haste_fhir_operation_error::OperationOutcomeError> {
246        match self {
247            PGConnection::Pool(pool, _) => delete_tenant(pool, id.as_ref()).await,
248            PGConnection::Transaction(tx, _, _) => {
249                let mut tx = tx.lock().await;
250                delete_tenant(&mut **tx, id.as_ref()).await
251            }
252        }
253    }
254
255    async fn search(
256        &self,
257        _tenant: &TenantId,
258        claims: &TenantSearchClaims,
259    ) -> Result<Vec<Tenant>, OperationOutcomeError> {
260        match self {
261            PGConnection::Pool(pool, _) => search_tenant(pool, claims).await,
262            PGConnection::Transaction(tx, _, _) => {
263                let mut tx = tx.lock().await;
264                search_tenant(&mut **tx, claims).await
265            }
266        }
267    }
268}
269
270#[cfg(test)]
271mod tests {
272    use super::validate_tenant_customization;
273
274    #[test]
275    fn free_tenants_cannot_set_branding() {
276        let result =
277            validate_tenant_customization("free", Some(&"Example Health".to_string()), None, None);
278
279        assert!(result.is_err());
280    }
281
282    #[test]
283    fn paid_tenants_can_set_branding() {
284        let result = validate_tenant_customization(
285            "professional",
286            Some(&"Example Health".to_string()),
287            Some(&vec![1, 2, 3]),
288            Some(&"image/png".to_string()),
289        );
290
291        assert!(result.is_ok());
292    }
293}