1use super::{
4 map_db_error, parse_datetime_row, parse_decimal_row, parse_enum_row, parse_uuid_row,
5 with_immediate_transaction,
6};
7use chrono::Utc;
8use r2d2::Pool;
9use r2d2_sqlite::SqliteConnectionManager;
10use rusqlite::OptionalExtension;
11use rust_decimal::Decimal;
12use stateset_core::{
13 ApplyPrepayment, CommerceError, CreatePrepayment, CurrencyCode, Prepayment,
14 PrepaymentApplication, PrepaymentApplicationId, PrepaymentFilter, PrepaymentId,
15 PrepaymentRepository, PrepaymentStatus, PrepaymentTargetType, Result,
16};
17
18#[derive(Debug)]
19pub struct SqlitePrepaymentRepository {
20 pool: Pool<SqliteConnectionManager>,
21}
22
23impl SqlitePrepaymentRepository {
24 #[must_use]
25 pub const fn new(pool: Pool<SqliteConnectionManager>) -> Self {
26 Self { pool }
27 }
28
29 fn conn(&self) -> Result<r2d2::PooledConnection<SqliteConnectionManager>> {
30 self.pool.get().map_err(|e| CommerceError::DatabaseError(e.to_string()))
31 }
32
33 fn row_to_prepayment(row: &rusqlite::Row<'_>) -> rusqlite::Result<Prepayment> {
34 Ok(Prepayment {
35 id: parse_uuid_row(&row.get::<_, String>("id")?, "prepayment", "id")?.into(),
36 number: row.get("number")?,
37 supplier_id: parse_uuid_row(
38 &row.get::<_, String>("supplier_id")?,
39 "prepayment",
40 "supplier_id",
41 )?,
42 amount: parse_decimal_row(&row.get::<_, String>("amount")?, "prepayment", "amount")?,
43 remaining: parse_decimal_row(
44 &row.get::<_, String>("remaining")?,
45 "prepayment",
46 "remaining",
47 )?,
48 currency: parse_enum_row::<CurrencyCode>(
49 &row.get::<_, String>("currency")?,
50 "prepayment",
51 "currency",
52 )?,
53 status: parse_enum_row::<PrepaymentStatus>(
54 &row.get::<_, String>("status")?,
55 "prepayment",
56 "status",
57 )?,
58 method: row.get("method")?,
59 reference: row.get("reference")?,
60 memo: row.get("memo")?,
61 created_at: parse_datetime_row(
62 &row.get::<_, String>("created_at")?,
63 "prepayment",
64 "created_at",
65 )?,
66 updated_at: parse_datetime_row(
67 &row.get::<_, String>("updated_at")?,
68 "prepayment",
69 "updated_at",
70 )?,
71 })
72 }
73
74 fn row_to_app(row: &rusqlite::Row<'_>) -> rusqlite::Result<PrepaymentApplication> {
75 Ok(PrepaymentApplication {
76 id: parse_uuid_row(&row.get::<_, String>("id")?, "prepayment_app", "id")?.into(),
77 prepayment_id: parse_uuid_row(
78 &row.get::<_, String>("prepayment_id")?,
79 "prepayment_app",
80 "prepayment_id",
81 )?
82 .into(),
83 target_type: parse_enum_row::<PrepaymentTargetType>(
84 &row.get::<_, String>("target_type")?,
85 "prepayment_app",
86 "target_type",
87 )?,
88 target_id: parse_uuid_row(
89 &row.get::<_, String>("target_id")?,
90 "prepayment_app",
91 "target_id",
92 )?,
93 amount: parse_decimal_row(
94 &row.get::<_, String>("amount")?,
95 "prepayment_app",
96 "amount",
97 )?,
98 reversed: row.get::<_, i32>("reversed")? != 0,
99 created_at: parse_datetime_row(
100 &row.get::<_, String>("created_at")?,
101 "prepayment_app",
102 "created_at",
103 )?,
104 })
105 }
106
107 fn fetch(tx: &rusqlite::Connection, id: &str) -> rusqlite::Result<Prepayment> {
108 tx.query_row("SELECT * FROM prepayments WHERE id = ?", [id], Self::row_to_prepayment)
109 }
110
111 fn conflict(msg: &str) -> rusqlite::Error {
112 rusqlite::Error::ToSqlConversionFailure(Box::new(CommerceError::Conflict(msg.to_string())))
113 }
114
115 fn validation(msg: &str) -> rusqlite::Error {
116 rusqlite::Error::ToSqlConversionFailure(Box::new(CommerceError::ValidationError(
117 msg.to_string(),
118 )))
119 }
120
121 fn status_for(remaining: Decimal, current: PrepaymentStatus) -> PrepaymentStatus {
122 match current {
123 PrepaymentStatus::Cancelled | PrepaymentStatus::Refunded => current,
124 _ if remaining <= Decimal::ZERO => PrepaymentStatus::Applied,
125 _ => PrepaymentStatus::Open,
126 }
127 }
128}
129
130impl PrepaymentRepository for SqlitePrepaymentRepository {
131 fn create(&self, input: CreatePrepayment) -> Result<Prepayment> {
132 if input.amount <= Decimal::ZERO {
133 return Err(CommerceError::ValidationError(
134 "prepayment amount must be positive".into(),
135 ));
136 }
137 let id = PrepaymentId::new();
138 let id_str = id.to_string();
139 let now_str = Utc::now().to_rfc3339();
140 let number = format!("PRE-{}", &id_str[..8]);
141 let currency = input.currency.unwrap_or(CurrencyCode::USD);
142 with_immediate_transaction(&self.pool, |tx| {
143 tx.execute(
144 "INSERT INTO prepayments (id, number, supplier_id, amount, remaining, currency, status, method, reference, memo, created_at, updated_at)
145 VALUES (?, ?, ?, ?, ?, ?, 'open', ?, ?, ?, ?, ?)",
146 rusqlite::params![
147 &id_str,
148 &number,
149 input.supplier_id.to_string(),
150 input.amount.to_string(),
151 input.amount.to_string(),
152 currency.to_string(),
153 &input.method,
154 &input.reference,
155 &input.memo,
156 &now_str,
157 &now_str,
158 ],
159 )?;
160 Self::fetch(tx, &id_str)
161 })
162 }
163
164 fn get(&self, id: PrepaymentId) -> Result<Option<Prepayment>> {
165 let conn = self.conn()?;
166 match Self::fetch(&conn, &id.to_string()) {
167 Ok(p) => Ok(Some(p)),
168 Err(rusqlite::Error::QueryReturnedNoRows) => Ok(None),
169 Err(e) => Err(map_db_error(e)),
170 }
171 }
172
173 fn list(&self, filter: PrepaymentFilter) -> Result<Vec<Prepayment>> {
174 let conn = self.conn()?;
175 let mut sql = "SELECT * FROM prepayments WHERE 1=1".to_string();
176 let mut params: Vec<Box<dyn rusqlite::types::ToSql>> = vec![];
177 if let Some(supplier) = filter.supplier_id {
178 sql.push_str(" AND supplier_id = ?");
179 params.push(Box::new(supplier.to_string()));
180 }
181 if let Some(status) = filter.status {
182 sql.push_str(" AND status = ?");
183 params.push(Box::new(status.to_string()));
184 }
185 sql.push_str(" ORDER BY created_at DESC");
186 crate::sqlite::append_limit_offset(&mut sql, filter.limit, filter.offset);
187 let param_refs: Vec<&dyn rusqlite::types::ToSql> =
188 params.iter().map(|p| p.as_ref()).collect();
189 let mut stmt = conn.prepare(&sql).map_err(map_db_error)?;
190 let rows = stmt
191 .query_map(param_refs.as_slice(), Self::row_to_prepayment)
192 .map_err(map_db_error)?
193 .collect::<std::result::Result<Vec<_>, _>>()
194 .map_err(map_db_error)?;
195 Ok(rows)
196 }
197
198 fn apply(&self, id: PrepaymentId, input: ApplyPrepayment) -> Result<Prepayment> {
199 let id_str = id.to_string();
200 let now = Utc::now().to_rfc3339();
201 with_immediate_transaction(&self.pool, |tx| {
202 let prepayment =
203 Self::fetch(tx, &id_str).optional()?.ok_or(rusqlite::Error::QueryReturnedNoRows)?;
204 if prepayment.status != PrepaymentStatus::Open {
205 return Err(Self::conflict("prepayment is not open for application"));
206 }
207 if input.amount <= Decimal::ZERO {
208 return Err(Self::validation("application amount must be positive"));
209 }
210 if input.amount > prepayment.remaining {
211 return Err(Self::validation("application exceeds remaining balance"));
212 }
213 let new_remaining = prepayment.remaining - input.amount;
214 let new_status = Self::status_for(new_remaining, prepayment.status);
215 tx.execute(
216 "INSERT INTO prepayment_applications (id, prepayment_id, target_type, target_id, amount, reversed, created_at)
217 VALUES (?, ?, ?, ?, ?, 0, ?)",
218 rusqlite::params![
219 PrepaymentApplicationId::new().to_string(),
220 &id_str,
221 input.target_type.to_string(),
222 input.target_id.to_string(),
223 input.amount.to_string(),
224 &now,
225 ],
226 )?;
227 tx.execute(
228 "UPDATE prepayments SET remaining = ?, status = ?, updated_at = ? WHERE id = ?",
229 rusqlite::params![new_remaining.to_string(), new_status.to_string(), &now, &id_str],
230 )?;
231 Self::fetch(tx, &id_str)
232 })
233 }
234
235 fn list_applications(&self, id: PrepaymentId) -> Result<Vec<PrepaymentApplication>> {
236 let conn = self.conn()?;
237 let mut stmt = conn
238 .prepare(
239 "SELECT * FROM prepayment_applications WHERE prepayment_id = ? ORDER BY created_at",
240 )
241 .map_err(map_db_error)?;
242 let rows = stmt
243 .query_map([id.to_string()], Self::row_to_app)
244 .map_err(map_db_error)?
245 .collect::<std::result::Result<Vec<_>, _>>()
246 .map_err(map_db_error)?;
247 Ok(rows)
248 }
249
250 fn reverse_application(
251 &self,
252 id: PrepaymentId,
253 application_id: PrepaymentApplicationId,
254 ) -> Result<Prepayment> {
255 let id_str = id.to_string();
256 let app_str = application_id.to_string();
257 let now = Utc::now().to_rfc3339();
258 with_immediate_transaction(&self.pool, |tx| {
259 let row: Option<(String, i32)> = tx
260 .query_row(
261 "SELECT amount, reversed FROM prepayment_applications WHERE id = ? AND prepayment_id = ?",
262 rusqlite::params![&app_str, &id_str],
263 |r| Ok((r.get(0)?, r.get(1)?)),
264 )
265 .optional()?;
266 let (amount_str, reversed) = row.ok_or(rusqlite::Error::QueryReturnedNoRows)?;
267 if reversed != 0 {
268 return Err(Self::conflict("application already reversed"));
269 }
270 let amount: Decimal = amount_str.parse().unwrap_or(Decimal::ZERO);
271 let prepayment = Self::fetch(tx, &id_str)?;
272 if prepayment.status == PrepaymentStatus::Refunded {
273 return Err(Self::conflict("cannot reverse against a refunded prepayment"));
274 }
275 let new_remaining = prepayment.remaining + amount;
276 let new_status = Self::status_for(new_remaining, prepayment.status);
277 tx.execute("UPDATE prepayment_applications SET reversed = 1 WHERE id = ?", [&app_str])?;
278 tx.execute(
279 "UPDATE prepayments SET remaining = ?, status = ?, updated_at = ? WHERE id = ?",
280 rusqlite::params![new_remaining.to_string(), new_status.to_string(), &now, &id_str],
281 )?;
282 Self::fetch(tx, &id_str)
283 })
284 }
285
286 fn refund(&self, id: PrepaymentId) -> Result<Prepayment> {
287 let id_str = id.to_string();
288 let now = Utc::now().to_rfc3339();
289 with_immediate_transaction(&self.pool, |tx| {
290 let prepayment = Self::fetch(tx, &id_str)?;
291 if prepayment.status == PrepaymentStatus::Cancelled {
292 return Err(Self::conflict("cannot refund a cancelled prepayment"));
293 }
294 tx.execute(
296 "UPDATE prepayments SET remaining = '0', status = 'refunded', updated_at = ? WHERE id = ?",
297 rusqlite::params![&now, &id_str],
298 )?;
299 Self::fetch(tx, &id_str)
300 })
301 }
302}
303
304#[cfg(test)]
305mod tests {
306 use super::*;
307 use crate::DatabaseConfig;
308 use crate::sqlite::SqliteDatabase;
309 use rust_decimal_macros::dec;
310 use uuid::Uuid;
311
312 fn test_repo() -> SqlitePrepaymentRepository {
313 let db = SqliteDatabase::new(&DatabaseConfig::in_memory()).expect("in-memory db");
314 SqlitePrepaymentRepository::new(db.pool().clone())
315 }
316
317 fn new_prepayment(repo: &SqlitePrepaymentRepository, amount: Decimal) -> Prepayment {
318 repo.create(CreatePrepayment {
319 supplier_id: Uuid::new_v4(),
320 amount,
321 currency: Some(CurrencyCode::USD),
322 method: Some("wire".into()),
323 reference: Some("WIRE-123".into()),
324 memo: None,
325 })
326 .expect("create prepayment")
327 }
328
329 fn apply(amount: Decimal) -> ApplyPrepayment {
330 ApplyPrepayment {
331 target_type: PrepaymentTargetType::Bill,
332 target_id: Uuid::new_v4(),
333 amount,
334 }
335 }
336
337 #[test]
338 fn create_rejects_non_positive() {
339 let repo = test_repo();
340 assert!(
341 repo.create(CreatePrepayment {
342 supplier_id: Uuid::new_v4(),
343 amount: dec!(0),
344 currency: None,
345 method: None,
346 reference: None,
347 memo: None,
348 })
349 .is_err()
350 );
351 }
352
353 #[test]
354 fn apply_decrements_and_marks_applied() {
355 let repo = test_repo();
356 let p = new_prepayment(&repo, dec!(100));
357 let after = repo.apply(p.id, apply(dec!(40))).expect("apply");
358 assert_eq!(after.remaining, dec!(60));
359 assert_eq!(after.status, PrepaymentStatus::Open);
360 let after = repo.apply(p.id, apply(dec!(60))).expect("apply rest");
361 assert_eq!(after.remaining, dec!(0));
362 assert_eq!(after.status, PrepaymentStatus::Applied);
363 assert_eq!(repo.list_applications(p.id).expect("apps").len(), 2);
364 }
365
366 #[test]
367 fn apply_over_balance_rejected() {
368 let repo = test_repo();
369 let p = new_prepayment(&repo, dec!(50));
370 assert!(repo.apply(p.id, apply(dec!(60))).is_err());
371 }
372
373 #[test]
374 fn reverse_restores_balance() {
375 let repo = test_repo();
376 let p = new_prepayment(&repo, dec!(100));
377 repo.apply(p.id, apply(dec!(100))).expect("apply");
378 let apps = repo.list_applications(p.id).expect("apps");
379 let reversed = repo.reverse_application(p.id, apps[0].id).expect("reverse");
380 assert_eq!(reversed.remaining, dec!(100));
381 assert_eq!(reversed.status, PrepaymentStatus::Open);
382 assert!(repo.reverse_application(p.id, apps[0].id).is_err());
383 }
384
385 #[test]
386 fn refund_closes_prepayment() {
387 let repo = test_repo();
388 let p = new_prepayment(&repo, dec!(100));
389 repo.apply(p.id, apply(dec!(30))).expect("apply");
390 let refunded = repo.refund(p.id).expect("refund");
391 assert_eq!(refunded.status, PrepaymentStatus::Refunded);
392 assert_eq!(refunded.remaining, dec!(0));
393 assert!(repo.apply(p.id, apply(dec!(10))).is_err());
395 }
396}