1mod 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#[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#[derive(Clone)]
44pub enum Filter {
45 Eq { column: String, value: SqlValue },
47 In {
49 column: String,
50 values: Vec<SqlValue>,
51 },
52 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 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#[derive(Clone)]
99pub struct Db {
100 conn: DatabaseConnection,
101}
102
103impl Db {
104 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 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 pub fn connection(&self) -> &DatabaseConnection {
147 &self.conn
148 }
149
150 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 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 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 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 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 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 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 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 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 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 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)?; s.push_str(&format!(" - '{name}'"));
472 }
473 }
474 Ok(s)
475 }
476}
477
478fn 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
489fn 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 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}