//! SQLite implementations of the repository ports. use async_trait::async_trait; use chrono::{DateTime, Utc}; use domain::auth::{AuthEvent, RefreshToken}; use domain::ports::{AuditLog, RefreshTokenRepository, UserRepository}; use domain::user::{Role, User, UserUpdate}; use domain::DomainError; use sqlx::sqlite::SqliteRow; use sqlx::Row; use uuid::Uuid; use crate::DbPool; fn storage(e: sqlx::Error) -> DomainError { DomainError::Storage(e.to_string()) } fn parse_ts(s: &str) -> DateTime { DateTime::parse_from_rfc3339(s) .map(|d| d.with_timezone(&Utc)) .unwrap_or_default() } fn user_from_row(r: &SqliteRow) -> User { User { id: r.get::("id"), email: r.get("email"), display_name: r.get("display_name"), password_hash: r.get("password_hash"), role: Role::parse(r.get::("role").as_str()).unwrap_or(Role::User), is_active: r.get::("is_active"), created_at: parse_ts(r.get::("created_at").as_str()), } } const USER_COLS: &str = "id, email, display_name, password_hash, role, is_active, created_at"; pub struct SqliteUsers(pub DbPool); #[async_trait] impl UserRepository for SqliteUsers { async fn find_by_id(&self, id: Uuid) -> Result, DomainError> { sqlx::query(&format!("SELECT {USER_COLS} FROM users WHERE id = ?")) .bind(id) .fetch_optional(&self.0) .await .map(|r| r.as_ref().map(user_from_row)) .map_err(storage) } async fn find_by_email(&self, email: &str) -> Result, DomainError> { sqlx::query(&format!("SELECT {USER_COLS} FROM users WHERE email = ?")) .bind(email) .fetch_optional(&self.0) .await .map(|r| r.as_ref().map(user_from_row)) .map_err(storage) } async fn list(&self) -> Result, DomainError> { sqlx::query(&format!("SELECT {USER_COLS} FROM users ORDER BY email")) .fetch_all(&self.0) .await .map(|rows| rows.iter().map(user_from_row).collect()) .map_err(storage) } async fn count(&self) -> Result { let n: i64 = sqlx::query_scalar("SELECT COUNT(*) FROM users") .fetch_one(&self.0) .await .map_err(storage)?; Ok(n as u64) } async fn count_active_admins(&self) -> Result { let n: i64 = sqlx::query_scalar("SELECT COUNT(*) FROM users WHERE role = 'admin' AND is_active = 1") .fetch_one(&self.0) .await .map_err(storage)?; Ok(n as u64) } async fn insert(&self, u: &User) -> Result<(), DomainError> { sqlx::query(&format!( "INSERT INTO users ({USER_COLS}) VALUES (?, ?, ?, ?, ?, ?, ?)" )) .bind(u.id) .bind(&u.email) .bind(&u.display_name) .bind(&u.password_hash) .bind(u.role.as_str()) .bind(u.is_active) .bind(u.created_at.to_rfc3339()) .execute(&self.0) .await .map(|_| ()) .map_err(storage) } async fn update(&self, id: Uuid, up: &UserUpdate) -> Result { let res = sqlx::query( "UPDATE users SET display_name = COALESCE(?, display_name), role = COALESCE(?, role), \ is_active = COALESCE(?, is_active) WHERE id = ?", ) .bind(&up.display_name) .bind(up.role.map(Role::as_str)) .bind(up.is_active) .bind(id) .execute(&self.0) .await .map_err(storage)?; if res.rows_affected() == 0 { return Err(DomainError::NotFound); } self.find_by_id(id).await?.ok_or(DomainError::NotFound) } async fn set_password_hash(&self, id: Uuid, hash: &str) -> Result<(), DomainError> { let res = sqlx::query("UPDATE users SET password_hash = ? WHERE id = ?") .bind(hash) .bind(id) .execute(&self.0) .await .map_err(storage)?; (res.rows_affected() > 0) .then_some(()) .ok_or(DomainError::NotFound) } } pub struct SqliteRefreshTokens(pub DbPool); #[async_trait] impl RefreshTokenRepository for SqliteRefreshTokens { async fn insert(&self, t: &RefreshToken) -> Result<(), DomainError> { sqlx::query("INSERT INTO refresh_tokens (id, user_id, family, token_hash, expires_at, revoked) VALUES (?, ?, ?, ?, ?, ?)") .bind(t.id) .bind(t.user_id) .bind(t.family) .bind(&t.token_hash) .bind(t.expires_at.to_rfc3339()) .bind(t.revoked) .execute(&self.0) .await .map(|_| ()) .map_err(storage) } async fn find_by_hash(&self, hash: &str) -> Result, DomainError> { sqlx::query("SELECT id, user_id, family, token_hash, expires_at, revoked FROM refresh_tokens WHERE token_hash = ?") .bind(hash) .fetch_optional(&self.0) .await .map(|row| { row.map(|r| RefreshToken { id: r.get("id"), user_id: r.get("user_id"), family: r.get("family"), token_hash: r.get("token_hash"), expires_at: parse_ts(r.get::("expires_at").as_str()), revoked: r.get("revoked"), }) }) .map_err(storage) } async fn revoke(&self, id: Uuid) -> Result<(), DomainError> { sqlx::query("UPDATE refresh_tokens SET revoked = 1 WHERE id = ?") .bind(id) .execute(&self.0) .await .map(|_| ()) .map_err(storage) } async fn revoke_family(&self, family: Uuid) -> Result<(), DomainError> { sqlx::query("UPDATE refresh_tokens SET revoked = 1 WHERE family = ?") .bind(family) .execute(&self.0) .await .map(|_| ()) .map_err(storage) } } pub struct SqliteAuditLog(pub DbPool); #[async_trait] impl AuditLog for SqliteAuditLog { async fn record(&self, e: &AuthEvent) -> Result<(), DomainError> { sqlx::query("INSERT INTO auth_events (user_id, email, kind, ip, at) VALUES (?, ?, ?, ?, ?)") .bind(e.user_id) .bind(&e.email) .bind(e.kind.as_str()) .bind(&e.ip) .bind(e.at.to_rfc3339()) .execute(&self.0) .await .map(|_| ()) .map_err(storage) } } #[cfg(test)] mod tests { use super::*; #[tokio::test] async fn users_roundtrip_and_update() { let pool = crate::connect("sqlite::memory:").await.unwrap(); let repo = SqliteUsers(pool); let u = User { id: Uuid::new_v4(), email: "a@x.de".into(), display_name: "A".into(), password_hash: "h".into(), role: Role::Admin, is_active: true, created_at: Utc::now(), }; repo.insert(&u).await.unwrap(); assert_eq!( repo.find_by_email("a@x.de").await.unwrap().unwrap().id, u.id ); assert_eq!(repo.count_active_admins().await.unwrap(), 1); let updated = repo .update( u.id, &UserUpdate { is_active: Some(false), ..Default::default() }, ) .await .unwrap(); assert!(!updated.is_active); assert_eq!(updated.display_name, "A"); assert_eq!( repo.update(Uuid::new_v4(), &UserUpdate::default()) .await .unwrap_err(), DomainError::NotFound ); } #[tokio::test] async fn refresh_tokens_revoke_by_family() { let pool = crate::connect("sqlite::memory:").await.unwrap(); let users = SqliteUsers(pool.clone()); let u = User { id: Uuid::new_v4(), email: "a@x.de".into(), display_name: "A".into(), password_hash: "h".into(), role: Role::User, is_active: true, created_at: Utc::now(), }; users.insert(&u).await.unwrap(); let repo = SqliteRefreshTokens(pool); let family = Uuid::new_v4(); for h in ["h1", "h2"] { repo.insert(&RefreshToken { id: Uuid::new_v4(), user_id: u.id, family, token_hash: h.into(), expires_at: Utc::now() + chrono::Duration::days(1), revoked: false, }) .await .unwrap(); } repo.revoke_family(family).await.unwrap(); assert!(repo.find_by_hash("h1").await.unwrap().unwrap().revoked); assert!(repo.find_by_hash("h2").await.unwrap().unwrap().revoked); assert!(repo.find_by_hash("nope").await.unwrap().is_none()); } }