Skip to main content

stateset_db/sqlite/
edi_documents.rs

1//! SQLite implementation of the EDI document repository
2
3use super::{
4    map_db_error, parse_datetime_row, parse_enum_row, parse_uuid_row, with_immediate_transaction,
5};
6use chrono::Utc;
7use r2d2::Pool;
8use r2d2_sqlite::SqliteConnectionManager;
9use stateset_core::{
10    CommerceError, CreateEdiDocument, EdiAggregateSummary, EdiCount, EdiDirection, EdiDocument,
11    EdiDocumentFilter, EdiDocumentId, EdiDocumentRepository, EdiStatus, Result,
12};
13
14#[derive(Debug)]
15pub struct SqliteEdiDocumentRepository {
16    pool: Pool<SqliteConnectionManager>,
17}
18
19impl SqliteEdiDocumentRepository {
20    #[must_use]
21    pub const fn new(pool: Pool<SqliteConnectionManager>) -> Self {
22        Self { pool }
23    }
24
25    fn conn(&self) -> Result<r2d2::PooledConnection<SqliteConnectionManager>> {
26        self.pool.get().map_err(|e| CommerceError::DatabaseError(e.to_string()))
27    }
28
29    fn row_to_doc(row: &rusqlite::Row<'_>) -> rusqlite::Result<EdiDocument> {
30        Ok(EdiDocument {
31            id: parse_uuid_row(&row.get::<_, String>("id")?, "edi_document", "id")?.into(),
32            document_type: row.get("document_type")?,
33            direction: parse_enum_row::<EdiDirection>(
34                &row.get::<_, String>("direction")?,
35                "edi_document",
36                "direction",
37            )?,
38            status: parse_enum_row::<EdiStatus>(
39                &row.get::<_, String>("status")?,
40                "edi_document",
41                "status",
42            )?,
43            partner: row.get("partner")?,
44            reference: row.get("reference")?,
45            payload: row.get("payload")?,
46            error_message: row.get("error_message")?,
47            created_at: parse_datetime_row(
48                &row.get::<_, String>("created_at")?,
49                "edi_document",
50                "created_at",
51            )?,
52            updated_at: parse_datetime_row(
53                &row.get::<_, String>("updated_at")?,
54                "edi_document",
55                "updated_at",
56            )?,
57        })
58    }
59}
60
61impl EdiDocumentRepository for SqliteEdiDocumentRepository {
62    fn create(&self, input: CreateEdiDocument) -> Result<EdiDocument> {
63        let id = EdiDocumentId::new();
64        let id_str = id.to_string();
65        let now = Utc::now().to_rfc3339();
66        with_immediate_transaction(&self.pool, |tx| {
67            tx.execute(
68                "INSERT INTO edi_documents (id, document_type, direction, status, partner, reference, payload, created_at, updated_at)
69                 VALUES (?, ?, ?, 'pending', ?, ?, ?, ?, ?)",
70                rusqlite::params![
71                    &id_str,
72                    &input.document_type,
73                    input.direction.to_string(),
74                    &input.partner,
75                    &input.reference,
76                    &input.payload,
77                    &now,
78                    &now,
79                ],
80            )?;
81            tx.query_row("SELECT * FROM edi_documents WHERE id = ?", [&id_str], Self::row_to_doc)
82        })
83    }
84
85    fn get(&self, id: EdiDocumentId) -> Result<Option<EdiDocument>> {
86        let conn = self.conn()?;
87        match conn.query_row(
88            "SELECT * FROM edi_documents WHERE id = ?",
89            [id.to_string()],
90            Self::row_to_doc,
91        ) {
92            Ok(d) => Ok(Some(d)),
93            Err(rusqlite::Error::QueryReturnedNoRows) => Ok(None),
94            Err(e) => Err(map_db_error(e)),
95        }
96    }
97
98    fn list(&self, filter: EdiDocumentFilter) -> Result<Vec<EdiDocument>> {
99        let conn = self.conn()?;
100        let mut sql = "SELECT * FROM edi_documents WHERE 1=1".to_string();
101        let mut params: Vec<Box<dyn rusqlite::types::ToSql>> = vec![];
102        if let Some(ref t) = filter.document_type {
103            sql.push_str(" AND document_type = ?");
104            params.push(Box::new(t.clone()));
105        }
106        if let Some(direction) = filter.direction {
107            sql.push_str(" AND direction = ?");
108            params.push(Box::new(direction.to_string()));
109        }
110        if let Some(status) = filter.status {
111            sql.push_str(" AND status = ?");
112            params.push(Box::new(status.to_string()));
113        }
114        if let Some(ref partner) = filter.partner {
115            sql.push_str(" AND partner = ?");
116            params.push(Box::new(partner.clone()));
117        }
118        sql.push_str(" ORDER BY created_at DESC");
119        crate::sqlite::append_limit_offset(&mut sql, filter.limit, filter.offset);
120        let param_refs: Vec<&dyn rusqlite::types::ToSql> =
121            params.iter().map(|p| p.as_ref()).collect();
122        let mut stmt = conn.prepare(&sql).map_err(map_db_error)?;
123        let rows = stmt
124            .query_map(param_refs.as_slice(), Self::row_to_doc)
125            .map_err(map_db_error)?
126            .collect::<std::result::Result<Vec<_>, _>>()
127            .map_err(map_db_error)?;
128        Ok(rows)
129    }
130
131    fn set_status(
132        &self,
133        id: EdiDocumentId,
134        status: EdiStatus,
135        error_message: Option<String>,
136    ) -> Result<EdiDocument> {
137        let id_str = id.to_string();
138        let now = Utc::now().to_rfc3339();
139        with_immediate_transaction(&self.pool, |tx| {
140            tx.execute(
141                "UPDATE edi_documents SET status = ?, error_message = ?, updated_at = ? WHERE id = ?",
142                rusqlite::params![status.to_string(), &error_message, &now, &id_str],
143            )?;
144            tx.query_row("SELECT * FROM edi_documents WHERE id = ?", [&id_str], Self::row_to_doc)
145        })
146    }
147
148    fn summary(&self) -> Result<EdiAggregateSummary> {
149        let conn = self.conn()?;
150        let total: u64 = conn
151            .query_row("SELECT COUNT(*) FROM edi_documents", [], |r| r.get::<_, i64>(0))
152            .map_err(map_db_error)? as u64;
153
154        let collect = |sql: &str| -> Result<Vec<EdiCount>> {
155            let mut stmt = conn.prepare(sql).map_err(map_db_error)?;
156            let rows = stmt
157                .query_map([], |r| {
158                    Ok(EdiCount { key: r.get::<_, String>(0)?, count: r.get::<_, i64>(1)? as u64 })
159                })
160                .map_err(map_db_error)?
161                .collect::<std::result::Result<Vec<_>, _>>()
162                .map_err(map_db_error)?;
163            Ok(rows)
164        };
165
166        let by_status =
167            collect("SELECT status, COUNT(*) FROM edi_documents GROUP BY status ORDER BY status")?;
168        let by_type = collect(
169            "SELECT document_type, COUNT(*) FROM edi_documents GROUP BY document_type ORDER BY document_type",
170        )?;
171        Ok(EdiAggregateSummary { total, by_status, by_type })
172    }
173}
174
175#[cfg(test)]
176mod tests {
177    use super::*;
178    use crate::DatabaseConfig;
179    use crate::sqlite::SqliteDatabase;
180
181    fn test_repo() -> SqliteEdiDocumentRepository {
182        let db = SqliteDatabase::new(&DatabaseConfig::in_memory()).expect("in-memory db");
183        SqliteEdiDocumentRepository::new(db.pool().clone())
184    }
185
186    fn create(
187        repo: &SqliteEdiDocumentRepository,
188        doc_type: &str,
189        dir: EdiDirection,
190    ) -> EdiDocument {
191        repo.create(CreateEdiDocument {
192            document_type: doc_type.into(),
193            direction: dir,
194            partner: Some("ACME-EDI".into()),
195            reference: Some("PO-1001".into()),
196            payload: Some("ISA*00*...".into()),
197        })
198        .expect("create")
199    }
200
201    #[test]
202    fn create_get_and_default_status() {
203        let repo = test_repo();
204        let d = create(&repo, "850", EdiDirection::Inbound);
205        assert_eq!(d.status, EdiStatus::Pending);
206        let fetched = repo.get(d.id).expect("get").expect("found");
207        assert_eq!(fetched.document_type, "850");
208    }
209
210    #[test]
211    fn set_status_records_error() {
212        let repo = test_repo();
213        let d = create(&repo, "810", EdiDirection::Outbound);
214        let errored =
215            repo.set_status(d.id, EdiStatus::Error, Some("malformed segment".into())).expect("set");
216        assert_eq!(errored.status, EdiStatus::Error);
217        assert_eq!(errored.error_message.as_deref(), Some("malformed segment"));
218    }
219
220    #[test]
221    fn list_filters() {
222        let repo = test_repo();
223        create(&repo, "850", EdiDirection::Inbound);
224        create(&repo, "856", EdiDirection::Outbound);
225        let inbound = repo
226            .list(EdiDocumentFilter {
227                direction: Some(EdiDirection::Inbound),
228                ..Default::default()
229            })
230            .expect("list");
231        assert_eq!(inbound.len(), 1);
232        let by_type = repo
233            .list(EdiDocumentFilter { document_type: Some("856".into()), ..Default::default() })
234            .expect("list");
235        assert_eq!(by_type.len(), 1);
236    }
237
238    #[test]
239    fn summary_groups_counts() {
240        let repo = test_repo();
241        create(&repo, "850", EdiDirection::Inbound);
242        create(&repo, "850", EdiDirection::Inbound);
243        let d = create(&repo, "810", EdiDirection::Outbound);
244        repo.set_status(d.id, EdiStatus::Error, None).expect("set");
245
246        let summary = repo.summary().expect("summary");
247        assert_eq!(summary.total, 3);
248        // by_type: 810 -> 1, 850 -> 2
249        let t850 = summary.by_type.iter().find(|c| c.key == "850").unwrap();
250        assert_eq!(t850.count, 2);
251        // by_status: error -> 1, pending -> 2
252        let err = summary.by_status.iter().find(|c| c.key == "error").unwrap();
253        assert_eq!(err.count, 1);
254    }
255}