1use super::{
6 map_db_error, parse_datetime_row, parse_decimal_row, parse_enum_row, parse_uuid_opt_row,
7 parse_uuid_row, with_immediate_transaction,
8};
9use chrono::Utc;
10use r2d2::Pool;
11use r2d2_sqlite::SqliteConnectionManager;
12use stateset_core::{
13 CommerceError, ConversionRuleType, CreateUnitClass, CreateUnitConversionRule,
14 CreateUnitOfMeasure, Result, UnitClass, UnitClassId, UnitConversionRule, UnitConversionRuleId,
15 UnitOfMeasure, UnitOfMeasureFilter, UnitOfMeasureId, UnitOfMeasureRepository,
16};
17
18#[derive(Debug)]
19pub struct SqliteUnitOfMeasureRepository {
20 pool: Pool<SqliteConnectionManager>,
21}
22
23impl SqliteUnitOfMeasureRepository {
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_class(row: &rusqlite::Row<'_>) -> rusqlite::Result<UnitClass> {
34 Ok(UnitClass {
35 id: parse_uuid_row(&row.get::<_, String>("id")?, "unit_class", "id")?.into(),
36 name: row.get("name")?,
37 description: row.get("description")?,
38 base_uom_id: parse_uuid_opt_row(
39 row.get::<_, Option<String>>("base_uom_id")?,
40 "unit_class",
41 "base_uom_id",
42 )?
43 .map(Into::into),
44 created_at: parse_datetime_row(
45 &row.get::<_, String>("created_at")?,
46 "unit_class",
47 "created_at",
48 )?,
49 updated_at: parse_datetime_row(
50 &row.get::<_, String>("updated_at")?,
51 "unit_class",
52 "updated_at",
53 )?,
54 })
55 }
56
57 fn row_to_uom(row: &rusqlite::Row<'_>) -> rusqlite::Result<UnitOfMeasure> {
58 Ok(UnitOfMeasure {
59 id: parse_uuid_row(&row.get::<_, String>("id")?, "uom", "id")?.into(),
60 unit_class_id: parse_uuid_row(
61 &row.get::<_, String>("unit_class_id")?,
62 "uom",
63 "unit_class_id",
64 )?
65 .into(),
66 name: row.get("name")?,
67 abbreviation: row.get("abbreviation")?,
68 factor: parse_decimal_row(&row.get::<_, String>("factor")?, "uom", "factor")?,
69 is_base: row.get::<_, i32>("is_base")? != 0,
70 created_at: parse_datetime_row(
71 &row.get::<_, String>("created_at")?,
72 "uom",
73 "created_at",
74 )?,
75 updated_at: parse_datetime_row(
76 &row.get::<_, String>("updated_at")?,
77 "uom",
78 "updated_at",
79 )?,
80 })
81 }
82
83 fn row_to_rule(row: &rusqlite::Row<'_>) -> rusqlite::Result<UnitConversionRule> {
84 Ok(UnitConversionRule {
85 id: parse_uuid_row(&row.get::<_, String>("id")?, "conversion_rule", "id")?.into(),
86 rule_type: parse_enum_row::<ConversionRuleType>(
87 &row.get::<_, String>("rule_type")?,
88 "conversion_rule",
89 "rule_type",
90 )?,
91 product_id: parse_uuid_opt_row(
92 row.get::<_, Option<String>>("product_id")?,
93 "conversion_rule",
94 "product_id",
95 )?
96 .map(Into::into),
97 from_uom_id: parse_uuid_row(
98 &row.get::<_, String>("from_uom_id")?,
99 "conversion_rule",
100 "from_uom_id",
101 )?
102 .into(),
103 to_uom_id: parse_uuid_row(
104 &row.get::<_, String>("to_uom_id")?,
105 "conversion_rule",
106 "to_uom_id",
107 )?
108 .into(),
109 factor: parse_decimal_row(
110 &row.get::<_, String>("factor")?,
111 "conversion_rule",
112 "factor",
113 )?,
114 created_at: parse_datetime_row(
115 &row.get::<_, String>("created_at")?,
116 "conversion_rule",
117 "created_at",
118 )?,
119 updated_at: parse_datetime_row(
120 &row.get::<_, String>("updated_at")?,
121 "conversion_rule",
122 "updated_at",
123 )?,
124 })
125 }
126}
127
128impl UnitOfMeasureRepository for SqliteUnitOfMeasureRepository {
129 fn create_class(&self, input: CreateUnitClass) -> Result<UnitClass> {
130 let id = UnitClassId::new();
131 let id_str = id.to_string();
132 let now_str = Utc::now().to_rfc3339();
133 with_immediate_transaction(&self.pool, |tx| {
134 tx.execute(
135 "INSERT INTO unit_classes (id, name, description, created_at, updated_at)
136 VALUES (?, ?, ?, ?, ?)",
137 rusqlite::params![&id_str, &input.name, &input.description, &now_str, &now_str],
138 )?;
139 tx.query_row("SELECT * FROM unit_classes WHERE id = ?", [&id_str], Self::row_to_class)
140 })
141 }
142
143 fn list_classes(&self) -> Result<Vec<UnitClass>> {
144 let conn = self.conn()?;
145 let mut stmt =
146 conn.prepare("SELECT * FROM unit_classes ORDER BY name").map_err(map_db_error)?;
147 let rows = stmt
148 .query_map([], Self::row_to_class)
149 .map_err(map_db_error)?
150 .collect::<std::result::Result<Vec<_>, _>>()
151 .map_err(map_db_error)?;
152 Ok(rows)
153 }
154
155 fn delete_class(&self, id: UnitClassId) -> Result<()> {
156 let conn = self.conn()?;
157 let referenced: i64 = conn
158 .query_row(
159 "SELECT COUNT(*) FROM units_of_measure WHERE unit_class_id = ?",
160 [id.to_string()],
161 |r| r.get(0),
162 )
163 .map_err(map_db_error)?;
164 if referenced > 0 {
165 return Err(CommerceError::Conflict("unit class still has units of measure".into()));
166 }
167 conn.execute("DELETE FROM unit_classes WHERE id = ?", [id.to_string()])
168 .map_err(map_db_error)?;
169 Ok(())
170 }
171
172 fn create_uom(&self, input: CreateUnitOfMeasure) -> Result<UnitOfMeasure> {
173 let id = UnitOfMeasureId::new();
174 let id_str = id.to_string();
175 let now_str = Utc::now().to_rfc3339();
176 with_immediate_transaction(&self.pool, |tx| {
177 tx.execute(
178 "INSERT INTO units_of_measure (id, unit_class_id, name, abbreviation, factor, is_base, created_at, updated_at)
179 VALUES (?, ?, ?, ?, ?, 0, ?, ?)",
180 rusqlite::params![
181 &id_str,
182 input.unit_class_id.to_string(),
183 &input.name,
184 &input.abbreviation,
185 input.factor.to_string(),
186 &now_str,
187 &now_str,
188 ],
189 )?;
190 tx.query_row("SELECT * FROM units_of_measure WHERE id = ?", [&id_str], Self::row_to_uom)
191 })
192 }
193
194 fn list_uoms(&self, filter: UnitOfMeasureFilter) -> Result<Vec<UnitOfMeasure>> {
195 let conn = self.conn()?;
196 let (mut sql, param): (String, Option<String>) = match filter.class_id {
197 Some(c) => (
198 "SELECT * FROM units_of_measure WHERE unit_class_id = ? ORDER BY name".to_string(),
199 Some(c.to_string()),
200 ),
201 None => ("SELECT * FROM units_of_measure ORDER BY name".to_string(), None),
202 };
203 crate::sqlite::append_limit_offset(&mut sql, filter.limit, filter.offset);
204 let mut stmt = conn.prepare(&sql).map_err(map_db_error)?;
205 let rows = if let Some(p) = param {
206 stmt.query_map([p], Self::row_to_uom)
207 } else {
208 stmt.query_map([], Self::row_to_uom)
209 }
210 .map_err(map_db_error)?
211 .collect::<std::result::Result<Vec<_>, _>>()
212 .map_err(map_db_error)?;
213 Ok(rows)
214 }
215
216 fn set_base_uom(&self, id: UnitOfMeasureId) -> Result<UnitOfMeasure> {
217 let id_str = id.to_string();
218 let now_str = Utc::now().to_rfc3339();
219 with_immediate_transaction(&self.pool, |tx| {
220 let class_id: String = tx.query_row(
221 "SELECT unit_class_id FROM units_of_measure WHERE id = ?",
222 [&id_str],
223 |r| r.get(0),
224 )?;
225 tx.execute(
227 "UPDATE units_of_measure SET is_base = 0, updated_at = ? WHERE unit_class_id = ?",
228 rusqlite::params![&now_str, &class_id],
229 )?;
230 tx.execute(
231 "UPDATE units_of_measure SET is_base = 1, updated_at = ? WHERE id = ?",
232 rusqlite::params![&now_str, &id_str],
233 )?;
234 tx.execute(
235 "UPDATE unit_classes SET base_uom_id = ?, updated_at = ? WHERE id = ?",
236 rusqlite::params![&id_str, &now_str, &class_id],
237 )?;
238 tx.query_row("SELECT * FROM units_of_measure WHERE id = ?", [&id_str], Self::row_to_uom)
239 })
240 }
241
242 fn delete_uom(&self, id: UnitOfMeasureId) -> Result<()> {
243 let conn = self.conn()?;
244 let referenced: i64 = conn
245 .query_row(
246 "SELECT COUNT(*) FROM unit_conversion_rules WHERE from_uom_id = ?1 OR to_uom_id = ?1",
247 [id.to_string()],
248 |r| r.get(0),
249 )
250 .map_err(map_db_error)?;
251 if referenced > 0 {
252 return Err(CommerceError::Conflict(
253 "unit of measure is still referenced by a conversion rule".into(),
254 ));
255 }
256 conn.execute("DELETE FROM units_of_measure WHERE id = ?", [id.to_string()])
257 .map_err(map_db_error)?;
258 Ok(())
259 }
260
261 fn create_rule(&self, input: CreateUnitConversionRule) -> Result<UnitConversionRule> {
262 match input.rule_type {
264 ConversionRuleType::Sku if input.product_id.is_none() => {
265 return Err(CommerceError::ValidationError(
266 "SKU conversion rules require a product_id".into(),
267 ));
268 }
269 ConversionRuleType::System if input.product_id.is_some() => {
270 return Err(CommerceError::ValidationError(
271 "SYSTEM conversion rules must not carry a product_id".into(),
272 ));
273 }
274 _ => {}
275 }
276 let id = UnitConversionRuleId::new();
277 let id_str = id.to_string();
278 let now_str = Utc::now().to_rfc3339();
279 with_immediate_transaction(&self.pool, |tx| {
280 tx.execute(
281 "INSERT INTO unit_conversion_rules (id, rule_type, product_id, from_uom_id, to_uom_id, factor, created_at, updated_at)
282 VALUES (?, ?, ?, ?, ?, ?, ?, ?)",
283 rusqlite::params![
284 &id_str,
285 input.rule_type.to_string(),
286 input.product_id.map(|p| p.to_string()),
287 input.from_uom_id.to_string(),
288 input.to_uom_id.to_string(),
289 input.factor.to_string(),
290 &now_str,
291 &now_str,
292 ],
293 )?;
294 tx.query_row(
295 "SELECT * FROM unit_conversion_rules WHERE id = ?",
296 [&id_str],
297 Self::row_to_rule,
298 )
299 })
300 }
301
302 fn list_rules(&self) -> Result<Vec<UnitConversionRule>> {
303 let conn = self.conn()?;
304 let mut stmt = conn
305 .prepare("SELECT * FROM unit_conversion_rules ORDER BY created_at DESC")
306 .map_err(map_db_error)?;
307 let rows = stmt
308 .query_map([], Self::row_to_rule)
309 .map_err(map_db_error)?
310 .collect::<std::result::Result<Vec<_>, _>>()
311 .map_err(map_db_error)?;
312 Ok(rows)
313 }
314
315 fn delete_rule(&self, id: UnitConversionRuleId) -> Result<()> {
316 let conn = self.conn()?;
317 conn.execute("DELETE FROM unit_conversion_rules WHERE id = ?", [id.to_string()])
318 .map_err(map_db_error)?;
319 Ok(())
320 }
321}
322
323#[cfg(test)]
324mod tests {
325 use super::*;
326 use crate::DatabaseConfig;
327 use crate::sqlite::SqliteDatabase;
328 use rust_decimal_macros::dec;
329 use stateset_core::ProductId;
330
331 fn test_repo() -> SqliteUnitOfMeasureRepository {
332 let db = SqliteDatabase::new(&DatabaseConfig::in_memory()).expect("in-memory db");
333 SqliteUnitOfMeasureRepository::new(db.pool().clone())
334 }
335
336 #[test]
337 fn class_uom_lifecycle() {
338 let repo = test_repo();
339 let class = repo
340 .create_class(CreateUnitClass { name: "Weight".into(), description: None })
341 .expect("class");
342 let g = repo
343 .create_uom(CreateUnitOfMeasure {
344 unit_class_id: class.id,
345 name: "Gram".into(),
346 abbreviation: "g".into(),
347 factor: dec!(1),
348 })
349 .expect("uom");
350 let _kg = repo
351 .create_uom(CreateUnitOfMeasure {
352 unit_class_id: class.id,
353 name: "Kilogram".into(),
354 abbreviation: "kg".into(),
355 factor: dec!(1000),
356 })
357 .expect("uom2");
358
359 assert_eq!(
360 repo.list_uoms(UnitOfMeasureFilter { class_id: Some(class.id), ..Default::default() })
361 .expect("list")
362 .len(),
363 2
364 );
365
366 let based = repo.set_base_uom(g.id).expect("base");
367 assert!(based.is_base);
368
369 assert!(repo.delete_class(class.id).is_err());
371 }
372
373 #[test]
374 fn list_uoms_applies_limit_and_offset() {
375 let repo = test_repo();
376 let class = repo
377 .create_class(CreateUnitClass { name: "Length".into(), description: None })
378 .expect("class");
379 for (name, abbr) in [("Meter", "m"), ("Centimeter", "cm"), ("Kilometer", "km")] {
380 repo.create_uom(CreateUnitOfMeasure {
381 unit_class_id: class.id,
382 name: name.into(),
383 abbreviation: abbr.into(),
384 factor: dec!(1),
385 })
386 .expect("uom");
387 }
388
389 let all = repo.list_uoms(UnitOfMeasureFilter::default()).expect("list all");
390 assert_eq!(all.len(), 3);
391
392 let page = repo
393 .list_uoms(UnitOfMeasureFilter { limit: Some(2), ..Default::default() })
394 .expect("limited");
395 assert_eq!(page.len(), 2, "limit must bound the result set");
396 assert_eq!(page[0].name, all[0].name);
397
398 let rest = repo
399 .list_uoms(UnitOfMeasureFilter {
400 limit: Some(2),
401 offset: Some(2),
402 ..Default::default()
403 })
404 .expect("offset");
405 assert_eq!(rest.len(), 1);
406 assert_eq!(rest[0].name, all[2].name);
407 }
408
409 #[test]
410 fn sku_rule_requires_product() {
411 let repo = test_repo();
412 let res = repo.create_rule(CreateUnitConversionRule {
413 rule_type: ConversionRuleType::Sku,
414 product_id: None,
415 from_uom_id: UnitOfMeasureId::new(),
416 to_uom_id: UnitOfMeasureId::new(),
417 factor: dec!(2),
418 });
419 assert!(res.is_err());
420 }
421
422 #[test]
423 fn system_rule_rejects_product() {
424 let repo = test_repo();
425 let res = repo.create_rule(CreateUnitConversionRule {
426 rule_type: ConversionRuleType::System,
427 product_id: Some(ProductId::new()),
428 from_uom_id: UnitOfMeasureId::new(),
429 to_uom_id: UnitOfMeasureId::new(),
430 factor: dec!(2),
431 });
432 assert!(res.is_err());
433 }
434
435 #[test]
436 fn create_and_list_rules() {
437 let repo = test_repo();
438 let rule = repo
439 .create_rule(CreateUnitConversionRule {
440 rule_type: ConversionRuleType::System,
441 product_id: None,
442 from_uom_id: UnitOfMeasureId::new(),
443 to_uom_id: UnitOfMeasureId::new(),
444 factor: dec!(12),
445 })
446 .expect("rule");
447 assert_eq!(rule.factor, dec!(12));
448 assert_eq!(repo.list_rules().expect("list").len(), 1);
449 repo.delete_rule(rule.id).expect("delete");
450 assert_eq!(repo.list_rules().expect("list").len(), 0);
451 }
452}