1use super::{
4 map_db_error, parse_datetime_row, parse_enum_row, parse_json_row, parse_uuid_row,
5 with_immediate_transaction,
6};
7use chrono::Utc;
8use r2d2::Pool;
9use r2d2_sqlite::SqliteConnectionManager;
10use stateset_core::{
11 CommerceError, CreateFraudAssessment, CreateFraudRule, FraudAssessment, FraudAssessmentFilter,
12 FraudDecision, FraudRepository, FraudRule, FraudRuleFilter, FraudRuleId, FraudSignal, OrderId,
13 Result, UpdateFraudRule,
14};
15
16#[derive(Debug)]
17pub struct SqliteFraudRepository {
18 pool: Pool<SqliteConnectionManager>,
19}
20
21impl SqliteFraudRepository {
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_assessment(row: &rusqlite::Row<'_>) -> rusqlite::Result<FraudAssessment> {
32 let signals_json: String = row.get("signals")?;
33 let signals: Vec<FraudSignal> =
34 parse_json_row(&signals_json, "fraud_assessment", "signals")?;
35
36 Ok(FraudAssessment {
37 order_id: parse_uuid_row(
38 &row.get::<_, String>("order_id")?,
39 "fraud_assessment",
40 "order_id",
41 )?
42 .into(),
43 risk_score: row.get("risk_score")?,
44 signals,
45 decision: parse_enum_row(
46 &row.get::<_, String>("decision")?,
47 "fraud_assessment",
48 "decision",
49 )?,
50 reviewed_by: row.get("reviewed_by")?,
51 review_notes: row.get("review_notes")?,
52 created_at: parse_datetime_row(
53 &row.get::<_, String>("created_at")?,
54 "fraud_assessment",
55 "created_at",
56 )?,
57 updated_at: parse_datetime_row(
58 &row.get::<_, String>("updated_at")?,
59 "fraud_assessment",
60 "updated_at",
61 )?,
62 })
63 }
64
65 fn row_to_rule(row: &rusqlite::Row<'_>) -> rusqlite::Result<FraudRule> {
66 Ok(FraudRule {
67 id: parse_uuid_row(&row.get::<_, String>("id")?, "fraud_rule", "id")?.into(),
68 name: row.get("name")?,
69 description: row.get("description")?,
70 signal_type: parse_enum_row(
71 &row.get::<_, String>("signal_type")?,
72 "fraud_rule",
73 "signal_type",
74 )?,
75 threshold: row.get("threshold")?,
76 action: parse_enum_row(&row.get::<_, String>("action")?, "fraud_rule", "action")?,
77 enabled: row.get::<_, i32>("enabled")? != 0,
78 created_at: parse_datetime_row(
79 &row.get::<_, String>("created_at")?,
80 "fraud_rule",
81 "created_at",
82 )?,
83 updated_at: parse_datetime_row(
84 &row.get::<_, String>("updated_at")?,
85 "fraud_rule",
86 "updated_at",
87 )?,
88 })
89 }
90}
91
92impl FraudRepository for SqliteFraudRepository {
93 fn create_assessment(&self, input: CreateFraudAssessment) -> Result<FraudAssessment> {
94 let now = Utc::now();
95 let now_str = now.to_rfc3339();
96 let order_id_str = input.order_id.to_string();
97
98 let signals: Vec<FraudSignal> = input
99 .signals
100 .into_iter()
101 .map(|s| FraudSignal {
102 order_id: input.order_id,
103 signal_type: s.signal_type,
104 score: s.score,
105 details: s.details,
106 detected_at: now,
107 })
108 .collect();
109
110 let risk_score = FraudAssessment::calculate_risk_score(&signals);
111 let decision =
112 if risk_score >= 0.8 { FraudDecision::Review } else { FraudDecision::Accept };
113
114 let signals_json = serde_json::to_string(&signals)
115 .map_err(|e| CommerceError::DatabaseError(e.to_string()))?;
116
117 with_immediate_transaction(&self.pool, |tx| {
118 tx.execute(
119 "INSERT INTO fraud_assessments (order_id, risk_score, signals, decision, reviewed_by, review_notes, created_at, updated_at)
120 VALUES (?, ?, ?, ?, NULL, NULL, ?, ?)",
121 rusqlite::params![
122 &order_id_str,
123 risk_score,
124 &signals_json,
125 decision.to_string(),
126 &now_str,
127 &now_str,
128 ],
129 )?;
130
131 tx.query_row(
132 "SELECT * FROM fraud_assessments WHERE order_id = ?",
133 [&order_id_str],
134 Self::row_to_assessment,
135 )
136 })
137 }
138
139 fn get_assessment(&self, order_id: OrderId) -> Result<Option<FraudAssessment>> {
140 let conn = self.conn()?;
141 match conn.query_row(
142 "SELECT * FROM fraud_assessments WHERE order_id = ?",
143 [order_id.to_string()],
144 Self::row_to_assessment,
145 ) {
146 Ok(a) => Ok(Some(a)),
147 Err(rusqlite::Error::QueryReturnedNoRows) => Ok(None),
148 Err(e) => Err(map_db_error(e)),
149 }
150 }
151
152 fn list_assessments(&self, filter: FraudAssessmentFilter) -> Result<Vec<FraudAssessment>> {
153 let conn = self.conn()?;
154 let mut sql = "SELECT * FROM fraud_assessments WHERE 1=1".to_string();
155 let mut params: Vec<Box<dyn rusqlite::types::ToSql>> = vec![];
156
157 if let Some(decision) = filter.decision {
158 sql.push_str(" AND decision = ?");
159 params.push(Box::new(decision.to_string()));
160 }
161 if let Some(min_score) = filter.min_risk_score {
162 sql.push_str(" AND risk_score >= ?");
163 params.push(Box::new(min_score));
164 }
165 if filter.unreviewed_only == Some(true) {
166 sql.push_str(" AND reviewed_by IS NULL");
167 }
168
169 sql.push_str(" ORDER BY created_at DESC");
170
171 crate::sqlite::append_limit_offset(&mut sql, filter.limit, filter.offset);
172
173 let param_refs: Vec<&dyn rusqlite::types::ToSql> =
174 params.iter().map(|p| p.as_ref()).collect();
175 let mut stmt = conn.prepare(&sql).map_err(map_db_error)?;
176 let rows = stmt
177 .query_map(param_refs.as_slice(), Self::row_to_assessment)
178 .map_err(map_db_error)?
179 .collect::<std::result::Result<Vec<_>, _>>()
180 .map_err(map_db_error)?;
181 Ok(rows)
182 }
183
184 fn review_assessment(
185 &self,
186 order_id: OrderId,
187 decision: FraudDecision,
188 reviewer: String,
189 notes: Option<String>,
190 ) -> Result<FraudAssessment> {
191 let order_id_str = order_id.to_string();
192 let now_str = Utc::now().to_rfc3339();
193
194 with_immediate_transaction(&self.pool, |tx| {
195 tx.execute(
196 "UPDATE fraud_assessments SET decision = ?, reviewed_by = ?, review_notes = ?, updated_at = ? WHERE order_id = ?",
197 rusqlite::params![
198 decision.to_string(),
199 &reviewer,
200 ¬es,
201 &now_str,
202 &order_id_str,
203 ],
204 )?;
205
206 tx.query_row(
207 "SELECT * FROM fraud_assessments WHERE order_id = ?",
208 [&order_id_str],
209 Self::row_to_assessment,
210 )
211 })
212 }
213
214 fn create_rule(&self, input: CreateFraudRule) -> Result<FraudRule> {
215 let id = FraudRuleId::new();
216 let now = Utc::now();
217 let id_str = id.to_string();
218 let now_str = now.to_rfc3339();
219
220 with_immediate_transaction(&self.pool, |tx| {
221 tx.execute(
222 "INSERT INTO fraud_rules (id, name, description, signal_type, threshold, action, enabled, created_at, updated_at)
223 VALUES (?, ?, ?, ?, ?, ?, 1, ?, ?)",
224 rusqlite::params![
225 &id_str,
226 &input.name,
227 &input.description,
228 input.signal_type.to_string(),
229 input.threshold,
230 input.action.to_string(),
231 &now_str,
232 &now_str,
233 ],
234 )?;
235
236 tx.query_row("SELECT * FROM fraud_rules WHERE id = ?", [&id_str], Self::row_to_rule)
237 })
238 }
239
240 fn get_rule(&self, id: FraudRuleId) -> Result<Option<FraudRule>> {
241 let conn = self.conn()?;
242 match conn.query_row(
243 "SELECT * FROM fraud_rules WHERE id = ?",
244 [id.to_string()],
245 Self::row_to_rule,
246 ) {
247 Ok(r) => Ok(Some(r)),
248 Err(rusqlite::Error::QueryReturnedNoRows) => Ok(None),
249 Err(e) => Err(map_db_error(e)),
250 }
251 }
252
253 fn update_rule(&self, id: FraudRuleId, input: UpdateFraudRule) -> Result<FraudRule> {
254 let id_str = id.to_string();
255 let now_str = Utc::now().to_rfc3339();
256
257 with_immediate_transaction(&self.pool, |tx| {
258 let mut sets = vec!["updated_at = ?".to_string()];
259 let mut params: Vec<Box<dyn rusqlite::types::ToSql>> = vec![Box::new(now_str.clone())];
260
261 if let Some(ref name) = input.name {
262 sets.push("name = ?".into());
263 params.push(Box::new(name.clone()));
264 }
265 if let Some(ref desc) = input.description {
266 sets.push("description = ?".into());
267 params.push(Box::new(desc.clone()));
268 }
269 if let Some(threshold) = input.threshold {
270 sets.push("threshold = ?".into());
271 params.push(Box::new(threshold));
272 }
273 if let Some(action) = input.action {
274 sets.push("action = ?".into());
275 params.push(Box::new(action.to_string()));
276 }
277 if let Some(enabled) = input.enabled {
278 sets.push("enabled = ?".into());
279 params.push(Box::new(enabled as i32));
280 }
281
282 let sql = format!("UPDATE fraud_rules SET {} WHERE id = ?", sets.join(", "));
283 params.push(Box::new(id_str.clone()));
284
285 let param_refs: Vec<&dyn rusqlite::types::ToSql> =
286 params.iter().map(|p| p.as_ref()).collect();
287 tx.execute(&sql, param_refs.as_slice())?;
288
289 tx.query_row("SELECT * FROM fraud_rules WHERE id = ?", [&id_str], Self::row_to_rule)
290 })
291 }
292
293 fn list_rules(&self, filter: FraudRuleFilter) -> Result<Vec<FraudRule>> {
294 let conn = self.conn()?;
295 let mut sql = "SELECT * FROM fraud_rules WHERE 1=1".to_string();
296 let mut params: Vec<Box<dyn rusqlite::types::ToSql>> = vec![];
297
298 if let Some(signal_type) = filter.signal_type {
299 sql.push_str(" AND signal_type = ?");
300 params.push(Box::new(signal_type.to_string()));
301 }
302 if let Some(action) = filter.action {
303 sql.push_str(" AND action = ?");
304 params.push(Box::new(action.to_string()));
305 }
306 if let Some(enabled) = filter.enabled {
307 sql.push_str(" AND enabled = ?");
308 params.push(Box::new(enabled as i32));
309 }
310
311 sql.push_str(" ORDER BY created_at DESC");
312
313 crate::sqlite::append_limit_offset(&mut sql, filter.limit, filter.offset);
314
315 let param_refs: Vec<&dyn rusqlite::types::ToSql> =
316 params.iter().map(|p| p.as_ref()).collect();
317 let mut stmt = conn.prepare(&sql).map_err(map_db_error)?;
318 let rows = stmt
319 .query_map(param_refs.as_slice(), Self::row_to_rule)
320 .map_err(map_db_error)?
321 .collect::<std::result::Result<Vec<_>, _>>()
322 .map_err(map_db_error)?;
323 Ok(rows)
324 }
325
326 fn delete_rule(&self, id: FraudRuleId) -> Result<()> {
327 let conn = self.conn()?;
328 conn.execute("DELETE FROM fraud_rules WHERE id = ?", [id.to_string()])
329 .map_err(map_db_error)?;
330 Ok(())
331 }
332
333 fn get_active_rules(&self) -> Result<Vec<FraudRule>> {
334 let conn = self.conn()?;
335 let mut stmt = conn
336 .prepare("SELECT * FROM fraud_rules WHERE enabled = 1 ORDER BY created_at DESC")
337 .map_err(map_db_error)?;
338 let rows = stmt
339 .query_map([], Self::row_to_rule)
340 .map_err(map_db_error)?
341 .collect::<std::result::Result<Vec<_>, _>>()
342 .map_err(map_db_error)?;
343 Ok(rows)
344 }
345}
346
347#[cfg(test)]
348mod tests {
349 use super::*;
350 use crate::DatabaseConfig;
351 use crate::sqlite::SqliteDatabase;
352 use stateset_core::{CreateFraudSignal, FraudSignalType};
353
354 fn test_repo() -> SqliteFraudRepository {
355 let db = SqliteDatabase::new(&DatabaseConfig::in_memory()).expect("in-memory db");
356 let conn = db.conn().expect("conn");
357 conn.execute_batch(
358 "CREATE TABLE IF NOT EXISTS fraud_assessments (
359 order_id TEXT PRIMARY KEY,
360 risk_score REAL NOT NULL DEFAULT 0.0,
361 signals TEXT NOT NULL DEFAULT '[]',
362 decision TEXT NOT NULL DEFAULT 'accept',
363 reviewed_by TEXT,
364 review_notes TEXT,
365 created_at TEXT NOT NULL DEFAULT (datetime('now')),
366 updated_at TEXT NOT NULL DEFAULT (datetime('now'))
367 );
368 CREATE TABLE IF NOT EXISTS fraud_rules (
369 id TEXT PRIMARY KEY,
370 name TEXT NOT NULL,
371 description TEXT,
372 signal_type TEXT NOT NULL,
373 threshold REAL NOT NULL DEFAULT 0.5,
374 action TEXT NOT NULL DEFAULT 'review',
375 enabled INTEGER NOT NULL DEFAULT 1,
376 created_at TEXT NOT NULL DEFAULT (datetime('now')),
377 updated_at TEXT NOT NULL DEFAULT (datetime('now'))
378 );",
379 )
380 .expect("create tables");
381 SqliteFraudRepository::new(db.pool().clone())
382 }
383
384 #[test]
385 fn create_and_get_assessment() {
386 let repo = test_repo();
387 let order_id = OrderId::new();
388 let assessment = repo
389 .create_assessment(CreateFraudAssessment {
390 order_id,
391 signals: vec![CreateFraudSignal {
392 signal_type: FraudSignalType::VelocitySpike,
393 score: 0.6,
394 details: "High order velocity".into(),
395 }],
396 })
397 .expect("create assessment");
398
399 assert_eq!(assessment.order_id, order_id);
400 assert!((assessment.risk_score - 0.6).abs() < f64::EPSILON);
401 assert_eq!(assessment.signals.len(), 1);
402
403 let fetched = repo.get_assessment(order_id).expect("get").expect("found");
404 assert_eq!(fetched.order_id, order_id);
405 }
406
407 #[test]
408 fn create_and_delete_rule() {
409 let repo = test_repo();
410 let rule = repo
411 .create_rule(CreateFraudRule {
412 name: "Velocity check".into(),
413 description: Some("Block fast orders".into()),
414 signal_type: FraudSignalType::VelocitySpike,
415 threshold: 0.8,
416 action: FraudDecision::Reject,
417 })
418 .expect("create rule");
419
420 assert_eq!(rule.name, "Velocity check");
421 assert!(rule.enabled);
422
423 let fetched = repo.get_rule(rule.id).expect("get").expect("found");
424 assert_eq!(fetched.id, rule.id);
425
426 repo.delete_rule(rule.id).expect("delete");
427 assert!(repo.get_rule(rule.id).expect("get after delete").is_none());
428 }
429}