Skip to main content

stateset_db/sqlite/
companies.rs

1//! SQLite implementation of the company (B2B account) repository
2
3use super::{
4    map_db_error, parse_datetime_row, parse_decimal_row, parse_enum_row, parse_json_row,
5    parse_uuid_row, with_immediate_transaction,
6};
7use chrono::Utc;
8use r2d2::Pool;
9use r2d2_sqlite::SqliteConnectionManager;
10use stateset_core::{
11    CommerceError, Company, CompanyFilter, CompanyId, CompanyPriceOverride, CompanyRepository,
12    CompanyShippingAddress, CompanyStatus, Contact, ContactId, CreateCompany, CreateContact,
13    CurrencyCode, Result, UpdateCompany,
14};
15
16#[derive(Debug)]
17pub struct SqliteCompanyRepository {
18    pool: Pool<SqliteConnectionManager>,
19}
20
21impl SqliteCompanyRepository {
22    #[must_use]
23    pub const fn new(pool: Pool<SqliteConnectionManager>) -> Self {
24        Self { pool }
25    }
26
27    fn conn(&self) -> Result<r2d2::PooledConnection<SqliteConnectionManager>> {
28        self.pool.get().map_err(|e| CommerceError::DatabaseError(e.to_string()))
29    }
30
31    fn row_to_company(row: &rusqlite::Row<'_>) -> rusqlite::Result<Company> {
32        let tags_json: String = row.get("tags")?;
33        let metadata_json: String = row.get("metadata")?;
34        Ok(Company {
35            id: parse_uuid_row(&row.get::<_, String>("id")?, "company", "id")?.into(),
36            name: row.get("name")?,
37            reference: row.get("reference")?,
38            email: row.get("email")?,
39            phone: row.get("phone")?,
40            currency: parse_enum_row::<CurrencyCode>(
41                &row.get::<_, String>("currency")?,
42                "company",
43                "currency",
44            )?,
45            payment_terms_days: row.get("payment_terms_days")?,
46            status: parse_enum_row::<CompanyStatus>(
47                &row.get::<_, String>("status")?,
48                "company",
49                "status",
50            )?,
51            tags: parse_json_row(&tags_json, "company", "tags")?,
52            metadata: parse_json_row(&metadata_json, "company", "metadata")?,
53            created_at: parse_datetime_row(
54                &row.get::<_, String>("created_at")?,
55                "company",
56                "created_at",
57            )?,
58            updated_at: parse_datetime_row(
59                &row.get::<_, String>("updated_at")?,
60                "company",
61                "updated_at",
62            )?,
63        })
64    }
65
66    fn row_to_address(row: &rusqlite::Row<'_>) -> rusqlite::Result<CompanyShippingAddress> {
67        Ok(CompanyShippingAddress {
68            id: parse_uuid_row(&row.get::<_, String>("id")?, "company_address", "id")?.into(),
69            company_id: parse_uuid_row(
70                &row.get::<_, String>("company_id")?,
71                "company_address",
72                "company_id",
73            )?
74            .into(),
75            label: row.get("label")?,
76            name: row.get("name")?,
77            line1: row.get("line1")?,
78            line2: row.get("line2")?,
79            city: row.get("city")?,
80            region: row.get("region")?,
81            postal_code: row.get("postal_code")?,
82            country: row.get("country")?,
83            is_default: row.get::<_, i32>("is_default")? != 0,
84            created_at: parse_datetime_row(
85                &row.get::<_, String>("created_at")?,
86                "company_address",
87                "created_at",
88            )?,
89            updated_at: parse_datetime_row(
90                &row.get::<_, String>("updated_at")?,
91                "company_address",
92                "updated_at",
93            )?,
94        })
95    }
96
97    fn row_to_contact(row: &rusqlite::Row<'_>) -> rusqlite::Result<Contact> {
98        let company_ids_json: String = row.get("company_ids")?;
99        Ok(Contact {
100            id: parse_uuid_row(&row.get::<_, String>("id")?, "contact", "id")?.into(),
101            first_name: row.get("first_name")?,
102            last_name: row.get("last_name")?,
103            email: row.get("email")?,
104            phone: row.get("phone")?,
105            title: row.get("title")?,
106            company_ids: parse_json_row(&company_ids_json, "contact", "company_ids")?,
107            portal_enabled: row.get::<_, i32>("portal_enabled")? != 0,
108            is_active: row.get::<_, i32>("is_active")? != 0,
109            created_at: parse_datetime_row(
110                &row.get::<_, String>("created_at")?,
111                "contact",
112                "created_at",
113            )?,
114            updated_at: parse_datetime_row(
115                &row.get::<_, String>("updated_at")?,
116                "contact",
117                "updated_at",
118            )?,
119        })
120    }
121
122    fn row_to_override(row: &rusqlite::Row<'_>) -> rusqlite::Result<CompanyPriceOverride> {
123        Ok(CompanyPriceOverride {
124            company_id: parse_uuid_row(
125                &row.get::<_, String>("company_id")?,
126                "price_override",
127                "company_id",
128            )?
129            .into(),
130            product_id: parse_uuid_row(
131                &row.get::<_, String>("product_id")?,
132                "price_override",
133                "product_id",
134            )?
135            .into(),
136            price: parse_decimal_row(&row.get::<_, String>("price")?, "price_override", "price")?,
137            currency: parse_enum_row::<CurrencyCode>(
138                &row.get::<_, String>("currency")?,
139                "price_override",
140                "currency",
141            )?,
142            created_at: parse_datetime_row(
143                &row.get::<_, String>("created_at")?,
144                "price_override",
145                "created_at",
146            )?,
147            updated_at: parse_datetime_row(
148                &row.get::<_, String>("updated_at")?,
149                "price_override",
150                "updated_at",
151            )?,
152        })
153    }
154}
155
156impl CompanyRepository for SqliteCompanyRepository {
157    fn create(&self, input: CreateCompany) -> Result<Company> {
158        let id = CompanyId::new();
159        let id_str = id.to_string();
160        let now_str = Utc::now().to_rfc3339();
161        let currency = input.currency.unwrap_or(CurrencyCode::USD);
162        let tags_json = serde_json::to_string(&input.tags)
163            .map_err(|e| CommerceError::DatabaseError(e.to_string()))?;
164        let metadata_json = serde_json::to_string(&input.metadata)
165            .map_err(|e| CommerceError::DatabaseError(e.to_string()))?;
166
167        with_immediate_transaction(&self.pool, |tx| {
168            tx.execute(
169                "INSERT INTO companies (id, name, reference, email, phone, currency, payment_terms_days, status, tags, metadata, created_at, updated_at)
170                 VALUES (?, ?, ?, ?, ?, ?, ?, 'active', ?, ?, ?, ?)",
171                rusqlite::params![
172                    &id_str,
173                    &input.name,
174                    &input.reference,
175                    &input.email,
176                    &input.phone,
177                    currency.to_string(),
178                    input.payment_terms_days,
179                    &tags_json,
180                    &metadata_json,
181                    &now_str,
182                    &now_str,
183                ],
184            )?;
185            tx.query_row("SELECT * FROM companies WHERE id = ?", [&id_str], Self::row_to_company)
186        })
187    }
188
189    fn get(&self, id: CompanyId) -> Result<Option<Company>> {
190        let conn = self.conn()?;
191        match conn.query_row(
192            "SELECT * FROM companies WHERE id = ?",
193            [id.to_string()],
194            Self::row_to_company,
195        ) {
196            Ok(c) => Ok(Some(c)),
197            Err(rusqlite::Error::QueryReturnedNoRows) => Ok(None),
198            Err(e) => Err(map_db_error(e)),
199        }
200    }
201
202    fn update(&self, id: CompanyId, input: UpdateCompany) -> Result<Company> {
203        let id_str = id.to_string();
204        let now_str = Utc::now().to_rfc3339();
205
206        with_immediate_transaction(&self.pool, |tx| {
207            let mut sets = vec!["updated_at = ?".to_string()];
208            let mut params: Vec<Box<dyn rusqlite::types::ToSql>> = vec![Box::new(now_str.clone())];
209
210            if let Some(ref name) = input.name {
211                sets.push("name = ?".into());
212                params.push(Box::new(name.clone()));
213            }
214            if let Some(ref reference) = input.reference {
215                sets.push("reference = ?".into());
216                params.push(Box::new(reference.clone()));
217            }
218            if let Some(ref email) = input.email {
219                sets.push("email = ?".into());
220                params.push(Box::new(email.clone()));
221            }
222            if let Some(ref phone) = input.phone {
223                sets.push("phone = ?".into());
224                params.push(Box::new(phone.clone()));
225            }
226            if let Some(currency) = input.currency {
227                sets.push("currency = ?".into());
228                params.push(Box::new(currency.to_string()));
229            }
230            if let Some(terms) = input.payment_terms_days {
231                sets.push("payment_terms_days = ?".into());
232                params.push(Box::new(terms));
233            }
234            if let Some(status) = input.status {
235                sets.push("status = ?".into());
236                params.push(Box::new(status.to_string()));
237            }
238            if let Some(ref tags) = input.tags {
239                let json = serde_json::to_string(tags).map_err(|e| {
240                    rusqlite::Error::ToSqlConversionFailure(Box::new(CommerceError::DatabaseError(
241                        e.to_string(),
242                    )))
243                })?;
244                sets.push("tags = ?".into());
245                params.push(Box::new(json));
246            }
247            if let Some(ref metadata) = input.metadata {
248                let json = serde_json::to_string(metadata).map_err(|e| {
249                    rusqlite::Error::ToSqlConversionFailure(Box::new(CommerceError::DatabaseError(
250                        e.to_string(),
251                    )))
252                })?;
253                sets.push("metadata = ?".into());
254                params.push(Box::new(json));
255            }
256
257            let sql = format!("UPDATE companies SET {} WHERE id = ?", sets.join(", "));
258            params.push(Box::new(id_str.clone()));
259            let param_refs: Vec<&dyn rusqlite::types::ToSql> =
260                params.iter().map(|p| p.as_ref()).collect();
261            tx.execute(&sql, param_refs.as_slice())?;
262
263            tx.query_row("SELECT * FROM companies WHERE id = ?", [&id_str], Self::row_to_company)
264        })
265    }
266
267    fn list(&self, filter: CompanyFilter) -> Result<Vec<Company>> {
268        let conn = self.conn()?;
269        let mut sql = "SELECT * FROM companies WHERE 1=1".to_string();
270        let mut params: Vec<Box<dyn rusqlite::types::ToSql>> = vec![];
271
272        if let Some(status) = filter.status {
273            sql.push_str(" AND status = ?");
274            params.push(Box::new(status.to_string()));
275        }
276        if let Some(ref search) = filter.search {
277            sql.push_str(" AND (name LIKE ? ESCAPE '\\' OR reference LIKE ? ESCAPE '\\' OR email LIKE ? ESCAPE '\\')");
278            let pat = format!("%{}%", super::escape_like(search));
279            params.push(Box::new(pat.clone()));
280            params.push(Box::new(pat.clone()));
281            params.push(Box::new(pat));
282        }
283        sql.push_str(" ORDER BY name ASC");
284        crate::sqlite::append_limit_offset(&mut sql, filter.limit, filter.offset);
285
286        let param_refs: Vec<&dyn rusqlite::types::ToSql> =
287            params.iter().map(|p| p.as_ref()).collect();
288        let mut stmt = conn.prepare(&sql).map_err(map_db_error)?;
289        let rows = stmt
290            .query_map(param_refs.as_slice(), Self::row_to_company)
291            .map_err(map_db_error)?
292            .collect::<std::result::Result<Vec<_>, _>>()
293            .map_err(map_db_error)?;
294        Ok(rows)
295    }
296
297    fn delete(&self, id: CompanyId) -> Result<()> {
298        let conn = self.conn()?;
299        conn.execute("DELETE FROM companies WHERE id = ?", [id.to_string()])
300            .map_err(map_db_error)?;
301        Ok(())
302    }
303
304    fn list_addresses(&self, id: CompanyId) -> Result<Vec<CompanyShippingAddress>> {
305        let conn = self.conn()?;
306        let mut stmt = conn
307            .prepare("SELECT * FROM company_shipping_addresses WHERE company_id = ? ORDER BY is_default DESC, created_at ASC")
308            .map_err(map_db_error)?;
309        let rows = stmt
310            .query_map([id.to_string()], Self::row_to_address)
311            .map_err(map_db_error)?
312            .collect::<std::result::Result<Vec<_>, _>>()
313            .map_err(map_db_error)?;
314        Ok(rows)
315    }
316
317    fn list_price_overrides(&self, id: CompanyId) -> Result<Vec<CompanyPriceOverride>> {
318        let conn = self.conn()?;
319        let mut stmt = conn
320            .prepare("SELECT * FROM company_price_overrides WHERE company_id = ?")
321            .map_err(map_db_error)?;
322        let rows = stmt
323            .query_map([id.to_string()], Self::row_to_override)
324            .map_err(map_db_error)?
325            .collect::<std::result::Result<Vec<_>, _>>()
326            .map_err(map_db_error)?;
327        Ok(rows)
328    }
329
330    fn create_contact(&self, input: CreateContact) -> Result<Contact> {
331        if input.company_ids.is_empty() {
332            return Err(CommerceError::ValidationError(
333                "a contact must be linked to at least one company".into(),
334            ));
335        }
336        let id = ContactId::new();
337        let id_str = id.to_string();
338        let now_str = Utc::now().to_rfc3339();
339        let company_ids_json = serde_json::to_string(&input.company_ids)
340            .map_err(|e| CommerceError::DatabaseError(e.to_string()))?;
341
342        with_immediate_transaction(&self.pool, |tx| {
343            tx.execute(
344                "INSERT INTO contacts (id, first_name, last_name, email, phone, title, company_ids, portal_enabled, is_active, created_at, updated_at)
345                 VALUES (?, ?, ?, ?, ?, ?, ?, 0, 1, ?, ?)",
346                rusqlite::params![
347                    &id_str,
348                    &input.first_name,
349                    &input.last_name,
350                    &input.email,
351                    &input.phone,
352                    &input.title,
353                    &company_ids_json,
354                    &now_str,
355                    &now_str,
356                ],
357            )?;
358            tx.query_row("SELECT * FROM contacts WHERE id = ?", [&id_str], Self::row_to_contact)
359        })
360    }
361
362    fn get_contact(&self, id: ContactId) -> Result<Option<Contact>> {
363        let conn = self.conn()?;
364        match conn.query_row(
365            "SELECT * FROM contacts WHERE id = ?",
366            [id.to_string()],
367            Self::row_to_contact,
368        ) {
369            Ok(c) => Ok(Some(c)),
370            Err(rusqlite::Error::QueryReturnedNoRows) => Ok(None),
371            Err(e) => Err(map_db_error(e)),
372        }
373    }
374
375    fn list_contacts(&self, company_id: CompanyId) -> Result<Vec<Contact>> {
376        let conn = self.conn()?;
377        // company_ids is a JSON array of UUID strings; match by substring.
378        let mut stmt = conn
379            .prepare("SELECT * FROM contacts WHERE is_active = 1 AND company_ids LIKE ? ORDER BY first_name")
380            .map_err(map_db_error)?;
381        let needle = format!("%\"{company_id}\"%");
382        let rows = stmt
383            .query_map([needle], Self::row_to_contact)
384            .map_err(map_db_error)?
385            .collect::<std::result::Result<Vec<_>, _>>()
386            .map_err(map_db_error)?;
387        Ok(rows)
388    }
389}
390
391#[cfg(test)]
392mod tests {
393    use super::*;
394    use crate::DatabaseConfig;
395    use crate::sqlite::SqliteDatabase;
396
397    fn test_repo() -> SqliteCompanyRepository {
398        let db = SqliteDatabase::new(&DatabaseConfig::in_memory()).expect("in-memory db");
399        SqliteCompanyRepository::new(db.pool().clone())
400    }
401
402    fn new_company(repo: &SqliteCompanyRepository, name: &str) -> Company {
403        repo.create(CreateCompany {
404            name: name.into(),
405            reference: Some("ACME-1".into()),
406            email: Some("ap@acme.test".into()),
407            phone: None,
408            currency: Some(CurrencyCode::USD),
409            payment_terms_days: Some(30),
410            tags: vec!["wholesale".into()],
411            metadata: serde_json::Value::Null,
412        })
413        .expect("create company")
414    }
415
416    #[test]
417    fn create_get_update() {
418        let repo = test_repo();
419        let c = new_company(&repo, "Acme Inc");
420        assert_eq!(c.payment_terms_days, Some(30));
421        let fetched = repo.get(c.id).expect("get").expect("found");
422        assert_eq!(fetched.name, "Acme Inc");
423
424        let updated = repo
425            .update(c.id, UpdateCompany { name: Some("Acme LLC".into()), ..Default::default() })
426            .expect("update");
427        assert_eq!(updated.name, "Acme LLC");
428        assert_eq!(updated.payment_terms_days, Some(30));
429    }
430
431    #[test]
432    fn list_search_and_status_filter() {
433        let repo = test_repo();
434        new_company(&repo, "Globex");
435        new_company(&repo, "Initech");
436        let all = repo.list(CompanyFilter::default()).expect("list");
437        assert_eq!(all.len(), 2);
438        let found = repo
439            .list(CompanyFilter { search: Some("Glob".into()), ..Default::default() })
440            .expect("search");
441        assert_eq!(found.len(), 1);
442        assert_eq!(found[0].name, "Globex");
443    }
444
445    #[test]
446    fn contacts_link_and_list() {
447        let repo = test_repo();
448        let c = new_company(&repo, "Acme");
449        let contact = repo
450            .create_contact(CreateContact {
451                first_name: "Ada".into(),
452                last_name: Some("Byron".into()),
453                email: None,
454                phone: None,
455                title: Some("Buyer".into()),
456                company_ids: vec![c.id],
457            })
458            .expect("create contact");
459        assert!(contact.belongs_to(c.id));
460        let listed = repo.list_contacts(c.id).expect("list contacts");
461        assert_eq!(listed.len(), 1);
462        assert_eq!(listed[0].display_name(), "Ada Byron");
463    }
464
465    #[test]
466    fn contact_requires_company() {
467        let repo = test_repo();
468        let res = repo.create_contact(CreateContact {
469            first_name: "Solo".into(),
470            last_name: None,
471            email: None,
472            phone: None,
473            title: None,
474            company_ids: vec![],
475        });
476        assert!(res.is_err());
477    }
478}