taler-rust

GNU Taler code in Rust. Largely core banking integrations.
Log | Files | Refs | Submodules | README | LICENSE

db.rs (11977B)


      1 /*
      2   This file is part of TALER
      3   Copyright (C) 2024-2026 Taler Systems SA
      4 
      5   TALER is free software; you can redistribute it and/or modify it under the
      6   terms of the GNU Affero General Public License as published by the Free Software
      7   Foundation; either version 3, or (at your option) any later version.
      8 
      9   TALER is distributed in the hope that it will be useful, but WITHOUT ANY
     10   WARRANTY; without even the implied warranty of MERCHANTABILITY or FITNESS FOR
     11   A PARTICULAR PURPOSE.  See the GNU Affero General Public License for more details.
     12 
     13   You should have received a copy of the GNU Affero General Public License along with
     14   TALER; see the file COPYING.  If not, see <http://www.gnu.org/licenses/>
     15 */
     16 
     17 use std::{str::FromStr, time::Duration};
     18 
     19 use jiff::{
     20     Timestamp,
     21     civil::{Date, Time},
     22     tz::TimeZone,
     23 };
     24 use sqlx::{
     25     Decode, Error, PgPool, Postgres, QueryBuilder, Row, Type,
     26     error::BoxDynError,
     27     postgres::PgRow,
     28     query::{Query, QueryScalar},
     29 };
     30 use taler_common::{
     31     api::params::{History, Page, Pooling},
     32     types::{
     33         amount::{Amount, Currency, Decimal},
     34         iban::IBAN,
     35         payto::PaytoURI,
     36         utils::date_to_utc_ts,
     37     },
     38 };
     39 use tokio::sync::watch::{self};
     40 use url::Url;
     41 
     42 pub type PgQueryBuilder<'b> = QueryBuilder<'b, Postgres>;
     43 
     44 /* ------ Serialization ----- */
     45 
     46 pub trait PgError {
     47     const PG_SERIALIZATION_FAILURE: &str = "40001";
     48     const PG_DEADLOCK_DETECTED: &str = "40P01";
     49     const PG_UNIQUE_VIOLATION: &str = "23505";
     50     const PG_FOREIGN_KEY_VIOLATION: &str = "23503";
     51 
     52     fn is_retryable_err(&self) -> bool;
     53     fn is_unique_err(&self) -> bool;
     54     fn is_fk_err(&self) -> bool;
     55 }
     56 
     57 impl PgError for sqlx::error::Error {
     58     fn is_retryable_err(&self) -> bool {
     59         if let sqlx::Error::Database(e) = self {
     60             return matches!(
     61                 e.downcast_ref::<sqlx::postgres::PgDatabaseError>().code(),
     62                 Self::PG_SERIALIZATION_FAILURE | Self::PG_DEADLOCK_DETECTED
     63             );
     64         }
     65         false
     66     }
     67 
     68     fn is_unique_err(&self) -> bool {
     69         if let sqlx::Error::Database(e) = self {
     70             return e.downcast_ref::<sqlx::postgres::PgDatabaseError>().code()
     71                 == Self::PG_UNIQUE_VIOLATION;
     72         }
     73         false
     74     }
     75 
     76     fn is_fk_err(&self) -> bool {
     77         if let sqlx::Error::Database(e) = self {
     78             return e.downcast_ref::<sqlx::postgres::PgDatabaseError>().code()
     79                 == Self::PG_FOREIGN_KEY_VIOLATION;
     80         }
     81         false
     82     }
     83 }
     84 
     85 #[macro_export]
     86 macro_rules! serialized {
     87     ($logic:expr) => {{
     88         use $crate::db::PgError;
     89         let mut attempts = 0;
     90         const MAX_RETRIES: u32 = 5;
     91 
     92         loop {
     93             let res: sqlx::Result<_, sqlx::Error> = $logic.await;
     94             if let Err(e) = &res
     95                 && e.is_retryable_err()
     96                 && attempts < MAX_RETRIES
     97             {
     98                 attempts += 1;
     99                 tokio::task::yield_now().await;
    100                 continue;
    101             }
    102             break res;
    103         }
    104     }};
    105 }
    106 
    107 /* ----- Routines ------ */
    108 
    109 pub async fn page<'a, 'b, R: Send + Unpin>(
    110     db: &PgPool,
    111     params: &Page,
    112     id_col: &str,
    113     prepare: impl Fn() -> QueryBuilder<'a, Postgres> + Copy,
    114     map: impl Fn(PgRow) -> Result<R, Error> + Send + Copy,
    115 ) -> Result<Vec<R>, Error> {
    116     serialized!(async {
    117         let mut builder = prepare();
    118         if let Some(offset) = params.offset {
    119             builder
    120                 .push(format_args!(
    121                     " {id_col} {}",
    122                     if params.backward() { '<' } else { '>' }
    123                 ))
    124                 .push_bind(offset);
    125         } else {
    126             builder.push("TRUE");
    127         }
    128         builder.push(format_args!(
    129             " ORDER BY {id_col} {} LIMIT ",
    130             if params.backward() { "DESC" } else { "ASC" }
    131         ));
    132         builder
    133             .push_bind(params.limit())
    134             .build()
    135             .try_map(map)
    136             .fetch_all(db)
    137             .await
    138     })
    139 }
    140 
    141 pub async fn pooling<R, N, F: Future<Output = sqlx::Result<R>>>(
    142     params: &Pooling,
    143     listen: impl FnOnce() -> watch::Receiver<N>,
    144     filter: impl FnMut(&N) -> bool,
    145     mut load: impl FnMut() -> F,
    146     mut check: impl FnMut(&R) -> bool,
    147 ) -> Result<R, Error> {
    148     let timeout = params.timeout_ms.unwrap_or_default();
    149     if timeout > 0 {
    150         let mut listener = listen();
    151         let init = load().await?;
    152         // Long polling if we found no transactions
    153         if !check(&init) {
    154             tokio::time::timeout(Duration::from_millis(timeout), async {
    155                 listener.wait_for(filter).await.ok();
    156             })
    157             .await
    158             .ok();
    159             // Whether pooling worked or not we load some fresh data from the database
    160             load().await
    161         } else {
    162             Ok(init)
    163         }
    164     } else {
    165         load().await
    166     }
    167 }
    168 
    169 pub async fn history<T: Send + Unpin>(
    170     db: &PgPool,
    171     id_col: &str,
    172     params: &History,
    173     listen: impl FnOnce() -> watch::Receiver<i64>,
    174     prepare: impl Fn() -> QueryBuilder<'static, Postgres> + Copy,
    175     map: impl Fn(PgRow) -> Result<T, Error> + Send + Copy,
    176 ) -> Result<Vec<T>, Error> {
    177     let load = async || page(db, &params.page, id_col, prepare, map).await;
    178     // When going backward there is always at least one transaction or none
    179     let poll = if params.page.limit < 0 {
    180         &Pooling::default()
    181     } else {
    182         &params.pooling
    183     };
    184     pooling(
    185         poll,
    186         listen,
    187         |id| *id > params.page.offset.unwrap_or_default(),
    188         load,
    189         |init| !init.is_empty(),
    190     )
    191     .await
    192 }
    193 
    194 /* ----- Bind ----- */
    195 
    196 pub trait BindHelper {
    197     fn bind_timestamp(self, timestamp: &Timestamp) -> Self;
    198     fn bind_date(self, date: &Date) -> Self;
    199 }
    200 
    201 impl<'q> BindHelper for Query<'q, Postgres, <Postgres as sqlx::Database>::Arguments<'q>> {
    202     fn bind_timestamp(self, timestamp: &Timestamp) -> Self {
    203         self.bind(timestamp.as_microsecond())
    204     }
    205 
    206     fn bind_date(self, date: &Date) -> Self {
    207         self.bind_timestamp(&date_to_utc_ts(date))
    208     }
    209 }
    210 
    211 impl<'q, T> BindHelper
    212     for QueryScalar<'q, Postgres, T, <Postgres as sqlx::Database>::Arguments<'q>>
    213 {
    214     fn bind_timestamp(self, timestamp: &Timestamp) -> Self {
    215         self.bind(timestamp.as_microsecond())
    216     }
    217 
    218     fn bind_date(self, date: &Date) -> Self {
    219         self.bind_timestamp(&date_to_utc_ts(date))
    220     }
    221 }
    222 
    223 /* ----- Get ----- */
    224 
    225 pub trait TypeHelper {
    226     fn try_get_map<
    227         'r,
    228         I: sqlx::ColumnIndex<Self>,
    229         T: Decode<'r, Postgres> + Type<Postgres>,
    230         E: Into<BoxDynError>,
    231         R,
    232         M: FnOnce(T) -> Result<R, E>,
    233     >(
    234         &'r self,
    235         index: I,
    236         map: M,
    237     ) -> sqlx::Result<R>;
    238     fn try_get_opt_map<
    239         'r,
    240         I: sqlx::ColumnIndex<Self>,
    241         T: Decode<'r, Postgres> + Type<Postgres>,
    242         E: Into<BoxDynError>,
    243         R,
    244         M: FnOnce(T) -> Result<R, E>,
    245     >(
    246         &'r self,
    247         index: I,
    248         map: M,
    249     ) -> sqlx::Result<Option<R>> {
    250         self.try_get_map(index, |it: Option<T>| it.map(map).transpose())
    251     }
    252     fn try_get_parse<I: sqlx::ColumnIndex<Self>, E: Into<BoxDynError>, T: FromStr<Err = E>>(
    253         &self,
    254         index: I,
    255     ) -> sqlx::Result<T> {
    256         self.try_get_map(index, |s: &str| s.parse())
    257     }
    258     fn try_get_opt_parse<I: sqlx::ColumnIndex<Self>, E: Into<BoxDynError>, T: FromStr<Err = E>>(
    259         &self,
    260         index: I,
    261     ) -> sqlx::Result<Option<T>> {
    262         self.try_get_map(index, |s: Option<&str>| s.map(|s| s.parse()).transpose())
    263     }
    264     fn try_get_timestamp<I: sqlx::ColumnIndex<Self>>(&self, index: I) -> sqlx::Result<Timestamp> {
    265         self.try_get_map(index, |micros| {
    266             jiff::Timestamp::from_microsecond(micros)
    267                 .map_err(|e| format!("expected timestamp micros got overflowing {micros}: {e}"))
    268         })
    269     }
    270     fn try_get_opt_timestamp<I: sqlx::ColumnIndex<Self>>(
    271         &self,
    272         index: I,
    273     ) -> sqlx::Result<Option<Timestamp>> {
    274         self.try_get_map(index, |micros: Option<i64>| {
    275             if let Some(micros) = micros {
    276                 Some(jiff::Timestamp::from_microsecond(micros).map_err(|e| {
    277                     format!("expected timestamp micros got overflowing {micros}: {e}")
    278                 }))
    279                 .transpose()
    280             } else {
    281                 Ok(None)
    282             }
    283         })
    284     }
    285     fn try_get_date<I: sqlx::ColumnIndex<Self>>(&self, index: I) -> sqlx::Result<Date> {
    286         let timestamp = self.try_get_timestamp(index)?;
    287         let zoned = timestamp.to_zoned(TimeZone::UTC);
    288         assert_eq!(zoned.time(), Time::midnight());
    289         Ok(zoned.date())
    290     }
    291     fn try_get_u16<I: sqlx::ColumnIndex<Self>>(&self, index: I) -> sqlx::Result<u16> {
    292         self.try_get_map(index, |signed: i16| signed.try_into())
    293     }
    294     fn try_get_opt_u16<I: sqlx::ColumnIndex<Self>>(&self, index: I) -> sqlx::Result<Option<u16>> {
    295         self.try_get_opt_map(index, |signed: i16| signed.try_into())
    296     }
    297     fn try_get_u32<I: sqlx::ColumnIndex<Self>>(&self, index: I) -> sqlx::Result<u32> {
    298         self.try_get_map(index, |signed: i32| signed.try_into())
    299     }
    300     fn try_get_opt_u32<I: sqlx::ColumnIndex<Self>>(&self, index: I) -> sqlx::Result<Option<u32>> {
    301         self.try_get_opt_map(index, |signed: i32| signed.try_into())
    302     }
    303     fn try_get_u64<I: sqlx::ColumnIndex<Self>>(&self, index: I) -> sqlx::Result<u64> {
    304         self.try_get_map(index, |signed: i64| signed.try_into())
    305     }
    306     fn try_get_opt_u64<I: sqlx::ColumnIndex<Self>>(&self, index: I) -> sqlx::Result<Option<u64>> {
    307         self.try_get_opt_map(index, |signed: i64| signed.try_into())
    308     }
    309     fn try_get_url<I: sqlx::ColumnIndex<Self>>(&self, index: I) -> sqlx::Result<Url> {
    310         self.try_get_parse(index)
    311     }
    312     fn try_get_payto<I: sqlx::ColumnIndex<Self>>(&self, index: I) -> sqlx::Result<PaytoURI> {
    313         self.try_get_parse(index)
    314     }
    315     fn try_get_opt_payto<I: sqlx::ColumnIndex<Self>>(
    316         &self,
    317         index: I,
    318     ) -> sqlx::Result<Option<PaytoURI>> {
    319         self.try_get_opt_parse(index)
    320     }
    321     fn try_get_iban<I: sqlx::ColumnIndex<Self>>(&self, index: I) -> sqlx::Result<IBAN> {
    322         self.try_get_parse(index)
    323     }
    324     fn try_get_amount<I: sqlx::ColumnIndex<Self>>(
    325         &self,
    326         index: I,
    327         currency: &Currency,
    328     ) -> sqlx::Result<Amount>;
    329     fn try_get_opt_amount<I: sqlx::ColumnIndex<Self>>(
    330         &self,
    331         index: I,
    332         currency: &Currency,
    333     ) -> sqlx::Result<Option<Amount>>;
    334 
    335     /** Flag consider NULL and false to be the same */
    336     fn try_get_flag<I: sqlx::ColumnIndex<Self>>(&self, index: I) -> sqlx::Result<bool>;
    337 }
    338 
    339 impl TypeHelper for PgRow {
    340     fn try_get_map<
    341         'r,
    342         I: sqlx::ColumnIndex<Self>,
    343         T: Decode<'r, Postgres> + Type<Postgres>,
    344         E: Into<BoxDynError>,
    345         R,
    346         M: FnOnce(T) -> Result<R, E>,
    347     >(
    348         &'r self,
    349         index: I,
    350         map: M,
    351     ) -> sqlx::Result<R> {
    352         let primitive: T = self.try_get(&index)?;
    353         map(primitive).map_err(|source| sqlx::Error::ColumnDecode {
    354             index: format!("{index:?}"),
    355             source: source.into(),
    356         })
    357     }
    358 
    359     fn try_get_amount<I: sqlx::ColumnIndex<Self>>(
    360         &self,
    361         index: I,
    362         currency: &Currency,
    363     ) -> sqlx::Result<Amount> {
    364         let decimal: Decimal = self.try_get(index)?;
    365         Ok(Amount::new_decimal(currency, decimal))
    366     }
    367 
    368     fn try_get_opt_amount<I: sqlx::ColumnIndex<Self>>(
    369         &self,
    370         index: I,
    371         currency: &Currency,
    372     ) -> sqlx::Result<Option<Amount>> {
    373         let decimal: Option<Decimal> = self.try_get(index)?;
    374         Ok(decimal.map(|decimal| Amount::new_decimal(currency, decimal)))
    375     }
    376 
    377     fn try_get_flag<I: sqlx::ColumnIndex<Self>>(&self, index: I) -> sqlx::Result<bool> {
    378         let opt_bool: Option<bool> = self.try_get(index)?;
    379         Ok(opt_bool.unwrap_or(false))
    380     }
    381 }