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, ¶ms.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 ¶ms.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 }