Skip to main content

stateset_db/sqlite/
units_of_measure.rs

1//! SQLite implementation of the units-of-measure repository
2//!
3//! Covers unit classes, units of measure, and conversion rules.
4
5use 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            // Unset any current base in the class, set the new one.
226            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        // Enforce the rule_type / product_id invariant.
263        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        // class can't be deleted while UOMs exist
370        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}