1use 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 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}