Skip to main content

apiplant_db/
lib.rs

1//! # apiplant-db
2//!
3//! The database layer. It has two jobs:
4//!
5//! * **Migrations** ([`migrate`]) — make Postgres match the resource schemas.
6//! * **CRUD** ([`Db`]) — build parameterised statements for a [`Resource`] at
7//!   runtime and hand rows back as plain JSON.
8//!
9//! Rows come back as JSON by letting Postgres do the conversion (`to_jsonb` /
10//! `jsonb_agg`), so the executor only ever extracts a single JSON column and
11//! never needs a compile-time entity for a table it only learned about from a
12//! TOML file. Values always travel as `$n` bind parameters; only validated,
13//! double-quoted identifiers are ever interpolated.
14
15mod ident;
16pub mod migrate;
17pub mod value;
18
19use apiplant_core::Resource;
20use sea_orm::sea_query::Value as SqlValue;
21use sea_orm::{
22    ConnectOptions, ConnectionTrait, Database, DatabaseBackend, DatabaseConnection, Statement,
23};
24use uuid::Uuid;
25
26use ident::quote_ident;
27pub use migrate::migrate;
28
29/// Database errors.
30#[derive(thiserror::Error, Debug)]
31pub enum Error {
32    #[error("database: {0}")]
33    Db(#[from] sea_orm::DbErr),
34    #[error("schema: {0}")]
35    Schema(String),
36    #[error("bad input: {0}")]
37    BadInput(String),
38}
39
40/// An extra predicate applied to a query: equality (owner/org scoping,
41/// `?field=` filters) or membership (`id IN (…)`, e.g. "organisations you belong
42/// to"). Column names are always validated and quoted; values are always bound.
43#[derive(Clone)]
44pub enum Filter {
45    /// `column = value`.
46    Eq { column: String, value: SqlValue },
47    /// `column IN (values…)`. An empty set matches no rows.
48    In {
49        column: String,
50        values: Vec<SqlValue>,
51    },
52    /// `column ILIKE '%value%'` — a case-insensitive substring match, which is
53    /// what a search box means by "search". The pattern's own wildcards are
54    /// escaped, so a term containing `%` looks for a per-cent sign.
55    Contains { column: String, value: String },
56}
57
58impl Filter {
59    pub fn eq(column: impl Into<String>, value: impl Into<SqlValue>) -> Self {
60        Filter::Eq {
61            column: column.into(),
62            value: value.into(),
63        }
64    }
65
66    pub fn in_(column: impl Into<String>, values: Vec<SqlValue>) -> Self {
67        Filter::In {
68            column: column.into(),
69            values,
70        }
71    }
72
73    /// Convenience: `column IN (…uuids)`.
74    pub fn in_uuids(column: impl Into<String>, ids: Vec<Uuid>) -> Self {
75        Filter::In {
76            column: column.into(),
77            values: ids.into_iter().map(SqlValue::from).collect(),
78        }
79    }
80
81    pub fn contains(column: impl Into<String>, value: impl Into<String>) -> Self {
82        Filter::Contains {
83            column: column.into(),
84            value: value.into(),
85        }
86    }
87
88    fn column(&self) -> &str {
89        match self {
90            Filter::Eq { column, .. }
91            | Filter::In { column, .. }
92            | Filter::Contains { column, .. } => column,
93        }
94    }
95}
96
97/// A connection pool plus the dynamic CRUD executor.
98#[derive(Clone)]
99pub struct Db {
100    conn: DatabaseConnection,
101}
102
103impl Db {
104    /// Open a pool against the given Postgres URL, creating the database first
105    /// if it does not exist yet.
106    ///
107    /// A fresh checkout pointed at a running Postgres would otherwise fail with
108    /// `database "…" does not exist` before migrations ever get a chance to
109    /// run, so on that specific error we connect to the `postgres` maintenance
110    /// database on the same server, `CREATE DATABASE`, and retry once. Any
111    /// other failure (bad credentials, no server) is returned untouched.
112    pub async fn connect(url: &str, max_connections: u32) -> Result<Self, Error> {
113        match Self::open(url, max_connections).await {
114            Ok(db) => Ok(db),
115            Err(err) if is_missing_database(&err) => {
116                let Some((admin_url, name)) = maintenance_url(url) else {
117                    return Err(err);
118                };
119                tracing::info!("database `{name}` does not exist; creating it");
120                let admin = Self::open(&admin_url, 1).await?;
121                // Another worker starting at the same time may win the race and
122                // create it first, which is fine: what matters is whether the
123                // database is there on the retry, so a failed CREATE is only
124                // reported if the retry also fails.
125                let created = admin
126                    .raw_json(&format!("CREATE DATABASE {}", quote_ident(&name)?), &[])
127                    .await;
128                match (Self::open(url, max_connections).await, created) {
129                    (Ok(db), _) => Ok(db),
130                    (Err(_), Err(create_err)) => Err(create_err),
131                    (Err(open_err), Ok(_)) => Err(open_err),
132                }
133            }
134            Err(err) => Err(err),
135        }
136    }
137
138    async fn open(url: &str, max_connections: u32) -> Result<Self, Error> {
139        let mut opt = ConnectOptions::new(url.to_owned());
140        opt.max_connections(max_connections).sqlx_logging(false);
141        let conn = Database::connect(opt).await?;
142        Ok(Db { conn })
143    }
144
145    /// Access the underlying connection (used by [`migrate`]).
146    pub fn connection(&self) -> &DatabaseConnection {
147        &self.conn
148    }
149
150    // --- CRUD -------------------------------------------------------------
151
152    /// `GET /<resource>` — a JSON array of rows, newest first.
153    pub async fn list(
154        &self,
155        r: &Resource,
156        filters: &[Filter],
157        limit: i64,
158        offset: i64,
159    ) -> Result<serde_json::Value, Error> {
160        let table = quote_ident(&r.table_name())?;
161        let (where_sql, mut params, n) = self.build_where(filters)?;
162        let order = if r.meta.timestamps {
163            "ORDER BY created_at DESC"
164        } else {
165            ""
166        };
167        let limit_ph = format!("${}", n);
168        let offset_ph = format!("${}", n + 1);
169        params.push(SqlValue::from(limit));
170        params.push(SqlValue::from(offset));
171
172        let hidden = self.hidden_subtraction(r)?;
173        let sql = format!(
174            "SELECT coalesce(jsonb_agg(to_jsonb(t){hidden}), '[]'::jsonb) AS result \
175             FROM (SELECT * FROM {table} {where_sql} {order} LIMIT {limit_ph} OFFSET {offset_ph}) t"
176        );
177        let row = self
178            .conn
179            .query_one(Statement::from_sql_and_values(
180                DatabaseBackend::Postgres,
181                sql,
182                params,
183            ))
184            .await?
185            .ok_or_else(|| Error::Db(sea_orm::DbErr::Custom("no aggregate row".into())))?;
186        Ok(row.try_get::<serde_json::Value>("", "result")?)
187    }
188
189    /// `GET /<resource>/<id>` — one row or `None`.
190    pub async fn get(
191        &self,
192        r: &Resource,
193        id: Uuid,
194        filters: &[Filter],
195    ) -> Result<Option<serde_json::Value>, Error> {
196        let table = quote_ident(&r.table_name())?;
197        let mut all = vec![Filter::eq("id", id)];
198        all.extend_from_slice(filters);
199        let (where_sql, params, _) = self.build_where(&all)?;
200        let hidden = self.hidden_subtraction(r)?;
201        let sql = format!(
202            "SELECT to_jsonb(t){hidden} AS result FROM (SELECT * FROM {table} {where_sql} LIMIT 1) t"
203        );
204        let row = self
205            .conn
206            .query_one(Statement::from_sql_and_values(
207                DatabaseBackend::Postgres,
208                sql,
209                params,
210            ))
211            .await?;
212        match row {
213            Some(row) => Ok(Some(row.try_get::<serde_json::Value>("", "result")?)),
214            None => Ok(None),
215        }
216    }
217
218    /// `POST /<resource>` — insert and return the created row.
219    pub async fn create(
220        &self,
221        r: &Resource,
222        data: &serde_json::Map<String, serde_json::Value>,
223    ) -> Result<serde_json::Value, Error> {
224        let table = quote_ident(&r.table_name())?;
225        let mut cols = Vec::new();
226        let mut placeholders = Vec::new();
227        let mut params: Vec<SqlValue> = Vec::new();
228        let mut n = 1;
229        for (name, field) in &r.fields {
230            if let Some(v) = data.get(name) {
231                cols.push(quote_ident(name)?);
232                placeholders.push(format!("${n}"));
233                params.push(value::json_to_sql(field.ty, v).map_err(Error::BadInput)?);
234                n += 1;
235            }
236        }
237
238        let hidden = self.hidden_subtraction(r)?;
239        let returning = format!("RETURNING (to_jsonb({table}.*){hidden}) AS result");
240        let sql = if cols.is_empty() {
241            format!("INSERT INTO {table} DEFAULT VALUES {returning}")
242        } else {
243            format!(
244                "INSERT INTO {table} ({}) VALUES ({}) {returning}",
245                cols.join(", "),
246                placeholders.join(", ")
247            )
248        };
249        let row = self
250            .conn
251            .query_one(Statement::from_sql_and_values(
252                DatabaseBackend::Postgres,
253                sql,
254                params,
255            ))
256            .await?
257            .ok_or_else(|| Error::Db(sea_orm::DbErr::Custom("insert returned no row".into())))?;
258        Ok(row.try_get::<serde_json::Value>("", "result")?)
259    }
260
261    /// `PATCH /<resource>/<id>` — update present fields, return the new row.
262    pub async fn update(
263        &self,
264        r: &Resource,
265        id: Uuid,
266        data: &serde_json::Map<String, serde_json::Value>,
267        filters: &[Filter],
268    ) -> Result<Option<serde_json::Value>, Error> {
269        let table = quote_ident(&r.table_name())?;
270        let mut assignments = Vec::new();
271        let mut params: Vec<SqlValue> = Vec::new();
272        let mut n = 1;
273        for (name, field) in &r.fields {
274            if let Some(v) = data.get(name) {
275                assignments.push(format!("{} = ${n}", quote_ident(name)?));
276                params.push(value::json_to_sql(field.ty, v).map_err(Error::BadInput)?);
277                n += 1;
278            }
279        }
280        if r.meta.timestamps {
281            assignments.push("updated_at = now()".to_string());
282        }
283        if assignments.is_empty() {
284            return self.get(r, id, filters).await;
285        }
286
287        let mut where_parts = vec![format!("{} = ${n}", quote_ident("id")?)];
288        params.push(SqlValue::from(id));
289        n += 1;
290        for f in filters {
291            where_parts.push(Self::render_filter(f, &mut params, &mut n)?);
292        }
293
294        let hidden = self.hidden_subtraction(r)?;
295        let sql = format!(
296            "UPDATE {table} SET {} WHERE {} RETURNING (to_jsonb({table}.*){hidden}) AS result",
297            assignments.join(", "),
298            where_parts.join(" AND "),
299        );
300        let row = self
301            .conn
302            .query_one(Statement::from_sql_and_values(
303                DatabaseBackend::Postgres,
304                sql,
305                params,
306            ))
307            .await?;
308        match row {
309            Some(row) => Ok(Some(row.try_get::<serde_json::Value>("", "result")?)),
310            None => Ok(None),
311        }
312    }
313
314    /// `DELETE /<resource>/<id>` — returns whether a row was removed.
315    pub async fn delete(&self, r: &Resource, id: Uuid, filters: &[Filter]) -> Result<bool, Error> {
316        let table = quote_ident(&r.table_name())?;
317        let mut all = vec![Filter::eq("id", id)];
318        all.extend_from_slice(filters);
319        let (where_sql, params, _) = self.build_where(&all)?;
320        let res = self
321            .conn
322            .execute(Statement::from_sql_and_values(
323                DatabaseBackend::Postgres,
324                format!("DELETE FROM {table} {where_sql}"),
325                params,
326            ))
327            .await?;
328        Ok(res.rows_affected() > 0)
329    }
330
331    /// Fetch multiple rows of a resource by id (used for relation expansion).
332    /// `filters` carry the caller's authorization scope, so an expansion can
333    /// never reach a row a direct read would have refused. Returns a JSON array
334    /// with hidden fields stripped; order is unspecified.
335    pub async fn fetch_by_ids(
336        &self,
337        r: &Resource,
338        ids: &[Uuid],
339        filters: &[Filter],
340    ) -> Result<serde_json::Value, Error> {
341        if ids.is_empty() {
342            return Ok(serde_json::Value::Array(Vec::new()));
343        }
344        let table = quote_ident(&r.table_name())?;
345        let mut all = vec![Filter::in_uuids("id", ids.to_vec())];
346        all.extend_from_slice(filters);
347        let (where_sql, params, _) = self.build_where(&all)?;
348        let hidden = self.hidden_subtraction(r)?;
349        let sql = format!(
350            "SELECT coalesce(jsonb_agg(to_jsonb(t){hidden}), '[]'::jsonb) AS result \
351             FROM (SELECT * FROM {table} {where_sql}) t"
352        );
353        let row = self
354            .conn
355            .query_one(Statement::from_sql_and_values(
356                DatabaseBackend::Postgres,
357                sql,
358                params,
359            ))
360            .await?
361            .ok_or_else(|| Error::Db(sea_orm::DbErr::Custom("no aggregate row".into())))?;
362        Ok(row.try_get::<serde_json::Value>("", "result")?)
363    }
364
365    /// Raw query bridge used by function `.so`s. `SELECT`/`WITH` statements come
366    /// back as a JSON array of rows; anything else returns `{"rows_affected":n}`.
367    pub async fn raw_json(
368        &self,
369        sql: &str,
370        params: &[serde_json::Value],
371    ) -> Result<serde_json::Value, Error> {
372        let vals: Vec<SqlValue> = params.iter().map(value::json_param).collect();
373        let head = sql.trim_start();
374        let is_query = (head.len() >= 6 && head[..6].eq_ignore_ascii_case("select"))
375            || (head.len() >= 4 && head[..4].eq_ignore_ascii_case("with"));
376
377        if is_query {
378            let wrapped =
379                format!("SELECT coalesce(jsonb_agg(t), '[]'::jsonb) AS result FROM ({sql}) t");
380            let row = self
381                .conn
382                .query_one(Statement::from_sql_and_values(
383                    DatabaseBackend::Postgres,
384                    wrapped,
385                    vals,
386                ))
387                .await?
388                .ok_or_else(|| Error::Db(sea_orm::DbErr::Custom("no aggregate row".into())))?;
389            Ok(row.try_get::<serde_json::Value>("", "result")?)
390        } else {
391            let res = self
392                .conn
393                .execute(Statement::from_sql_and_values(
394                    DatabaseBackend::Postgres,
395                    sql.to_string(),
396                    vals,
397                ))
398                .await?;
399            Ok(serde_json::json!({ "rows_affected": res.rows_affected() }))
400        }
401    }
402
403    // --- helpers ----------------------------------------------------------
404
405    /// Build a `WHERE …` clause from filters; returns the SQL, the bound values,
406    /// and the next free parameter index.
407    fn build_where(&self, filters: &[Filter]) -> Result<(String, Vec<SqlValue>, usize), Error> {
408        if filters.is_empty() {
409            return Ok((String::new(), Vec::new(), 1));
410        }
411        let mut parts = Vec::new();
412        let mut params = Vec::new();
413        let mut n = 1;
414        for f in filters {
415            parts.push(Self::render_filter(f, &mut params, &mut n)?);
416        }
417        Ok((format!("WHERE {}", parts.join(" AND ")), params, n))
418    }
419
420    /// Render one filter to SQL, appending its bound values to `params` and
421    /// advancing the `$n` counter.
422    fn render_filter(
423        f: &Filter,
424        params: &mut Vec<SqlValue>,
425        n: &mut usize,
426    ) -> Result<String, Error> {
427        let col = quote_ident(f.column())?;
428        Ok(match f {
429            Filter::Eq { value, .. } => {
430                let part = format!("{col} = ${n}");
431                params.push(value.clone());
432                *n += 1;
433                part
434            }
435            Filter::In { values, .. } => {
436                if values.is_empty() {
437                    return Ok("false".to_string());
438                }
439                let placeholders: Vec<String> = values
440                    .iter()
441                    .map(|v| {
442                        let p = format!("${n}");
443                        params.push(v.clone());
444                        *n += 1;
445                        p
446                    })
447                    .collect();
448                format!("{col} IN ({})", placeholders.join(", "))
449            }
450            Filter::Contains { value, .. } => {
451                // The term is bound, so it cannot be SQL — but it is still a
452                // LIKE *pattern*, and an unescaped `%` would match everything.
453                let escaped = value
454                    .replace('\\', "\\\\")
455                    .replace('%', "\\%")
456                    .replace('_', "\\_");
457                let part = format!("{col}::text ILIKE ${n}");
458                params.push(SqlValue::from(format!("%{escaped}%")));
459                *n += 1;
460                part
461            }
462        })
463    }
464
465    /// `- 'col'` fragments that strip hidden fields from a `to_jsonb` result.
466    fn hidden_subtraction(&self, r: &Resource) -> Result<String, Error> {
467        let mut s = String::new();
468        for (name, field) in &r.fields {
469            if field.hidden {
470                quote_ident(name)?; // validate the identifier before embedding
471                s.push_str(&format!(" - '{name}'"));
472            }
473        }
474        Ok(s)
475    }
476}
477
478/// Does this error mean "the database in the URL isn't there"?
479///
480/// sea-orm flattens the sqlx error into its message, so the SQLSTATE for
481/// `invalid_catalog_name` (`3D000`) is matched on text — that code only ever
482/// means a missing database.
483fn is_missing_database(err: &Error) -> bool {
484    let Error::Db(err) = err else { return false };
485    let msg = err.to_string();
486    msg.contains("3D000") || msg.contains("does not exist")
487}
488
489/// Split a Postgres URL into (same server, `postgres` database) and the database
490/// name it asked for. Returns `None` when the URL names no database, in which
491/// case there is nothing to create.
492fn maintenance_url(url: &str) -> Option<(String, String)> {
493    let (before_query, query) = match url.find(['?', '#']) {
494        Some(i) => (&url[..i], &url[i..]),
495        None => (url, ""),
496    };
497    // Skip the `scheme://` so its slashes aren't mistaken for the path.
498    let authority_start = before_query.find("://")? + 3;
499    let slash = authority_start + before_query[authority_start..].find('/')?;
500    let name = &before_query[slash + 1..];
501    if name.is_empty() || name.contains('/') {
502        return None;
503    }
504    Some((
505        format!("{}/postgres{query}", &before_query[..slash]),
506        name.to_string(),
507    ))
508}
509
510#[cfg(test)]
511mod connect_tests {
512    use super::maintenance_url;
513
514    #[test]
515    fn swaps_the_database_name() {
516        assert_eq!(
517            maintenance_url("postgres://user:pw@127.0.0.1:5432/apiplant"),
518            Some((
519                "postgres://user:pw@127.0.0.1:5432/postgres".into(),
520                "apiplant".into()
521            ))
522        );
523    }
524
525    #[test]
526    fn keeps_query_parameters() {
527        assert_eq!(
528            maintenance_url("postgres://localhost/app?sslmode=require"),
529            Some((
530                "postgres://localhost/postgres?sslmode=require".into(),
531                "app".into()
532            ))
533        );
534    }
535
536    #[test]
537    fn no_database_in_url() {
538        assert_eq!(maintenance_url("postgres://localhost"), None);
539        assert_eq!(maintenance_url("postgres://localhost/"), None);
540    }
541}