Skip to main content

stateset_db/sqlite/
fraud.rs

1//! SQLite implementation of fraud detection repository
2
3use 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                    &notes,
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}