diff --git a/crates/tranquil-comms/Cargo.toml b/crates/tranquil-comms/Cargo.toml index 87cdc1e..aa8cb13 100644 --- a/crates/tranquil-comms/Cargo.toml +++ b/crates/tranquil-comms/Cargo.toml @@ -14,5 +14,6 @@ serde_json = { workspace = true } sqlx = { workspace = true } thiserror = { workspace = true } tokio = { workspace = true } +tranquil-db-traits = { workspace = true } urlencoding = { workspace = true } uuid = { workspace = true } diff --git a/crates/tranquil-comms/src/types.rs b/crates/tranquil-comms/src/types.rs index cb14ce0..2af7bd7 100644 --- a/crates/tranquil-comms/src/types.rs +++ b/crates/tranquil-comms/src/types.rs @@ -1,63 +1,6 @@ -use chrono::{DateTime, Utc}; -use serde::{Deserialize, Serialize}; use uuid::Uuid; -#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize, sqlx::Type)] -#[serde(rename_all = "lowercase")] -#[sqlx(type_name = "comms_channel", rename_all = "lowercase")] -pub enum CommsChannel { - Email, - Discord, - Telegram, - Signal, -} - -#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, sqlx::Type)] -#[serde(rename_all = "lowercase")] -#[sqlx(type_name = "comms_status", rename_all = "lowercase")] -pub enum CommsStatus { - Pending, - Processing, - Sent, - Failed, -} - -#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, sqlx::Type)] -#[serde(rename_all = "snake_case")] -#[sqlx(type_name = "comms_type", rename_all = "snake_case")] -pub enum CommsType { - Welcome, - EmailVerification, - PasswordReset, - EmailUpdate, - AccountDeletion, - AdminEmail, - PlcOperation, - TwoFactorCode, - PasskeyRecovery, - LegacyLoginAlert, - MigrationVerification, -} - -#[derive(Debug, Clone)] -pub struct QueuedComms { - pub id: Uuid, - pub user_id: Uuid, - pub channel: CommsChannel, - pub comms_type: CommsType, - pub status: CommsStatus, - pub recipient: String, - pub subject: Option, - pub body: String, - pub metadata: Option, - pub attempts: i32, - pub max_attempts: i32, - pub last_error: Option, - pub created_at: DateTime, - pub updated_at: DateTime, - pub scheduled_for: DateTime, - pub processed_at: Option>, -} +pub use tranquil_db_traits::{CommsChannel, CommsStatus, CommsType, QueuedComms}; pub struct NewComms { pub user_id: Uuid, diff --git a/crates/tranquil-pds/Cargo.toml b/crates/tranquil-pds/Cargo.toml index ded273b..81f3b46 100644 --- a/crates/tranquil-pds/Cargo.toml +++ b/crates/tranquil-pds/Cargo.toml @@ -15,6 +15,8 @@ tranquil-scopes = { workspace = true } tranquil-auth = { workspace = true } tranquil-oauth = { workspace = true } tranquil-comms = { workspace = true } +tranquil-db = { workspace = true } +tranquil-db-traits = { workspace = true } aes-gcm = { workspace = true } backon = { workspace = true } diff --git a/crates/tranquil-pds/src/api/actor/preferences.rs b/crates/tranquil-pds/src/api/actor/preferences.rs index 44d9d17..16213e1 100644 --- a/crates/tranquil-pds/src/api/actor/preferences.rs +++ b/crates/tranquil-pds/src/api/actor/preferences.rs @@ -38,23 +38,13 @@ pub async fn get_preferences( ) -> Response { let auth_user = auth.0; let has_full_access = auth_user.permissions().has_full_access(); - let user_id: uuid::Uuid = - match sqlx::query_scalar!("SELECT id FROM users WHERE did = $1", &*auth_user.did) - .fetch_optional(&state.db) - .await - { - Ok(Some(id)) => id, - _ => { - return ApiError::InternalError(Some("User not found".into())).into_response(); - } - }; - let prefs_result = sqlx::query!( - "SELECT name, value_json FROM account_preferences WHERE user_id = $1", - user_id - ) - .fetch_all(&state.db) - .await; - let prefs = match prefs_result { + let user_id: uuid::Uuid = match state.user_repo.get_id_by_did(&auth_user.did).await { + Ok(Some(id)) => id, + _ => { + return ApiError::InternalError(Some("User not found".into())).into_response(); + } + }; + let prefs = match state.infra_repo.get_account_preferences(user_id).await { Ok(rows) => rows, Err(_) => { return ApiError::InternalError(Some("Failed to fetch preferences".into())) @@ -64,21 +54,20 @@ pub async fn get_preferences( let mut personal_details_pref: Option = None; let mut preferences: Vec = prefs .into_iter() - .filter(|row| { - row.name == APP_BSKY_NAMESPACE - || row.name.starts_with(&format!("{}.", APP_BSKY_NAMESPACE)) + .filter(|(name, _)| { + name == APP_BSKY_NAMESPACE || name.starts_with(&format!("{}.", APP_BSKY_NAMESPACE)) }) - .filter_map(|row| { - if row.name == DECLARED_AGE_PREF { + .filter_map(|(name, value_json)| { + if name == DECLARED_AGE_PREF { return None; } - if row.name == PERSONAL_DETAILS_PREF { + if name == PERSONAL_DETAILS_PREF { if !has_full_access { return None; } - personal_details_pref = serde_json::from_value(row.value_json.clone()).ok(); + personal_details_pref = serde_json::from_value(value_json.clone()).ok(); } - serde_json::from_value(row.value_json).ok() + serde_json::from_value(value_json).ok() }) .collect(); if let Some(age) = personal_details_pref @@ -109,16 +98,12 @@ pub async fn put_preferences( ) -> Response { let auth_user = auth.0; let has_full_access = auth_user.permissions().has_full_access(); - let user_id: uuid::Uuid = - match sqlx::query_scalar!("SELECT id FROM users WHERE did = $1", &*auth_user.did) - .fetch_optional(&state.db) - .await - { - Ok(Some(id)) => id, - _ => { - return ApiError::InternalError(Some("User not found".into())).into_response(); - } - }; + let user_id: uuid::Uuid = match state.user_repo.get_id_by_did(&auth_user.did).await { + Ok(Some(id)) => id, + _ => { + return ApiError::InternalError(Some("User not found".into())).into_response(); + } + }; if input.preferences.len() > MAX_PREFERENCES_COUNT { return ApiError::InvalidRequest(format!( "Too many preferences: {} exceeds limit of {}", @@ -195,50 +180,24 @@ pub async fn put_preferences( )) .into_response(); } - let mut tx = match state.db.begin().await { - Ok(tx) => tx, - Err(_) => { - return ApiError::InternalError(Some("Failed to start transaction".into())) - .into_response(); - } - }; - let delete_result = sqlx::query!( - "DELETE FROM account_preferences WHERE user_id = $1 AND (name = $2 OR name LIKE $3)", - user_id, - APP_BSKY_NAMESPACE, - format!("{}.%", APP_BSKY_NAMESPACE) - ) - .execute(&mut *tx) - .await; - if delete_result.is_err() { - let _ = tx.rollback().await; - return ApiError::InternalError(Some("Failed to clear preferences".into())).into_response(); - } - for pref in input.preferences { - let pref_type = match pref.get("$type").and_then(|t| t.as_str()) { - Some(t) => t, - None => continue, - }; - if pref_type == DECLARED_AGE_PREF { - continue; - } - let insert_result = sqlx::query!( - "INSERT INTO account_preferences (user_id, name, value_json) VALUES ($1, $2, $3)", - user_id, - pref_type, - pref - ) - .execute(&mut *tx) - .await; - if insert_result.is_err() { - let _ = tx.rollback().await; - return ApiError::InternalError(Some("Failed to save preference".into())) - .into_response(); - } - } - if tx.commit().await.is_err() { - return ApiError::InternalError(Some("Failed to commit transaction".into())) - .into_response(); + let prefs_to_save: Vec<(String, Value)> = input + .preferences + .into_iter() + .filter_map(|pref| { + let pref_type = pref.get("$type").and_then(|t| t.as_str())?; + if pref_type == DECLARED_AGE_PREF { + return None; + } + Some((pref_type.to_string(), pref)) + }) + .collect(); + + if let Err(_) = state + .infra_repo + .replace_namespace_preferences(user_id, APP_BSKY_NAMESPACE, prefs_to_save) + .await + { + return ApiError::InternalError(Some("Failed to save preferences".into())).into_response(); } StatusCode::OK.into_response() } diff --git a/crates/tranquil-pds/src/api/admin/account/delete.rs b/crates/tranquil-pds/src/api/admin/account/delete.rs index 412cf78..78ad075 100644 --- a/crates/tranquil-pds/src/api/admin/account/delete.rs +++ b/crates/tranquil-pds/src/api/admin/account/delete.rs @@ -22,10 +22,7 @@ pub async fn delete_account( Json(input): Json, ) -> Response { let did = &input.did; - let user = sqlx::query!("SELECT id, handle FROM users WHERE did = $1", did.as_str()) - .fetch_optional(&state.db) - .await; - let (user_id, handle) = match user { + let (user_id, handle) = match state.user_repo.get_id_and_handle_by_did(did).await { Ok(Some(row)) => (row.id, row.handle), Ok(None) => { return ApiError::AccountNotFound.into_response(); @@ -35,100 +32,13 @@ pub async fn delete_account( return ApiError::InternalError(None).into_response(); } }; - let mut tx = match state.db.begin().await { - Ok(tx) => tx, - Err(e) => { - error!("Failed to begin transaction for account deletion: {:?}", e); - return ApiError::InternalError(None).into_response(); - } - }; - if let Err(e) = sqlx::query!("DELETE FROM session_tokens WHERE did = $1", did.as_str()) - .execute(&mut *tx) + if let Err(e) = state + .user_repo + .admin_delete_account_complete(user_id, did) .await { - error!("Failed to delete session tokens for {}: {:?}", did, e); - return ApiError::InternalError(Some("Failed to delete session tokens".into())) - .into_response(); - } - if let Err(e) = sqlx::query!("DELETE FROM used_refresh_tokens WHERE session_id IN (SELECT id FROM session_tokens WHERE did = $1)", did.as_str()) - .execute(&mut *tx) - .await - { - error!("Failed to delete used refresh tokens for {}: {:?}", did, e); - } - if let Err(e) = sqlx::query!("DELETE FROM records WHERE repo_id = $1", user_id) - .execute(&mut *tx) - .await - { - error!("Failed to delete records for user {}: {:?}", user_id, e); - return ApiError::InternalError(Some("Failed to delete records".into())).into_response(); - } - if let Err(e) = sqlx::query!("DELETE FROM repos WHERE user_id = $1", user_id) - .execute(&mut *tx) - .await - { - error!("Failed to delete repos for user {}: {:?}", user_id, e); - return ApiError::InternalError(Some("Failed to delete repos".into())).into_response(); - } - if let Err(e) = sqlx::query!("DELETE FROM blobs WHERE created_by_user = $1", user_id) - .execute(&mut *tx) - .await - { - error!("Failed to delete blobs for user {}: {:?}", user_id, e); - return ApiError::InternalError(Some("Failed to delete blobs".into())).into_response(); - } - if let Err(e) = sqlx::query!("DELETE FROM app_passwords WHERE user_id = $1", user_id) - .execute(&mut *tx) - .await - { - error!( - "Failed to delete app passwords for user {}: {:?}", - user_id, e - ); - return ApiError::InternalError(Some("Failed to delete app passwords".into())) - .into_response(); - } - if let Err(e) = sqlx::query!( - "DELETE FROM invite_code_uses WHERE used_by_user = $1", - user_id - ) - .execute(&mut *tx) - .await - { - error!( - "Failed to delete invite code uses for user {}: {:?}", - user_id, e - ); - } - if let Err(e) = sqlx::query!( - "DELETE FROM invite_codes WHERE created_by_user = $1", - user_id - ) - .execute(&mut *tx) - .await - { - error!( - "Failed to delete invite codes for user {}: {:?}", - user_id, e - ); - } - if let Err(e) = sqlx::query!("DELETE FROM user_keys WHERE user_id = $1", user_id) - .execute(&mut *tx) - .await - { - error!("Failed to delete user keys for user {}: {:?}", user_id, e); - return ApiError::InternalError(Some("Failed to delete user keys".into())).into_response(); - } - if let Err(e) = sqlx::query!("DELETE FROM users WHERE id = $1", user_id) - .execute(&mut *tx) - .await - { - error!("Failed to delete user {}: {:?}", user_id, e); - return ApiError::InternalError(Some("Failed to delete user".into())).into_response(); - } - if let Err(e) = tx.commit().await { - error!("Failed to commit account deletion transaction: {:?}", e); - return ApiError::InternalError(Some("Failed to commit deletion".into())).into_response(); + error!("Failed to delete account {}: {:?}", did, e); + return ApiError::InternalError(Some("Failed to delete account".into())).into_response(); } if let Err(e) = crate::api::repo::record::sequence_account_event(&state, did, false, Some("deleted")).await diff --git a/crates/tranquil-pds/src/api/admin/account/email.rs b/crates/tranquil-pds/src/api/admin/account/email.rs index 890a273..837f6b9 100644 --- a/crates/tranquil-pds/src/api/admin/account/email.rs +++ b/crates/tranquil-pds/src/api/admin/account/email.rs @@ -35,22 +35,8 @@ pub async fn send_email( if content.is_empty() { return ApiError::InvalidRequest("content is required".into()).into_response(); } - let user = sqlx::query!( - "SELECT id, email, handle FROM users WHERE did = $1", - input.recipient_did.as_str() - ) - .fetch_optional(&state.db) - .await; - let (user_id, email, handle) = match user { - Ok(Some(row)) => { - let email = match row.email { - Some(e) => e, - None => { - return ApiError::NoEmail.into_response(); - } - }; - (row.id, email, row.handle) - } + let user = match state.user_repo.get_by_did(&input.recipient_did).await { + Ok(Some(row)) => row, Ok(None) => { return ApiError::AccountNotFound.into_response(); } @@ -59,19 +45,30 @@ pub async fn send_email( return ApiError::InternalError(None).into_response(); } }; + let email = match user.email { + Some(e) => e, + None => { + return ApiError::NoEmail.into_response(); + } + }; + let (user_id, handle) = (user.id, user.handle); let hostname = std::env::var("PDS_HOSTNAME").unwrap_or_else(|_| "localhost".to_string()); let subject = input .subject .clone() .unwrap_or_else(|| format!("Message from {}", hostname)); - let item = crate::comms::NewComms::email( - user_id, - crate::comms::CommsType::AdminEmail, - email, - subject, - content.to_string(), - ); - let result = crate::comms::enqueue_comms(&state.db, item).await; + let result = state + .infra_repo + .enqueue_comms( + Some(user_id), + tranquil_db_traits::CommsChannel::Email, + tranquil_db_traits::CommsType::AdminEmail, + &email, + Some(&subject), + content, + None, + ) + .await; match result { Ok(_) => { tracing::info!( diff --git a/crates/tranquil-pds/src/api/admin/account/info.rs b/crates/tranquil-pds/src/api/admin/account/info.rs index 59d80fc..16df6cf 100644 --- a/crates/tranquil-pds/src/api/admin/account/info.rs +++ b/crates/tranquil-pds/src/api/admin/account/info.rs @@ -9,6 +9,7 @@ use axum::{ response::{IntoResponse, Response}, }; use serde::{Deserialize, Serialize}; +use std::collections::HashMap; use tracing::error; #[derive(Deserialize)] @@ -43,8 +44,10 @@ pub struct InviteCodeInfo { pub code: String, pub available: i32, pub disabled: bool, - pub for_account: Did, - pub created_by: Did, + #[serde(skip_serializing_if = "Option::is_none")] + pub for_account: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub created_by: Option, pub created_at: String, pub uses: Vec, } @@ -67,120 +70,86 @@ pub async fn get_account_info( _auth: BearerAuthAdmin, Query(params): Query, ) -> Response { - let result = sqlx::query!( - r#" - SELECT id, did, handle, email, created_at, invites_disabled, email_verified, deactivated_at - FROM users - WHERE did = $1 - "#, - params.did.as_str() - ) - .fetch_optional(&state.db) - .await; - match result { - Ok(Some(row)) => { - let invited_by = get_invited_by(&state.db, row.id).await; - let invites = get_invites_for_user(&state.db, row.id).await; - ( - StatusCode::OK, - Json(AccountInfo { - did: row.did.into(), - handle: row.handle.into(), - email: row.email, - indexed_at: row.created_at.to_rfc3339(), - invite_note: None, - invites_disabled: row.invites_disabled.unwrap_or(false), - email_confirmed_at: if row.email_verified { - Some(row.created_at.to_rfc3339()) - } else { - None - }, - deactivated_at: row.deactivated_at.map(|dt| dt.to_rfc3339()), - invited_by, - invites, - }), - ) - .into_response() - } - Ok(None) => ApiError::AccountNotFound.into_response(), + let account = match state.infra_repo.get_admin_account_info_by_did(¶ms.did).await { + Ok(Some(a)) => a, + Ok(None) => return ApiError::AccountNotFound.into_response(), Err(e) => { error!("DB error in get_account_info: {:?}", e); - ApiError::InternalError(None).into_response() + return ApiError::InternalError(None).into_response(); } - } -} + }; + + let invited_by = get_invited_by(&state, account.id).await; + let invites = get_invites_for_user(&state, account.id).await; -async fn get_invited_by(db: &sqlx::PgPool, user_id: uuid::Uuid) -> Option { - let use_row = sqlx::query!( - r#" - SELECT icu.code - FROM invite_code_uses icu - WHERE icu.used_by_user = $1 - LIMIT 1 - "#, - user_id + ( + StatusCode::OK, + Json(AccountInfo { + did: account.did, + handle: account.handle, + email: account.email, + indexed_at: account.created_at.to_rfc3339(), + invite_note: None, + invites_disabled: account.invites_disabled, + email_confirmed_at: if account.email_verified { + Some(account.created_at.to_rfc3339()) + } else { + None + }, + deactivated_at: account.deactivated_at.map(|dt| dt.to_rfc3339()), + invited_by, + invites, + }), ) - .fetch_optional(db) - .await - .ok()??; - get_invite_code_info(db, &use_row.code).await + .into_response() } -async fn get_invites_for_user( - db: &sqlx::PgPool, - user_id: uuid::Uuid, -) -> Option> { - let invite_codes = sqlx::query!( - r#" - SELECT ic.code, ic.available_uses, ic.disabled, ic.for_account, ic.created_at, u.did as created_by - FROM invite_codes ic - JOIN users u ON ic.created_by_user = u.id - WHERE ic.created_by_user = $1 - "#, - user_id - ) - .fetch_all(db) - .await - .ok()?; +async fn get_invited_by(state: &AppState, user_id: uuid::Uuid) -> Option { + let code = state + .infra_repo + .get_invite_code_used_by_user(user_id) + .await + .ok()??; + + get_invite_code_info(state, &code).await +} + +async fn get_invites_for_user(state: &AppState, user_id: uuid::Uuid) -> Option> { + let invite_codes = state + .infra_repo + .get_invites_created_by_user(user_id) + .await + .ok()?; if invite_codes.is_empty() { return None; } let code_strings: Vec = invite_codes.iter().map(|ic| ic.code.clone()).collect(); - let mut uses_by_code: std::collections::HashMap> = - std::collections::HashMap::new(); - sqlx::query!( - r#" - SELECT icu.code, u.did as used_by, icu.used_at - FROM invite_code_uses icu - JOIN users u ON icu.used_by_user = u.id - WHERE icu.code = ANY($1) - "#, - &code_strings - ) - .fetch_all(db) - .await - .ok()? - .into_iter() - .for_each(|r| { - uses_by_code - .entry(r.code) - .or_default() - .push(InviteCodeUseInfo { - used_by: r.used_by.into(), - used_at: r.used_at.to_rfc3339(), + + let uses = state + .infra_repo + .get_invite_code_uses_batch(&code_strings) + .await + .ok()?; + + let uses_by_code: HashMap> = + uses.into_iter().fold(HashMap::new(), |mut acc, u| { + acc.entry(u.code.clone()).or_default().push(InviteCodeUseInfo { + used_by: u.used_by_did, + used_at: u.used_at.to_rfc3339(), }); - }); + acc + }); let invites: Vec = invite_codes .into_iter() .map(|ic| InviteCodeInfo { code: ic.code.clone(), available: ic.available_uses, - disabled: ic.disabled.unwrap_or(false), - for_account: ic.for_account.into(), - created_by: ic.created_by.into(), + disabled: ic.disabled, + for_account: ic.for_account, + created_by: ic.created_by, created_at: ic.created_at.to_rfc3339(), uses: uses_by_code.get(&ic.code).cloned().unwrap_or_default(), }) @@ -193,42 +162,27 @@ async fn get_invites_for_user( } } -async fn get_invite_code_info(db: &sqlx::PgPool, code: &str) -> Option { - let row = sqlx::query!( - r#" - SELECT ic.code, ic.available_uses, ic.disabled, ic.for_account, ic.created_at, u.did as created_by - FROM invite_codes ic - JOIN users u ON ic.created_by_user = u.id - WHERE ic.code = $1 - "#, - code - ) - .fetch_optional(db) - .await - .ok()??; - let uses = sqlx::query!( - r#" - SELECT u.did as used_by, icu.used_at - FROM invite_code_uses icu - JOIN users u ON icu.used_by_user = u.id - WHERE icu.code = $1 - "#, - code - ) - .fetch_all(db) - .await - .ok()?; +async fn get_invite_code_info(state: &AppState, code: &str) -> Option { + let info = state.infra_repo.get_invite_code_info(code).await.ok()??; + + let uses = state + .infra_repo + .get_invite_code_uses(code) + .await + .ok() + .unwrap_or_default(); + Some(InviteCodeInfo { - code: row.code, - available: row.available_uses, - disabled: row.disabled.unwrap_or(false), - for_account: row.for_account.into(), - created_by: row.created_by.into(), - created_at: row.created_at.to_rfc3339(), + code: info.code, + available: info.available_uses, + disabled: info.disabled, + for_account: info.for_account, + created_by: info.created_by, + created_at: info.created_at.to_rfc3339(), uses: uses .into_iter() .map(|u| InviteCodeUseInfo { - used_by: u.used_by.into(), + used_by: u.used_by_did, used_at: u.used_at.to_rfc3339(), }) .collect(), @@ -244,132 +198,108 @@ pub async fn get_account_infos( .into_iter() .filter(|d| !d.is_empty()) .collect(); + if dids.is_empty() { return ApiError::InvalidRequest("dids is required".into()).into_response(); } - let users = match sqlx::query!( - r#" - SELECT id, did, handle, email, created_at, invites_disabled, email_verified, deactivated_at - FROM users - WHERE did = ANY($1) - "#, - &dids - ) - .fetch_all(&state.db) - .await - { - Ok(rows) => rows, + + let dids_typed: Vec = dids + .iter() + .filter_map(|d| d.parse().ok()) + .collect(); + let accounts = match state.infra_repo.get_admin_account_infos_by_dids(&dids_typed).await { + Ok(accounts) => accounts, Err(e) => { error!("Failed to fetch account infos: {:?}", e); return ApiError::InternalError(None).into_response(); } }; - let user_ids: Vec = users.iter().map(|u| u.id).collect(); - - let all_invite_codes = sqlx::query!( - r#" - SELECT ic.code, ic.available_uses, ic.disabled, ic.for_account, ic.created_at, - ic.created_by_user, u.did as created_by - FROM invite_codes ic - JOIN users u ON ic.created_by_user = u.id - WHERE ic.created_by_user = ANY($1) - "#, - &user_ids - ) - .fetch_all(&state.db) - .await - .unwrap_or_default(); + let user_ids: Vec = accounts.iter().map(|u| u.id).collect(); - let all_codes: Vec = all_invite_codes.iter().map(|c| c.code.clone()).collect(); - let all_invite_uses = if !all_codes.is_empty() { - sqlx::query!( - r#" - SELECT icu.code, u.did as used_by, icu.used_at - FROM invite_code_uses icu - JOIN users u ON icu.used_by_user = u.id - WHERE icu.code = ANY($1) - "#, - &all_codes - ) - .fetch_all(&state.db) + let all_invite_codes = state + .infra_repo + .get_invite_codes_by_users(&user_ids) .await - .unwrap_or_default() + .unwrap_or_default(); + + let all_codes: Vec = all_invite_codes.iter().map(|(_, c)| c.code.clone()).collect(); + + let all_invite_uses = if !all_codes.is_empty() { + state + .infra_repo + .get_invite_code_uses_batch(&all_codes) + .await + .unwrap_or_default() } else { Vec::new() }; - let invited_by_map: std::collections::HashMap = sqlx::query!( - r#" - SELECT icu.used_by_user, icu.code - FROM invite_code_uses icu - WHERE icu.used_by_user = ANY($1) - "#, - &user_ids - ) - .fetch_all(&state.db) - .await - .unwrap_or_default() - .into_iter() - .map(|r| (r.used_by_user, r.code)) - .collect(); - - let uses_by_code: std::collections::HashMap> = + let invited_by_map: HashMap = state + .infra_repo + .get_invite_code_uses_by_users(&user_ids) + .await + .unwrap_or_default() + .into_iter() + .collect(); + + let uses_by_code: HashMap> = all_invite_uses .into_iter() - .fold(std::collections::HashMap::new(), |mut acc, u| { + .fold(HashMap::new(), |mut acc, u| { acc.entry(u.code.clone()).or_default().push(InviteCodeUseInfo { - used_by: u.used_by.into(), + used_by: u.used_by_did, used_at: u.used_at.to_rfc3339(), }); acc }); let (codes_by_user, code_info_map): ( - std::collections::HashMap>, - std::collections::HashMap, + HashMap>, + HashMap, ) = all_invite_codes.into_iter().fold( - (std::collections::HashMap::new(), std::collections::HashMap::new()), - |(mut by_user, mut by_code), ic| { + (HashMap::new(), HashMap::new()), + |(mut by_user, mut by_code), (user_id, ic)| { let info = InviteCodeInfo { code: ic.code.clone(), available: ic.available_uses, - disabled: ic.disabled.unwrap_or(false), - for_account: ic.for_account.into(), - created_by: ic.created_by.into(), + disabled: ic.disabled, + for_account: ic.for_account, + created_by: ic.created_by, created_at: ic.created_at.to_rfc3339(), uses: uses_by_code.get(&ic.code).cloned().unwrap_or_default(), }; by_code.insert(ic.code.clone(), info.clone()); - by_user.entry(ic.created_by_user).or_default().push(info); + by_user.entry(user_id).or_default().push(info); (by_user, by_code) }, ); - let infos: Vec = users + let infos: Vec = accounts .into_iter() - .map(|row| { + .map(|account| { let invited_by = invited_by_map - .get(&row.id) + .get(&account.id) .and_then(|code| code_info_map.get(code).cloned()); - let invites = codes_by_user.get(&row.id).cloned(); + let invites = codes_by_user.get(&account.id).cloned(); AccountInfo { - did: row.did.into(), - handle: row.handle.into(), - email: row.email, - indexed_at: row.created_at.to_rfc3339(), + did: account.did, + handle: account.handle, + email: account.email, + indexed_at: account.created_at.to_rfc3339(), invite_note: None, - invites_disabled: row.invites_disabled.unwrap_or(false), - email_confirmed_at: if row.email_verified { - Some(row.created_at.to_rfc3339()) + invites_disabled: account.invites_disabled, + email_confirmed_at: if account.email_verified { + Some(account.created_at.to_rfc3339()) } else { None }, - deactivated_at: row.deactivated_at.map(|dt| dt.to_rfc3339()), + deactivated_at: account.deactivated_at.map(|dt| dt.to_rfc3339()), invited_by, invites, } }) .collect(); + (StatusCode::OK, Json(GetAccountInfosOutput { infos })).into_response() } diff --git a/crates/tranquil-pds/src/api/admin/account/search.rs b/crates/tranquil-pds/src/api/admin/account/search.rs index 27b6e2b..5280f67 100644 --- a/crates/tranquil-pds/src/api/admin/account/search.rs +++ b/crates/tranquil-pds/src/api/admin/account/search.rs @@ -54,68 +54,37 @@ pub async fn search_accounts( Query(params): Query, ) -> Response { let limit = params.limit.clamp(1, 100); - let cursor_did = params.cursor.as_deref().unwrap_or(""); let email_filter = params.email.as_deref().map(|e| format!("%{}%", e)); let handle_filter = params.handle.as_deref().map(|h| format!("%{}%", h)); - let result = sqlx::query_as::< - _, - ( - String, - String, - Option, - chrono::DateTime, - bool, - Option>, - Option, - ), - >( - r#" - SELECT did, handle, email, created_at, email_verified, deactivated_at, invites_disabled - FROM users - WHERE did > $1 - AND ($2::text IS NULL OR email ILIKE $2) - AND ($3::text IS NULL OR handle ILIKE $3) - ORDER BY did ASC - LIMIT $4 - "#, - ) - .bind(cursor_did) - .bind(&email_filter) - .bind(&handle_filter) - .bind(limit + 1) - .fetch_all(&state.db) - .await; + let cursor_did: Option = params.cursor.as_ref().and_then(|c| c.parse().ok()); + let result = state + .user_repo + .search_accounts( + cursor_did.as_ref(), + email_filter.as_deref(), + handle_filter.as_deref(), + limit + 1, + ) + .await; match result { Ok(rows) => { let has_more = rows.len() > limit as usize; let accounts: Vec = rows .into_iter() .take(limit as usize) - .map( - |( - did, - handle, - email, - created_at, - email_verified, - deactivated_at, - invites_disabled, - )| { - AccountView { - did: did.clone().into(), - handle: handle.into(), - email, - indexed_at: created_at.to_rfc3339(), - email_confirmed_at: if email_verified { - Some(created_at.to_rfc3339()) - } else { - None - }, - deactivated_at: deactivated_at.map(|dt| dt.to_rfc3339()), - invites_disabled, - } + .map(|row| AccountView { + did: row.did.clone(), + handle: row.handle, + email: row.email, + indexed_at: row.created_at.to_rfc3339(), + email_confirmed_at: if row.email_verified { + Some(row.created_at.to_rfc3339()) + } else { + None }, - ) + deactivated_at: row.deactivated_at.map(|dt| dt.to_rfc3339()), + invites_disabled: row.invites_disabled, + }) .collect(); let next_cursor = if has_more { accounts.last().map(|a| a.did.to_string()) diff --git a/crates/tranquil-pds/src/api/admin/account/update.rs b/crates/tranquil-pds/src/api/admin/account/update.rs index 9dec226..26dbab4 100644 --- a/crates/tranquil-pds/src/api/admin/account/update.rs +++ b/crates/tranquil-pds/src/api/admin/account/update.rs @@ -27,16 +27,13 @@ pub async fn update_account_email( if account.is_empty() || email.is_empty() { return ApiError::InvalidRequest("account and email are required".into()).into_response(); } - let result = sqlx::query!("UPDATE users SET email = $1 WHERE did = $2", email, account) - .execute(&state.db) - .await; - match result { - Ok(r) => { - if r.rows_affected() == 0 { - return ApiError::AccountNotFound.into_response(); - } - EmptyResponse::ok().into_response() - } + let account_did: Did = match account.parse() { + Ok(d) => d, + Err(_) => return ApiError::InvalidDid("Invalid DID format".into()).into_response(), + }; + match state.user_repo.admin_update_email(&account_did, email).await { + Ok(0) => ApiError::AccountNotFound.into_response(), + Ok(_) => EmptyResponse::ok().into_response(), Err(e) => { error!("DB error updating email: {:?}", e); ApiError::InternalError(None).into_response() @@ -67,45 +64,35 @@ pub async fn update_account_handle( return ApiError::InvalidHandle(None).into_response(); } let hostname = std::env::var("PDS_HOSTNAME").unwrap_or_else(|_| "localhost".to_string()); + let hostname_for_handles = hostname.split(':').next().unwrap_or(&hostname); let handle = if !input_handle.contains('.') { - format!("{}.{}", input_handle, hostname) + format!("{}.{}", input_handle, hostname_for_handles) } else { input_handle.to_string() }; - let old_handle = sqlx::query_scalar!("SELECT handle FROM users WHERE did = $1", did.as_str()) - .fetch_optional(&state.db) + let old_handle = state + .user_repo + .get_handle_by_did(did) .await .ok() .flatten(); - let existing = sqlx::query!( - "SELECT id FROM users WHERE handle = $1 AND did != $2", - handle, - did.as_str() - ) - .fetch_optional(&state.db) - .await; - if let Ok(Some(_)) = existing { + let user_id = match state.user_repo.get_id_by_did(did).await { + Ok(Some(id)) => id, + _ => return ApiError::AccountNotFound.into_response(), + }; + let handle_for_check = Handle::new_unchecked(&handle); + if let Ok(true) = state.user_repo.check_handle_exists(&handle_for_check, user_id).await { return ApiError::HandleTaken.into_response(); } - let result = sqlx::query!( - "UPDATE users SET handle = $1 WHERE did = $2", - handle, - did.as_str() - ) - .execute(&state.db) - .await; - match result { - Ok(r) => { - if r.rows_affected() == 0 { - return ApiError::AccountNotFound.into_response(); - } + match state.user_repo.admin_update_handle(did, &handle_for_check).await { + Ok(0) => ApiError::AccountNotFound.into_response(), + Ok(_) => { if let Some(old) = old_handle { let _ = state.cache.delete(&format!("handle:{}", old)).await; } let _ = state.cache.delete(&format!("handle:{}", handle)).await; - let handle_typed = Handle::new_unchecked(&handle); if let Err(e) = - crate::api::repo::record::sequence_identity_event(&state, did, Some(&handle_typed)) + crate::api::repo::record::sequence_identity_event(&state, did, Some(&handle_for_check)) .await { warn!( @@ -114,7 +101,7 @@ pub async fn update_account_handle( ); } if let Err(e) = - crate::api::identity::did::update_plc_handle(&state, did.as_str(), &handle).await + crate::api::identity::did::update_plc_handle(&state, did, &handle_for_check).await { warn!("Failed to update PLC handle for admin handle update: {}", e); } @@ -150,20 +137,9 @@ pub async fn update_account_password( return ApiError::InternalError(None).into_response(); } }; - let result = sqlx::query!( - "UPDATE users SET password_hash = $1 WHERE did = $2", - password_hash, - did.as_str() - ) - .execute(&state.db) - .await; - match result { - Ok(r) => { - if r.rows_affected() == 0 { - return ApiError::AccountNotFound.into_response(); - } - EmptyResponse::ok().into_response() - } + match state.user_repo.admin_update_password(did, &password_hash).await { + Ok(0) => ApiError::AccountNotFound.into_response(), + Ok(_) => EmptyResponse::ok().into_response(), Err(e) => { error!("DB error updating password: {:?}", e); ApiError::InternalError(None).into_response() diff --git a/crates/tranquil-pds/src/api/admin/config.rs b/crates/tranquil-pds/src/api/admin/config.rs index a281b3f..a1e31d3 100644 --- a/crates/tranquil-pds/src/api/admin/config.rs +++ b/crates/tranquil-pds/src/api/admin/config.rs @@ -4,6 +4,7 @@ use crate::state::AppState; use axum::{Json, extract::State}; use serde::{Deserialize, Serialize}; use tracing::error; +use tranquil_types::CidLink; #[derive(Serialize)] #[serde(rename_all = "camelCase")] @@ -42,14 +43,25 @@ fn is_valid_hex_color(s: &str) -> bool { pub async fn get_server_config( State(state): State, ) -> Result, ApiError> { - let rows: Vec<(String, String)> = sqlx::query_as( - "SELECT key, value FROM server_config WHERE key IN ('server_name', 'primary_color', 'primary_color_dark', 'secondary_color', 'secondary_color_dark', 'logo_cid')" - ) - .fetch_all(&state.db) - .await?; - - let config_map: std::collections::HashMap = - rows.into_iter().collect(); + let keys = &[ + "server_name", + "primary_color", + "primary_color_dark", + "secondary_color", + "secondary_color_dark", + "logo_cid", + ]; + + let rows = state + .infra_repo + .get_server_configs(keys) + .await + .map_err(|e| { + error!("DB error fetching server config: {:?}", e); + ApiError::InternalError(None) + })?; + + let config_map: std::collections::HashMap = rows.into_iter().collect(); Ok(Json(ServerConfigResponse { server_name: config_map @@ -64,26 +76,6 @@ pub async fn get_server_config( })) } -async fn upsert_config(db: &sqlx::PgPool, key: &str, value: &str) -> Result<(), sqlx::Error> { - sqlx::query( - "INSERT INTO server_config (key, value, updated_at) VALUES ($1, $2, NOW()) - ON CONFLICT (key) DO UPDATE SET value = $2, updated_at = NOW()", - ) - .bind(key) - .bind(value) - .execute(db) - .await?; - Ok(()) -} - -async fn delete_config(db: &sqlx::PgPool, key: &str) -> Result<(), sqlx::Error> { - sqlx::query("DELETE FROM server_config WHERE key = $1") - .bind(key) - .execute(db) - .await?; - Ok(()) -} - pub async fn update_server_config( State(state): State, _admin: BearerAuthAdmin, @@ -96,14 +88,35 @@ pub async fn update_server_config( "Server name must be 1-100 characters".into(), )); } - upsert_config(&state.db, "server_name", trimmed).await?; + state + .infra_repo + .upsert_server_config("server_name", trimmed) + .await + .map_err(|e| { + error!("DB error upserting server_name: {:?}", e); + ApiError::InternalError(None) + })?; } if let Some(ref color) = req.primary_color { if color.is_empty() { - delete_config(&state.db, "primary_color").await?; + state + .infra_repo + .delete_server_config("primary_color") + .await + .map_err(|e| { + error!("DB error deleting primary_color: {:?}", e); + ApiError::InternalError(None) + })?; } else if is_valid_hex_color(color) { - upsert_config(&state.db, "primary_color", color).await?; + state + .infra_repo + .upsert_server_config("primary_color", color) + .await + .map_err(|e| { + error!("DB error upserting primary_color: {:?}", e); + ApiError::InternalError(None) + })?; } else { return Err(ApiError::InvalidRequest( "Invalid primary color format (expected #RRGGBB)".into(), @@ -113,9 +126,23 @@ pub async fn update_server_config( if let Some(ref color) = req.primary_color_dark { if color.is_empty() { - delete_config(&state.db, "primary_color_dark").await?; + state + .infra_repo + .delete_server_config("primary_color_dark") + .await + .map_err(|e| { + error!("DB error deleting primary_color_dark: {:?}", e); + ApiError::InternalError(None) + })?; } else if is_valid_hex_color(color) { - upsert_config(&state.db, "primary_color_dark", color).await?; + state + .infra_repo + .upsert_server_config("primary_color_dark", color) + .await + .map_err(|e| { + error!("DB error upserting primary_color_dark: {:?}", e); + ApiError::InternalError(None) + })?; } else { return Err(ApiError::InvalidRequest( "Invalid primary dark color format (expected #RRGGBB)".into(), @@ -125,9 +152,23 @@ pub async fn update_server_config( if let Some(ref color) = req.secondary_color { if color.is_empty() { - delete_config(&state.db, "secondary_color").await?; + state + .infra_repo + .delete_server_config("secondary_color") + .await + .map_err(|e| { + error!("DB error deleting secondary_color: {:?}", e); + ApiError::InternalError(None) + })?; } else if is_valid_hex_color(color) { - upsert_config(&state.db, "secondary_color", color).await?; + state + .infra_repo + .upsert_server_config("secondary_color", color) + .await + .map_err(|e| { + error!("DB error upserting secondary_color: {:?}", e); + ApiError::InternalError(None) + })?; } else { return Err(ApiError::InvalidRequest( "Invalid secondary color format (expected #RRGGBB)".into(), @@ -137,9 +178,23 @@ pub async fn update_server_config( if let Some(ref color) = req.secondary_color_dark { if color.is_empty() { - delete_config(&state.db, "secondary_color_dark").await?; + state + .infra_repo + .delete_server_config("secondary_color_dark") + .await + .map_err(|e| { + error!("DB error deleting secondary_color_dark: {:?}", e); + ApiError::InternalError(None) + })?; } else if is_valid_hex_color(color) { - upsert_config(&state.db, "secondary_color_dark", color).await?; + state + .infra_repo + .upsert_server_config("secondary_color_dark", color) + .await + .map_err(|e| { + error!("DB error upserting secondary_color_dark: {:?}", e); + ApiError::InternalError(None) + })?; } else { return Err(ApiError::InvalidRequest( "Invalid secondary dark color format (expected #RRGGBB)".into(), @@ -148,10 +203,12 @@ pub async fn update_server_config( } if let Some(ref logo_cid) = req.logo_cid { - let old_logo_cid: Option = - sqlx::query_scalar("SELECT value FROM server_config WHERE key = 'logo_cid'") - .fetch_optional(&state.db) - .await?; + let old_logo_cid = state + .infra_repo + .get_server_config("logo_cid") + .await + .ok() + .flatten(); let should_delete_old = match (&old_logo_cid, logo_cid.is_empty()) { (Some(old), true) => Some(old.clone()), @@ -159,27 +216,38 @@ pub async fn update_server_config( _ => None, }; - if let Some(old_cid) = should_delete_old - && let Ok(Some(blob)) = - sqlx::query!("SELECT storage_key FROM blobs WHERE cid = $1", old_cid) - .fetch_optional(&state.db) - .await - { - if let Err(e) = state.blob_store.delete(&blob.storage_key).await { - error!("Failed to delete old logo blob from storage: {:?}", e); - } - if let Err(e) = sqlx::query!("DELETE FROM blobs WHERE cid = $1", old_cid) - .execute(&state.db) - .await + if let Some(old_cid_str) = should_delete_old { + let old_cid = CidLink::new_unchecked(old_cid_str); + if let Ok(Some(storage_key)) = + state.infra_repo.get_blob_storage_key_by_cid(&old_cid).await { - error!("Failed to delete old logo blob record: {:?}", e); + if let Err(e) = state.blob_store.delete(&storage_key).await { + error!("Failed to delete old logo blob from storage: {:?}", e); + } + if let Err(e) = state.infra_repo.delete_blob_by_cid(&old_cid).await { + error!("Failed to delete old logo blob record: {:?}", e); + } } } if logo_cid.is_empty() { - delete_config(&state.db, "logo_cid").await?; + state + .infra_repo + .delete_server_config("logo_cid") + .await + .map_err(|e| { + error!("DB error deleting logo_cid: {:?}", e); + ApiError::InternalError(None) + })?; } else { - upsert_config(&state.db, "logo_cid", logo_cid).await?; + state + .infra_repo + .upsert_server_config("logo_cid", logo_cid) + .await + .map_err(|e| { + error!("DB error upserting logo_cid: {:?}", e); + ApiError::InternalError(None) + })?; } } diff --git a/crates/tranquil-pds/src/api/admin/invite.rs b/crates/tranquil-pds/src/api/admin/invite.rs index a87b1d1..ad9193e 100644 --- a/crates/tranquil-pds/src/api/admin/invite.rs +++ b/crates/tranquil-pds/src/api/admin/invite.rs @@ -10,6 +10,7 @@ use axum::{ }; use serde::{Deserialize, Serialize}; use tracing::error; +use tranquil_db_traits::InviteCodeSortOrder; #[derive(Deserialize)] #[serde(rename_all = "camelCase")] @@ -24,20 +25,22 @@ pub async fn disable_invite_codes( Json(input): Json, ) -> Response { if let Some(codes) = &input.codes { - let _ = sqlx::query!( - "UPDATE invite_codes SET disabled = TRUE WHERE code = ANY($1)", - codes as &[String] - ) - .execute(&state.db) - .await; + if let Err(e) = state.infra_repo.disable_invite_codes_by_code(codes).await { + error!("DB error disabling invite codes: {:?}", e); + } } if let Some(accounts) = &input.accounts { - let _ = sqlx::query!( - "UPDATE invite_codes SET disabled = TRUE WHERE created_by_user IN (SELECT id FROM users WHERE did = ANY($1))", - accounts as &[String] - ) - .execute(&state.db) - .await; + let accounts_typed: Vec = accounts + .iter() + .filter_map(|a| a.parse().ok()) + .collect(); + if let Err(e) = state + .infra_repo + .disable_invite_codes_by_account(&accounts_typed) + .await + { + error!("DB error disabling invite codes by account: {:?}", e); + } } EmptyResponse::ok().into_response() } @@ -81,59 +84,16 @@ pub async fn get_invite_codes( Query(params): Query, ) -> Response { let limit = params.limit.unwrap_or(100).clamp(1, 500); - let sort = params.sort.as_deref().unwrap_or("recent"); - let order_clause = match sort { - "usage" => "available_uses DESC", - _ => "created_at DESC", + let sort_order = match params.sort.as_deref() { + Some("usage") => InviteCodeSortOrder::Usage, + _ => InviteCodeSortOrder::Recent, }; - let codes_result = if let Some(cursor) = ¶ms.cursor { - sqlx::query_as::< - _, - ( - String, - i32, - Option, - uuid::Uuid, - chrono::DateTime, - ), - >(&format!( - r#" - SELECT ic.code, ic.available_uses, ic.disabled, ic.created_by_user, ic.created_at - FROM invite_codes ic - WHERE ic.created_at < (SELECT created_at FROM invite_codes WHERE code = $1) - ORDER BY {} - LIMIT $2 - "#, - order_clause - )) - .bind(cursor) - .bind(limit) - .fetch_all(&state.db) - .await - } else { - sqlx::query_as::< - _, - ( - String, - i32, - Option, - uuid::Uuid, - chrono::DateTime, - ), - >(&format!( - r#" - SELECT ic.code, ic.available_uses, ic.disabled, ic.created_by_user, ic.created_at - FROM invite_codes ic - ORDER BY {} - LIMIT $1 - "#, - order_clause - )) - .bind(limit) - .fetch_all(&state.db) + + let codes_rows = match state + .infra_repo + .list_invite_codes(params.cursor.as_deref(), limit, sort_order) .await - }; - let codes_rows = match codes_result { + { Ok(rows) => rows, Err(e) => { error!("DB error fetching invite codes: {:?}", e); @@ -141,72 +101,58 @@ pub async fn get_invite_codes( } }; - let user_ids: Vec = codes_rows.iter().map(|(_, _, _, uid, _)| *uid).collect(); - let code_strings: Vec = codes_rows.iter().map(|(c, _, _, _, _)| c.clone()).collect(); + let user_ids: Vec = codes_rows.iter().map(|r| r.created_by_user).collect(); + let code_strings: Vec = codes_rows.iter().map(|r| r.code.clone()).collect(); - let mut creator_dids: std::collections::HashMap = - std::collections::HashMap::new(); - sqlx::query!( - "SELECT id, did FROM users WHERE id = ANY($1)", - &user_ids - ) - .fetch_all(&state.db) - .await - .unwrap_or_default() - .into_iter() - .for_each(|r| { - creator_dids.insert(r.id, r.did); - }); - - let mut uses_by_code: std::collections::HashMap> = - std::collections::HashMap::new(); - if !code_strings.is_empty() { - sqlx::query!( - r#" - SELECT icu.code, u.did, icu.used_at - FROM invite_code_uses icu - JOIN users u ON icu.used_by_user = u.id - WHERE icu.code = ANY($1) - ORDER BY icu.used_at DESC - "#, - &code_strings - ) - .fetch_all(&state.db) + let creator_dids: std::collections::HashMap = state + .infra_repo + .get_user_dids_by_ids(&user_ids) .await .unwrap_or_default() .into_iter() - .for_each(|r| { - uses_by_code - .entry(r.code) - .or_default() - .push(InviteCodeUseInfo { - used_by: r.did, - used_at: r.used_at.to_rfc3339(), + .collect(); + + let uses_by_code: std::collections::HashMap> = if code_strings + .is_empty() + { + std::collections::HashMap::new() + } else { + state + .infra_repo + .get_invite_code_uses_batch(&code_strings) + .await + .unwrap_or_default() + .into_iter() + .fold(std::collections::HashMap::new(), |mut acc, u| { + acc.entry(u.code.clone()).or_default().push(InviteCodeUseInfo { + used_by: u.used_by_did.to_string(), + used_at: u.used_at.to_rfc3339(), }); - }); - } + acc + }) + }; let codes: Vec = codes_rows .iter() - .map(|(code, available_uses, disabled, created_by_user, created_at)| { + .map(|r| { let creator_did = creator_dids - .get(created_by_user) - .cloned() + .get(&r.created_by_user) + .map(|d| d.to_string()) .unwrap_or_else(|| "unknown".to_string()); InviteCodeInfo { - code: code.clone(), - available: *available_uses, - disabled: disabled.unwrap_or(false), + code: r.code.clone(), + available: r.available_uses, + disabled: r.disabled.unwrap_or(false), for_account: creator_did.clone(), created_by: creator_did, - created_at: created_at.to_rfc3339(), - uses: uses_by_code.get(code).cloned().unwrap_or_default(), + created_at: r.created_at.to_rfc3339(), + uses: uses_by_code.get(&r.code).cloned().unwrap_or_default(), } }) .collect(); let next_cursor = if codes_rows.len() == limit as usize { - codes_rows.last().map(|(code, _, _, _, _)| code.clone()) + codes_rows.last().map(|r| r.code.clone()) } else { None }; @@ -234,19 +180,13 @@ pub async fn disable_account_invites( if account.is_empty() { return ApiError::InvalidRequest("account is required".into()).into_response(); } - let result = sqlx::query!( - "UPDATE users SET invites_disabled = TRUE WHERE did = $1", - account - ) - .execute(&state.db) - .await; - match result { - Ok(r) => { - if r.rows_affected() == 0 { - return ApiError::AccountNotFound.into_response(); - } - EmptyResponse::ok().into_response() - } + let account_did: tranquil_types::Did = match account.parse() { + Ok(d) => d, + Err(_) => return ApiError::InvalidDid("Invalid DID format".into()).into_response(), + }; + match state.user_repo.set_invites_disabled(&account_did, true).await { + Ok(true) => EmptyResponse::ok().into_response(), + Ok(false) => ApiError::AccountNotFound.into_response(), Err(e) => { error!("DB error disabling account invites: {:?}", e); ApiError::InternalError(None).into_response() @@ -268,19 +208,13 @@ pub async fn enable_account_invites( if account.is_empty() { return ApiError::InvalidRequest("account is required".into()).into_response(); } - let result = sqlx::query!( - "UPDATE users SET invites_disabled = FALSE WHERE did = $1", - account - ) - .execute(&state.db) - .await; - match result { - Ok(r) => { - if r.rows_affected() == 0 { - return ApiError::AccountNotFound.into_response(); - } - EmptyResponse::ok().into_response() - } + let account_did: tranquil_types::Did = match account.parse() { + Ok(d) => d, + Err(_) => return ApiError::InvalidDid("Invalid DID format".into()).into_response(), + }; + match state.user_repo.set_invites_disabled(&account_did, false).await { + Ok(true) => EmptyResponse::ok().into_response(), + Ok(false) => ApiError::AccountNotFound.into_response(), Err(e) => { error!("DB error enabling account invites: {:?}", e); ApiError::InternalError(None).into_response() diff --git a/crates/tranquil-pds/src/api/admin/server_stats.rs b/crates/tranquil-pds/src/api/admin/server_stats.rs index 12d7b72..2e2624c 100644 --- a/crates/tranquil-pds/src/api/admin/server_stats.rs +++ b/crates/tranquil-pds/src/api/admin/server_stats.rs @@ -17,42 +17,10 @@ pub struct ServerStatsResponse { } pub async fn get_server_stats(State(state): State, _auth: BearerAuthAdmin) -> Response { - let user_count: i64 = match sqlx::query_scalar!("SELECT COUNT(*) FROM users") - .fetch_one(&state.db) - .await - { - Ok(Some(count)) => count, - Ok(None) => 0, - Err(_) => 0, - }; - - let repo_count: i64 = match sqlx::query_scalar!("SELECT COUNT(*) FROM repos") - .fetch_one(&state.db) - .await - { - Ok(Some(count)) => count, - Ok(None) => 0, - Err(_) => 0, - }; - - let record_count: i64 = match sqlx::query_scalar!("SELECT COUNT(*) FROM records") - .fetch_one(&state.db) - .await - { - Ok(Some(count)) => count, - Ok(None) => 0, - Err(_) => 0, - }; - - let blob_storage_bytes: i64 = - match sqlx::query_scalar!("SELECT COALESCE(SUM(size_bytes), 0)::BIGINT FROM blobs") - .fetch_one(&state.db) - .await - { - Ok(Some(bytes)) => bytes, - Ok(None) => 0, - Err(_) => 0, - }; + let user_count = state.user_repo.count_users().await.unwrap_or(0); + let repo_count = state.repo_repo.count_repos().await.unwrap_or(0); + let record_count = state.repo_repo.count_all_records().await.unwrap_or(0); + let blob_storage_bytes = state.blob_repo.sum_blob_storage().await.unwrap_or(0); Json(ServerStatsResponse { user_count, diff --git a/crates/tranquil-pds/src/api/admin/status.rs b/crates/tranquil-pds/src/api/admin/status.rs index e8391e1..88cc747 100644 --- a/crates/tranquil-pds/src/api/admin/status.rs +++ b/crates/tranquil-pds/src/api/admin/status.rs @@ -1,7 +1,7 @@ use crate::api::error::ApiError; use crate::auth::BearerAuthAdmin; use crate::state::AppState; -use crate::types::Did; +use crate::types::{CidLink, Did}; use axum::{ Json, extract::{Query, State}, @@ -41,20 +41,18 @@ pub async fn get_subject_status( if params.did.is_none() && params.uri.is_none() && params.blob.is_none() { return ApiError::InvalidRequest("Must provide did, uri, or blob".into()).into_response(); } - if let Some(did) = ¶ms.did { - let user = sqlx::query!( - "SELECT did, deactivated_at, takedown_ref FROM users WHERE did = $1", - did - ) - .fetch_optional(&state.db) - .await; - match user { - Ok(Some(row)) => { - let deactivated = row.deactivated_at.map(|_| StatusAttr { + if let Some(did_str) = ¶ms.did { + let did: Did = match did_str.parse() { + Ok(d) => d, + Err(_) => return ApiError::InvalidDid("Invalid DID format".into()).into_response(), + }; + match state.user_repo.get_status_by_did(&did).await { + Ok(Some(status)) => { + let deactivated = status.deactivated_at.map(|_| StatusAttr { applied: true, r#ref: None, }); - let takedown = row.takedown_ref.as_ref().map(|r| StatusAttr { + let takedown = status.takedown_ref.as_ref().map(|r| StatusAttr { applied: true, r#ref: Some(r.clone()), }); @@ -63,7 +61,7 @@ pub async fn get_subject_status( Json(SubjectStatus { subject: json!({ "$type": "com.atproto.admin.defs#repoRef", - "did": row.did + "did": did_str }), takedown, deactivated, @@ -80,16 +78,14 @@ pub async fn get_subject_status( } } } - if let Some(uri) = ¶ms.uri { - let record = sqlx::query!( - "SELECT r.id, r.takedown_ref FROM records r WHERE r.record_cid = $1", - uri - ) - .fetch_optional(&state.db) - .await; - match record { - Ok(Some(row)) => { - let takedown = row.takedown_ref.as_ref().map(|r| StatusAttr { + if let Some(uri_str) = ¶ms.uri { + let cid: CidLink = match uri_str.parse() { + Ok(c) => c, + Err(_) => return ApiError::InvalidRequest("Invalid CID format".into()).into_response(), + }; + match state.repo_repo.get_record_by_cid(&cid).await { + Ok(Some(record)) => { + let takedown = record.takedown_ref.as_ref().map(|r| StatusAttr { applied: true, r#ref: Some(r.clone()), }); @@ -98,8 +94,8 @@ pub async fn get_subject_status( Json(SubjectStatus { subject: json!({ "$type": "com.atproto.repo.strongRef", - "uri": uri, - "cid": uri + "uri": uri_str, + "cid": uri_str }), takedown, deactivated: None, @@ -116,7 +112,11 @@ pub async fn get_subject_status( } } } - if let Some(blob_cid) = ¶ms.blob { + if let Some(blob_cid_str) = ¶ms.blob { + let blob_cid: CidLink = match blob_cid_str.parse() { + Ok(c) => c, + Err(_) => return ApiError::InvalidRequest("Invalid CID format".into()).into_response(), + }; let did = match ¶ms.did { Some(d) => d, None => { @@ -124,15 +124,9 @@ pub async fn get_subject_status( .into_response(); } }; - let blob = sqlx::query!( - "SELECT cid, takedown_ref FROM blobs WHERE cid = $1", - blob_cid - ) - .fetch_optional(&state.db) - .await; - match blob { - Ok(Some(row)) => { - let takedown = row.takedown_ref.as_ref().map(|r| StatusAttr { + match state.blob_repo.get_blob_with_takedown(&blob_cid).await { + Ok(Some(blob)) => { + let takedown = blob.takedown_ref.as_ref().map(|r| StatusAttr { applied: true, r#ref: Some(r.clone()), }); @@ -142,7 +136,7 @@ pub async fn get_subject_status( subject: json!({ "$type": "com.atproto.admin.defs#repoBlobRef", "did": did, - "cid": row.cid + "cid": blob.cid }), takedown, deactivated: None, @@ -187,26 +181,16 @@ pub async fn update_subject_status( let did_str = input.subject.get("did").and_then(|d| d.as_str()); if let Some(did_str) = did_str { let did = Did::new_unchecked(did_str); - let mut tx = match state.db.begin().await { - Ok(tx) => tx, - Err(e) => { - error!("Failed to begin transaction: {:?}", e); - return ApiError::InternalError(None).into_response(); - } - }; if let Some(takedown) = &input.takedown { let takedown_ref = if takedown.applied { - takedown.r#ref.clone() + takedown.r#ref.as_deref() } else { None }; - if let Err(e) = sqlx::query!( - "UPDATE users SET takedown_ref = $1 WHERE did = $2", - takedown_ref, - did.as_str() - ) - .execute(&mut *tx) - .await + if let Err(e) = state + .user_repo + .set_user_takedown(&did, takedown_ref) + .await { error!("Failed to update user takedown status for {}: {:?}", did, e); return ApiError::InternalError(Some( @@ -217,19 +201,9 @@ pub async fn update_subject_status( } if let Some(deactivated) = &input.deactivated { let result = if deactivated.applied { - sqlx::query!( - "UPDATE users SET deactivated_at = NOW() WHERE did = $1", - did.as_str() - ) - .execute(&mut *tx) - .await + state.user_repo.deactivate_account(&did, None).await } else { - sqlx::query!( - "UPDATE users SET deactivated_at = NULL WHERE did = $1", - did.as_str() - ) - .execute(&mut *tx) - .await + state.user_repo.activate_account(&did).await }; if let Err(e) = result { error!( @@ -242,10 +216,6 @@ pub async fn update_subject_status( .into_response(); } } - if let Err(e) = tx.commit().await { - error!("Failed to commit transaction: {:?}", e); - return ApiError::InternalError(None).into_response(); - } if let Some(takedown) = &input.takedown { let status = if takedown.applied { Some("takendown") @@ -280,11 +250,7 @@ pub async fn update_subject_status( warn!("Failed to sequence account event for deactivation: {}", e); } } - if let Ok(Some(handle)) = - sqlx::query_scalar!("SELECT handle FROM users WHERE did = $1", did.as_str()) - .fetch_optional(&state.db) - .await - { + if let Ok(Some(handle)) = state.user_repo.get_handle_by_did(&did).await { let _ = state.cache.delete(&format!("handle:{}", handle)).await; } return ( @@ -304,25 +270,26 @@ pub async fn update_subject_status( } } Some("com.atproto.repo.strongRef") => { - let uri = input.subject.get("uri").and_then(|u| u.as_str()); - if let Some(uri) = uri { + let uri_str = input.subject.get("uri").and_then(|u| u.as_str()); + if let Some(uri_str) = uri_str { + let cid: CidLink = match uri_str.parse() { + Ok(c) => c, + Err(_) => return ApiError::InvalidRequest("Invalid CID format".into()).into_response(), + }; if let Some(takedown) = &input.takedown { let takedown_ref = if takedown.applied { - takedown.r#ref.clone() + takedown.r#ref.as_deref() } else { None }; - if let Err(e) = sqlx::query!( - "UPDATE records SET takedown_ref = $1 WHERE record_cid = $2", - takedown_ref, - uri - ) - .execute(&state.db) - .await + if let Err(e) = state + .repo_repo + .set_record_takedown(&cid, takedown_ref) + .await { error!( "Failed to update record takedown status for {}: {:?}", - uri, e + uri_str, e ); return ApiError::InternalError(Some( "Failed to update takedown status".into(), @@ -344,23 +311,24 @@ pub async fn update_subject_status( } } Some("com.atproto.admin.defs#repoBlobRef") => { - let cid = input.subject.get("cid").and_then(|c| c.as_str()); - if let Some(cid) = cid { + let cid_str = input.subject.get("cid").and_then(|c| c.as_str()); + if let Some(cid_str) = cid_str { + let cid: CidLink = match cid_str.parse() { + Ok(c) => c, + Err(_) => return ApiError::InvalidRequest("Invalid CID format".into()).into_response(), + }; if let Some(takedown) = &input.takedown { let takedown_ref = if takedown.applied { - takedown.r#ref.clone() + takedown.r#ref.as_deref() } else { None }; - if let Err(e) = sqlx::query!( - "UPDATE blobs SET takedown_ref = $1 WHERE cid = $2", - takedown_ref, - cid - ) - .execute(&state.db) - .await + if let Err(e) = state + .blob_repo + .update_blob_takedown(&cid, takedown_ref) + .await { - error!("Failed to update blob takedown status for {}: {:?}", cid, e); + error!("Failed to update blob takedown status for {}: {:?}", cid_str, e); return ApiError::InternalError(Some( "Failed to update takedown status".into(), )) diff --git a/crates/tranquil-pds/src/api/age_assurance.rs b/crates/tranquil-pds/src/api/age_assurance.rs index a6b9766..ced1298 100644 --- a/crates/tranquil-pds/src/api/age_assurance.rs +++ b/crates/tranquil-pds/src/api/age_assurance.rs @@ -43,7 +43,8 @@ async fn get_account_created_at(state: &AppState, headers: &HeaderMap) -> Option let http_uri = "/"; let auth_user = match validate_token_with_dpop( - &state.db, + state.user_repo.as_ref(), + state.oauth_repo.as_ref(), &extracted.token, extracted.is_dpop, dpop_proof, @@ -64,22 +65,18 @@ async fn get_account_created_at(state: &AppState, headers: &HeaderMap) -> Option } }; - let row = match sqlx::query!( - "SELECT created_at FROM users WHERE did = $1", - &auth_user.did - ) - .fetch_optional(&state.db) - .await - { - Ok(r) => { - tracing::debug!(?r, "age assurance: query result"); - r + match state.user_repo.get_by_did(&auth_user.did).await { + Ok(Some(user)) => { + tracing::debug!(created_at = ?user.created_at, "age assurance: got user"); + Some(user.created_at.to_rfc3339()) + } + Ok(None) => { + tracing::debug!("age assurance: user not found"); + None } Err(e) => { tracing::warn!(?e, "age assurance: query failed"); - return None; + None } - }; - - row.map(|r| r.created_at.to_rfc3339()) + } } diff --git a/crates/tranquil-pds/src/api/backup.rs b/crates/tranquil-pds/src/api/backup.rs index 643fdc6..59f05b2 100644 --- a/crates/tranquil-pds/src/api/backup.rs +++ b/crates/tranquil-pds/src/api/backup.rs @@ -14,6 +14,7 @@ use cid::Cid; use serde::{Deserialize, Serialize}; use serde_json::json; use std::str::FromStr; +use tranquil_db::{BackupRepository, OldBackupInfo}; use tracing::{error, info, warn}; #[derive(Serialize)] @@ -35,14 +36,8 @@ pub struct ListBackupsOutput { } pub async fn list_backups(State(state): State, auth: BearerAuth) -> Response { - let user = match sqlx::query!( - "SELECT id, backup_enabled FROM users WHERE did = $1", - auth.0.did.as_str() - ) - .fetch_optional(&state.db) - .await - { - Ok(Some(u)) => u, + let (user_id, backup_enabled) = match state.backup_repo.get_user_backup_status(&auth.0.did).await { + Ok(Some(status)) => status, Ok(None) => { return ApiError::AccountNotFound.into_response(); } @@ -52,18 +47,7 @@ pub async fn list_backups(State(state): State, auth: BearerAuth) -> Re } }; - let backups = match sqlx::query!( - r#" - SELECT id, repo_rev, repo_root_cid, block_count, size_bytes, created_at - FROM account_backups - WHERE user_id = $1 - ORDER BY created_at DESC - "#, - user.id - ) - .fetch_all(&state.db) - .await - { + let backups = match state.backup_repo.list_backups_for_user(user_id).await { Ok(rows) => rows, Err(e) => { error!("DB error fetching backups: {:?}", e); @@ -87,7 +71,7 @@ pub async fn list_backups(State(state): State, auth: BearerAuth) -> Re StatusCode::OK, Json(ListBackupsOutput { backups: backup_list, - backup_enabled: user.backup_enabled, + backup_enabled, }), ) .into_response() @@ -110,19 +94,7 @@ pub async fn get_backup( } }; - let backup = match sqlx::query!( - r#" - SELECT ab.storage_key, ab.repo_rev - FROM account_backups ab - JOIN users u ON u.id = ab.user_id - WHERE ab.id = $1 AND u.did = $2 - "#, - backup_id, - auth.0.did.as_str() - ) - .fetch_optional(&state.db) - .await - { + let backup_info = match state.backup_repo.get_backup_storage_info(backup_id, &auth.0.did).await { Ok(Some(b)) => b, Ok(None) => { return ApiError::BackupNotFound.into_response(); @@ -140,7 +112,7 @@ pub async fn get_backup( } }; - let car_bytes = match backup_storage.get_backup(&backup.storage_key).await { + let car_bytes = match backup_storage.get_backup(&backup_info.storage_key).await { Ok(bytes) => bytes, Err(e) => { error!("Failed to fetch backup from storage: {:?}", e); @@ -155,7 +127,7 @@ pub async fn get_backup( (axum::http::header::CONTENT_TYPE, "application/vnd.ipld.car"), ( axum::http::header::CONTENT_DISPOSITION, - &format!("attachment; filename=\"{}.car\"", backup.repo_rev), + &format!("attachment; filename=\"{}.car\"", backup_info.repo_rev), ), ], car_bytes, @@ -180,18 +152,7 @@ pub async fn create_backup(State(state): State, auth: BearerAuth) -> R } }; - let user = match sqlx::query!( - r#" - SELECT u.id, u.did, u.backup_enabled, u.deactivated_at, r.repo_root_cid, r.repo_rev - FROM users u - JOIN repos r ON r.user_id = u.id - WHERE u.did = $1 - "#, - auth.0.did.as_str() - ) - .fetch_optional(&state.db) - .await - { + let user = match state.backup_repo.get_user_for_backup(&auth.0.did).await { Ok(Some(u)) => u, Ok(None) => { return ApiError::AccountNotFound.into_response(); @@ -221,7 +182,7 @@ pub async fn create_backup(State(state): State, auth: BearerAuth) -> R }; let car_bytes = - match generate_full_backup(&state.db, &state.block_store, user.id, &head_cid).await { + match generate_full_backup(state.repo_repo.as_ref(), &state.block_store, user.id, &head_cid).await { Ok(bytes) => bytes, Err(e) => { error!("Failed to generate CAR: {:?}", e); @@ -244,22 +205,14 @@ pub async fn create_backup(State(state): State, auth: BearerAuth) -> R } }; - let backup_id = match sqlx::query_scalar!( - r#" - INSERT INTO account_backups (user_id, storage_key, repo_root_cid, repo_rev, block_count, size_bytes) - VALUES ($1, $2, $3, $4, $5, $6) - RETURNING id - "#, + let backup_id = match state.backup_repo.insert_backup( user.id, - storage_key, - user.repo_root_cid, - repo_rev, + &storage_key, + &user.repo_root_cid, + &repo_rev, block_count, - size_bytes - ) - .fetch_one(&state.db) - .await - { + size_bytes, + ).await { Ok(id) => id, Err(e) => { error!("DB error inserting backup: {:?}", e); @@ -282,7 +235,7 @@ pub async fn create_backup(State(state): State, auth: BearerAuth) -> R ); let retention = BackupStorage::retention_count(); - if let Err(e) = cleanup_old_backups(&state.db, backup_storage, user.id, retention).await { + if let Err(e) = cleanup_old_backups(state.backup_repo.as_ref(), backup_storage, user.id, retention).await { warn!(did = %user.did, error = %e, "Failed to cleanup old backups after manual backup"); } @@ -299,25 +252,15 @@ pub async fn create_backup(State(state): State, auth: BearerAuth) -> R } async fn cleanup_old_backups( - db: &sqlx::PgPool, + backup_repo: &dyn BackupRepository, backup_storage: &BackupStorage, user_id: uuid::Uuid, retention_count: u32, ) -> Result<(), String> { - let old_backups = sqlx::query!( - r#" - SELECT id, storage_key - FROM account_backups - WHERE user_id = $1 - ORDER BY created_at DESC - OFFSET $2 - "#, - user_id, - retention_count as i64 - ) - .fetch_all(db) - .await - .map_err(|e| format!("DB error fetching old backups: {}", e))?; + let old_backups: Vec = backup_repo + .get_old_backups(user_id, retention_count as i64) + .await + .map_err(|e| format!("DB error fetching old backups: {}", e))?; for backup in old_backups { if let Err(e) = backup_storage.delete_backup(&backup.storage_key).await { @@ -329,8 +272,8 @@ async fn cleanup_old_backups( continue; } - sqlx::query!("DELETE FROM account_backups WHERE id = $1", backup.id) - .execute(db) + backup_repo + .delete_backup(backup.id) .await .map_err(|e| format!("Failed to delete old backup record: {}", e))?; } @@ -355,19 +298,7 @@ pub async fn delete_backup( } }; - let backup = match sqlx::query!( - r#" - SELECT ab.id, ab.storage_key, u.deactivated_at - FROM account_backups ab - JOIN users u ON u.id = ab.user_id - WHERE ab.id = $1 AND u.did = $2 - "#, - backup_id, - auth.0.did.as_str() - ) - .fetch_optional(&state.db) - .await - { + let backup = match state.backup_repo.get_backup_for_deletion(backup_id, &auth.0.did).await { Ok(Some(b)) => b, Ok(None) => { return ApiError::BackupNotFound.into_response(); @@ -392,10 +323,7 @@ pub async fn delete_backup( ); } - if let Err(e) = sqlx::query!("DELETE FROM account_backups WHERE id = $1", backup.id) - .execute(&state.db) - .await - { + if let Err(e) = state.backup_repo.delete_backup(backup.id).await { error!("DB error deleting backup: {:?}", e); return ApiError::InternalError(Some("Failed to delete backup".into())).into_response(); } @@ -416,14 +344,8 @@ pub async fn set_backup_enabled( auth: BearerAuth, Json(input): Json, ) -> Response { - let user = match sqlx::query!( - "SELECT deactivated_at FROM users WHERE did = $1", - auth.0.did.as_str() - ) - .fetch_optional(&state.db) - .await - { - Ok(Some(u)) => u, + let deactivated_at = match state.backup_repo.get_user_deactivated_status(&auth.0.did).await { + Ok(Some(status)) => status, Ok(None) => { return ApiError::AccountNotFound.into_response(); } @@ -433,18 +355,11 @@ pub async fn set_backup_enabled( } }; - if user.deactivated_at.is_some() { + if deactivated_at.is_some() { return ApiError::AccountDeactivated.into_response(); } - if let Err(e) = sqlx::query!( - "UPDATE users SET backup_enabled = $1 WHERE did = $2", - input.enabled, - auth.0.did.as_str() - ) - .execute(&state.db) - .await - { + if let Err(e) = state.backup_repo.update_backup_enabled(&auth.0.did, input.enabled).await { error!("DB error updating backup_enabled: {:?}", e); return ApiError::InternalError(Some("Failed to update setting".into())).into_response(); } @@ -455,11 +370,8 @@ pub async fn set_backup_enabled( } pub async fn export_blobs(State(state): State, auth: BearerAuth) -> Response { - let user = match sqlx::query!("SELECT id FROM users WHERE did = $1", auth.0.did.as_str()) - .fetch_optional(&state.db) - .await - { - Ok(Some(u)) => u, + let user_id = match state.backup_repo.get_user_id_by_did(&auth.0.did).await { + Ok(Some(id)) => id, Ok(None) => { return ApiError::AccountNotFound.into_response(); } @@ -469,18 +381,7 @@ pub async fn export_blobs(State(state): State, auth: BearerAuth) -> Re } }; - let blobs = match sqlx::query!( - r#" - SELECT DISTINCT b.cid, b.storage_key, b.mime_type - FROM blobs b - JOIN record_blobs rb ON rb.blob_cid = b.cid - WHERE rb.repo_id = $1 - "#, - user.id - ) - .fetch_all(&state.db) - .await - { + let blobs = match state.backup_repo.get_blobs_for_export(user_id).await { Ok(rows) => rows, Err(e) => { error!("DB error fetching blobs: {:?}", e); diff --git a/crates/tranquil-pds/src/api/delegation.rs b/crates/tranquil-pds/src/api/delegation.rs index 2dd7cd7..67c998d 100644 --- a/crates/tranquil-pds/src/api/delegation.rs +++ b/crates/tranquil-pds/src/api/delegation.rs @@ -1,8 +1,7 @@ use crate::api::error::ApiError; use crate::api::repo::record::utils::create_signed_commit; use crate::auth::BearerAuth; -use crate::delegation::{self, DelegationActionType}; -use crate::oauth::db as oauth_db; +use crate::delegation::{DelegationActionType, SCOPE_PRESETS, scopes}; use crate::state::{AppState, RateLimitKind}; use crate::types::{Did, Handle, Nsid, Rkey}; use crate::util::extract_client_ip; @@ -35,7 +34,11 @@ pub struct ListControllersResponse { } pub async fn list_controllers(State(state): State, auth: BearerAuth) -> Response { - let controllers = match delegation::get_delegations_for_account(&state.db, &auth.0.did).await { + let controllers = match state + .delegation_repo + .get_delegations_for_account(&auth.0.did) + .await + { Ok(c) => c, Err(e) => { tracing::error!("Failed to list controllers: {:?}", e); @@ -49,7 +52,7 @@ pub async fn list_controllers(State(state): State, auth: BearerAuth) - .into_iter() .map(|c| ControllerInfo { did: c.did.into(), - handle: c.handle, + handle: c.handle.into(), granted_scopes: c.granted_scopes, granted_at: c.granted_at, is_active: c.is_active, @@ -70,23 +73,23 @@ pub async fn add_controller( auth: BearerAuth, Json(input): Json, ) -> Response { - if let Err(e) = delegation::scopes::validate_delegation_scopes(&input.granted_scopes) { + if let Err(e) = scopes::validate_delegation_scopes(&input.granted_scopes) { return ApiError::InvalidScopes(e).into_response(); } - let controller_exists: bool = sqlx::query_scalar!( - r#"SELECT EXISTS(SELECT 1 FROM users WHERE did = $1) as "exists!""#, - input.controller_did.as_str() - ) - .fetch_one(&state.db) - .await - .unwrap_or(false); + let controller_exists = state + .user_repo + .get_by_did(&input.controller_did) + .await + .ok() + .flatten() + .is_some(); if !controller_exists { return ApiError::ControllerNotFound.into_response(); } - match delegation::controls_any_accounts(&state.db, &auth.0.did).await { + match state.delegation_repo.controls_any_accounts(&auth.0.did).await { Ok(true) => { return ApiError::InvalidDelegation( "Cannot add controllers to an account that controls other accounts".into(), @@ -101,7 +104,11 @@ pub async fn add_controller( Ok(false) => {} } - match delegation::has_any_controllers(&state.db, &input.controller_did).await { + match state + .delegation_repo + .has_any_controllers(&input.controller_did) + .await + { Ok(true) => { return ApiError::InvalidDelegation( "Cannot add a controlled account as a controller".into(), @@ -116,29 +123,31 @@ pub async fn add_controller( Ok(false) => {} } - match delegation::create_delegation( - &state.db, - &auth.0.did, - &input.controller_did, - &input.granted_scopes, - &auth.0.did, - ) - .await + match state + .delegation_repo + .create_delegation( + &auth.0.did, + &input.controller_did, + &input.granted_scopes, + &auth.0.did, + ) + .await { Ok(_) => { - let _ = delegation::log_delegation_action( - &state.db, - &auth.0.did, - &auth.0.did, - Some(&input.controller_did), - DelegationActionType::GrantCreated, - Some(serde_json::json!({ - "granted_scopes": input.granted_scopes - })), - None, - None, - ) - .await; + let _ = state + .delegation_repo + .log_delegation_action( + &auth.0.did, + &auth.0.did, + Some(&input.controller_did), + DelegationActionType::GrantCreated, + Some(serde_json::json!({ + "granted_scopes": input.granted_scopes + })), + None, + None, + ) + .await; ( StatusCode::OK, @@ -165,45 +174,39 @@ pub async fn remove_controller( auth: BearerAuth, Json(input): Json, ) -> Response { - match delegation::revoke_delegation(&state.db, &auth.0.did, &input.controller_did, &auth.0.did) + match state + .delegation_repo + .revoke_delegation(&auth.0.did, &input.controller_did, &auth.0.did) .await { Ok(true) => { - let revoked_app_passwords = sqlx::query_scalar!( - r#"DELETE FROM app_passwords - WHERE user_id = (SELECT id FROM users WHERE did = $1) - AND created_by_controller_did = $2 - RETURNING id"#, - &auth.0.did, - input.controller_did.as_str() - ) - .fetch_all(&state.db) - .await - .map(|r| r.len()) - .unwrap_or(0); - - let revoked_oauth_tokens = oauth_db::revoke_tokens_for_controller( - &state.db, - &auth.0.did, - &input.controller_did, - ) - .await - .unwrap_or(0); - - let _ = delegation::log_delegation_action( - &state.db, - &auth.0.did, - &auth.0.did, - Some(&input.controller_did), - DelegationActionType::GrantRevoked, - Some(serde_json::json!({ - "revoked_app_passwords": revoked_app_passwords, - "revoked_oauth_tokens": revoked_oauth_tokens - })), - None, - None, - ) - .await; + let revoked_app_passwords = state + .session_repo + .delete_app_passwords_by_controller(&auth.0.did, &input.controller_did) + .await + .unwrap_or(0) as usize; + + let revoked_oauth_tokens = state + .oauth_repo + .revoke_tokens_for_controller(&auth.0.did, &input.controller_did) + .await + .unwrap_or(0); + + let _ = state + .delegation_repo + .log_delegation_action( + &auth.0.did, + &auth.0.did, + Some(&input.controller_did), + DelegationActionType::GrantRevoked, + Some(serde_json::json!({ + "revoked_app_passwords": revoked_app_passwords, + "revoked_oauth_tokens": revoked_oauth_tokens + })), + None, + None, + ) + .await; ( StatusCode::OK, @@ -232,32 +235,30 @@ pub async fn update_controller_scopes( auth: BearerAuth, Json(input): Json, ) -> Response { - if let Err(e) = delegation::scopes::validate_delegation_scopes(&input.granted_scopes) { + if let Err(e) = scopes::validate_delegation_scopes(&input.granted_scopes) { return ApiError::InvalidScopes(e).into_response(); } - match delegation::update_delegation_scopes( - &state.db, - &auth.0.did, - &input.controller_did, - &input.granted_scopes, - ) - .await + match state + .delegation_repo + .update_delegation_scopes(&auth.0.did, &input.controller_did, &input.granted_scopes) + .await { Ok(true) => { - let _ = delegation::log_delegation_action( - &state.db, - &auth.0.did, - &auth.0.did, - Some(&input.controller_did), - DelegationActionType::ScopesModified, - Some(serde_json::json!({ - "new_scopes": input.granted_scopes - })), - None, - None, - ) - .await; + let _ = state + .delegation_repo + .log_delegation_action( + &auth.0.did, + &auth.0.did, + Some(&input.controller_did), + DelegationActionType::ScopesModified, + Some(serde_json::json!({ + "new_scopes": input.granted_scopes + })), + None, + None, + ) + .await; ( StatusCode::OK, @@ -291,7 +292,11 @@ pub struct ListControlledAccountsResponse { } pub async fn list_controlled_accounts(State(state): State, auth: BearerAuth) -> Response { - let accounts = match delegation::get_accounts_controlled_by(&state.db, &auth.0.did).await { + let accounts = match state + .delegation_repo + .get_accounts_controlled_by(&auth.0.did) + .await + { Ok(a) => a, Err(e) => { tracing::error!("Failed to list controlled accounts: {:?}", e); @@ -305,7 +310,7 @@ pub async fn list_controlled_accounts(State(state): State, auth: Beare .into_iter() .map(|a| DelegatedAccountInfo { did: a.did.into(), - handle: a.handle, + handle: a.handle.into(), granted_scopes: a.granted_scopes, granted_at: a.granted_at, }) @@ -352,19 +357,21 @@ pub async fn get_audit_log( let limit = params.limit.clamp(1, 100); let offset = params.offset.max(0); - let entries = - match delegation::audit::get_audit_log_for_account(&state.db, &auth.0.did, limit, offset) - .await - { - Ok(e) => e, - Err(e) => { - tracing::error!("Failed to get audit log: {:?}", e); - return ApiError::InternalError(Some("Failed to get audit log".into())) - .into_response(); - } - }; + let entries = match state + .delegation_repo + .get_audit_log_for_account(&auth.0.did, limit, offset) + .await + { + Ok(e) => e, + Err(e) => { + tracing::error!("Failed to get audit log: {:?}", e); + return ApiError::InternalError(Some("Failed to get audit log".into())).into_response(); + } + }; - let total = delegation::audit::count_audit_log_entries(&state.db, &auth.0.did) + let total = state + .delegation_repo + .count_audit_log_entries(&auth.0.did) .await .unwrap_or_default(); @@ -401,7 +408,7 @@ pub struct GetScopePresetsResponse { pub async fn get_scope_presets() -> Response { Json(GetScopePresetsResponse { - presets: delegation::SCOPE_PRESETS + presets: SCOPE_PRESETS .iter() .map(|p| ScopePresetInfo { name: p.name, @@ -448,11 +455,11 @@ pub async fn create_delegated_account( .into_response(); } - if let Err(e) = delegation::scopes::validate_delegation_scopes(&input.controller_scopes) { + if let Err(e) = scopes::validate_delegation_scopes(&input.controller_scopes) { return ApiError::InvalidScopes(e).into_response(); } - match delegation::has_any_controllers(&state.db, &auth.0.did).await { + match state.delegation_repo.has_any_controllers(&auth.0.did).await { Ok(true) => { return ApiError::InvalidDelegation( "Cannot create delegated accounts from a controlled account".into(), @@ -468,7 +475,8 @@ pub async fn create_delegated_account( } let hostname = std::env::var("PDS_HOSTNAME").unwrap_or_else(|_| "localhost".to_string()); - let pds_suffix = format!(".{}", hostname); + let hostname_for_handles = hostname.split(':').next().unwrap_or(&hostname); + let pds_suffix = format!(".{}", hostname_for_handles); let handle = if !input.handle.contains('.') || input.handle.ends_with(&pds_suffix) { let handle_to_validate = if input.handle.ends_with(&pds_suffix) { @@ -480,7 +488,7 @@ pub async fn create_delegated_account( &input.handle }; match crate::api::validation::validate_short_handle(handle_to_validate) { - Ok(h) => format!("{}.{}", h, hostname), + Ok(h) => format!("{}.{}", h, hostname_for_handles), Err(e) => { return ApiError::InvalidRequest(e.to_string()).into_response(); } @@ -501,17 +509,9 @@ pub async fn create_delegated_account( } if let Some(ref code) = input.invite_code { - let valid = sqlx::query_scalar!( - "SELECT available_uses > 0 AND NOT disabled FROM invite_codes WHERE code = $1", - code - ) - .fetch_optional(&state.db) - .await - .ok() - .flatten() - .unwrap_or(Some(false)); + let valid = state.infra_repo.is_invite_code_valid(code).await.unwrap_or(false); - if valid != Some(true) { + if !valid { return ApiError::InvalidInviteCode.into_response(); } } else { @@ -572,44 +572,6 @@ pub async fn create_delegated_account( let handle = Handle::new_unchecked(&handle); info!(did = %did, handle = %handle, controller = %&auth.0.did, "Created DID for delegated account"); - let mut tx = match state.db.begin().await { - Ok(tx) => tx, - Err(e) => { - error!("Error starting transaction: {:?}", e); - return ApiError::InternalError(None).into_response(); - } - }; - - let user_insert: Result<(uuid::Uuid,), _> = sqlx::query_as( - r#"INSERT INTO users ( - handle, email, did, password_hash, password_required, - account_type, preferred_comms_channel - ) VALUES ($1, $2, $3, NULL, FALSE, 'delegated'::account_type, 'email'::comms_channel) RETURNING id"#, - ) - .bind(handle.as_str()) - .bind(&email) - .bind(did.as_str()) - .fetch_one(&mut *tx) - .await; - - let user_id = match user_insert { - Ok((id,)) => id, - Err(e) => { - if let Some(db_err) = e.as_database_error() - && db_err.code().as_deref() == Some("23505") - { - let constraint = db_err.constraint().unwrap_or(""); - if constraint.contains("handle") { - return ApiError::HandleNotAvailable(None).into_response(); - } else if constraint.contains("email") { - return ApiError::EmailTaken.into_response(); - } - } - error!("Error inserting user: {:?}", e); - return ApiError::InternalError(None).into_response(); - } - }; - let encrypted_key_bytes = match crate::config::encrypt_key(&secret_key_bytes) { Ok(bytes) => bytes, Err(e) => { @@ -618,34 +580,6 @@ pub async fn create_delegated_account( } }; - if let Err(e) = sqlx::query!( - "INSERT INTO user_keys (user_id, key_bytes, encryption_version, encrypted_at) VALUES ($1, $2, $3, NOW())", - user_id, - &encrypted_key_bytes[..], - crate::config::ENCRYPTION_VERSION - ) - .execute(&mut *tx) - .await - { - error!("Error inserting user key: {:?}", e); - return ApiError::InternalError(None).into_response(); - } - - if let Err(e) = sqlx::query!( - r#"INSERT INTO account_delegations (delegated_did, controller_did, granted_scopes, granted_by) - VALUES ($1, $2, $3, $4)"#, - did.as_str(), - auth.0.did.as_str(), - input.controller_scopes, - auth.0.did.as_str() - ) - .execute(&mut *tx) - .await - { - error!("Error creating initial delegation: {:?}", e); - return ApiError::InternalError(None).into_response(); - } - let mst = Mst::new(Arc::new(state.block_store.clone())); let mst_root = match mst.persist().await { Ok(c) => c, @@ -670,58 +604,35 @@ pub async fn create_delegated_account( return ApiError::InternalError(None).into_response(); } }; - let commit_cid_str = commit_cid.to_string(); - let rev_str = rev.as_ref().to_string(); - if let Err(e) = sqlx::query!( - "INSERT INTO repos (user_id, repo_root_cid, repo_rev) VALUES ($1, $2, $3)", - user_id, - commit_cid_str, - rev_str - ) - .execute(&mut *tx) - .await - { - error!("Error inserting repo: {:?}", e); - return ApiError::InternalError(None).into_response(); - } let genesis_block_cids = vec![mst_root.to_bytes(), commit_cid.to_bytes()]; - if let Err(e) = sqlx::query!( - r#" - INSERT INTO user_blocks (user_id, block_cid) - SELECT $1, block_cid FROM UNNEST($2::bytea[]) AS t(block_cid) - ON CONFLICT (user_id, block_cid) DO NOTHING - "#, - user_id, - &genesis_block_cids - ) - .execute(&mut *tx) - .await - { - error!("Error inserting user_blocks: {:?}", e); - return ApiError::InternalError(None).into_response(); - } - - if let Some(ref code) = input.invite_code { - let _ = sqlx::query!( - "UPDATE invite_codes SET available_uses = available_uses - 1 WHERE code = $1", - code - ) - .execute(&mut *tx) - .await; - let _ = sqlx::query!( - "INSERT INTO invite_code_uses (code, used_by_user) VALUES ($1, $2)", - code, - user_id - ) - .execute(&mut *tx) - .await; - } + let create_input = tranquil_db_traits::CreateDelegatedAccountInput { + handle: handle.clone(), + email: email.clone(), + did: did.clone(), + controller_did: auth.0.did.clone(), + controller_scopes: input.controller_scopes.clone(), + encrypted_key_bytes, + encryption_version: crate::config::ENCRYPTION_VERSION, + commit_cid: commit_cid.to_string(), + repo_rev: rev.as_ref().to_string(), + genesis_block_cids, + invite_code: input.invite_code.clone(), + }; - if let Err(e) = tx.commit().await { - error!("Error committing transaction: {:?}", e); - return ApiError::InternalError(None).into_response(); - } + let _user_id = match state.user_repo.create_delegated_account(&create_input).await { + Ok(id) => id, + Err(tranquil_db_traits::CreateAccountError::HandleTaken) => { + return ApiError::HandleNotAvailable(None).into_response(); + } + Err(tranquil_db_traits::CreateAccountError::EmailTaken) => { + return ApiError::EmailTaken.into_response(); + } + Err(e) => { + error!("Error creating delegated account: {:?}", e); + return ApiError::InternalError(None).into_response(); + } + }; if let Err(e) = crate::api::repo::record::sequence_identity_event(&state, &did, Some(&handle)).await @@ -751,20 +662,21 @@ pub async fn create_delegated_account( warn!("Failed to create default profile for {}: {}", did, e); } - let _ = delegation::log_delegation_action( - &state.db, - &did, - &auth.0.did, - Some(&auth.0.did), - DelegationActionType::GrantCreated, - Some(json!({ - "account_created": true, - "granted_scopes": input.controller_scopes - })), - None, - None, - ) - .await; + let _ = state + .delegation_repo + .log_delegation_action( + &did, + &auth.0.did, + Some(&auth.0.did), + DelegationActionType::GrantCreated, + Some(json!({ + "account_created": true, + "granted_scopes": input.controller_scopes + })), + None, + None, + ) + .await; info!(did = %did, handle = %handle, controller = %&auth.0.did, "Delegated account created"); diff --git a/crates/tranquil-pds/src/api/error.rs b/crates/tranquil-pds/src/api/error.rs index add3c24..9ca3c2d 100644 --- a/crates/tranquil-pds/src/api/error.rs +++ b/crates/tranquil-pds/src/api/error.rs @@ -158,7 +158,8 @@ impl ApiError { Self::RepoTakendown | Self::RepoDeactivated | Self::RepoNotFound(_) => { StatusCode::BAD_REQUEST } - Self::InvalidSwap(_) | Self::TotpAlreadyEnabled => StatusCode::CONFLICT, + Self::TotpAlreadyEnabled => StatusCode::CONFLICT, + Self::InvalidSwap(_) => StatusCode::BAD_REQUEST, Self::InvalidRequest(_) | Self::InvalidHandle(_) | Self::HandleNotAvailable(_) @@ -474,22 +475,14 @@ impl From for ApiError { crate::auth::TokenValidationError::OAuthTokenExpired => { Self::OAuthExpiredToken(Some("Token has expired".to_string())) } - } - } -} - -impl From for ApiError { - fn from(e: crate::util::DbLookupError) -> Self { - match e { - crate::util::DbLookupError::NotFound => Self::AccountNotFound, - crate::util::DbLookupError::DatabaseError(db_err) => { - tracing::error!("Database error: {:?}", db_err); - Self::DatabaseError + crate::auth::TokenValidationError::InvalidToken => { + Self::AuthenticationFailed(Some("Invalid token format".to_string())) } } } } + impl From for ApiError { fn from(e: crate::auth::extractor::AuthError) -> Self { match e { @@ -644,6 +637,15 @@ impl From for ApiError { } } +pub fn parse_did(s: &str) -> Result { + s.parse() + .map_err(|_| ApiError::InvalidDid("Invalid DID format".into()).into_response()) +} + +pub fn parse_did_option(s: Option<&str>) -> Result, Response> { + s.map(parse_did).transpose() +} + pub struct AtpJson(pub T); impl FromRequest for AtpJson diff --git a/crates/tranquil-pds/src/api/identity/account.rs b/crates/tranquil-pds/src/api/identity/account.rs index 4e9cb38..d409b31 100644 --- a/crates/tranquil-pds/src/api/identity/account.rs +++ b/crates/tranquil-pds/src/api/identity/account.rs @@ -243,32 +243,20 @@ pub async fn create_account( }) }; let hostname = std::env::var("PDS_HOSTNAME").unwrap_or_else(|_| "localhost".to_string()); + let hostname_for_handles = hostname.split(':').next().unwrap_or(&hostname); let pds_endpoint = format!("https://{}", hostname); - let suffix = format!(".{}", hostname); + let suffix = format!(".{}", hostname_for_handles); let handle = if input.handle.ends_with(&suffix) { - format!("{}.{}", validated_short_handle, hostname) + format!("{}.{}", validated_short_handle, hostname_for_handles) } else if input.handle.contains('.') { validated_short_handle.clone() } else { - format!("{}.{}", validated_short_handle, hostname) + format!("{}.{}", validated_short_handle, hostname_for_handles) }; let (secret_key_bytes, reserved_key_id): (Vec, Option) = if let Some(signing_key_did) = &input.signing_key { - let reserved = sqlx::query!( - r#" - SELECT id, private_key_bytes - FROM reserved_signing_keys - WHERE public_key_did_key = $1 - AND used_at IS NULL - AND expires_at > NOW() - FOR UPDATE - "#, - signing_key_did - ) - .fetch_optional(&state.db) - .await; - match reserved { - Ok(Some(row)) => (row.private_key_bytes, Some(row.id)), + match state.infra_repo.get_reserved_signing_key(signing_key_did).await { + Ok(Some(key)) => (key.private_key_bytes, Some(key.id)), Ok(None) => { return ApiError::InvalidSigningKey.into_response(); } @@ -294,7 +282,7 @@ pub async fn create_account( if !crate::api::server::meta::is_self_hosted_did_web_enabled() { return ApiError::SelfHostedDidWebDisabled.into_response(); } - let subdomain_host = format!("{}.{}", input.handle, hostname); + let subdomain_host = format!("{}.{}", input.handle, hostname_for_handles); let encoded_subdomain = subdomain_host.replace(':', "%3A"); let self_hosted_did = format!("did:web:{}", encoded_subdomain); info!(did = %self_hosted_did, "Creating self-hosted did:web account (subdomain)"); @@ -414,55 +402,17 @@ pub async fn create_account( } } }; - let mut tx = match state.db.begin().await { - Ok(tx) => tx, - Err(e) => { - error!("Error starting transaction: {:?}", e); - return ApiError::InternalError(None).into_response(); - } - }; if is_migration { - let existing_account: Option<(uuid::Uuid, String, Option>)> = - sqlx::query_as("SELECT id, handle, deactivated_at FROM users WHERE did = $1 FOR UPDATE") - .bind(&did) - .fetch_optional(&mut *tx) - .await - .unwrap_or(None); - if let Some((account_id, old_handle, deactivated_at)) = existing_account { - if deactivated_at.is_some() { - info!(did = %did, old_handle = %old_handle, new_handle = %handle, "Preparing existing account for inbound migration"); - let update_result: Result<_, sqlx::Error> = - sqlx::query("UPDATE users SET handle = $1 WHERE id = $2") - .bind(&handle) - .bind(account_id) - .execute(&mut *tx) - .await; - if let Err(e) = update_result { - if let Some(db_err) = e.as_database_error() - && db_err - .constraint() - .map(|c| c.contains("handle")) - .unwrap_or(false) - { - return ApiError::HandleTaken.into_response(); - } - error!("Error reactivating account: {:?}", e); - return ApiError::InternalError(None).into_response(); - } - if let Err(e) = tx.commit().await { - error!("Error committing reactivation: {:?}", e); - return ApiError::InternalError(None).into_response(); - } - let key_row: Option<(Vec, i32)> = sqlx::query_as( - "SELECT key_bytes, encryption_version FROM user_keys WHERE user_id = $1", - ) - .bind(account_id) - .fetch_optional(&state.db) - .await - .unwrap_or(None); - let secret_key_bytes = match key_row { - Some((key_bytes, encryption_version)) => { - match crate::config::decrypt_key(&key_bytes, Some(encryption_version)) { + let reactivate_input = tranquil_db_traits::MigrationReactivationInput { + did: Did::new_unchecked(&did), + new_handle: Handle::new_unchecked(&handle), + }; + match state.user_repo.reactivate_migration_account(&reactivate_input).await { + Ok(reactivated) => { + info!(did = %did, old_handle = %reactivated.old_handle, new_handle = %handle, "Preparing existing account for inbound migration"); + let secret_key_bytes = match state.user_repo.get_user_key_by_id(reactivated.user_id).await { + Ok(Some(key_info)) => { + match crate::config::decrypt_key(&key_info.key_bytes, key_info.encryption_version) { Ok(k) => k, Err(e) => { error!("Error decrypting key for reactivated account: {:?}", e); @@ -470,7 +420,7 @@ pub async fn create_account( } } } - None => { + _ => { error!("No signing key found for reactivated account"); return ApiError::InternalError(Some( "Account signing key not found".into(), @@ -496,17 +446,19 @@ pub async fn create_account( return ApiError::InternalError(None).into_response(); } }; - let session_result: Result<_, sqlx::Error> = sqlx::query( - "INSERT INTO session_tokens (did, access_jti, refresh_jti, access_expires_at, refresh_expires_at) VALUES ($1, $2, $3, $4, $5)", - ) - .bind(&did) - .bind(&access_meta.jti) - .bind(&refresh_meta.jti) - .bind(access_meta.expires_at) - .bind(refresh_meta.expires_at) - .execute(&state.db) - .await; - if let Err(e) = session_result { + let session_data = tranquil_db_traits::SessionTokenCreate { + did: Did::new_unchecked(&did), + access_jti: access_meta.jti.clone(), + refresh_jti: refresh_meta.jti.clone(), + access_expires_at: access_meta.expires_at, + refresh_expires_at: refresh_meta.expires_at, + legacy_login: false, + mfa_verified: false, + scope: None, + controller_did: None, + app_password_name: None, + }; + if let Err(e) = state.session_repo.create_session(&session_data).await { error!("Error creating session: {:?}", e); return ApiError::InternalError(None).into_response(); } @@ -514,7 +466,7 @@ pub async fn create_account( axum::http::StatusCode::OK, Json(CreateAccountOutput { handle: handle.clone().into(), - did: did.clone().into(), + did: Did::new_unchecked(&did), did_doc: state.did_resolver.resolve_did_document(&did).await, access_jwt: access_meta.token, refresh_jwt: refresh_meta.token, @@ -523,20 +475,34 @@ pub async fn create_account( }), ) .into_response(); - } else { + } + Err(tranquil_db_traits::MigrationReactivationError::NotFound) => { + } + Err(tranquil_db_traits::MigrationReactivationError::NotDeactivated) => { return ApiError::AccountAlreadyExists.into_response(); } + Err(tranquil_db_traits::MigrationReactivationError::HandleTaken) => { + return ApiError::HandleTaken.into_response(); + } + Err(e) => { + error!("Error reactivating migration account: {:?}", e); + return ApiError::InternalError(None).into_response(); + } } } - let exists_result: Option<(i32,)> = - sqlx::query_as("SELECT 1 FROM users WHERE handle = $1 AND deactivated_at IS NULL") - .bind(&handle) - .fetch_optional(&mut *tx) - .await - .unwrap_or(None); - if exists_result.is_some() { + + let handle_typed = Handle::new_unchecked(&handle); + let handle_available = match state.user_repo.check_handle_available_for_new_account(&handle_typed).await { + Ok(available) => available, + Err(e) => { + error!("Error checking handle availability: {:?}", e); + return ApiError::InternalError(None).into_response(); + } + }; + if !handle_available { return ApiError::HandleTaken.into_response(); } + let invite_code_required = std::env::var("INVITE_CODE_REQUIRED") .map(|v| v == "true" || v == "1") .unwrap_or(false); @@ -552,37 +518,18 @@ pub async fn create_account( if let Some(code) = &input.invite_code && !code.trim().is_empty() { - let invite_query = sqlx::query!( - "SELECT available_uses FROM invite_codes WHERE code = $1 FOR UPDATE", - code - ) - .fetch_optional(&mut *tx) - .await; - match invite_query { - Ok(Some(row)) => { - if row.available_uses <= 0 { - return ApiError::InvalidInviteCode.into_response(); - } - let update_invite = sqlx::query!( - "UPDATE invite_codes SET available_uses = available_uses - 1 WHERE code = $1", - code - ) - .execute(&mut *tx) - .await; - if let Err(e) = update_invite { - error!("Error updating invite code: {:?}", e); - return ApiError::InternalError(None).into_response(); - } - } - Ok(None) => { - return ApiError::InvalidInviteCode.into_response(); - } + let valid = match state.user_repo.check_and_consume_invite_code(code).await { + Ok(v) => v, Err(e) => { error!("Error checking invite code: {:?}", e); return ApiError::InternalError(None).into_response(); } + }; + if !valid { + return ApiError::InvalidInviteCode.into_response(); } } + if let Err(e) = validate_password(&input.password) { return ApiError::InvalidRequest(e.to_string()).into_response(); } @@ -600,74 +547,12 @@ pub async fn create_account( return ApiError::InternalError(None).into_response(); } }; - let is_first_user = sqlx::query_scalar!("SELECT COUNT(*) as count FROM users") - .fetch_one(&mut *tx) - .await - .map(|c| c.unwrap_or(0) == 0) - .unwrap_or(false); + let deactivated_at: Option> = if is_migration || is_did_web_byod { Some(chrono::Utc::now()) } else { None }; - let user_insert: Result<(uuid::Uuid,), _> = sqlx::query_as( - r#"INSERT INTO users ( - handle, email, did, password_hash, - preferred_comms_channel, - discord_id, telegram_username, signal_number, - is_admin, deactivated_at, email_verified - ) VALUES ($1, $2, $3, $4, $5::comms_channel, $6, $7, $8, $9, $10, $11) RETURNING id"#, - ) - .bind(&handle) - .bind(&email) - .bind(&did) - .bind(&password_hash) - .bind(verification_channel) - .bind( - input - .discord_id - .as_deref() - .map(|s| s.trim()) - .filter(|s| !s.is_empty()), - ) - .bind( - input - .telegram_username - .as_deref() - .map(|s| s.trim()) - .filter(|s| !s.is_empty()), - ) - .bind( - input - .signal_number - .as_deref() - .map(|s| s.trim()) - .filter(|s| !s.is_empty()), - ) - .bind(is_first_user) - .bind(deactivated_at) - .bind(false) - .fetch_one(&mut *tx) - .await; - let user_id = match user_insert { - Ok((id,)) => id, - Err(e) => { - if let Some(db_err) = e.as_database_error() - && db_err.code().as_deref() == Some("23505") - { - let constraint = db_err.constraint().unwrap_or(""); - if constraint.contains("handle") || constraint.contains("users_handle") { - return ApiError::HandleNotAvailable(None).into_response(); - } else if constraint.contains("email") || constraint.contains("users_email") { - return ApiError::EmailTaken.into_response(); - } else if constraint.contains("did") || constraint.contains("users_did") { - return ApiError::AccountAlreadyExists.into_response(); - } - } - error!("Error inserting user: {:?}", e); - return ApiError::InternalError(None).into_response(); - } - }; let encrypted_key_bytes = match crate::config::encrypt_key(&secret_key_bytes) { Ok(enc) => enc, @@ -676,30 +561,7 @@ pub async fn create_account( return ApiError::InternalError(None).into_response(); } }; - let key_insert = sqlx::query!( - "INSERT INTO user_keys (user_id, key_bytes, encryption_version, encrypted_at) VALUES ($1, $2, $3, NOW())", - user_id, - &encrypted_key_bytes[..], - crate::config::ENCRYPTION_VERSION - ) - .execute(&mut *tx) - .await; - if let Err(e) = key_insert { - error!("Error inserting user key: {:?}", e); - return ApiError::InternalError(None).into_response(); - } - if let Some(key_id) = reserved_key_id { - let mark_used = sqlx::query!( - "UPDATE reserved_signing_keys SET used_at = NOW() WHERE id = $1", - key_id - ) - .execute(&mut *tx) - .await; - if let Err(e) = mark_used { - error!("Error marking reserved key as used: {:?}", e); - return ApiError::InternalError(None).into_response(); - } - } + let mst = Mst::new(Arc::new(state.block_store.clone())); let mst_root = match mst.persist().await { Ok(c) => c, @@ -727,71 +589,60 @@ pub async fn create_account( }; let commit_cid_str = commit_cid.to_string(); let rev_str = rev.as_ref().to_string(); - let repo_insert = sqlx::query!( - "INSERT INTO repos (user_id, repo_root_cid, repo_rev) VALUES ($1, $2, $3)", - user_id, - commit_cid_str, - rev_str - ) - .execute(&mut *tx) - .await; - if let Err(e) = repo_insert { - error!("Error initializing repo: {:?}", e); - return ApiError::InternalError(None).into_response(); - } let genesis_block_cids = vec![mst_root.to_bytes(), commit_cid.to_bytes()]; - if let Err(e) = sqlx::query!( - r#" - INSERT INTO user_blocks (user_id, block_cid) - SELECT $1, block_cid FROM UNNEST($2::bytea[]) AS t(block_cid) - ON CONFLICT (user_id, block_cid) DO NOTHING - "#, - user_id, - &genesis_block_cids - ) - .execute(&mut *tx) - .await - { - error!("Error inserting user_blocks: {:?}", e); - return ApiError::InternalError(None).into_response(); - } - if let Some(code) = &input.invite_code - && !code.trim().is_empty() - { - let use_insert = sqlx::query!( - "INSERT INTO invite_code_uses (code, used_by_user) VALUES ($1, $2)", - code, - user_id - ) - .execute(&mut *tx) - .await; - if let Err(e) = use_insert { - error!("Error recording invite usage: {:?}", e); - return ApiError::InternalError(None).into_response(); - } - } - if std::env::var("PDS_AGE_ASSURANCE_OVERRIDE").is_ok() { - let birthdate_pref = json!({ + + let birthdate_pref = std::env::var("PDS_AGE_ASSURANCE_OVERRIDE").ok().map(|_| { + json!({ "$type": "app.bsky.actor.defs#personalDetailsPref", "birthDate": "1998-05-06T00:00:00.000Z" - }); - if let Err(e) = sqlx::query!( - "INSERT INTO account_preferences (user_id, name, value_json) VALUES ($1, $2, $3) - ON CONFLICT (user_id, name) DO NOTHING", - user_id, - "app.bsky.actor.defs#personalDetailsPref", - birthdate_pref - ) - .execute(&mut *tx) - .await - { - warn!("Failed to set default birthdate preference: {:?}", e); + }) + }); + + let preferred_comms_channel = match verification_channel { + "email" => tranquil_db_traits::CommsChannel::Email, + "discord" => tranquil_db_traits::CommsChannel::Discord, + "telegram" => tranquil_db_traits::CommsChannel::Telegram, + "signal" => tranquil_db_traits::CommsChannel::Signal, + _ => tranquil_db_traits::CommsChannel::Email, + }; + + let create_input = tranquil_db_traits::CreatePasswordAccountInput { + handle: Handle::new_unchecked(&handle), + email: email.clone(), + did: Did::new_unchecked(&did), + password_hash, + preferred_comms_channel, + discord_id: input.discord_id.as_deref().map(|s| s.trim()).filter(|s| !s.is_empty()).map(String::from), + telegram_username: input.telegram_username.as_deref().map(|s| s.trim()).filter(|s| !s.is_empty()).map(String::from), + signal_number: input.signal_number.as_deref().map(|s| s.trim()).filter(|s| !s.is_empty()).map(String::from), + deactivated_at, + encrypted_key_bytes, + encryption_version: crate::config::ENCRYPTION_VERSION, + reserved_key_id, + commit_cid: commit_cid_str.clone(), + repo_rev: rev_str.clone(), + genesis_block_cids, + invite_code: input.invite_code.clone(), + birthdate_pref, + }; + + let create_result = match state.user_repo.create_password_account(&create_input).await { + Ok(r) => r, + Err(tranquil_db_traits::CreateAccountError::HandleTaken) => { + return ApiError::HandleNotAvailable(None).into_response(); } - } - if let Err(e) = tx.commit().await { - error!("Error committing transaction: {:?}", e); - return ApiError::InternalError(None).into_response(); - } + Err(tranquil_db_traits::CreateAccountError::EmailTaken) => { + return ApiError::EmailTaken.into_response(); + } + Err(tranquil_db_traits::CreateAccountError::DidExists) => { + return ApiError::AccountAlreadyExists.into_response(); + } + Err(e) => { + error!("Error creating password account: {:?}", e); + return ApiError::InternalError(None).into_response(); + } + }; + let user_id = create_result.user_id; if !is_migration && !is_did_web_byod { let did_typed = Did::new_unchecked(&did); let handle_typed = Handle::new_unchecked(&handle); @@ -858,13 +709,13 @@ pub async fn create_account( ); let formatted_token = crate::auth::verification_token::format_token_for_display(&verification_token); - if let Err(e) = crate::comms::enqueue_signup_verification( - &state.db, + if let Err(e) = crate::comms::comms_repo::enqueue_signup_verification( + state.infra_repo.as_ref(), user_id, verification_channel, recipient, &formatted_token, - None, + &hostname, ) .await { @@ -877,8 +728,9 @@ pub async fn create_account( } else if let Some(ref user_email) = email { let token = crate::auth::verification_token::generate_migration_token(&did, user_email); let formatted_token = crate::auth::verification_token::format_token_for_display(&token); - if let Err(e) = crate::comms::enqueue_migration_verification( - &state.db, + if let Err(e) = crate::comms::comms_repo::enqueue_migration_verification( + state.user_repo.as_ref(), + state.infra_repo.as_ref(), user_id, user_email, &formatted_token, @@ -906,17 +758,19 @@ pub async fn create_account( return ApiError::InternalError(None).into_response(); } }; - if let Err(e) = sqlx::query!( - "INSERT INTO session_tokens (did, access_jti, refresh_jti, access_expires_at, refresh_expires_at) VALUES ($1, $2, $3, $4, $5)", - did, - access_meta.jti, - refresh_meta.jti, - access_meta.expires_at, - refresh_meta.expires_at - ) - .execute(&state.db) - .await - { + let session_data = tranquil_db_traits::SessionTokenCreate { + did: Did::new_unchecked(&did), + access_jti: access_meta.jti.clone(), + refresh_jti: refresh_meta.jti.clone(), + access_expires_at: access_meta.expires_at, + refresh_expires_at: refresh_meta.expires_at, + legacy_login: false, + mfa_verified: false, + scope: None, + controller_did: None, + app_password_name: None, + }; + if let Err(e) = state.session_repo.create_session(&session_data).await { error!("createAccount: Error creating session: {:?}", e); return ApiError::InternalError(None).into_response(); } @@ -934,7 +788,7 @@ pub async fn create_account( StatusCode::OK, Json(CreateAccountOutput { handle: handle.clone().into(), - did: did.into(), + did: Did::new_unchecked(&did), did_doc, access_jwt: access_meta.token, refresh_jwt: refresh_meta.token, diff --git a/crates/tranquil-pds/src/api/identity/did.rs b/crates/tranquil-pds/src/api/identity/did.rs index a0cb1f3..0fe670d 100644 --- a/crates/tranquil-pds/src/api/identity/did.rs +++ b/crates/tranquil-pds/src/api/identity/did.rs @@ -34,17 +34,19 @@ pub async fn resolve_handle( State(state): State, Query(params): Query, ) -> Response { - let handle = params.handle.trim(); - if handle.is_empty() { + let handle_str = params.handle.trim(); + if handle_str.is_empty() { return ApiError::InvalidRequest("handle is required".into()).into_response(); } - let cache_key = format!("handle:{}", handle); + let cache_key = format!("handle:{}", handle_str); if let Some(did) = state.cache.get(&cache_key).await { return DidResponse::response(did).into_response(); } - let user = sqlx::query!("SELECT did FROM users WHERE handle = $1", handle) - .fetch_optional(&state.db) - .await; + let handle: Handle = match handle_str.parse() { + Ok(h) => h, + Err(_) => return ApiError::InvalidHandle(Some("Invalid handle format".into())).into_response(), + }; + let user = state.user_repo.get_by_handle(&handle).await; match user { Ok(Some(row)) => { let _ = state @@ -53,7 +55,7 @@ pub async fn resolve_handle( .await; DidResponse::response(row.did).into_response() } - Ok(None) => match crate::handle::resolve_handle(handle).await { + Ok(None) => match crate::handle::resolve_handle(handle.as_str()).await { Ok(did) => { let _ = state .cache @@ -130,15 +132,14 @@ pub async fn well_known_did(State(state): State, headers: HeaderMap) - } async fn serve_subdomain_did_doc(state: &AppState, handle: &str, hostname: &str) -> Response { - let full_handle = format!("{}.{}", handle, hostname); - let user = sqlx::query!( - "SELECT id, did, migrated_to_pds FROM users WHERE handle = $1", - full_handle - ) - .fetch_optional(&state.db) - .await; - let (user_id, did, migrated_to_pds) = match user { - Ok(Some(row)) => (row.id, row.did, row.migrated_to_pds), + let hostname_for_handles = hostname.split(':').next().unwrap_or(hostname); + let full_handle = format!("{}.{}", handle, hostname_for_handles); + let full_handle_typed: Handle = match full_handle.parse() { + Ok(h) => h, + Err(_) => return ApiError::InvalidHandle(Some("Invalid handle format".into())).into_response(), + }; + let user = match state.user_repo.get_did_web_info_by_handle(&full_handle_typed).await { + Ok(Some(u)) => u, Ok(None) => { return ApiError::NotFoundMsg("User not found".into()).into_response(); } @@ -147,10 +148,11 @@ async fn serve_subdomain_did_doc(state: &AppState, handle: &str, hostname: &str) return ApiError::InternalError(None).into_response(); } }; + let (user_id, did, migrated_to_pds) = (user.id, user.did, user.migrated_to_pds); if !did.starts_with("did:web:") { return ApiError::NotFoundMsg("User is not did:web".into()).into_response(); } - let subdomain_host = format!("{}.{}", handle, hostname); + let subdomain_host = format!("{}.{}", handle, hostname_for_handles); let encoded_subdomain = subdomain_host.replace(':', "%3A"); let expected_self_hosted = format!("did:web:{}", encoded_subdomain); if did != expected_self_hosted { @@ -158,14 +160,12 @@ async fn serve_subdomain_did_doc(state: &AppState, handle: &str, hostname: &str) .into_response(); } - let overrides = sqlx::query!( - "SELECT verification_methods, also_known_as FROM did_web_overrides WHERE user_id = $1", - user_id - ) - .fetch_optional(&state.db) - .await - .ok() - .flatten(); + let overrides = state + .user_repo + .get_did_web_overrides(user_id) + .await + .ok() + .flatten(); let service_endpoint = migrated_to_pds.unwrap_or_else(|| format!("https://{}", hostname)); @@ -204,20 +204,14 @@ async fn serve_subdomain_did_doc(state: &AppState, handle: &str, hostname: &str) .into_response(); } - let key_row = sqlx::query!( - "SELECT key_bytes, encryption_version FROM user_keys WHERE user_id = $1", - user_id - ) - .fetch_optional(&state.db) - .await; - let key_bytes: Vec = match key_row { - Ok(Some(row)) => match crate::config::decrypt_key(&row.key_bytes, row.encryption_version) { - Ok(k) => k, - Err(_) => { - return ApiError::InternalError(None).into_response(); - } - }, - _ => { + let key_info = match state.user_repo.get_user_key_by_id(user_id).await { + Ok(Some(k)) => k, + Ok(None) => return ApiError::InternalError(None).into_response(), + Err(_) => return ApiError::InternalError(None).into_response(), + }; + let key_bytes: Vec = match crate::config::decrypt_key(&key_info.key_bytes, key_info.encryption_version) { + Ok(k) => k, + Err(_) => { return ApiError::InternalError(None).into_response(); } }; @@ -264,15 +258,14 @@ async fn serve_subdomain_did_doc(state: &AppState, handle: &str, hostname: &str) pub async fn user_did_doc(State(state): State, Path(handle): Path) -> Response { let hostname = std::env::var("PDS_HOSTNAME").unwrap_or_else(|_| "localhost".to_string()); - let full_handle = format!("{}.{}", handle, hostname); - let user = sqlx::query!( - "SELECT id, did, migrated_to_pds FROM users WHERE handle = $1", - full_handle - ) - .fetch_optional(&state.db) - .await; - let (user_id, did, migrated_to_pds) = match user { - Ok(Some(row)) => (row.id, row.did, row.migrated_to_pds), + let hostname_for_handles = hostname.split(':').next().unwrap_or(&hostname); + let full_handle = format!("{}.{}", handle, hostname_for_handles); + let full_handle_typed: Handle = match full_handle.parse() { + Ok(h) => h, + Err(_) => return ApiError::InvalidHandle(Some("Invalid handle format".into())).into_response(), + }; + let user = match state.user_repo.get_did_web_info_by_handle(&full_handle_typed).await { + Ok(Some(u)) => u, Ok(None) => { return ApiError::NotFoundMsg("User not found".into()).into_response(); } @@ -281,12 +274,13 @@ pub async fn user_did_doc(State(state): State, Path(handle): Path, Path(handle): Path, Path(handle): Path = match key_row { - Ok(Some(row)) => match crate::config::decrypt_key(&row.key_bytes, row.encryption_version) { - Ok(k) => k, - Err(_) => { - return ApiError::InternalError(None).into_response(); - } - }, - _ => { + let key_info = match state.user_repo.get_user_key_by_id(user_id).await { + Ok(Some(k)) => k, + Ok(None) => return ApiError::InternalError(None).into_response(), + Err(_) => return ApiError::InternalError(None).into_response(), + }; + let key_bytes: Vec = match crate::config::decrypt_key(&key_info.key_bytes, key_info.encryption_version) { + Ok(k) => k, + Err(_) => { return ApiError::InternalError(None).into_response(); } }; @@ -404,7 +390,8 @@ pub async fn verify_did_web( handle: &str, expected_signing_key: Option<&str>, ) -> Result<(), String> { - let subdomain_host = format!("{}.{}", handle, hostname); + let hostname_for_handles = hostname.split(':').next().unwrap_or(hostname); + let subdomain_host = format!("{}.{}", handle, hostname_for_handles); let encoded_subdomain = subdomain_host.replace(':', "%3A"); let expected_subdomain_did = format!("did:web:{}", encoded_subdomain); if did == expected_subdomain_did { @@ -527,15 +514,10 @@ pub async fn get_recommended_did_credentials( auth: BearerAuthAllowDeactivated, ) -> Response { let auth_user = auth.0; - let user = match sqlx::query!( - "SELECT handle FROM users u JOIN user_keys k ON u.id = k.user_id WHERE u.did = $1", - &auth_user.did - ) - .fetch_optional(&state.db) - .await - { - Ok(Some(row)) => row, - _ => return ApiError::InternalError(None).into_response(), + let handle = match state.user_repo.get_handle_by_did(&auth_user.did).await { + Ok(Some(h)) => h, + Ok(None) => return ApiError::InternalError(None).into_response(), + Err(_) => return ApiError::InternalError(None).into_response(), }; let key_bytes = match auth_user.key_bytes { Some(kb) => kb, @@ -571,7 +553,7 @@ pub async fn get_recommended_did_credentials( StatusCode::OK, Json(GetRecommendedDidCredentialsOutput { rotation_keys, - also_known_as: vec![format!("at://{}", user.handle)], + also_known_as: vec![format!("at://{}", handle)], verification_methods: VerificationMethods { atproto: did_key }, services: Services { atproto_pds: AtprotoPds { @@ -619,12 +601,10 @@ pub async fn update_handle( return ApiError::RateLimitExceeded(Some("Daily handle update limit exceeded.".into())) .into_response(); } - let user_row = match sqlx::query!("SELECT id, handle FROM users WHERE did = $1", did.as_str()) - .fetch_optional(&state.db) - .await - { + let user_row = match state.user_repo.get_id_and_handle_by_did(&did).await { Ok(Some(row)) => row, - _ => return ApiError::InternalError(None).into_response(), + Ok(None) => return ApiError::InternalError(None).into_response(), + Err(_) => return ApiError::InternalError(None).into_response(), }; let user_id = user_row.id; let current_handle = user_row.handle; @@ -657,8 +637,9 @@ pub async fn update_handle( .into_response(); } let hostname = std::env::var("PDS_HOSTNAME").unwrap_or_else(|_| "localhost".to_string()); - let suffix = format!(".{}", hostname); - let is_service_domain = crate::handle::is_service_domain_handle(&new_handle, &hostname); + let hostname_for_handles = hostname.split(':').next().unwrap_or(&hostname); + let suffix = format!(".{}", hostname_for_handles); + let is_service_domain = crate::handle::is_service_domain_handle(&new_handle, hostname_for_handles); let handle = if is_service_domain { let short_part = if new_handle.ends_with(&suffix) { new_handle.strip_suffix(&suffix).unwrap_or(&new_handle) @@ -668,7 +649,7 @@ pub async fn update_handle( let full_handle = if new_handle.ends_with(&suffix) { new_handle.clone() } else { - format!("{}.{}", new_handle, hostname) + format!("{}.{}", new_handle, hostname_for_handles) }; if full_handle == current_handle { let handle_typed = Handle::new_unchecked(&full_handle); @@ -727,23 +708,18 @@ pub async fn update_handle( } new_handle.clone() }; - let existing = sqlx::query!( - "SELECT id FROM users WHERE handle = $1 AND id != $2", - handle, - user_id - ) - .fetch_optional(&state.db) - .await; - if let Ok(Some(_)) = existing { + let handle_typed: Handle = match handle.parse() { + Ok(h) => h, + Err(_) => return ApiError::InvalidHandle(Some("Invalid handle format".into())).into_response(), + }; + let handle_exists = match state.user_repo.check_handle_exists(&handle_typed, user_id).await { + Ok(exists) => exists, + Err(_) => return ApiError::InternalError(None).into_response(), + }; + if handle_exists { return ApiError::HandleTaken.into_response(); } - let result = sqlx::query!( - "UPDATE users SET handle = $1 WHERE id = $2", - handle, - user_id - ) - .execute(&state.db) - .await; + let result = state.user_repo.update_handle(user_id, &handle_typed).await; match result { Ok(_) => { if !current_handle.is_empty() { @@ -753,14 +729,13 @@ pub async fn update_handle( .await; } let _ = state.cache.delete(&format!("handle:{}", handle)).await; - let handle_typed = Handle::new_unchecked(&handle); if let Err(e) = crate::api::repo::record::sequence_identity_event(&state, &did, Some(&handle_typed)) .await { warn!("Failed to sequence identity event for handle update: {}", e); } - if let Err(e) = update_plc_handle(&state, &did, &handle).await { + if let Err(e) = update_plc_handle(&state, &did, &handle_typed).await { warn!("Failed to update PLC handle: {}", e); } EmptyResponse::ok().into_response() @@ -774,22 +749,13 @@ pub async fn update_handle( pub async fn update_plc_handle( state: &AppState, - did: &str, - new_handle: &str, + did: &crate::types::Did, + new_handle: &Handle, ) -> Result<(), Box> { - if !did.starts_with("did:plc:") { + if !did.as_str().starts_with("did:plc:") { return Ok(()); } - let user_row = sqlx::query!( - r#"SELECT u.id, uk.key_bytes, uk.encryption_version - FROM users u - JOIN user_keys uk ON u.id = uk.user_id - WHERE u.did = $1"#, - did - ) - .fetch_optional(&state.db) - .await?; - let user_row = match user_row { + let user_row = match state.user_repo.get_user_with_key_by_did(did).await? { Some(r) => r, None => return Ok(()), }; @@ -810,12 +776,14 @@ pub async fn well_known_atproto_did(State(state): State, headers: Head Some(h) => h, None => return (StatusCode::BAD_REQUEST, "Missing host header").into_response(), }; - let handle = host.split(':').next().unwrap_or(host); - let user = sqlx::query!("SELECT did FROM users WHERE handle = $1", handle) - .fetch_optional(&state.db) - .await; + let handle_str = host.split(':').next().unwrap_or(host); + let handle: Handle = match handle_str.parse() { + Ok(h) => h, + Err(_) => return (StatusCode::BAD_REQUEST, "Invalid handle format").into_response(), + }; + let user = state.user_repo.get_by_handle(&handle).await; match user { - Ok(Some(row)) => row.did.into_response(), + Ok(Some(row)) => row.did.to_string().into_response(), Ok(None) => (StatusCode::NOT_FOUND, "Handle not found").into_response(), Err(e) => { error!("DB error in well-known atproto-did: {:?}", e); diff --git a/crates/tranquil-pds/src/api/identity/plc/request.rs b/crates/tranquil-pds/src/api/identity/plc/request.rs index 96e9737..4700d9f 100644 --- a/crates/tranquil-pds/src/api/identity/plc/request.rs +++ b/crates/tranquil-pds/src/api/identity/plc/request.rs @@ -25,43 +25,34 @@ pub async fn request_plc_operation_signature( ) { return e; } - let user = match sqlx::query!("SELECT id FROM users WHERE did = $1", &auth_user.did) - .fetch_optional(&state.db) - .await - { - Ok(Some(row)) => row, + let user_id = match state.user_repo.get_id_by_did(&auth_user.did).await { + Ok(Some(id)) => id, Ok(None) => return ApiError::AccountNotFound.into_response(), Err(e) => { error!("DB error: {:?}", e); return ApiError::InternalError(None).into_response(); } }; - let _ = sqlx::query!( - "DELETE FROM plc_operation_tokens WHERE user_id = $1 OR expires_at < NOW()", - user.id - ) - .execute(&state.db) - .await; + let _ = state.infra_repo.delete_plc_tokens_for_user(user_id).await; let plc_token = generate_plc_token(); let expires_at = Utc::now() + Duration::minutes(10); - if let Err(e) = sqlx::query!( - r#" - INSERT INTO plc_operation_tokens (user_id, token, expires_at) - VALUES ($1, $2, $3) - "#, - user.id, - plc_token, - expires_at - ) - .execute(&state.db) - .await + if let Err(e) = state + .infra_repo + .insert_plc_token(user_id, &plc_token, expires_at) + .await { error!("Failed to create PLC token: {:?}", e); return ApiError::InternalError(None).into_response(); } let hostname = std::env::var("PDS_HOSTNAME").unwrap_or_else(|_| "localhost".to_string()); - if let Err(e) = - crate::comms::enqueue_plc_operation(&state.db, user.id, &plc_token, &hostname).await + if let Err(e) = crate::comms::comms_repo::enqueue_plc_operation( + state.user_repo.as_ref(), + state.infra_repo.as_ref(), + user_id, + &plc_token, + &hostname, + ) + .await { warn!("Failed to enqueue PLC operation notification: {:?}", e); } diff --git a/crates/tranquil-pds/src/api/identity/plc/sign.rs b/crates/tranquil-pds/src/api/identity/plc/sign.rs index fcb9782..4bbace7 100644 --- a/crates/tranquil-pds/src/api/identity/plc/sign.rs +++ b/crates/tranquil-pds/src/api/identity/plc/sign.rs @@ -67,24 +67,16 @@ pub async fn sign_plc_operation( .into_response(); } }; - let user = match sqlx::query!("SELECT id FROM users WHERE did = $1", did) - .fetch_optional(&state.db) - .await - { - Ok(Some(row)) => row, - _ => { - return ApiError::AccountNotFound.into_response(); + let user_id = match state.user_repo.get_id_by_did(did).await { + Ok(Some(id)) => id, + Ok(None) => return ApiError::AccountNotFound.into_response(), + Err(e) => { + error!("DB error: {:?}", e); + return ApiError::InternalError(None).into_response(); } }; - let token_row = match sqlx::query!( - "SELECT id, expires_at FROM plc_operation_tokens WHERE user_id = $1 AND token = $2", - user.id, - token - ) - .fetch_optional(&state.db) - .await - { - Ok(Some(row)) => row, + let token_expiry = match state.infra_repo.get_plc_token_expiry(user_id, token).await { + Ok(Some(expiry)) => expiry, Ok(None) => { return ApiError::InvalidToken(Some("Invalid or expired token".into())).into_response(); } @@ -93,27 +85,20 @@ pub async fn sign_plc_operation( return ApiError::InternalError(None).into_response(); } }; - if Utc::now() > token_row.expires_at { - let _ = sqlx::query!( - "DELETE FROM plc_operation_tokens WHERE id = $1", - token_row.id - ) - .execute(&state.db) - .await; + if Utc::now() > token_expiry { + let _ = state.infra_repo.delete_plc_token(user_id, token).await; return ApiError::ExpiredToken(Some("Token has expired".into())).into_response(); } - let key_row = match sqlx::query!( - "SELECT key_bytes, encryption_version FROM user_keys WHERE user_id = $1", - user.id - ) - .fetch_optional(&state.db) - .await - { + let key_row = match state.user_repo.get_user_key_by_id(user_id).await { Ok(Some(row)) => row, - _ => { + Ok(None) => { return ApiError::InternalError(Some("User signing key not found".into())) .into_response(); } + Err(e) => { + error!("DB error: {:?}", e); + return ApiError::InternalError(None).into_response(); + } }; let key_bytes = match crate::config::decrypt_key(&key_row.key_bytes, key_row.encryption_version) { @@ -179,12 +164,7 @@ pub async fn sign_plc_operation( return ApiError::InternalError(None).into_response(); } }; - let _ = sqlx::query!( - "DELETE FROM plc_operation_tokens WHERE id = $1", - token_row.id - ) - .execute(&state.db) - .await; + let _ = state.infra_repo.delete_plc_token(user_id, token).await; info!("Signed PLC operation for user {}", did); ( StatusCode::OK, diff --git a/crates/tranquil-pds/src/api/identity/plc/submit.rs b/crates/tranquil-pds/src/api/identity/plc/submit.rs index 2a772d5..ce4d586 100644 --- a/crates/tranquil-pds/src/api/identity/plc/submit.rs +++ b/crates/tranquil-pds/src/api/identity/plc/submit.rs @@ -44,27 +44,24 @@ pub async fn submit_plc_operation( let op = &input.operation; let hostname = std::env::var("PDS_HOSTNAME").unwrap_or_else(|_| "localhost".to_string()); let public_url = format!("https://{}", hostname); - let user = match sqlx::query!("SELECT id, handle FROM users WHERE did = $1", did) - .fetch_optional(&state.db) - .await - { - Ok(Some(row)) => row, - _ => { - return ApiError::AccountNotFound.into_response(); + let user = match state.user_repo.get_id_and_handle_by_did(did).await { + Ok(Some(u)) => u, + Ok(None) => return ApiError::AccountNotFound.into_response(), + Err(e) => { + error!("DB error: {:?}", e); + return ApiError::InternalError(None).into_response(); } }; - let key_row = match sqlx::query!( - "SELECT key_bytes, encryption_version FROM user_keys WHERE user_id = $1", - user.id - ) - .fetch_optional(&state.db) - .await - { + let key_row = match state.user_repo.get_user_key_by_id(user.id).await { Ok(Some(row)) => row, - _ => { + Ok(None) => { return ApiError::InternalError(Some("User signing key not found".into())) .into_response(); } + Err(e) => { + error!("DB error: {:?}", e); + return ApiError::InternalError(None).into_response(); + } }; let key_bytes = match crate::config::decrypt_key(&key_row.key_bytes, key_row.encryption_version) { @@ -139,19 +136,13 @@ pub async fn submit_plc_operation( { return ApiError::from(e).into_response(); } - match sqlx::query!( - "INSERT INTO repo_seq (did, event_type, handle) VALUES ($1, 'identity', $2) RETURNING seq", - did, - user.handle - ) - .fetch_one(&state.db) - .await + match state + .repo_repo + .insert_identity_event(did, Some(&user.handle)) + .await { - Ok(row) => { - if let Err(e) = sqlx::query(&format!("NOTIFY repo_updates, '{}'", row.seq)) - .execute(&state.db) - .await - { + Ok(seq) => { + if let Err(e) = state.repo_repo.notify_update(seq).await { warn!("Failed to notify identity event: {:?}", e); } } diff --git a/crates/tranquil-pds/src/api/moderation/mod.rs b/crates/tranquil-pds/src/api/moderation/mod.rs index 0d3dce9..3aac97f 100644 --- a/crates/tranquil-pds/src/api/moderation/mod.rs +++ b/crates/tranquil-pds/src/api/moderation/mod.rs @@ -72,18 +72,9 @@ async fn proxy_to_report_service( let key_bytes = match &auth_user.key_bytes { Some(kb) => kb.clone(), None => { - match sqlx::query_as::<_, (Vec, Option)>( - "SELECT k.key_bytes, k.encryption_version - FROM users u - JOIN user_keys k ON u.id = k.user_id - WHERE u.did = $1", - ) - .bind(&auth_user.did) - .fetch_optional(&state.db) - .await - { - Ok(Some((key_bytes_enc, encryption_version))) => { - match crate::config::decrypt_key(&key_bytes_enc, encryption_version) { + match state.user_repo.get_with_key_by_did(&auth_user.did).await { + Ok(Some(user_with_key)) => { + match crate::config::decrypt_key(&user_with_key.key_bytes, user_with_key.encryption_version) { Ok(key) => key, Err(e) => { error!(error = ?e, "Failed to decrypt user key for report service auth"); @@ -185,7 +176,7 @@ async fn proxy_to_report_service( async fn create_report_locally( state: &AppState, - did: &str, + did: &crate::types::Did, is_takendown: bool, input: CreateReportInput, ) -> Response { @@ -214,19 +205,14 @@ async fn create_report_locally( let report_id = (uuid::Uuid::now_v7().as_u128() & 0x7FFF_FFFF_FFFF_FFFF) as i64; let subject_json = json!(input.subject); - let insert = sqlx::query!( - "INSERT INTO reports (id, reason_type, reason, subject_json, reported_by_did, created_at) VALUES ($1, $2, $3, $4, $5, $6)", + if let Err(e) = state.infra_repo.insert_report( report_id, - input.reason_type, - input.reason, + &input.reason_type, + input.reason.as_deref(), subject_json, did, - created_at - ) - .execute(&state.db) - .await; - - if let Err(e) = insert { + created_at, + ).await { error!("Failed to insert report: {:?}", e); return ApiError::InternalError(None).into_response(); } diff --git a/crates/tranquil-pds/src/api/notification_prefs.rs b/crates/tranquil-pds/src/api/notification_prefs.rs index 1e9854b..dfe1484 100644 --- a/crates/tranquil-pds/src/api/notification_prefs.rs +++ b/crates/tranquil-pds/src/api/notification_prefs.rs @@ -8,7 +8,6 @@ use axum::{ }; use serde::{Deserialize, Serialize}; use serde_json::json; -use sqlx::Row; use tracing::info; #[derive(Serialize)] @@ -26,47 +25,22 @@ pub struct NotificationPrefsResponse { pub async fn get_notification_prefs(State(state): State, auth: BearerAuth) -> Response { let user = auth.0; - let row = match sqlx::query( - r#" - SELECT - email, - preferred_comms_channel::text as channel, - discord_id, - discord_verified, - telegram_username, - telegram_verified, - signal_number, - signal_verified - FROM users - WHERE did = $1 - "#, - ) - .bind(&user.did) - .fetch_one(&state.db) - .await - { - Ok(r) => r, + let prefs = match state.user_repo.get_notification_prefs(&user.did).await { + Ok(Some(p)) => p, + Ok(None) => return ApiError::AccountNotFound.into_response(), Err(e) => { return ApiError::InternalError(Some(format!("Database error: {}", e))).into_response(); } }; - let email: String = row.get("email"); - let channel: String = row.get("channel"); - let discord_id: Option = row.get("discord_id"); - let discord_verified: bool = row.get("discord_verified"); - let telegram_username: Option = row.get("telegram_username"); - let telegram_verified: bool = row.get("telegram_verified"); - let signal_number: Option = row.get("signal_number"); - let signal_verified: bool = row.get("signal_verified"); Json(NotificationPrefsResponse { - preferred_channel: channel, - email, - discord_id, - discord_verified, - telegram_username, - telegram_verified, - signal_number, - signal_verified, + preferred_channel: prefs.preferred_channel, + email: prefs.email, + discord_id: prefs.discord_id, + discord_verified: prefs.discord_verified, + telegram_username: prefs.telegram_username, + telegram_verified: prefs.telegram_verified, + signal_number: prefs.signal_number, + signal_verified: prefs.signal_verified, }) .into_response() } @@ -91,37 +65,15 @@ pub struct GetNotificationHistoryResponse { pub async fn get_notification_history(State(state): State, auth: BearerAuth) -> Response { let user = auth.0; - let user_id: uuid::Uuid = - match sqlx::query_scalar!("SELECT id FROM users WHERE did = $1", &user.did) - .fetch_one(&state.db) - .await - { - Ok(id) => id, - Err(e) => { - return ApiError::InternalError(Some(format!("Database error: {}", e))) - .into_response(); - } - }; + let user_id: uuid::Uuid = match state.user_repo.get_id_by_did(&user.did).await { + Ok(Some(id)) => id, + Ok(None) => return ApiError::AccountNotFound.into_response(), + Err(e) => { + return ApiError::InternalError(Some(format!("Database error: {}", e))).into_response(); + } + }; - let rows = match sqlx::query!( - r#" - SELECT - created_at, - channel as "channel: String", - comms_type as "comms_type: String", - status as "status: String", - subject, - body - FROM comms_queue - WHERE user_id = $1 - ORDER BY created_at DESC - LIMIT 50 - "#, - user_id - ) - .fetch_all(&state.db) - .await - { + let rows = match state.infra_repo.get_notification_history(user_id, 50).await { Ok(r) => r, Err(e) => { return ApiError::InternalError(Some(format!("Database error: {}", e))).into_response(); @@ -181,7 +133,7 @@ pub struct UpdateNotificationPrefsResponse { } pub async fn request_channel_verification( - db: &sqlx::PgPool, + state: &AppState, user_id: uuid::Uuid, did: &str, channel: &str, @@ -195,8 +147,8 @@ pub async fn request_channel_verification( if channel == "email" { let hostname = std::env::var("PDS_HOSTNAME").unwrap_or_else(|_| "localhost".to_string()); let handle_str = handle.unwrap_or("user"); - crate::comms::enqueue_email_update( - db, + crate::comms::comms_repo::enqueue_email_update( + state.infra_repo.as_ref(), user_id, identifier, handle_str, @@ -206,20 +158,25 @@ pub async fn request_channel_verification( .await .map_err(|e| format!("Failed to enqueue email notification: {}", e))?; } else { - sqlx::query!( - r#" - INSERT INTO comms_queue (user_id, channel, comms_type, recipient, subject, body, metadata) - VALUES ($1, $2::comms_channel, 'channel_verification', $3, 'Verify your channel', $4, $5) - "#, - user_id, - channel as _, - identifier, - format!("Your verification code is: {}", formatted_token), - json!({"code": formatted_token}) - ) - .execute(db) - .await - .map_err(|e| format!("Failed to enqueue notification: {}", e))?; + let comms_channel = match channel { + "discord" => tranquil_db_traits::CommsChannel::Discord, + "telegram" => tranquil_db_traits::CommsChannel::Telegram, + "signal" => tranquil_db_traits::CommsChannel::Signal, + _ => return Err("Invalid channel".to_string()), + }; + state + .infra_repo + .enqueue_comms( + Some(user_id), + comms_channel, + tranquil_db_traits::CommsType::ChannelVerification, + identifier, + Some("Verify your channel"), + &format!("Your verification code is: {}", formatted_token), + Some(json!({"code": formatted_token})), + ) + .await + .map_err(|e| format!("Failed to enqueue notification: {}", e))?; } Ok(token) @@ -232,14 +189,13 @@ pub async fn update_notification_prefs( ) -> Response { let user = auth.0; - let user_row = match sqlx::query!( - "SELECT id, handle, email FROM users WHERE did = $1", - &user.did - ) - .fetch_one(&state.db) - .await + let user_row = match state + .user_repo + .get_id_handle_email_by_did(&user.did) + .await { - Ok(row) => row, + Ok(Some(row)) => row, + Ok(None) => return ApiError::AccountNotFound.into_response(), Err(e) => { return ApiError::InternalError(Some(format!("Database error: {}", e))).into_response(); } @@ -259,13 +215,10 @@ pub async fn update_notification_prefs( ) .into_response(); } - if let Err(e) = sqlx::query( - r#"UPDATE users SET preferred_comms_channel = $1::comms_channel, updated_at = NOW() WHERE did = $2"# - ) - .bind(channel) - .bind(&user.did) - .execute(&state.db) - .await + if let Err(e) = state + .user_repo + .update_preferred_comms_channel(&user.did, channel) + .await { return ApiError::InternalError(Some(format!("Database error: {}", e))).into_response(); } @@ -285,20 +238,17 @@ pub async fn update_notification_prefs( if current_email.as_ref().map(|e| e.to_lowercase()) == Some(email_clean.clone()) { info!(did = %user.did, "Email unchanged, skipping"); } else { - let exists = sqlx::query!( - "SELECT 1 as one FROM users WHERE LOWER(email) = $1 AND id != $2", - email_clean, - user_id - ) - .fetch_optional(&state.db) - .await; - - if let Ok(Some(_)) = exists { - return ApiError::EmailTaken.into_response(); + match state.user_repo.check_email_exists(&email_clean, user_id).await { + Ok(true) => return ApiError::EmailTaken.into_response(), + Err(e) => { + return ApiError::InternalError(Some(format!("Database error: {}", e))) + .into_response(); + } + Ok(false) => {} } if let Err(e) = request_channel_verification( - &state.db, + &state, user_id, &user.did, "email", @@ -316,21 +266,15 @@ pub async fn update_notification_prefs( if let Some(ref discord_id) = input.discord_id { if discord_id.is_empty() { - if let Err(e) = sqlx::query!( - "UPDATE users SET discord_id = NULL, discord_verified = FALSE, updated_at = NOW() WHERE id = $1", - user_id - ) - .execute(&state.db) - .await - { - return ApiError::InternalError(Some(format!("Database error: {}", e))).into_response(); + if let Err(e) = state.user_repo.clear_discord(user_id).await { + return ApiError::InternalError(Some(format!("Database error: {}", e))) + .into_response(); } info!(did = %user.did, "Cleared Discord ID"); } else { - if let Err(e) = request_channel_verification( - &state.db, user_id, &user.did, "discord", discord_id, None, - ) - .await + if let Err(e) = + request_channel_verification(&state, user_id, &user.did, "discord", discord_id, None) + .await { return ApiError::InternalError(Some(e)).into_response(); } @@ -342,19 +286,14 @@ pub async fn update_notification_prefs( if let Some(ref telegram) = input.telegram_username { let telegram_clean = telegram.trim_start_matches('@'); if telegram_clean.is_empty() { - if let Err(e) = sqlx::query!( - "UPDATE users SET telegram_username = NULL, telegram_verified = FALSE, updated_at = NOW() WHERE id = $1", - user_id - ) - .execute(&state.db) - .await - { - return ApiError::InternalError(Some(format!("Database error: {}", e))).into_response(); + if let Err(e) = state.user_repo.clear_telegram(user_id).await { + return ApiError::InternalError(Some(format!("Database error: {}", e))) + .into_response(); } info!(did = %user.did, "Cleared Telegram username"); } else { if let Err(e) = request_channel_verification( - &state.db, + &state, user_id, &user.did, "telegram", @@ -372,19 +311,14 @@ pub async fn update_notification_prefs( if let Some(ref signal) = input.signal_number { if signal.is_empty() { - if let Err(e) = sqlx::query!( - "UPDATE users SET signal_number = NULL, signal_verified = FALSE, updated_at = NOW() WHERE id = $1", - user_id - ) - .execute(&state.db) - .await - { - return ApiError::InternalError(Some(format!("Database error: {}", e))).into_response(); + if let Err(e) = state.user_repo.clear_signal(user_id).await { + return ApiError::InternalError(Some(format!("Database error: {}", e))) + .into_response(); } info!(did = %user.did, "Cleared Signal number"); } else { if let Err(e) = - request_channel_verification(&state.db, user_id, &user.did, "signal", signal, None) + request_channel_verification(&state, user_id, &user.did, "signal", signal, None) .await { return ApiError::InternalError(Some(e)).into_response(); diff --git a/crates/tranquil-pds/src/api/proxy.rs b/crates/tranquil-pds/src/api/proxy.rs index cb91b90..40d0797 100644 --- a/crates/tranquil-pds/src/api/proxy.rs +++ b/crates/tranquil-pds/src/api/proxy.rs @@ -225,7 +225,8 @@ async fn proxy_handler( let http_uri = crate::util::build_full_url(&uri.to_string()); match crate::auth::validate_token_with_dpop( - &state.db, + state.user_repo.as_ref(), + state.oauth_repo.as_ref(), &token, extracted.is_dpop, dpop_proof, diff --git a/crates/tranquil-pds/src/api/repo/blob.rs b/crates/tranquil-pds/src/api/repo/blob.rs index fed3d0a..0b4c1ee 100644 --- a/crates/tranquil-pds/src/api/repo/blob.rs +++ b/crates/tranquil-pds/src/api/repo/blob.rs @@ -1,7 +1,8 @@ use crate::api::error::ApiError; use crate::auth::{BearerAuthAllowDeactivated, ServiceTokenVerifier, is_service_token}; -use crate::delegation::{self, DelegationActionType}; +use crate::delegation::DelegationActionType; use crate::state::AppState; +use crate::types::{CidLink, Did}; use crate::util::get_max_blob_size; use axum::body::Body; use axum::{ @@ -55,7 +56,7 @@ pub async fn upload_blob( let is_service_auth = is_service_token(&token); - let (did, _is_migration, controller_did) = if is_service_auth { + let (did, _is_migration, controller_did): (Did, bool, Option) = if is_service_auth { debug!("Verifying service token for blob upload"); let verifier = ServiceTokenVerifier::new(); match verifier @@ -64,7 +65,11 @@ pub async fn upload_blob( { Ok(claims) => { debug!("Service token verified for DID: {}", claims.iss); - (claims.iss, false, None) + let did: Did = match claims.iss.parse() { + Ok(d) => d, + Err(_) => return ApiError::InvalidDid("Invalid DID format".into()).into_response(), + }; + (did, false, None) } Err(e) => { error!("Service token verification failed: {:?}", e); @@ -82,7 +87,8 @@ pub async fn upload_blob( std::env::var("PDS_HOSTNAME").unwrap_or_else(|_| "localhost".to_string()) ); match crate::auth::validate_token_with_dpop( - &state.db, + state.user_repo.as_ref(), + state.oauth_repo.as_ref(), &token, extracted.is_dpop, dpop_proof, @@ -105,17 +111,15 @@ pub async fn upload_blob( ) { return e; } - let deactivated = sqlx::query_scalar!( - "SELECT deactivated_at FROM users WHERE did = $1", - &user.did - ) - .fetch_optional(&state.db) - .await - .ok() - .flatten() - .flatten(); - let ctrl_did = user.controller_did.map(|d| d.to_string()); - (user.did.to_string(), deactivated.is_some(), ctrl_did) + let deactivated = state + .user_repo + .get_status_by_did(&user.did) + .await + .ok() + .flatten() + .and_then(|s| s.deactivated_at); + let ctrl_did = user.controller_did.clone(); + (user.did, deactivated.is_some(), ctrl_did) } Err(_) => { return ApiError::AuthenticationFailed(None).into_response(); @@ -123,7 +127,9 @@ pub async fn upload_blob( } }; - if crate::util::is_account_migrated(&state.db, &did) + if state + .user_repo + .is_account_migrated(&did) .await .unwrap_or(false) { @@ -135,11 +141,8 @@ pub async fn upload_blob( .and_then(|h| h.to_str().ok()) .unwrap_or("application/octet-stream"); - let user_query = sqlx::query!("SELECT id FROM users WHERE did = $1", did) - .fetch_optional(&state.db) - .await; - let user_id = match user_query { - Ok(Some(row)) => row.id, + let user_id = match state.user_repo.get_id_by_did(&did).await { + Ok(Some(id)) => id, _ => { return ApiError::InternalError(None).into_response(); } @@ -192,6 +195,7 @@ pub async fn upload_blob( }; let cid = Cid::new_v1(0x55, multihash); let cid_str = cid.to_string(); + let cid_link: CidLink = CidLink::new_unchecked(&cid_str); let storage_key = format!("blobs/{}", cid_str); info!( @@ -199,27 +203,11 @@ pub async fn upload_blob( size, cid_str ); - let mut tx = match state.db.begin().await { - Ok(tx) => tx, - Err(e) => { - let _ = state.blob_store.delete(&temp_key).await; - error!("Failed to begin transaction: {:?}", e); - return ApiError::InternalError(None).into_response(); - } - }; - - let insert = sqlx::query!( - "INSERT INTO blobs (cid, mime_type, size_bytes, created_by_user, storage_key) VALUES ($1, $2, $3, $4, $5) ON CONFLICT (cid) DO NOTHING RETURNING cid", - cid_str, - mime_type, - size as i64, - user_id, - storage_key - ) - .fetch_optional(&mut *tx) - .await; - - let was_inserted = match insert { + let was_inserted = match state + .blob_repo + .insert_blob(&cid_link, &mime_type, size as i64, user_id, &storage_key) + .await + { Ok(Some(_)) => true, Ok(None) => false, Err(e) => { @@ -229,41 +217,33 @@ pub async fn upload_blob( } }; - if was_inserted && let Err(e) = state.blob_store.copy(&temp_key, &storage_key).await { - let _ = state.blob_store.delete(&temp_key).await; - error!("Failed to copy blob to final location: {:?}", e); - return ApiError::InternalError(Some("Failed to store blob".into())).into_response(); + if was_inserted { + if let Err(e) = state.blob_store.copy(&temp_key, &storage_key).await { + let _ = state.blob_store.delete(&temp_key).await; + error!("Failed to copy blob to final location: {:?}", e); + return ApiError::InternalError(Some("Failed to store blob".into())).into_response(); + } } let _ = state.blob_store.delete(&temp_key).await; - if let Err(e) = tx.commit().await { - error!("Failed to commit blob transaction: {:?}", e); - if was_inserted && let Err(cleanup_err) = state.blob_store.delete(&storage_key).await { - error!( - "Failed to cleanup orphaned blob {}: {:?}", - storage_key, cleanup_err - ); - } - return ApiError::InternalError(None).into_response(); - } - if let Some(ref controller) = controller_did { - let _ = delegation::log_delegation_action( - &state.db, - &did, - controller, - Some(controller), - DelegationActionType::BlobUpload, - Some(json!({ - "cid": cid_str, - "mime_type": mime_type, - "size": size - })), - None, - None, - ) - .await; + let _ = state + .delegation_repo + .log_delegation_action( + &did, + controller, + Some(controller), + DelegationActionType::BlobUpload, + Some(json!({ + "cid": cid_str, + "mime_type": mime_type, + "size": size + })), + None, + None, + ) + .await; } Json(json!({ @@ -305,47 +285,35 @@ pub async fn list_missing_blobs( Query(params): Query, ) -> Response { let auth_user = auth.0; - let did = auth_user.did; - let user_query = sqlx::query!("SELECT id FROM users WHERE did = $1", did.as_str()) - .fetch_optional(&state.db) - .await; - let user_id = match user_query { - Ok(Some(row)) => row.id, - _ => { + let did = &auth_user.did; + let user = match state.user_repo.get_by_did(did).await { + Ok(Some(u)) => u, + Ok(None) => return ApiError::InternalError(None).into_response(), + Err(e) => { + error!("DB error fetching user: {:?}", e); return ApiError::InternalError(None).into_response(); } }; let limit = params.limit.unwrap_or(500).clamp(1, 1000); - let cursor_cid = params.cursor.as_deref().unwrap_or(""); - let missing_query = sqlx::query!( - r#" - SELECT rb.blob_cid, rb.record_uri - FROM record_blobs rb - LEFT JOIN blobs b ON rb.blob_cid = b.cid - WHERE rb.repo_id = $1 AND b.cid IS NULL AND rb.blob_cid > $2 - ORDER BY rb.blob_cid - LIMIT $3 - "#, - user_id, - cursor_cid, - limit + 1 - ) - .fetch_all(&state.db) - .await; - let rows = match missing_query { - Ok(r) => r, + let cursor = params.cursor.as_deref(); + let missing = match state + .blob_repo + .list_missing_blobs(user.id, cursor, limit + 1) + .await + { + Ok(m) => m, Err(e) => { error!("DB error fetching missing blobs: {:?}", e); return ApiError::InternalError(None).into_response(); } }; - let has_more = rows.len() > limit as usize; - let blobs: Vec = rows + let has_more = missing.len() > limit as usize; + let blobs: Vec = missing .into_iter() .take(limit as usize) - .map(|row| RecordBlob { - cid: row.blob_cid, - record_uri: row.record_uri, + .map(|m| RecordBlob { + cid: m.blob_cid.to_string(), + record_uri: m.record_uri.to_string(), }) .collect(); let next_cursor = if has_more { diff --git a/crates/tranquil-pds/src/api/repo/import.rs b/crates/tranquil-pds/src/api/repo/import.rs index e53b5af..2b059f8 100644 --- a/crates/tranquil-pds/src/api/repo/import.rs +++ b/crates/tranquil-pds/src/api/repo/import.rs @@ -6,6 +6,7 @@ use crate::state::AppState; use crate::sync::import::{ImportError, apply_import, parse_car}; use crate::sync::verify::CarVerifier; use crate::types::Did; +use tranquil_types::{AtUri, CidLink}; use axum::{ body::Bytes, extract::State, @@ -45,13 +46,7 @@ pub async fn import_repo( } let auth_user = auth.0; let did = &auth_user.did; - let user = match sqlx::query!( - "SELECT id, handle, deactivated_at, takedown_ref FROM users WHERE did = $1", - did - ) - .fetch_optional(&state.db) - .await - { + let user = match state.user_repo.get_by_did(did).await { Ok(Some(row)) => row, Ok(None) => { return ApiError::AccountNotFound.into_response(); @@ -190,46 +185,38 @@ pub async fn import_repo( .ok() .and_then(|s| s.parse().ok()) .unwrap_or(DEFAULT_MAX_BLOCKS); - match apply_import(&state.db, user_id, root, blocks.clone(), max_blocks).await { + match apply_import(&state.repo_repo, user_id, root, blocks.clone(), max_blocks).await { Ok(import_result) => { info!( "Successfully imported {} records for user {}", import_result.records.len(), did ); - let blob_refs: Vec<(String, String)> = import_result + let blob_refs: Vec<(AtUri, CidLink)> = import_result .records .iter() .flat_map(|record| { - let record_uri = format!("at://{}/{}/{}", did, record.collection, record.rkey); + let record_uri = AtUri::from_parts(did.as_str(), &record.collection, &record.rkey); record .blob_refs .iter() - .map(move |blob_ref| (record_uri.clone(), blob_ref.cid.clone())) + .map(move |blob_ref| (record_uri.clone(), CidLink::new_unchecked(blob_ref.cid.clone()))) }) .collect(); if !blob_refs.is_empty() { - let (record_uris, blob_cids): (Vec, Vec) = + let (record_uris, blob_cids): (Vec, Vec) = blob_refs.into_iter().unzip(); - match sqlx::query!( - r#" - INSERT INTO record_blobs (repo_id, record_uri, blob_cid) - SELECT $1, * FROM UNNEST($2::text[], $3::text[]) - ON CONFLICT (repo_id, record_uri, blob_cid) DO NOTHING - "#, - user_id, - &record_uris, - &blob_cids - ) - .execute(&state.db) - .await + match state + .blob_repo + .insert_record_blobs(user_id, &record_uris, &blob_cids) + .await { - Ok(result) => { + Ok(()) => { info!( "Recorded {} blob references for imported repo", - result.rows_affected() + blob_cids.len() ); } Err(e) => { @@ -237,16 +224,7 @@ pub async fn import_repo( } } } - let key_row = match sqlx::query!( - r#"SELECT uk.key_bytes, uk.encryption_version - FROM user_keys uk - JOIN users u ON uk.user_id = u.id - WHERE u.did = $1"#, - did - ) - .fetch_optional(&state.db) - .await - { + let key_row = match state.user_repo.get_user_with_key_by_did(did).await { Ok(Some(row)) => row, Ok(None) => { error!("No signing key found for user {}", did); @@ -295,36 +273,26 @@ pub async fn import_repo( return ApiError::InternalError(None).into_response(); } }; - let new_root_str = new_root_cid.to_string(); - if let Err(e) = sqlx::query!( - "UPDATE repos SET repo_root_cid = $1, repo_rev = $2, updated_at = NOW() WHERE user_id = $3", - new_root_str, - &new_rev_str, - user_id - ) - .execute(&state.db) - .await + let new_root_cid_link = CidLink::new_unchecked(&new_root_cid.to_string()); + if let Err(e) = state + .repo_repo + .update_repo_root(user_id, &new_root_cid_link, &new_rev_str) + .await { error!("Failed to update repo root: {:?}", e); return ApiError::InternalError(None).into_response(); } let mut all_block_cids: Vec> = blocks.keys().map(|c| c.to_bytes()).collect(); all_block_cids.push(new_root_cid.to_bytes()); - if let Err(e) = sqlx::query!( - r#" - INSERT INTO user_blocks (user_id, block_cid) - SELECT $1, block_cid FROM UNNEST($2::bytea[]) AS t(block_cid) - ON CONFLICT (user_id, block_cid) DO NOTHING - "#, - user_id, - &all_block_cids - ) - .execute(&state.db) - .await + if let Err(e) = state + .repo_repo + .insert_user_blocks(user_id, &all_block_cids, &new_rev_str) + .await { error!("Failed to insert user_blocks: {:?}", e); return ApiError::InternalError(None).into_response(); } + let new_root_str = new_root_cid.to_string(); info!( "Created new commit for imported repo: cid={}, rev={}", new_root_str, new_rev_str @@ -338,15 +306,14 @@ pub async fn import_repo( "$type": "app.bsky.actor.defs#personalDetailsPref", "birthDate": "1998-05-06T00:00:00.000Z" }); - if let Err(e) = sqlx::query!( - "INSERT INTO account_preferences (user_id, name, value_json) VALUES ($1, $2, $3) - ON CONFLICT (user_id, name) DO NOTHING", - user_id, - "app.bsky.actor.defs#personalDetailsPref", - birthdate_pref - ) - .execute(&state.db) - .await + if let Err(e) = state + .infra_repo + .insert_account_preference_if_not_exists( + user_id, + "app.bsky.actor.defs#personalDetailsPref", + birthdate_pref, + ) + .await { warn!( "Failed to set default birthdate preference for migrated user: {:?}", @@ -397,37 +364,20 @@ async fn sequence_import_event( state: &AppState, did: &Did, commit_cid: &str, -) -> Result<(), sqlx::Error> { - let prev_cid: Option = None; - let prev_data_cid: Option = None; - let ops = serde_json::json!([]); - let blobs: Vec = vec![]; - let blocks_cids: Vec = vec![]; - let did_str = did.as_str(); - - let mut tx = state.db.begin().await?; - - let seq_row = sqlx::query!( - r#" - INSERT INTO repo_seq (did, event_type, commit_cid, prev_cid, prev_data_cid, ops, blobs, blocks_cids) - VALUES ($1, 'commit', $2, $3, $4, $5, $6, $7) - RETURNING seq - "#, - did_str, - commit_cid, - prev_cid, - prev_data_cid, - ops, - &blobs, - &blocks_cids - ) - .fetch_one(&mut *tx) - .await?; - - sqlx::query(&format!("NOTIFY repo_updates, '{}'", seq_row.seq)) - .execute(&mut *tx) - .await?; +) -> Result<(), tranquil_db::DbError> { + let data = tranquil_db::CommitEventData { + did: did.clone(), + event_type: "commit".to_string(), + commit_cid: Some(CidLink::new_unchecked(commit_cid)), + prev_cid: None, + ops: Some(serde_json::json!([])), + blobs: Some(vec![]), + blocks_cids: Some(vec![]), + prev_data_cid: None, + rev: None, + }; - tx.commit().await?; + let seq = state.repo_repo.insert_commit_event(&data).await?; + state.repo_repo.notify_update(seq).await?; Ok(()) } diff --git a/crates/tranquil-pds/src/api/repo/meta.rs b/crates/tranquil-pds/src/api/repo/meta.rs index 711a32b..ec5cfbc 100644 --- a/crates/tranquil-pds/src/api/repo/meta.rs +++ b/crates/tranquil-pds/src/api/repo/meta.rs @@ -19,28 +19,33 @@ pub async fn describe_repo( Query(input): Query, ) -> Response { let hostname = std::env::var("PDS_HOSTNAME").unwrap_or_else(|_| "localhost".to_string()); + let hostname_for_handles = hostname.split(':').next().unwrap_or(&hostname); let user_row = if input.repo.is_did() { - sqlx::query!( - "SELECT id, handle, did FROM users WHERE did = $1", - input.repo.as_str() - ) - .fetch_optional(&state.db) - .await - .map(|opt| opt.map(|r| (r.id, r.handle, r.did))) + let did: crate::types::Did = match input.repo.as_str().parse() { + Ok(d) => d, + Err(_) => return ApiError::InvalidRequest("Invalid DID format".into()).into_response(), + }; + state + .user_repo + .get_by_did(&did) + .await + .map(|opt| opt.map(|r| (r.id, r.handle, r.did))) } else { let repo_str = input.repo.as_str(); - let handle = if !repo_str.contains('.') { - format!("{}.{}", repo_str, hostname) + let handle_str = if !repo_str.contains('.') { + format!("{}.{}", repo_str, hostname_for_handles) } else { repo_str.to_string() }; - sqlx::query!( - "SELECT id, handle, did FROM users WHERE handle = $1", - handle - ) - .fetch_optional(&state.db) - .await - .map(|opt| opt.map(|r| (r.id, r.handle, r.did))) + let handle: crate::types::Handle = match handle_str.parse() { + Ok(h) => h, + Err(_) => return ApiError::InvalidRequest("Invalid handle format".into()).into_response(), + }; + state + .user_repo + .get_by_handle(&handle) + .await + .map(|opt| opt.map(|r| (r.id, r.handle, r.did))) }; let (user_id, handle, did) = match user_row { Ok(Some((id, handle, did))) => (id, handle, did), @@ -51,16 +56,11 @@ pub async fn describe_repo( return ApiError::InternalError(None).into_response(); } }; - let collections_query = sqlx::query!( - "SELECT DISTINCT collection FROM records WHERE repo_id = $1", - user_id - ) - .fetch_all(&state.db) - .await; - let collections: Vec = match collections_query { - Ok(rows) => rows.iter().map(|r| r.collection.clone()).collect(), - Err(_) => Vec::new(), - }; + let collections = state + .repo_repo + .list_collections(user_id) + .await + .unwrap_or_default(); let did_doc = json!({ "id": did, "alsoKnownAs": [format!("at://{}", handle)] diff --git a/crates/tranquil-pds/src/api/repo/record/batch.rs b/crates/tranquil-pds/src/api/repo/record/batch.rs index 95a62a8..206f9d9 100644 --- a/crates/tranquil-pds/src/api/repo/record/batch.rs +++ b/crates/tranquil-pds/src/api/repo/record/batch.rs @@ -1,9 +1,8 @@ use super::validation::validate_record_with_status; -use super::write::has_verified_comms_channel; use crate::api::error::ApiError; use crate::api::repo::record::utils::{CommitParams, RecordOp, commit_and_log, extract_blob_cids}; use crate::auth::BearerAuth; -use crate::delegation::{self, DelegationActionType}; +use crate::delegation::DelegationActionType; use crate::repo::tracking::TrackingBlockStore; use crate::state::AppState; use crate::types::{AtIdentifier, AtUri, Did, Nsid, Rkey}; @@ -280,16 +279,22 @@ pub async fn apply_writes( return ApiError::InvalidRepo("Repo does not match authenticated user".into()) .into_response(); } - if crate::util::is_account_migrated(&state.db, &did) + if state + .user_repo + .is_account_migrated(&did) .await .unwrap_or(false) { return ApiError::AccountMigrated.into_response(); } - let is_verified = has_verified_comms_channel(&state.db, &did) + let is_verified = state + .user_repo + .has_verified_comms_channel(&did) .await .unwrap_or(false); - let is_delegated = crate::delegation::is_delegated_account(&state.db, &did) + let is_delegated = state + .delegation_repo + .is_delegated_account(&did) .await .unwrap_or(false); if !is_verified && !is_delegated { @@ -373,20 +378,18 @@ pub async fn apply_writes( } } - let user_id: uuid::Uuid = - match sqlx::query_scalar!("SELECT id FROM users WHERE did = $1", did.as_str()) - .fetch_optional(&state.db) - .await - { - Ok(Some(id)) => id, - _ => return ApiError::InternalError(Some("User not found".into())).into_response(), - }; - let root_cid_str: String = match sqlx::query_scalar!( - "SELECT repo_root_cid FROM repos WHERE user_id = $1", - user_id - ) - .fetch_optional(&state.db) - .await + let user_id: uuid::Uuid = match state + .user_repo + .get_id_by_did(&did) + .await + { + Ok(Some(id)) => id, + _ => return ApiError::InternalError(Some("User not found".into())).into_response(), + }; + let root_cid_str = match state + .repo_repo + .get_repo_root_cid_by_user_id(user_id) + .await { Ok(Some(cid_str)) => cid_str, _ => return ApiError::InternalError(Some("Repo root not found".into())).into_response(), @@ -544,21 +547,22 @@ pub async fn apply_writes( }) .collect(); - let _ = delegation::log_delegation_action( - &state.db, - &did, - controller, - Some(controller), - DelegationActionType::RepoWrite, - Some(json!({ - "action": "apply_writes", - "count": input.writes.len(), - "writes": write_summary - })), - None, - None, - ) - .await; + let _ = state + .delegation_repo + .log_delegation_action( + &did, + controller, + Some(controller), + DelegationActionType::RepoWrite, + Some(json!({ + "action": "apply_writes", + "count": input.writes.len(), + "writes": write_summary + })), + None, + None, + ) + .await; } ( diff --git a/crates/tranquil-pds/src/api/repo/record/delete.rs b/crates/tranquil-pds/src/api/repo/record/delete.rs index e65019c..e8af11c 100644 --- a/crates/tranquil-pds/src/api/repo/record/delete.rs +++ b/crates/tranquil-pds/src/api/repo/record/delete.rs @@ -1,10 +1,10 @@ use crate::api::error::ApiError; use crate::api::repo::record::utils::{CommitParams, RecordOp, commit_and_log}; use crate::api::repo::record::write::{CommitInfo, prepare_repo_write}; -use crate::delegation::{self, DelegationActionType}; +use crate::delegation::DelegationActionType; use crate::repo::tracking::TrackingBlockStore; use crate::state::AppState; -use crate::types::{AtIdentifier, Nsid, Rkey}; +use crate::types::{AtIdentifier, AtUri, Nsid, Rkey}; use axum::{ Json, extract::State, @@ -183,21 +183,27 @@ pub async fn delete_record( }; if let Some(ref controller) = controller_did { - let _ = delegation::log_delegation_action( - &state.db, - &did, - controller, - Some(controller), - DelegationActionType::RepoWrite, - Some(json!({ - "action": "delete", - "collection": collection_for_audit, - "rkey": rkey_for_audit - })), - None, - None, - ) - .await; + let _ = state + .delegation_repo + .log_delegation_action( + &did, + controller, + Some(controller), + DelegationActionType::RepoWrite, + Some(json!({ + "action": "delete", + "collection": collection_for_audit, + "rkey": rkey_for_audit + })), + None, + None, + ) + .await; + } + + let deleted_uri = AtUri::from_parts(&did, &input.collection, &input.rkey); + if let Err(e) = state.backlink_repo.remove_backlinks_by_uri(&deleted_uri).await { + error!("Failed to remove backlinks for {}: {}", deleted_uri, e); } ( @@ -211,3 +217,115 @@ pub async fn delete_record( ) .into_response() } + +use crate::types::Did; +use uuid::Uuid; + +pub async fn delete_record_internal( + state: &AppState, + did: &Did, + user_id: Uuid, + collection: &Nsid, + rkey: &Rkey, +) -> Result<(), String> { + let root_cid_str = state + .repo_repo + .get_repo_root_cid_by_user_id(user_id) + .await + .map_err(|e| format!("DB error: {}", e))? + .ok_or_else(|| "Repo root not found".to_string())?; + + let current_root_cid = + Cid::from_str(root_cid_str.as_str()).map_err(|_| "Invalid repo root CID".to_string())?; + + let tracking_store = TrackingBlockStore::new(state.block_store.clone()); + let commit_bytes = tracking_store + .get(¤t_root_cid) + .await + .map_err(|e| format!("Failed to fetch commit: {:?}", e))? + .ok_or_else(|| "Commit block not found".to_string())?; + + let commit = Commit::from_cbor(&commit_bytes) + .map_err(|e| format!("Failed to parse commit: {:?}", e))?; + + let mst = Mst::load(Arc::new(tracking_store.clone()), commit.data, None); + let key = format!("{}/{}", collection, rkey); + + let prev_record_cid = mst + .get(&key) + .await + .map_err(|e| format!("MST get error: {:?}", e))?; + + let Some(prev_cid) = prev_record_cid else { + return Ok(()); + }; + + let new_mst = mst + .delete(&key) + .await + .map_err(|e| format!("Failed to delete from MST: {:?}", e))?; + + let new_mst_root = new_mst + .persist() + .await + .map_err(|e| format!("Failed to persist MST: {:?}", e))?; + + let op = RecordOp::Delete { + collection: collection.clone(), + rkey: rkey.clone(), + prev: Some(prev_cid), + }; + + let mut new_mst_blocks = std::collections::BTreeMap::new(); + let mut old_mst_blocks = std::collections::BTreeMap::new(); + + new_mst + .blocks_for_path(&key, &mut new_mst_blocks) + .await + .map_err(|e| format!("Failed to get new MST blocks: {:?}", e))?; + + mst.blocks_for_path(&key, &mut old_mst_blocks) + .await + .map_err(|e| format!("Failed to get old MST blocks: {:?}", e))?; + + let mut relevant_blocks = new_mst_blocks.clone(); + relevant_blocks.extend(old_mst_blocks.iter().map(|(k, v)| (*k, v.clone()))); + + let written_cids: Vec = tracking_store + .get_all_relevant_cids() + .into_iter() + .chain(relevant_blocks.keys().copied()) + .collect::>() + .into_iter() + .collect(); + + let written_cids_str: Vec = written_cids.iter().map(|c| c.to_string()).collect(); + + let obsolete_cids: Vec = std::iter::once(current_root_cid) + .chain( + old_mst_blocks + .keys() + .filter(|cid| !new_mst_blocks.contains_key(*cid)) + .copied(), + ) + .chain(std::iter::once(prev_cid)) + .collect(); + + commit_and_log( + state, + CommitParams { + did, + user_id, + current_root_cid: Some(current_root_cid), + prev_data_cid: Some(commit.data), + new_mst_root, + ops: vec![op], + blocks_cids: &written_cids_str, + blobs: &[], + obsolete_cids, + }, + ) + .await?; + + Ok(()) +} diff --git a/crates/tranquil-pds/src/api/repo/record/mod.rs b/crates/tranquil-pds/src/api/repo/record/mod.rs index f7c41ca..b9de2a1 100644 --- a/crates/tranquil-pds/src/api/repo/record/mod.rs +++ b/crates/tranquil-pds/src/api/repo/record/mod.rs @@ -6,7 +6,7 @@ pub mod validation; pub mod write; pub use batch::apply_writes; -pub use delete::{DeleteRecordInput, delete_record}; +pub use delete::{DeleteRecordInput, delete_record, delete_record_internal}; pub use read::{GetRecordInput, ListRecordsInput, ListRecordsOutput, get_record, list_records}; pub use utils::*; pub use write::{ diff --git a/crates/tranquil-pds/src/api/repo/record/read.rs b/crates/tranquil-pds/src/api/repo/record/read.rs index 6fcc37a..99a4e99 100644 --- a/crates/tranquil-pds/src/api/repo/record/read.rs +++ b/crates/tranquil-pds/src/api/repo/record/read.rs @@ -59,22 +59,33 @@ pub async fn get_record( Query(input): Query, ) -> Response { let hostname = std::env::var("PDS_HOSTNAME").unwrap_or_else(|_| "localhost".to_string()); + let hostname_for_handles = hostname.split(':').next().unwrap_or(&hostname); let user_id_opt = if input.repo.is_did() { - sqlx::query!("SELECT id FROM users WHERE did = $1", input.repo.as_str()) - .fetch_optional(&state.db) + let did: crate::types::Did = match input.repo.as_str().parse() { + Ok(d) => d, + Err(_) => return ApiError::InvalidRequest("Invalid DID format".into()).into_response(), + }; + state + .user_repo + .get_id_by_did(&did) .await - .map(|opt| opt.map(|r| r.id)) + .map_err(|_| ()) } else { let repo_str = input.repo.as_str(); - let handle = if !repo_str.contains('.') { - format!("{}.{}", repo_str, hostname) + let handle_str = if !repo_str.contains('.') { + format!("{}.{}", repo_str, hostname_for_handles) } else { repo_str.to_string() }; - sqlx::query!("SELECT id FROM users WHERE handle = $1", handle) - .fetch_optional(&state.db) + let handle: crate::types::Handle = match handle_str.parse() { + Ok(h) => h, + Err(_) => return ApiError::InvalidRequest("Invalid handle format".into()).into_response(), + }; + state + .user_repo + .get_id_by_handle(&handle) .await - .map(|opt| opt.map(|r| r.id)) + .map_err(|_| ()) }; let user_id: uuid::Uuid = match user_id_opt { Ok(Some(id)) => id, @@ -85,20 +96,17 @@ pub async fn get_record( return ApiError::InternalError(None).into_response(); } }; - let record_row = sqlx::query!( - "SELECT record_cid FROM records WHERE repo_id = $1 AND collection = $2 AND rkey = $3", - user_id, - input.collection.as_str(), - input.rkey.as_str() - ) - .fetch_optional(&state.db) - .await; - let record_cid_str: String = match record_row { - Ok(Some(row)) => row.record_cid, + let record_row = state + .repo_repo + .get_record_cid(user_id, &input.collection, &input.rkey) + .await; + let record_cid_link = match record_row { + Ok(Some(cid)) => cid, _ => { return ApiError::RecordNotFound.into_response(); } }; + let record_cid_str = record_cid_link.to_string(); if let Some(expected_cid) = &input.cid && &record_cid_str != expected_cid { @@ -152,22 +160,33 @@ pub async fn list_records( Query(input): Query, ) -> Response { let hostname = std::env::var("PDS_HOSTNAME").unwrap_or_else(|_| "localhost".to_string()); + let hostname_for_handles = hostname.split(':').next().unwrap_or(&hostname); let user_id_opt = if input.repo.is_did() { - sqlx::query!("SELECT id FROM users WHERE did = $1", input.repo.as_str()) - .fetch_optional(&state.db) + let did: crate::types::Did = match input.repo.as_str().parse() { + Ok(d) => d, + Err(_) => return ApiError::InvalidRequest("Invalid DID format".into()).into_response(), + }; + state + .user_repo + .get_id_by_did(&did) .await - .map(|opt| opt.map(|r| r.id)) + .map_err(|_| ()) } else { let repo_str = input.repo.as_str(); - let handle = if !repo_str.contains('.') { - format!("{}.{}", repo_str, hostname) + let handle_str = if !repo_str.contains('.') { + format!("{}.{}", repo_str, hostname_for_handles) } else { repo_str.to_string() }; - sqlx::query!("SELECT id FROM users WHERE handle = $1", handle) - .fetch_optional(&state.db) + let handle: crate::types::Handle = match handle_str.parse() { + Ok(h) => h, + Err(_) => return ApiError::InvalidRequest("Invalid handle format".into()).into_response(), + }; + state + .user_repo + .get_id_by_handle(&handle) .await - .map(|opt| opt.map(|r| r.id)) + .map_err(|_| ()) }; let user_id: uuid::Uuid = match user_id_opt { Ok(Some(id)) => id, @@ -181,67 +200,33 @@ pub async fn list_records( let limit = input.limit.unwrap_or(50).clamp(1, 100); let reverse = input.reverse.unwrap_or(false); let limit_i64 = limit as i64; - let order = if reverse { "ASC" } else { "DESC" }; - let rows_res: Result, sqlx::Error> = if let Some(cursor) = &input.cursor { - let comparator = if reverse { ">" } else { "<" }; - let query = format!( - "SELECT rkey, record_cid FROM records WHERE repo_id = $1 AND collection = $2 AND rkey {} $3 ORDER BY rkey {} LIMIT $4", - comparator, order - ); - sqlx::query_as(&query) - .bind(user_id) - .bind(input.collection.as_str()) - .bind(cursor) - .bind(limit_i64) - .fetch_all(&state.db) - .await - } else { - let mut conditions = vec!["repo_id = $1", "collection = $2"]; - let mut param_idx = 3; - if input.rkey_start.is_some() { - conditions.push("rkey > $3"); - param_idx += 1; - } - if input.rkey_end.is_some() { - conditions.push(if param_idx == 3 { - "rkey < $3" - } else { - "rkey < $4" - }); - param_idx += 1; - } - let limit_idx = param_idx; - let query = format!( - "SELECT rkey, record_cid FROM records WHERE {} ORDER BY rkey {} LIMIT ${}", - conditions.join(" AND "), - order, - limit_idx - ); - let mut query_builder = sqlx::query_as::<_, (String, String)>(&query) - .bind(user_id) - .bind(input.collection.as_str()); - if let Some(start) = &input.rkey_start { - query_builder = query_builder.bind(start.as_str()); - } - if let Some(end) = &input.rkey_end { - query_builder = query_builder.bind(end.as_str()); - } - query_builder.bind(limit_i64).fetch_all(&state.db).await - }; - let rows = match rows_res { + let cursor_rkey = input.cursor.as_ref().and_then(|c| c.parse::().ok()); + let rows = match state + .repo_repo + .list_records( + user_id, + &input.collection, + cursor_rkey.as_ref(), + limit_i64, + reverse, + input.rkey_start.as_ref(), + input.rkey_end.as_ref(), + ) + .await + { Ok(r) => r, Err(e) => { error!("Error listing records: {:?}", e); return ApiError::InternalError(None).into_response(); } }; - let last_rkey = rows.last().map(|(rkey, _)| rkey.clone()); + let last_rkey = rows.last().map(|r| r.rkey.to_string()); let parsed_rows: Vec<(Cid, String, String)> = rows .iter() - .filter_map(|(rkey, cid_str)| { - Cid::from_str(cid_str) + .filter_map(|row| { + Cid::from_str(row.record_cid.as_str()) .ok() - .map(|cid| (cid, rkey.clone(), cid_str.clone())) + .map(|cid| (cid, row.rkey.to_string(), row.record_cid.to_string())) }) .collect(); let cids: Vec = parsed_rows.iter().map(|(cid, _, _)| *cid).collect(); diff --git a/crates/tranquil-pds/src/api/repo/record/utils.rs b/crates/tranquil-pds/src/api/repo/record/utils.rs index 68f83e5..286a9cf 100644 --- a/crates/tranquil-pds/src/api/repo/record/utils.rs +++ b/crates/tranquil-pds/src/api/repo/record/utils.rs @@ -36,6 +36,41 @@ fn extract_blob_cids_recursive(value: &Value, blobs: &mut Vec) { } } +use tranquil_db_traits::Backlink; +use crate::types::AtUri; + +pub fn extract_backlinks(uri: &AtUri, record: &Value) -> Vec { + let record_type = record + .get("$type") + .and_then(|v| v.as_str()) + .unwrap_or_default(); + + match record_type { + "app.bsky.graph.follow" | "app.bsky.graph.block" => record + .get("subject") + .and_then(|v| v.as_str()) + .filter(|s| s.starts_with("did:")) + .map(|subject| vec![Backlink { + uri: uri.clone(), + path: "subject".to_string(), + link_to: subject.to_string(), + }]) + .unwrap_or_default(), + "app.bsky.feed.like" | "app.bsky.feed.repost" => record + .get("subject") + .and_then(|v| v.get("uri")) + .and_then(|v| v.as_str()) + .filter(|s| s.starts_with("at://")) + .map(|subject_uri| vec![Backlink { + uri: uri.clone(), + path: "subject.uri".to_string(), + link_to: subject_uri.to_string(), + }]) + .unwrap_or_default(), + _ => Vec::new(), + } +} + pub fn create_signed_commit( did: &Did, data: Cid, @@ -98,6 +133,8 @@ pub async fn commit_and_log( state: &AppState, params: CommitParams<'_>, ) -> Result { + use tranquil_db_traits::{ApplyCommitError, ApplyCommitInput, CommitEventData, RecordDelete, RecordUpsert}; + let CommitParams { did, user_id, @@ -109,13 +146,12 @@ pub async fn commit_and_log( blobs, obsolete_cids, } = params; - let key_row = sqlx::query!( - "SELECT key_bytes, encryption_version FROM user_keys WHERE user_id = $1", - user_id - ) - .fetch_one(&state.db) - .await - .map_err(|e| format!("Failed to fetch signing key: {}", e))?; + let key_row = state + .user_repo + .get_user_key_by_id(user_id) + .await + .map_err(|e| format!("Failed to fetch signing key: {}", e))? + .ok_or_else(|| "Signing key not found".to_string())?; let key_bytes = crate::config::decrypt_key(&key_row.key_bytes, key_row.encryption_version) .map_err(|e| format!("Failed to decrypt signing key: {}", e))?; let signing_key = @@ -129,184 +165,48 @@ pub async fn commit_and_log( .put(&new_commit_bytes) .await .map_err(|e| format!("Failed to save commit block: {:?}", e))?; - let mut tx = state - .db - .begin() - .await - .map_err(|e| format!("Failed to begin transaction: {}", e))?; - let lock_result = sqlx::query!( - "SELECT repo_root_cid FROM repos WHERE user_id = $1 FOR UPDATE NOWAIT", - user_id - ) - .fetch_optional(&mut *tx) - .await; - match lock_result { - Err(e) => { - if let Some(db_err) = e.as_database_error() - && db_err.code().as_deref() == Some("55P03") - { - return Err( - "ConcurrentModification: Another request is modifying this repo".to_string(), - ); - } - return Err(format!("Failed to acquire repo lock: {}", e)); - } - Ok(Some(row)) => { - if let Some(expected_root) = ¤t_root_cid - && row.repo_root_cid != expected_root.to_string() - { - return Err( - "ConcurrentModification: Repo has been modified since last read".to_string(), - ); - } - } - Ok(None) => { - return Err("Repo not found".to_string()); - } - } - let is_account_active = sqlx::query_scalar!( - "SELECT deactivated_at IS NULL FROM users WHERE id = $1", - user_id - ) - .fetch_optional(&mut *tx) - .await - .map_err(|e| format!("Failed to check account status: {}", e))? - .flatten() - .unwrap_or(false); - sqlx::query!( - "UPDATE repos SET repo_root_cid = $1, repo_rev = $2 WHERE user_id = $3", - new_root_cid.to_string(), - &rev_str, - user_id - ) - .execute(&mut *tx) - .await - .map_err(|e| format!("DB Error (repos): {}", e))?; + let mut all_block_cids: Vec> = blocks_cids .iter() .filter_map(|s| Cid::from_str(s).ok()) .map(|c| c.to_bytes()) .collect(); all_block_cids.push(new_root_cid.to_bytes()); - if !all_block_cids.is_empty() { - sqlx::query!( - r#" - INSERT INTO user_blocks (user_id, block_cid) - SELECT $1, block_cid FROM UNNEST($2::bytea[]) AS t(block_cid) - ON CONFLICT (user_id, block_cid) DO NOTHING - "#, - user_id, - &all_block_cids - ) - .execute(&mut *tx) - .await - .map_err(|e| format!("DB Error (user_blocks): {}", e))?; - } - if !obsolete_cids.is_empty() { - let obsolete_bytes: Vec> = obsolete_cids.iter().map(|c| c.to_bytes()).collect(); - sqlx::query!( - r#" - DELETE FROM user_blocks - WHERE user_id = $1 - AND block_cid = ANY($2) - "#, - user_id, - &obsolete_bytes as &[Vec] - ) - .execute(&mut *tx) - .await - .map_err(|e| format!("DB Error (user_blocks delete obsolete): {}", e))?; - } - let (upserts, deletes): (Vec<_>, Vec<_>) = ops - .iter() - .partition(|op| matches!(op, RecordOp::Create { .. } | RecordOp::Update { .. })); - let (upsert_collections, upsert_rkeys, upsert_cids): (Vec, Vec, Vec) = - upserts - .into_iter() - .filter_map(|op| match op { - RecordOp::Create { - collection, - rkey, - cid, + + let obsolete_bytes: Vec> = obsolete_cids.iter().map(|c| c.to_bytes()).collect(); + + let (record_upserts, record_deletes): (Vec, Vec) = ops.iter().fold( + (Vec::new(), Vec::new()), + |(mut upserts, mut deletes), op| { + match op { + RecordOp::Create { collection, rkey, cid } + | RecordOp::Update { collection, rkey, cid, .. } => { + upserts.push(RecordUpsert { + collection: collection.clone(), + rkey: rkey.clone(), + cid: crate::types::CidLink::new_unchecked(&cid.to_string()), + }); } - | RecordOp::Update { - collection, - rkey, - cid, - .. - } => Some((collection.to_string(), rkey.to_string(), cid.to_string())), - _ => None, - }) - .fold( - (Vec::new(), Vec::new(), Vec::new()), - |(mut cols, mut rkeys, mut cids), (c, r, ci)| { - cols.push(c); - rkeys.push(r); - cids.push(ci); - (cols, rkeys, cids) - }, - ); - let (delete_collections, delete_rkeys): (Vec, Vec) = deletes - .into_iter() - .filter_map(|op| match op { - RecordOp::Delete { - collection, rkey, .. - } => Some((collection.to_string(), rkey.to_string())), - _ => None, - }) - .unzip(); - if !upsert_collections.is_empty() { - sqlx::query!( - r#" - INSERT INTO records (repo_id, collection, rkey, record_cid, repo_rev) - SELECT $1, collection, rkey, record_cid, $5 - FROM UNNEST($2::text[], $3::text[], $4::text[]) AS t(collection, rkey, record_cid) - ON CONFLICT (repo_id, collection, rkey) DO UPDATE - SET record_cid = EXCLUDED.record_cid, repo_rev = EXCLUDED.repo_rev, created_at = NOW() - "#, - user_id, - &upsert_collections, - &upsert_rkeys, - &upsert_cids, - rev_str - ) - .execute(&mut *tx) - .await - .map_err(|e| format!("DB Error (records batch upsert): {}", e))?; - } - if !delete_collections.is_empty() { - sqlx::query!( - r#" - DELETE FROM records - WHERE repo_id = $1 - AND (collection, rkey) IN (SELECT * FROM UNNEST($2::text[], $3::text[])) - "#, - user_id, - &delete_collections, - &delete_rkeys - ) - .execute(&mut *tx) - .await - .map_err(|e| format!("DB Error (records batch delete): {}", e))?; - } - let ops_json = ops + RecordOp::Delete { collection, rkey, .. } => { + deletes.push(RecordDelete { + collection: collection.clone(), + rkey: rkey.clone(), + }); + } + } + (upserts, deletes) + }, + ); + + let ops_json: Vec = ops .iter() .map(|op| match op { - RecordOp::Create { - collection, - rkey, - cid, - } => json!({ + RecordOp::Create { collection, rkey, cid } => json!({ "action": "create", "path": format!("{}/{}", collection, rkey), "cid": cid.to_string() }), - RecordOp::Update { - collection, - rkey, - cid, - prev, - } => { + RecordOp::Update { collection, rkey, cid, prev } => { let mut obj = json!({ "action": "update", "path": format!("{}/{}", collection, rkey), @@ -317,11 +217,7 @@ pub async fn commit_and_log( } obj } - RecordOp::Delete { - collection, - rkey, - prev, - } => { + RecordOp::Delete { collection, rkey, prev } => { let mut obj = json!({ "action": "delete", "path": format!("{}/{}", collection, rkey), @@ -333,40 +229,49 @@ pub async fn commit_and_log( obj } }) - .collect::>(); - if is_account_active { - let event_type = "commit"; - let prev_cid_str = current_root_cid.map(|c| c.to_string()); - let prev_data_cid_str = prev_data_cid.map(|c| c.to_string()); - let seq_row = sqlx::query!( - r#" - INSERT INTO repo_seq (did, event_type, commit_cid, prev_cid, ops, blobs, blocks_cids, prev_data_cid) - VALUES ($1, $2, $3, $4, $5, $6, $7, $8) - RETURNING seq - "#, - did.as_str(), - event_type, - new_root_cid.to_string(), - prev_cid_str, - json!(ops_json), - blobs, - blocks_cids, - prev_data_cid_str, - ) - .fetch_one(&mut *tx) - .await - .map_err(|e| format!("DB Error (repo_seq): {}", e))?; - sqlx::query(&format!("NOTIFY repo_updates, '{}'", seq_row.seq)) - .execute(&mut *tx) - .await - .map_err(|e| format!("DB Error (notify): {}", e))?; - } - tx.commit() + .collect(); + + let commit_event = CommitEventData { + did: did.clone(), + event_type: "commit".to_string(), + commit_cid: Some(crate::types::CidLink::new_unchecked(&new_root_cid.to_string())), + prev_cid: current_root_cid.map(|c| crate::types::CidLink::new_unchecked(&c.to_string())), + ops: Some(json!(ops_json)), + blobs: Some(blobs.to_vec()), + blocks_cids: Some(blocks_cids.to_vec()), + prev_data_cid: prev_data_cid.map(|c| crate::types::CidLink::new_unchecked(&c.to_string())), + rev: Some(rev_str.clone()), + }; + + let input = ApplyCommitInput { + user_id, + did: did.clone(), + expected_root_cid: current_root_cid.map(|c| crate::types::CidLink::new_unchecked(&c.to_string())), + new_root_cid: crate::types::CidLink::new_unchecked(&new_root_cid.to_string()), + new_rev: rev_str.clone(), + new_block_cids: all_block_cids, + obsolete_block_cids: obsolete_bytes, + record_upserts, + record_deletes, + commit_event, + }; + + let result = state + .repo_repo + .apply_commit(input) .await - .map_err(|e| format!("Failed to commit transaction: {}", e))?; - if is_account_active { + .map_err(|e| match e { + ApplyCommitError::RepoNotFound => "Repo not found".to_string(), + ApplyCommitError::ConcurrentModification => { + "ConcurrentModification: Repo has been modified since last read".to_string() + } + ApplyCommitError::Database(msg) => format!("DB Error: {}", msg), + })?; + + if result.is_account_active { let _ = sequence_sync_event(state, did, &new_root_cid.to_string(), Some(&rev_str)).await; } + Ok(CommitResult { commit_cid: new_root_cid, rev: rev_str, @@ -382,21 +287,20 @@ pub async fn create_record_internal( use crate::repo::tracking::TrackingBlockStore; use jacquard_repo::mst::Mst; use std::sync::Arc; - let user_id: Uuid = sqlx::query_scalar!("SELECT id FROM users WHERE did = $1", did.as_str()) - .fetch_optional(&state.db) + let user_id: Uuid = state + .user_repo + .get_id_by_did(did) .await .map_err(|e| format!("DB error: {}", e))? .ok_or_else(|| "User not found".to_string())?; - let root_cid_str: String = sqlx::query_scalar!( - "SELECT repo_root_cid FROM repos WHERE user_id = $1", - user_id - ) - .fetch_optional(&state.db) - .await - .map_err(|e| format!("DB error: {}", e))? - .ok_or_else(|| "Repo not found".to_string())?; + let root_cid_link = state + .repo_repo + .get_repo_root_cid_by_user_id(user_id) + .await + .map_err(|e| format!("DB error: {}", e))? + .ok_or_else(|| "Repo not found".to_string())?; let current_root_cid = - Cid::from_str(&root_cid_str).map_err(|_| "Invalid repo root CID".to_string())?; + Cid::from_str(root_cid_link.as_str()).map_err(|_| "Invalid repo root CID".to_string())?; let tracking_store = TrackingBlockStore::new(state.block_store.clone()); let commit_bytes = tracking_store .get(¤t_root_cid) @@ -481,31 +385,11 @@ pub async fn sequence_identity_event( did: &Did, handle: Option<&Handle>, ) -> Result { - let mut tx = state - .db - .begin() + state + .repo_repo + .insert_identity_event(did, handle) .await - .map_err(|e| format!("Failed to begin transaction: {}", e))?; - let seq_row = sqlx::query!( - r#" - INSERT INTO repo_seq (did, event_type, handle) - VALUES ($1, 'identity', $2) - RETURNING seq - "#, - did.as_str(), - handle.map(|h| h.as_str()), - ) - .fetch_one(&mut *tx) - .await - .map_err(|e| format!("DB Error (repo_seq identity): {}", e))?; - sqlx::query(&format!("NOTIFY repo_updates, '{}'", seq_row.seq)) - .execute(&mut *tx) - .await - .map_err(|e| format!("DB Error (notify): {}", e))?; - tx.commit() - .await - .map_err(|e| format!("Failed to commit transaction: {}", e))?; - Ok(seq_row.seq) + .map_err(|e| format!("DB Error (identity event): {}", e)) } pub async fn sequence_account_event( state: &AppState, @@ -513,32 +397,11 @@ pub async fn sequence_account_event( active: bool, status: Option<&str>, ) -> Result { - let mut tx = state - .db - .begin() - .await - .map_err(|e| format!("Failed to begin transaction: {}", e))?; - let seq_row = sqlx::query!( - r#" - INSERT INTO repo_seq (did, event_type, active, status) - VALUES ($1, 'account', $2, $3) - RETURNING seq - "#, - did.as_str(), - active, - status, - ) - .fetch_one(&mut *tx) - .await - .map_err(|e| format!("DB Error (repo_seq account): {}", e))?; - sqlx::query(&format!("NOTIFY repo_updates, '{}'", seq_row.seq)) - .execute(&mut *tx) - .await - .map_err(|e| format!("DB Error (notify): {}", e))?; - tx.commit() + state + .repo_repo + .insert_account_event(did, active, status) .await - .map_err(|e| format!("Failed to commit transaction: {}", e))?; - Ok(seq_row.seq) + .map_err(|e| format!("DB Error (account event): {}", e)) } pub async fn sequence_sync_event( state: &AppState, @@ -546,32 +409,12 @@ pub async fn sequence_sync_event( commit_cid: &str, rev: Option<&str>, ) -> Result { - let mut tx = state - .db - .begin() - .await - .map_err(|e| format!("Failed to begin transaction: {}", e))?; - let seq_row = sqlx::query!( - r#" - INSERT INTO repo_seq (did, event_type, commit_cid, rev) - VALUES ($1, 'sync', $2, $3) - RETURNING seq - "#, - did.as_str(), - commit_cid, - rev, - ) - .fetch_one(&mut *tx) - .await - .map_err(|e| format!("DB Error (repo_seq sync): {}", e))?; - sqlx::query(&format!("NOTIFY repo_updates, '{}'", seq_row.seq)) - .execute(&mut *tx) - .await - .map_err(|e| format!("DB Error (notify): {}", e))?; - tx.commit() + let cid_link = crate::types::CidLink::new_unchecked(commit_cid); + state + .repo_repo + .insert_sync_event(did, &cid_link, rev) .await - .map_err(|e| format!("Failed to commit transaction: {}", e))?; - Ok(seq_row.seq) + .map_err(|e| format!("DB Error (sync event): {}", e)) } pub async fn sequence_genesis_commit( @@ -581,39 +424,11 @@ pub async fn sequence_genesis_commit( mst_root_cid: &Cid, rev: &str, ) -> Result { - let ops = serde_json::json!([]); - let blobs: Vec = vec![]; - let blocks_cids: Vec = vec![mst_root_cid.to_string(), commit_cid.to_string()]; - let prev_cid: Option<&str> = None; - let commit_cid_str = commit_cid.to_string(); - let mut tx = state - .db - .begin() - .await - .map_err(|e| format!("Failed to begin transaction: {}", e))?; - let seq_row = sqlx::query!( - r#" - INSERT INTO repo_seq (did, event_type, commit_cid, prev_cid, ops, blobs, blocks_cids, rev) - VALUES ($1, 'commit', $2, $3::TEXT, $4, $5, $6, $7) - RETURNING seq - "#, - did.as_str(), - commit_cid_str, - prev_cid, - ops, - &blobs, - &blocks_cids, - rev - ) - .fetch_one(&mut *tx) - .await - .map_err(|e| format!("DB Error (repo_seq genesis commit): {}", e))?; - sqlx::query(&format!("NOTIFY repo_updates, '{}'", seq_row.seq)) - .execute(&mut *tx) - .await - .map_err(|e| format!("DB Error (notify): {}", e))?; - tx.commit() + let commit_cid_link = crate::types::CidLink::new_unchecked(&commit_cid.to_string()); + let mst_root_cid_link = crate::types::CidLink::new_unchecked(&mst_root_cid.to_string()); + state + .repo_repo + .insert_genesis_commit_event(did, &commit_cid_link, &mst_root_cid_link, rev) .await - .map_err(|e| format!("Failed to commit transaction: {}", e))?; - Ok(seq_row.seq) + .map_err(|e| format!("DB Error (genesis commit event): {}", e)) } diff --git a/crates/tranquil-pds/src/api/repo/record/write.rs b/crates/tranquil-pds/src/api/repo/record/write.rs index 5d4ca17..ae3c6c8 100644 --- a/crates/tranquil-pds/src/api/repo/record/write.rs +++ b/crates/tranquil-pds/src/api/repo/record/write.rs @@ -1,7 +1,7 @@ use super::validation::validate_record_with_status; use crate::api::error::ApiError; -use crate::api::repo::record::utils::{CommitParams, RecordOp, commit_and_log, extract_blob_cids}; -use crate::delegation::{self, DelegationActionType}; +use crate::api::repo::record::utils::{CommitParams, RecordOp, commit_and_log, extract_backlinks, extract_blob_cids}; +use crate::delegation::DelegationActionType; use crate::repo::tracking::TrackingBlockStore; use crate::state::AppState; use crate::types::{AtIdentifier, AtUri, Did, Nsid, Rkey}; @@ -15,39 +15,11 @@ use cid::Cid; use jacquard_repo::{commit::Commit, mst::Mst, storage::BlockStore}; use serde::{Deserialize, Serialize}; use serde_json::json; -use sqlx::{PgPool, Row}; use std::str::FromStr; use std::sync::Arc; use tracing::error; use uuid::Uuid; -pub async fn has_verified_comms_channel(db: &PgPool, did: &Did) -> Result { - let row = sqlx::query( - r#" - SELECT - email_verified, - discord_verified, - telegram_verified, - signal_verified - FROM users - WHERE did = $1 - "#, - ) - .bind(did.as_str()) - .fetch_optional(db) - .await?; - match row { - Some(r) => { - let email_verified: bool = r.get("email_verified"); - let discord_verified: bool = r.get("discord_verified"); - let telegram_verified: bool = r.get("telegram_verified"); - let signal_verified: bool = r.get("signal_verified"); - Ok(email_verified || discord_verified || telegram_verified || signal_verified) - } - None => Ok(false), - } -} - pub struct RepoWriteAuth { pub did: Did, pub user_id: Uuid, @@ -70,7 +42,8 @@ pub async fn prepare_repo_write( .ok_or_else(|| ApiError::AuthenticationRequired.into_response())?; let dpop_proof = headers.get("DPoP").and_then(|h| h.to_str().ok()); let auth_user = crate::auth::validate_token_with_dpop( - &state.db, + state.user_repo.as_ref(), + state.oauth_repo.as_ref(), &extracted.token, extracted.is_dpop, dpop_proof, @@ -89,40 +62,45 @@ pub async fn prepare_repo_write( ApiError::InvalidRepo("Repo does not match authenticated user".into()).into_response(), ); } - if crate::util::is_account_migrated(&state.db, &auth_user.did) + if state + .user_repo + .is_account_migrated(&auth_user.did) .await .unwrap_or(false) { return Err(ApiError::AccountMigrated.into_response()); } - let is_verified = has_verified_comms_channel(&state.db, &auth_user.did) + let is_verified = state + .user_repo + .has_verified_comms_channel(&auth_user.did) .await .unwrap_or(false); - let is_delegated = crate::delegation::is_delegated_account(&state.db, &auth_user.did) + let is_delegated = state + .delegation_repo + .is_delegated_account(&auth_user.did) .await .unwrap_or(false); if !is_verified && !is_delegated { return Err(ApiError::AccountNotVerified.into_response()); } - let user_id = sqlx::query_scalar!("SELECT id FROM users WHERE did = $1", &auth_user.did) - .fetch_optional(&state.db) + let user_id = state + .user_repo + .get_id_by_did(&auth_user.did) .await .map_err(|e| { error!("DB error fetching user: {}", e); ApiError::InternalError(None).into_response() })? .ok_or_else(|| ApiError::InternalError(Some("User not found".into())).into_response())?; - let root_cid_str: String = sqlx::query_scalar!( - "SELECT repo_root_cid FROM repos WHERE user_id = $1", - user_id - ) - .fetch_optional(&state.db) - .await - .map_err(|e| { - error!("DB error fetching repo root: {}", e); - ApiError::InternalError(None).into_response() - })? - .ok_or_else(|| ApiError::InternalError(Some("Repo root not found".into())).into_response())?; + let root_cid_str = state + .repo_repo + .get_repo_root_cid_by_user_id(user_id) + .await + .map_err(|e| { + error!("DB error fetching repo root: {}", e); + ApiError::InternalError(None).into_response() + })? + .ok_or_else(|| ApiError::InternalError(Some("Repo root not found".into())).into_response())?; let current_root_cid = Cid::from_str(&root_cid_str).map_err(|_| { ApiError::InternalError(Some("Invalid repo root CID".into())).into_response() })?; @@ -200,16 +178,7 @@ pub async fn create_record( { return ApiError::InvalidSwap(Some("Repo has been modified".into())).into_response(); } - let tracking_store = TrackingBlockStore::new(state.block_store.clone()); - let commit_bytes = match tracking_store.get(¤t_root_cid).await { - Ok(Some(b)) => b, - _ => return ApiError::InternalError(Some("Commit block not found".into())).into_response(), - }; - let commit = match Commit::from_cbor(&commit_bytes) { - Ok(c) => c, - _ => return ApiError::InternalError(Some("Failed to parse commit".into())).into_response(), - }; - let mst = Mst::load(Arc::new(tracking_store.clone()), commit.data, None); + let validation_status = if input.validate == Some(false) { None } else { @@ -225,6 +194,79 @@ pub async fn create_record( } }; let rkey = input.rkey.unwrap_or_else(Rkey::generate); + + let tracking_store = TrackingBlockStore::new(state.block_store.clone()); + let commit_bytes = match tracking_store.get(¤t_root_cid).await { + Ok(Some(b)) => b, + _ => return ApiError::InternalError(Some("Commit block not found".into())).into_response(), + }; + let commit = match Commit::from_cbor(&commit_bytes) { + Ok(c) => c, + _ => return ApiError::InternalError(Some("Failed to parse commit".into())).into_response(), + }; + let mut mst = Mst::load(Arc::new(tracking_store.clone()), commit.data, None); + let initial_mst_root = commit.data; + + let mut ops: Vec = Vec::new(); + let mut conflict_uris_to_cleanup: Vec = Vec::new(); + let mut all_old_mst_blocks = std::collections::BTreeMap::new(); + + if input.validate != Some(false) { + let record_uri = AtUri::from_parts(&did, &input.collection, &rkey); + let backlinks = extract_backlinks(&record_uri, &input.record); + + if !backlinks.is_empty() { + let conflicts = match state + .backlink_repo + .get_backlink_conflicts(user_id, &input.collection, &backlinks) + .await + { + Ok(c) => c, + Err(e) => { + error!("Failed to check backlink conflicts: {}", e); + return ApiError::InternalError(None).into_response(); + } + }; + + for conflict_uri in conflicts { + let conflict_rkey = match conflict_uri.rkey() { + Some(r) => Rkey::from(r.to_string()), + None => continue, + }; + let conflict_collection = match conflict_uri.collection() { + Some(c) => Nsid::from(c.to_string()), + None => continue, + }; + let conflict_key = format!("{}/{}", conflict_collection, conflict_rkey); + + let prev_cid = match mst.get(&conflict_key).await { + Ok(Some(cid)) => cid, + Ok(None) => continue, + Err(_) => continue, + }; + + if mst.blocks_for_path(&conflict_key, &mut all_old_mst_blocks).await.is_err() { + error!("Failed to get old MST blocks for conflict {}", conflict_uri); + } + + mst = match mst.delete(&conflict_key).await { + Ok(m) => m, + Err(e) => { + error!("Failed to delete conflict from MST {}: {:?}", conflict_uri, e); + continue; + } + }; + + ops.push(RecordOp::Delete { + collection: conflict_collection, + rkey: conflict_rkey, + prev: Some(prev_cid), + }); + conflict_uris_to_cleanup.push(conflict_uri); + } + } + } + let record_ipld = crate::util::json_to_ipld(&input.record); let mut record_bytes = Vec::new(); if serde_ipld_dagcbor::to_writer(&mut record_bytes, &record_ipld).is_err() { @@ -238,6 +280,11 @@ pub async fn create_record( } }; let key = format!("{}/{}", input.collection, rkey); + + if mst.blocks_for_path(&key, &mut all_old_mst_blocks).await.is_err() { + error!("Failed to get old MST blocks for new record path"); + } + let new_mst = match mst.add(&key, record_cid).await { Ok(m) => m, _ => return ApiError::InternalError(Some("Failed to add to MST".into())).into_response(), @@ -246,13 +293,14 @@ pub async fn create_record( Ok(c) => c, _ => return ApiError::InternalError(Some("Failed to persist MST".into())).into_response(), }; - let op = RecordOp::Create { + + ops.push(RecordOp::Create { collection: input.collection.clone(), rkey: rkey.clone(), cid: record_cid, - }; + }); + let mut new_mst_blocks = std::collections::BTreeMap::new(); - let mut old_mst_blocks = std::collections::BTreeMap::new(); if new_mst .blocks_for_path(&key, &mut new_mst_blocks) .await @@ -261,17 +309,10 @@ pub async fn create_record( return ApiError::InternalError(Some("Failed to get new MST blocks for path".into())) .into_response(); } - if mst - .blocks_for_path(&key, &mut old_mst_blocks) - .await - .is_err() - { - return ApiError::InternalError(Some("Failed to get old MST blocks for path".into())) - .into_response(); - } + let mut relevant_blocks = new_mst_blocks.clone(); - relevant_blocks.extend(old_mst_blocks.iter().map(|(k, v)| (*k, v.clone()))); - relevant_blocks.insert(record_cid, bytes::Bytes::from(record_bytes)); + relevant_blocks.extend(all_old_mst_blocks.iter().map(|(k, v)| (*k, v.clone()))); + relevant_blocks.insert(record_cid, bytes::Bytes::new()); let written_cids: Vec = tracking_store .get_all_relevant_cids() .into_iter() @@ -283,21 +324,22 @@ pub async fn create_record( let blob_cids = extract_blob_cids(&input.record); let obsolete_cids: Vec = std::iter::once(current_root_cid) .chain( - old_mst_blocks + all_old_mst_blocks .keys() .filter(|cid| !new_mst_blocks.contains_key(*cid)) .copied(), ) .collect(); + let commit_result = match commit_and_log( &state, CommitParams { did: &did, user_id, current_root_cid: Some(current_root_cid), - prev_data_cid: Some(commit.data), + prev_data_cid: Some(initial_mst_root), new_mst_root, - ops: vec![op], + ops, blocks_cids: &written_cids_str, blobs: &blob_cids, obsolete_cids, @@ -312,28 +354,47 @@ pub async fn create_record( Err(e) => return ApiError::InternalError(Some(e)).into_response(), }; + for conflict_uri in conflict_uris_to_cleanup { + if let Err(e) = state.backlink_repo.remove_backlinks_by_uri(&conflict_uri).await { + error!("Failed to remove backlinks for {}: {}", conflict_uri, e); + } + } + if let Some(ref controller) = controller_did { - let _ = delegation::log_delegation_action( - &state.db, - &did, - controller, - Some(controller), - DelegationActionType::RepoWrite, - Some(json!({ - "action": "create", - "collection": input.collection, - "rkey": rkey - })), - None, - None, - ) - .await; + let _ = state + .delegation_repo + .log_delegation_action( + &did, + controller, + Some(controller), + DelegationActionType::RepoWrite, + Some(json!({ + "action": "create", + "collection": input.collection, + "rkey": rkey + })), + None, + None, + ) + .await; + } + + let created_uri = AtUri::from_parts(&did, &input.collection, &rkey); + let backlinks = extract_backlinks(&created_uri, &input.record); + if !backlinks.is_empty() { + if let Err(e) = state + .backlink_repo + .add_backlinks(user_id, &backlinks) + .await + { + error!("Failed to add backlinks for {}: {}", created_uri, e); + } } ( StatusCode::OK, Json(CreateRecordOutput { - uri: AtUri::from_parts(&did, &input.collection, &rkey), + uri: created_uri, cid: record_cid.to_string(), commit: CommitInfo { cid: commit_result.commit_cid.to_string(), @@ -574,21 +635,22 @@ pub async fn put_record( }; if let Some(ref controller) = controller_did { - let _ = delegation::log_delegation_action( - &state.db, - &did, - controller, - Some(controller), - DelegationActionType::RepoWrite, - Some(json!({ - "action": if is_update { "update" } else { "create" }, - "collection": input.collection, - "rkey": input.rkey - })), - None, - None, - ) - .await; + let _ = state + .delegation_repo + .log_delegation_action( + &did, + controller, + Some(controller), + DelegationActionType::RepoWrite, + Some(json!({ + "action": if is_update { "update" } else { "create" }, + "collection": input.collection, + "rkey": input.rkey + })), + None, + None, + ) + .await; } ( diff --git a/crates/tranquil-pds/src/api/server/account_status.rs b/crates/tranquil-pds/src/api/server/account_status.rs index 9775d50..3ce1c78 100644 --- a/crates/tranquil-pds/src/api/server/account_status.rs +++ b/crates/tranquil-pds/src/api/server/account_status.rs @@ -3,7 +3,7 @@ use crate::api::error::ApiError; use crate::cache::Cache; use crate::plc::PlcClient; use crate::state::AppState; -use crate::types::{Handle, PlainPassword}; +use crate::types::PlainPassword; use axum::{ Json, extract::State, @@ -54,7 +54,8 @@ pub async fn check_account_status( std::env::var("PDS_HOSTNAME").unwrap_or_else(|_| "localhost".to_string()) ); let did = match crate::auth::validate_token_with_dpop( - &state.db, + state.user_repo.as_ref(), + state.oauth_repo.as_ref(), &extracted.token, extracted.is_dpop, dpop_proof, @@ -68,43 +69,28 @@ pub async fn check_account_status( Ok(user) => user.did, Err(e) => return ApiError::from(e).into_response(), }; - let user_id = match sqlx::query_scalar!("SELECT id FROM users WHERE did = $1", did.as_str()) - .fetch_optional(&state.db) - .await - { + let user_id = match state.user_repo.get_id_by_did(&did).await { Ok(Some(id)) => id, _ => { return ApiError::InternalError(None).into_response(); } }; - let user_status = sqlx::query!( - "SELECT deactivated_at FROM users WHERE did = $1", - did.as_str() - ) - .fetch_optional(&state.db) - .await; - let deactivated_at = match user_status { - Ok(Some(row)) => row.deactivated_at, - _ => None, - }; - let repo_result = sqlx::query!( - "SELECT repo_root_cid, repo_rev FROM repos WHERE user_id = $1", - user_id - ) - .fetch_optional(&state.db) - .await; - let (repo_commit, repo_rev_from_db) = match repo_result { - Ok(Some(row)) => (row.repo_root_cid, row.repo_rev), - _ => (String::new(), None), - }; - let block_count: i64 = sqlx::query_scalar!( - "SELECT COUNT(*) FROM user_blocks WHERE user_id = $1", - user_id - ) - .fetch_one(&state.db) - .await - .unwrap_or(Some(0)) - .unwrap_or(0); + let is_active = state + .user_repo + .is_account_active_by_did(&did) + .await + .ok() + .flatten() + .unwrap_or(false); + let repo_info = state.repo_repo.get_repo(user_id).await.ok().flatten(); + let (repo_commit, repo_rev_from_db) = repo_info + .map(|r| (r.repo_root_cid.to_string(), r.repo_rev)) + .unwrap_or_else(|| (String::new(), None)); + let block_count: i64 = state + .repo_repo + .count_user_blocks(user_id) + .await + .unwrap_or(0); let repo_rev = if let Some(rev) = repo_rev_from_db { rev } else if !repo_commit.is_empty() { @@ -123,37 +109,22 @@ pub async fn check_account_status( } else { String::new() }; - let record_count: i64 = - sqlx::query_scalar!("SELECT COUNT(*) FROM records WHERE repo_id = $1", user_id) - .fetch_one(&state.db) - .await - .unwrap_or(Some(0)) - .unwrap_or(0); - let imported_blobs: i64 = sqlx::query_scalar!( - "SELECT COUNT(*) FROM blobs WHERE created_by_user = $1", - user_id - ) - .fetch_one(&state.db) - .await - .unwrap_or(Some(0)) - .unwrap_or(0); - let expected_blobs: i64 = sqlx::query_scalar!( - "SELECT COUNT(DISTINCT blob_cid) FROM record_blobs WHERE repo_id = $1", - user_id - ) - .fetch_one(&state.db) - .await - .unwrap_or(Some(0)) - .unwrap_or(0); - let valid_did = is_valid_did_for_service(&state.db, state.cache.clone(), did.as_str()).await; + let record_count: i64 = state.repo_repo.count_records(user_id).await.unwrap_or(0); + let imported_blobs: i64 = state.blob_repo.count_blobs_by_user(user_id).await.unwrap_or(0); + let expected_blobs: i64 = state + .blob_repo + .count_distinct_record_blobs(user_id) + .await + .unwrap_or(0); + let valid_did = is_valid_did_for_service(state.user_repo.as_ref(), state.cache.clone(), &did).await; ( StatusCode::OK, Json(CheckAccountStatusOutput { - activated: deactivated_at.is_none(), + activated: is_active, valid_did, repo_commit: repo_commit.clone(), repo_rev, - repo_blocks: block_count as i64, + repo_blocks: block_count, indexed_records: record_count, private_state_values: 0, expected_blobs, @@ -163,25 +134,25 @@ pub async fn check_account_status( .into_response() } -async fn is_valid_did_for_service(db: &sqlx::PgPool, cache: Arc, did: &str) -> bool { - assert_valid_did_document_for_service(db, cache, did, false) +async fn is_valid_did_for_service(user_repo: &dyn tranquil_db_traits::UserRepository, cache: Arc, did: &crate::types::Did) -> bool { + assert_valid_did_document_for_service(user_repo, cache, did, false) .await .is_ok() } async fn assert_valid_did_document_for_service( - db: &sqlx::PgPool, + user_repo: &dyn tranquil_db_traits::UserRepository, cache: Arc, - did: &str, + did: &crate::types::Did, with_retry: bool, ) -> Result<(), ApiError> { let hostname = std::env::var("PDS_HOSTNAME").unwrap_or_else(|_| "localhost".to_string()); let expected_endpoint = format!("https://{}", hostname); - if did.starts_with("did:plc:") { + if did.as_str().starts_with("did:plc:") { let max_attempts = if with_retry { 5 } else { 1 }; let cache_for_retry = cache.clone(); - let did_owned = did.to_string(); + let did_owned = did.as_str().to_string(); let expected_owned = expected_endpoint.clone(); let attempt_counter = Arc::new(AtomicUsize::new(0)); @@ -264,19 +235,16 @@ async fn assert_valid_did_document_for_service( .and_then(|v| v.get("atproto")) .and_then(|k| k.as_str()); - let user_row = sqlx::query!( - "SELECT uk.key_bytes, uk.encryption_version FROM user_keys uk JOIN users u ON uk.user_id = u.id WHERE u.did = $1", - did - ) - .fetch_optional(db) - .await - .map_err(|e| { - error!("Failed to fetch user key: {:?}", e); - ApiError::InternalError(None) - })?; + let user_key = user_repo + .get_user_key_by_did(&did) + .await + .map_err(|e| { + error!("Failed to fetch user key: {:?}", e); + ApiError::InternalError(None) + })?; - if let Some(row) = user_row { - let key_bytes = crate::config::decrypt_key(&row.key_bytes, row.encryption_version) + if let Some(key_info) = user_key { + let key_bytes = crate::config::decrypt_key(&key_info.key_bytes, key_info.encryption_version) .map_err(|e| { error!("Failed to decrypt user key: {}", e); ApiError::InternalError(None) @@ -297,7 +265,7 @@ async fn assert_valid_did_document_for_service( )); } } - } else if let Some(host_and_path) = did.strip_prefix("did:web:") { + } else if let Some(host_and_path) = did.as_str().strip_prefix("did:web:") { let client = crate::api::proxy_client::did_resolution_client(); let decoded = host_and_path.replace("%3A", ":"); let parts: Vec<&str> = decoded.split(':').collect(); @@ -374,7 +342,8 @@ pub async fn activate_account( std::env::var("PDS_HOSTNAME").unwrap_or_else(|_| "localhost".to_string()) ); let auth_user = match crate::auth::validate_token_with_dpop( - &state.db, + state.user_repo.as_ref(), + state.oauth_repo.as_ref(), &extracted.token, extracted.is_dpop, dpop_proof, @@ -414,7 +383,7 @@ pub async fn activate_account( ); let did_validation_start = std::time::Instant::now(); if let Err(e) = - assert_valid_did_document_for_service(&state.db, state.cache.clone(), did.as_str(), true) + assert_valid_did_document_for_service(state.user_repo.as_ref(), state.cache.clone(), &did, true) .await { info!( @@ -430,8 +399,9 @@ pub async fn activate_account( did_validation_start.elapsed() ); - let handle = sqlx::query_scalar!("SELECT handle FROM users WHERE did = $1", did.as_str()) - .fetch_optional(&state.db) + let handle = state + .user_repo + .get_handle_by_did(&did) .await .ok() .flatten(); @@ -439,12 +409,7 @@ pub async fn activate_account( "[MIGRATION] activateAccount: Activating account did={} handle={:?}", did, handle ); - let result = sqlx::query!( - "UPDATE users SET deactivated_at = NULL WHERE did = $1", - did.as_str() - ) - .execute(&state.db) - .await; + let result = state.user_repo.activate_account(&did).await; match result { Ok(_) => { info!( @@ -472,7 +437,7 @@ pub async fn activate_account( "[MIGRATION] activateAccount: Sequencing identity event for did={} handle={:?}", did, handle ); - let handle_typed = handle.as_ref().map(Handle::new_unchecked); + let handle_typed = handle.clone(); if let Err(e) = crate::api::repo::record::sequence_identity_event( &state, &did, @@ -487,20 +452,18 @@ pub async fn activate_account( } else { info!("[MIGRATION] activateAccount: Identity event sequenced successfully"); } - let repo_root = sqlx::query_scalar!( - "SELECT r.repo_root_cid FROM repos r JOIN users u ON r.user_id = u.id WHERE u.did = $1", - did.as_str() - ) - .fetch_optional(&state.db) - .await - .ok() - .flatten(); - if let Some(root_cid) = repo_root { + let repo_root = state + .repo_repo + .get_repo_root_by_did(&did) + .await + .ok() + .flatten(); + if let Some(root_cid_link) = repo_root { info!( "[MIGRATION] activateAccount: Sequencing sync event for did={} root_cid={}", - did, root_cid + did, root_cid_link ); - let rev = if let Ok(cid) = Cid::from_str(&root_cid) { + let rev = if let Ok(cid) = Cid::from_str(root_cid_link.as_str()) { if let Ok(Some(block)) = state.block_store.get(&cid).await { Commit::from_cbor(&block).ok().map(|c| c.rev().to_string()) } else { @@ -512,7 +475,7 @@ pub async fn activate_account( if let Err(e) = crate::api::repo::record::sequence_sync_event( &state, &did, - &root_cid, + root_cid_link.as_str(), rev.as_deref(), ) .await @@ -566,7 +529,8 @@ pub async fn deactivate_account( std::env::var("PDS_HOSTNAME").unwrap_or_else(|_| "localhost".to_string()) ); let auth_user = match crate::auth::validate_token_with_dpop( - &state.db, + state.user_repo.as_ref(), + state.oauth_repo.as_ref(), &extracted.token, extracted.is_dpop, dpop_proof, @@ -598,22 +562,20 @@ pub async fn deactivate_account( let did = auth_user.did; - let handle = sqlx::query_scalar!("SELECT handle FROM users WHERE did = $1", did.as_str()) - .fetch_optional(&state.db) + let handle = state + .user_repo + .get_handle_by_did(&did) .await .ok() .flatten(); - let result = sqlx::query!( - "UPDATE users SET deactivated_at = NOW(), delete_after = $2 WHERE did = $1", - did.as_str(), - delete_after - ) - .execute(&state.db) - .await; + let result = state + .user_repo + .deactivate_account(&did, delete_after) + .await; match result { - Ok(_) => { + Ok(true) => { if let Some(ref h) = handle { let _ = state.cache.delete(&format!("handle:{}", h)).await; } @@ -629,6 +591,9 @@ pub async fn deactivate_account( } EmptyResponse::ok().into_response() } + Ok(false) => { + EmptyResponse::ok().into_response() + } Err(e) => { error!("DB error deactivating account: {:?}", e); ApiError::InternalError(None).into_response() @@ -652,7 +617,8 @@ pub async fn request_account_delete( std::env::var("PDS_HOSTNAME").unwrap_or_else(|_| "localhost".to_string()) ); let validated = match crate::auth::validate_token_with_dpop( - &state.db, + state.user_repo.as_ref(), + state.oauth_repo.as_ref(), &extracted.token, extracted.is_dpop, dpop_proof, @@ -668,15 +634,12 @@ pub async fn request_account_delete( }; let did = validated.did.clone(); - if !crate::api::server::reauth::check_legacy_session_mfa(&state.db, did.as_str()).await { - return crate::api::server::reauth::legacy_mfa_required_response(&state.db, did.as_str()) + if !crate::api::server::reauth::check_legacy_session_mfa(&*state.session_repo, &did).await { + return crate::api::server::reauth::legacy_mfa_required_response(&*state.user_repo, &*state.session_repo, &did) .await; } - let user_id = match sqlx::query_scalar!("SELECT id FROM users WHERE did = $1", did.as_str()) - .fetch_optional(&state.db) - .await - { + let user_id = match state.user_repo.get_id_by_did(&did).await { Ok(Some(id)) => id, _ => { return ApiError::InternalError(None).into_response(); @@ -684,22 +647,23 @@ pub async fn request_account_delete( }; let confirmation_token = Uuid::new_v4().to_string(); let expires_at = Utc::now() + Duration::minutes(15); - let insert = sqlx::query!( - "INSERT INTO account_deletion_requests (token, did, expires_at) VALUES ($1, $2, $3)", - confirmation_token, - did.as_str(), - expires_at - ) - .execute(&state.db) - .await; - if let Err(e) = insert { + if let Err(e) = state + .infra_repo + .create_deletion_request(&confirmation_token, &did, expires_at) + .await + { error!("DB error creating deletion token: {:?}", e); return ApiError::InternalError(None).into_response(); } let hostname = std::env::var("PDS_HOSTNAME").unwrap_or_else(|_| "localhost".to_string()); - if let Err(e) = - crate::comms::enqueue_account_deletion(&state.db, user_id, &confirmation_token, &hostname) - .await + if let Err(e) = crate::comms::comms_repo::enqueue_account_deletion( + state.user_repo.as_ref(), + state.infra_repo.as_ref(), + user_id, + &confirmation_token, + &hostname, + ) + .await { warn!("Failed to enqueue account deletion notification: {:?}", e); } @@ -731,14 +695,8 @@ pub async fn delete_account( if token.is_empty() { return ApiError::InvalidToken(Some("token is required".into())).into_response(); } - let user = sqlx::query!( - "SELECT id, password_hash, handle FROM users WHERE did = $1", - did.as_str() - ) - .fetch_optional(&state.db) - .await; - let (user_id, password_hash, handle) = match user { - Ok(Some(row)) => (row.id, row.password_hash, row.handle), + let user = match state.user_repo.get_user_for_deletion(did).await { + Ok(Some(u)) => u, Ok(None) => { return ApiError::InvalidRequest("account not found".into()).into_response(); } @@ -747,6 +705,7 @@ pub async fn delete_account( return ApiError::InternalError(None).into_response(); } }; + let (user_id, password_hash, handle) = (user.id, user.password_hash, user.handle); let password_valid = if password_hash .as_ref() .map(|h| verify(password, h).unwrap_or(false)) @@ -754,28 +713,20 @@ pub async fn delete_account( { true } else { - let app_pass_rows = sqlx::query!( - "SELECT password_hash FROM app_passwords WHERE user_id = $1", - user_id - ) - .fetch_all(&state.db) - .await - .unwrap_or_default(); - app_pass_rows + let app_pass_hashes = state + .session_repo + .get_app_password_hashes_by_did(did) + .await + .unwrap_or_default(); + app_pass_hashes .iter() - .any(|row| verify(password, &row.password_hash).unwrap_or(false)) + .any(|h| verify(password, h).unwrap_or(false)) }; if !password_valid { return ApiError::AuthenticationFailed(Some("Invalid password".into())).into_response(); } - let deletion_request = sqlx::query!( - "SELECT did, expires_at FROM account_deletion_requests WHERE token = $1", - token - ) - .fetch_optional(&state.db) - .await; - let (token_did, expires_at) = match deletion_request { - Ok(Some(row)) => (row.did, row.expires_at), + let deletion_request = match state.infra_repo.get_deletion_request(token).await { + Ok(Some(req)) => req, Ok(None) => { return ApiError::InvalidToken(Some("Invalid or expired token".into())).into_response(); } @@ -784,96 +735,49 @@ pub async fn delete_account( return ApiError::InternalError(None).into_response(); } }; - if token_did != did.as_str() { + if &deletion_request.did != did { return ApiError::InvalidToken(Some("Token does not match account".into())).into_response(); } - if Utc::now() > expires_at { - let _ = sqlx::query!( - "DELETE FROM account_deletion_requests WHERE token = $1", - token - ) - .execute(&state.db) - .await; + if Utc::now() > deletion_request.expires_at { + let _ = state.infra_repo.delete_deletion_request(token).await; return ApiError::ExpiredToken(None).into_response(); } - let mut tx = match state.db.begin().await { - Ok(tx) => tx, - Err(e) => { - error!("Failed to begin transaction: {:?}", e); - return ApiError::InternalError(None).into_response(); - } - }; - let deletion_result: Result<(), sqlx::Error> = async { - sqlx::query!("DELETE FROM session_tokens WHERE did = $1", did) - .execute(&mut *tx) - .await?; - sqlx::query!("DELETE FROM records WHERE repo_id = $1", user_id) - .execute(&mut *tx) - .await?; - sqlx::query!("DELETE FROM repos WHERE user_id = $1", user_id) - .execute(&mut *tx) - .await?; - sqlx::query!("DELETE FROM blobs WHERE created_by_user = $1", user_id) - .execute(&mut *tx) - .await?; - sqlx::query!("DELETE FROM user_keys WHERE user_id = $1", user_id) - .execute(&mut *tx) - .await?; - sqlx::query!("DELETE FROM app_passwords WHERE user_id = $1", user_id) - .execute(&mut *tx) - .await?; - sqlx::query!("DELETE FROM account_deletion_requests WHERE did = $1", did) - .execute(&mut *tx) - .await?; - sqlx::query!("DELETE FROM users WHERE id = $1", user_id) - .execute(&mut *tx) - .await?; - Ok(()) + if let Err(e) = state + .user_repo + .delete_account_complete(user_id, did) + .await + { + error!("DB error deleting account: {:?}", e); + return ApiError::InternalError(None).into_response(); } + let account_seq = crate::api::repo::record::sequence_account_event( + &state, + did, + false, + Some("deleted"), + ) .await; - match deletion_result { - Ok(()) => { - if let Err(e) = tx.commit().await { - error!("Failed to commit account deletion transaction: {:?}", e); - return ApiError::InternalError(None).into_response(); - } - let account_seq = crate::api::repo::record::sequence_account_event( - &state, - did, - false, - Some("deleted"), - ) - .await; - match account_seq { - Ok(seq) => { - if let Err(e) = sqlx::query!( - "DELETE FROM repo_seq WHERE did = $1 AND seq != $2", - did, - seq - ) - .execute(&state.db) - .await - { - warn!( - "Failed to cleanup sequences for deleted account {}: {}", - did, e - ); - } - } - Err(e) => { - warn!( - "Failed to sequence account deletion event for {}: {}", - did, e - ); - } + match account_seq { + Ok(seq) => { + if let Err(e) = state + .repo_repo + .delete_sequences_except(did, seq) + .await + { + warn!( + "Failed to cleanup sequences for deleted account {}: {}", + did, e + ); } - let _ = state.cache.delete(&format!("handle:{}", handle)).await; - info!("Account {} deleted successfully", did); - EmptyResponse::ok().into_response() } Err(e) => { - error!("DB error deleting account, rolling back: {:?}", e); - ApiError::InternalError(None).into_response() + warn!( + "Failed to sequence account deletion event for {}: {}", + did, e + ); } } + let _ = state.cache.delete(&format!("handle:{}", handle)).await; + info!("Account {} deleted successfully", did); + EmptyResponse::ok().into_response() } diff --git a/crates/tranquil-pds/src/api/server/app_password.rs b/crates/tranquil-pds/src/api/server/app_password.rs index f1d883f..c949233 100644 --- a/crates/tranquil-pds/src/api/server/app_password.rs +++ b/crates/tranquil-pds/src/api/server/app_password.rs @@ -1,9 +1,8 @@ use crate::api::EmptyResponse; use crate::api::error::ApiError; use crate::auth::BearerAuth; -use crate::delegation::{self, DelegationActionType}; +use crate::delegation::{DelegationActionType, intersect_scopes}; use crate::state::{AppState, RateLimitKind}; -use crate::util::get_user_id_by_did; use axum::{ Json, extract::State, @@ -12,6 +11,7 @@ use axum::{ }; use serde::{Deserialize, Serialize}; use serde_json::json; +use tranquil_db_traits::AppPasswordCreate; use tracing::{error, warn}; #[derive(Serialize)] @@ -35,17 +35,16 @@ pub async fn list_app_passwords( State(state): State, BearerAuth(auth_user): BearerAuth, ) -> Response { - let user_id = match get_user_id_by_did(&state.db, &auth_user.did).await { - Ok(id) => id, - Err(e) => return ApiError::from(e).into_response(), + let user = match state.user_repo.get_by_did(&auth_user.did).await { + Ok(Some(u)) => u, + Ok(None) => return ApiError::AccountNotFound.into_response(), + Err(e) => { + error!("DB error getting user: {:?}", e); + return ApiError::InternalError(None).into_response(); + } }; - match sqlx::query!( - "SELECT name, created_at, privileged, scopes, created_by_controller_did FROM app_passwords WHERE user_id = $1 ORDER BY created_at DESC", - user_id - ) - .fetch_all(&state.db) - .await - { + + match state.session_repo.list_app_passwords(user.id).await { Ok(rows) => { let passwords: Vec = rows .iter() @@ -54,7 +53,7 @@ pub async fn list_app_passwords( created_at: row.created_at.to_rfc3339(), privileged: row.privileged, scopes: row.scopes.clone(), - created_by_controller: row.created_by_controller_did.clone(), + created_by_controller: row.created_by_controller_did.as_ref().map(|d| d.to_string()), }) .collect(); Json(ListAppPasswordsOutput { passwords }).into_response() @@ -98,34 +97,41 @@ pub async fn create_app_password( warn!(ip = %client_ip, "App password creation rate limit exceeded"); return ApiError::RateLimitExceeded(None).into_response(); } - let user_id = match get_user_id_by_did(&state.db, &auth_user.did).await { - Ok(id) => id, - Err(e) => return ApiError::from(e).into_response(), + + let user = match state.user_repo.get_by_did(&auth_user.did).await { + Ok(Some(u)) => u, + Ok(None) => return ApiError::AccountNotFound.into_response(), + Err(e) => { + error!("DB error getting user: {:?}", e); + return ApiError::InternalError(None).into_response(); + } }; + let name = input.name.trim(); if name.is_empty() { return ApiError::InvalidRequest("name is required".into()).into_response(); } - let existing = sqlx::query!( - "SELECT id FROM app_passwords WHERE user_id = $1 AND name = $2", - user_id, - name - ) - .fetch_optional(&state.db) - .await; - if let Ok(Some(_)) = existing { - return ApiError::DuplicateAppPassword.into_response(); + + match state.session_repo.get_app_password_by_name(user.id, name).await { + Ok(Some(_)) => return ApiError::DuplicateAppPassword.into_response(), + Err(e) => { + error!("DB error checking app password: {:?}", e); + return ApiError::InternalError(None).into_response(); + } + Ok(None) => {} } let (final_scopes, controller_did) = if let Some(ref controller) = auth_user.controller_did { - let grant = delegation::get_delegation(&state.db, &auth_user.did, controller) + let grant = state + .delegation_repo + .get_delegation(&auth_user.did, controller) .await .ok() .flatten(); let granted_scopes = grant.map(|g| g.granted_scopes).unwrap_or_default(); let requested = input.scopes.as_deref().unwrap_or("atproto"); - let intersected = delegation::intersect_scopes(requested, &granted_scopes); + let intersected = intersect_scopes(requested, &granted_scopes); if intersected.is_empty() && !granted_scopes.is_empty() { return ApiError::InsufficientScope(None).into_response(); @@ -152,6 +158,7 @@ pub async fn create_app_password( }) .collect::>() .join("-"); + let password_clone = password.clone(); let password_hash = match tokio::task::spawn_blocking(move || { bcrypt::hash(&password_clone, bcrypt::DEFAULT_COST) @@ -168,38 +175,38 @@ pub async fn create_app_password( return ApiError::InternalError(None).into_response(); } }; + let privileged = input.privileged.unwrap_or(false); let created_at = chrono::Utc::now(); - match sqlx::query!( - "INSERT INTO app_passwords (user_id, name, password_hash, created_at, privileged, scopes, created_by_controller_did) VALUES ($1, $2, $3, $4, $5, $6, $7)", - user_id, - name, + + let create_data = AppPasswordCreate { + user_id: user.id, + name: name.to_string(), password_hash, - created_at, privileged, - final_scopes, - controller_did.as_deref() - ) - .execute(&state.db) - .await - { + scopes: final_scopes.clone(), + created_by_controller_did: controller_did.clone(), + }; + + match state.session_repo.create_app_password(&create_data).await { Ok(_) => { if let Some(ref controller) = controller_did { - let _ = delegation::log_delegation_action( - &state.db, - &auth_user.did, - controller, - Some(controller), - DelegationActionType::AccountAction, - Some(json!({ - "action": "create_app_password", - "name": name, - "scopes": final_scopes - })), - None, - None, - ) - .await; + let _ = state + .delegation_repo + .log_delegation_action( + &auth_user.did, + controller, + Some(controller), + DelegationActionType::AccountAction, + Some(json!({ + "action": "create_app_password", + "name": name, + "scopes": final_scopes + })), + None, + None, + ) + .await; } Json(CreateAppPasswordOutput { name: name.to_string(), @@ -227,33 +234,35 @@ pub async fn revoke_app_password( BearerAuth(auth_user): BearerAuth, Json(input): Json, ) -> Response { - let user_id = match get_user_id_by_did(&state.db, &auth_user.did).await { - Ok(id) => id, - Err(e) => return ApiError::from(e).into_response(), + let user = match state.user_repo.get_by_did(&auth_user.did).await { + Ok(Some(u)) => u, + Ok(None) => return ApiError::AccountNotFound.into_response(), + Err(e) => { + error!("DB error getting user: {:?}", e); + return ApiError::InternalError(None).into_response(); + } }; + let name = input.name.trim(); if name.is_empty() { return ApiError::InvalidRequest("name is required".into()).into_response(); } - let sessions_to_invalidate = sqlx::query_scalar!( - "SELECT access_jti FROM session_tokens WHERE did = $1 AND app_password_name = $2", - &auth_user.did, - name - ) - .fetch_all(&state.db) - .await - .unwrap_or_default(); - if let Err(e) = sqlx::query!( - "DELETE FROM session_tokens WHERE did = $1 AND app_password_name = $2", - &auth_user.did, - name - ) - .execute(&state.db) - .await + + let sessions_to_invalidate = state + .session_repo + .get_session_jtis_by_app_password(&auth_user.did, name) + .await + .unwrap_or_default(); + + if let Err(e) = state + .session_repo + .delete_sessions_by_app_password(&auth_user.did, name) + .await { error!("DB error revoking sessions for app password: {:?}", e); return ApiError::InternalError(None).into_response(); } + futures::future::join_all(sessions_to_invalidate.iter().map(|jti| { let cache_key = format!("auth:session:{}:{}", &auth_user.did, jti); let cache = state.cache.clone(); @@ -262,16 +271,11 @@ pub async fn revoke_app_password( } })) .await; - if let Err(e) = sqlx::query!( - "DELETE FROM app_passwords WHERE user_id = $1 AND name = $2", - user_id, - name - ) - .execute(&state.db) - .await - { + + if let Err(e) = state.session_repo.delete_app_password(user.id, name).await { error!("DB error revoking app password: {:?}", e); return ApiError::InternalError(None).into_response(); } + EmptyResponse::ok().into_response() } diff --git a/crates/tranquil-pds/src/api/server/email.rs b/crates/tranquil-pds/src/api/server/email.rs index b86d6d1..afada3b 100644 --- a/crates/tranquil-pds/src/api/server/email.rs +++ b/crates/tranquil-pds/src/api/server/email.rs @@ -34,14 +34,7 @@ pub async fn request_email_update( return e; } - let did = auth.0.did.to_string(); - let user = match sqlx::query!( - "SELECT id, handle, email, email_verified FROM users WHERE did = $1", - did - ) - .fetch_optional(&state.db) - .await - { + let user = match state.user_repo.get_email_info_by_did(&auth.0.did).await { Ok(Some(row)) => row, Ok(None) => { return ApiError::AccountNotFound.into_response(); @@ -61,16 +54,21 @@ pub async fn request_email_update( if token_required { let code = crate::auth::verification_token::generate_channel_update_token( - &did, + &auth.0.did, "email_update", ¤t_email.to_lowercase(), ); let formatted_code = crate::auth::verification_token::format_token_for_display(&code); let hostname = std::env::var("PDS_HOSTNAME").unwrap_or_else(|_| "localhost".to_string()); - if let Err(e) = - crate::comms::enqueue_email_update_token(&state.db, user.id, &formatted_code, &hostname) - .await + if let Err(e) = crate::comms::comms_repo::enqueue_email_update_token( + state.user_repo.as_ref(), + state.infra_repo.as_ref(), + user.id, + &formatted_code, + &hostname, + ) + .await { warn!("Failed to enqueue email update notification: {:?}", e); } @@ -111,14 +109,8 @@ pub async fn confirm_email( return e; } - let did = auth.0.did.to_string(); - let user = match sqlx::query!( - "SELECT id, email, email_verified FROM users WHERE did = $1", - did - ) - .fetch_optional(&state.db) - .await - { + let did = &auth.0.did; + let user = match state.user_repo.get_email_info_by_did(did).await { Ok(Some(row)) => row, Ok(None) => { return ApiError::AccountNotFound.into_response(); @@ -154,7 +146,7 @@ pub async fn confirm_email( match verified { Ok(token_data) => { - if token_data.did != did { + if token_data.did != did.as_str() { return ApiError::InvalidToken(None).into_response(); } } @@ -166,14 +158,7 @@ pub async fn confirm_email( } } - let update = sqlx::query!( - "UPDATE users SET email_verified = TRUE, updated_at = NOW() WHERE id = $1", - user.id - ) - .execute(&state.db) - .await; - - if let Err(e) = update { + if let Err(e) = state.user_repo.set_email_verified(user.id, true).await { error!("DB error confirming email: {:?}", e); return ApiError::InternalError(None).into_response(); } @@ -207,14 +192,8 @@ pub async fn update_email( return e; } - let did = auth_user.did.to_string(); - let user = match sqlx::query!( - "SELECT id, email, email_verified FROM users WHERE did = $1", - did - ) - .fetch_optional(&state.db) - .await - { + let did = &auth_user.did; + let user = match state.user_repo.get_email_info_by_did(did).await { Ok(Some(row)) => row, Ok(None) => { return ApiError::AccountNotFound.into_response(); @@ -262,7 +241,7 @@ pub async fn update_email( match verified { Ok(token_data) => { - if token_data.did != did { + if token_data.did != did.as_str() { return ApiError::InvalidToken(None).into_response(); } } @@ -275,34 +254,12 @@ pub async fn update_email( } } - let exists = sqlx::query!( - "SELECT 1 as one FROM users WHERE LOWER(email) = $1 AND id != $2", - new_email, - user_id - ) - .fetch_optional(&state.db) - .await; - - if let Ok(Some(_)) = exists { + if let Ok(true) = state.user_repo.check_email_exists(&new_email, user_id).await { return ApiError::InvalidRequest("Email is already in use".into()).into_response(); } - let update: Result = sqlx::query!( - "UPDATE users SET email = $1, email_verified = FALSE, updated_at = NOW() WHERE id = $2", - new_email, - user_id - ) - .execute(&state.db) - .await; - - if let Err(e) = update { + if let Err(e) = state.user_repo.update_email(user_id, &new_email).await { error!("DB error updating email: {:?}", e); - if e.as_database_error() - .map(|db_err: &dyn sqlx::error::DatabaseError| db_err.is_unique_violation()) - .unwrap_or(false) - { - return ApiError::EmailTaken.into_response(); - } return ApiError::InternalError(None).into_response(); } @@ -310,29 +267,30 @@ pub async fn update_email( crate::auth::verification_token::generate_signup_token(&did, "email", &new_email); let formatted_token = crate::auth::verification_token::format_token_for_display(&verification_token); - if let Err(e) = crate::comms::enqueue_signup_verification( - &state.db, + let hostname = std::env::var("PDS_HOSTNAME").unwrap_or_else(|_| "localhost".to_string()); + if let Err(e) = crate::comms::comms_repo::enqueue_signup_verification( + state.infra_repo.as_ref(), user_id, "email", &new_email, &formatted_token, - None, + &hostname, ) .await { warn!("Failed to send verification email to new address: {:?}", e); } - match sqlx::query!( - "INSERT INTO account_preferences (user_id, name, value_json) VALUES ($1, 'email_auth_factor', $2) ON CONFLICT (user_id, name) DO UPDATE SET value_json = $2", - user_id, - json!(input.email_auth_factor.unwrap_or(false)) - ) - .execute(&state.db) - .await + if let Err(e) = state + .infra_repo + .upsert_account_preference( + user_id, + "email_auth_factor", + json!(input.email_auth_factor.unwrap_or(false)), + ) + .await { - Ok(_) => {} - Err(e) => warn!("Failed to update email_auth_factor preference: {}", e), + warn!("Failed to update email_auth_factor preference: {}", e); } info!("Email updated for user {}", user_id); @@ -357,15 +315,12 @@ pub async fn check_email_verified( return ApiError::RateLimitExceeded(None).into_response(); } - let user = sqlx::query!( - "SELECT email_verified FROM users WHERE email = $1 OR handle = $1", - input.identifier - ) - .fetch_optional(&state.db) - .await; - - match user { - Ok(Some(row)) => VerifiedResponse::response(row.email_verified).into_response(), + match state + .user_repo + .check_email_verified_by_identifier(&input.identifier) + .await + { + Ok(Some(verified)) => VerifiedResponse::response(verified).into_response(), Ok(None) => ApiError::AccountNotFound.into_response(), Err(e) => { error!("DB error checking email verified: {:?}", e); diff --git a/crates/tranquil-pds/src/api/server/invite.rs b/crates/tranquil-pds/src/api/server/invite.rs index 7a6160b..5c252c0 100644 --- a/crates/tranquil-pds/src/api/server/invite.rs +++ b/crates/tranquil-pds/src/api/server/invite.rs @@ -2,6 +2,7 @@ use crate::api::ApiError; use crate::auth::BearerAuth; use crate::auth::extractor::BearerAuthAdmin; use crate::state::AppState; +use crate::types::Did; use axum::{ Json, extract::State, @@ -50,27 +51,24 @@ pub async fn create_invite_code( return ApiError::InvalidRequest("useCount must be at least 1".into()).into_response(); } - let for_account = input - .for_account - .unwrap_or_else(|| auth_user.did.to_string()); + let for_account: Did = match &input.for_account { + Some(acct) => match acct.parse() { + Ok(d) => d, + Err(_) => return ApiError::InvalidDid("Invalid DID format".into()).into_response(), + }, + None => auth_user.did.clone(), + }; let code = gen_invite_code(); - match sqlx::query!( - "INSERT INTO invite_codes (code, available_uses, created_by_user, for_account) - SELECT $1, $2, id, $3 FROM users WHERE is_admin = true LIMIT 1", - code, - input.use_count, - for_account - ) - .execute(&state.db) - .await + match state + .infra_repo + .create_invite_code(&code, input.use_count, Some(&for_account)) + .await { - Ok(result) => { - if result.rows_affected() == 0 { - error!("No admin user found to create invite code"); - return ApiError::InternalError(None).into_response(); - } - Json(CreateInviteCodeOutput { code }).into_response() + Ok(true) => Json(CreateInviteCodeOutput { code }).into_response(), + Ok(false) => { + error!("No admin user found to create invite code"); + ApiError::InternalError(None).into_response() } Err(e) => { error!("DB error creating invite code: {:?}", e); @@ -108,45 +106,38 @@ pub async fn create_invite_codes( } let code_count = input.code_count.unwrap_or(1).max(1); - let for_accounts = input - .for_accounts - .filter(|v| !v.is_empty()) - .unwrap_or_else(|| vec![auth_user.did.to_string()]); - - let admin_user_id = - match sqlx::query_scalar!("SELECT id FROM users WHERE is_admin = true LIMIT 1") - .fetch_optional(&state.db) - .await - { - Ok(Some(id)) => id, - Ok(None) => { - error!("No admin user found to create invite codes"); - return ApiError::InternalError(None).into_response(); - } - Err(e) => { - error!("DB error looking up admin user: {:?}", e); - return ApiError::InternalError(None).into_response(); + let for_accounts: Vec = match &input.for_accounts { + Some(accounts) if !accounts.is_empty() => { + let parsed: Result, _> = accounts.iter().map(|a| a.parse()).collect(); + match parsed { + Ok(dids) => dids, + Err(_) => return ApiError::InvalidDid("Invalid DID format".into()).into_response(), } - }; + } + _ => vec![auth_user.did.clone()], + }; + + let admin_user_id = match state.user_repo.get_any_admin_user_id().await { + Ok(Some(id)) => id, + Ok(None) => { + error!("No admin user found to create invite codes"); + return ApiError::InternalError(None).into_response(); + } + Err(e) => { + error!("DB error looking up admin user: {:?}", e); + return ApiError::InternalError(None).into_response(); + } + }; let result = futures::future::try_join_all(for_accounts.into_iter().map(|account| { - let db = state.db.clone(); + let infra_repo = state.infra_repo.clone(); let use_count = input.use_count; async move { let codes: Vec = (0..code_count).map(|_| gen_invite_code()).collect(); - sqlx::query!( - r#" - INSERT INTO invite_codes (code, available_uses, created_by_user, for_account) - SELECT code, $2, $3, $4 FROM UNNEST($1::text[]) AS t(code) - "#, - &codes[..], - use_count, - admin_user_id, - account - ) - .execute(&db) - .await - .map(|_| AccountCodes { account, codes }) + infra_repo + .create_invite_codes_batch(&codes, use_count, admin_user_id, Some(&account)) + .await + .map(|_| AccountCodes { account: account.to_string(), codes }) } })) .await; @@ -203,78 +194,59 @@ pub async fn get_account_invite_codes( ) -> Response { let include_used = params.include_used.unwrap_or(true); - let codes_rows = match sqlx::query!( - r#" - SELECT - ic.code, - ic.available_uses, - ic.created_at, - ic.disabled, - ic.for_account, - (SELECT COUNT(*) FROM invite_code_uses icu WHERE icu.code = ic.code)::int as "use_count!" - FROM invite_codes ic - WHERE ic.for_account = $1 - ORDER BY ic.created_at DESC - "#, - &auth_user.did - ) - .fetch_all(&state.db) - .await + let codes_info = match state + .infra_repo + .get_invite_codes_for_account(&auth_user.did) + .await { - Ok(rows) => rows, + Ok(info) => info, Err(e) => { error!("DB error fetching invite codes: {:?}", e); return ApiError::InternalError(None).into_response(); } }; - let filtered_rows: Vec<_> = codes_rows + let filtered_codes: Vec<_> = codes_info .into_iter() - .filter(|row| { - let disabled = row.disabled.unwrap_or(false); - !disabled && (include_used || row.use_count < row.available_uses) - }) + .filter(|info| !info.disabled) .collect(); - let codes = futures::future::join_all(filtered_rows.into_iter().map(|row| { - let db = state.db.clone(); + let codes = futures::future::join_all(filtered_codes.into_iter().map(|info| { + let infra_repo = state.infra_repo.clone(); async move { - let uses = sqlx::query!( - r#" - SELECT u.did, u.handle, icu.used_at - FROM invite_code_uses icu - JOIN users u ON icu.used_by_user = u.id - WHERE icu.code = $1 - ORDER BY icu.used_at DESC - "#, - row.code - ) - .fetch_all(&db) - .await - .map(|use_rows| { - use_rows - .iter() - .map(|u| InviteCodeUse { - used_by: u.did.clone(), - used_by_handle: Some(u.handle.clone()), - used_at: u.used_at.to_rfc3339(), - }) - .collect() - }) - .unwrap_or_default(); + let uses = infra_repo + .get_invite_code_uses(&info.code) + .await + .map(|use_rows| { + use_rows + .into_iter() + .map(|u| InviteCodeUse { + used_by: u.used_by_did.to_string(), + used_by_handle: u.used_by_handle.map(|h| h.to_string()), + used_at: u.used_at.to_rfc3339(), + }) + .collect::>() + }) + .unwrap_or_default(); + + let use_count = uses.len() as i32; + if !include_used && use_count >= info.available_uses { + return None; + } - InviteCode { - code: row.code, - available: row.available_uses, + Some(InviteCode { + code: info.code, + available: info.available_uses, disabled: false, - for_account: row.for_account, - created_by: "admin".to_string(), - created_at: row.created_at.to_rfc3339(), + for_account: info.for_account.map(|d| d.to_string()).unwrap_or_default(), + created_by: info.created_by.map(|d| d.to_string()).unwrap_or_else(|| "admin".to_string()), + created_at: info.created_at.to_rfc3339(), uses, - } + }) } })) .await; + let codes: Vec = codes.into_iter().flatten().collect(); Json(GetAccountInviteCodesOutput { codes }).into_response() } diff --git a/crates/tranquil-pds/src/api/server/logo.rs b/crates/tranquil-pds/src/api/server/logo.rs index 91b02b4..e086d7b 100644 --- a/crates/tranquil-pds/src/api/server/logo.rs +++ b/crates/tranquil-pds/src/api/server/logo.rs @@ -9,31 +9,22 @@ use axum::{ use tracing::error; pub async fn get_logo(State(state): State) -> Response { - let logo_cid: Option = - match sqlx::query_scalar("SELECT value FROM server_config WHERE key = 'logo_cid'") - .fetch_optional(&state.db) - .await - { - Ok(cid) => cid, - Err(e) => { - error!("DB error fetching logo_cid: {:?}", e); - return StatusCode::INTERNAL_SERVER_ERROR.into_response(); - } - }; + let logo_cid = match state.infra_repo.get_server_config("logo_cid").await { + Ok(cid) => cid, + Err(e) => { + error!("DB error fetching logo_cid: {:?}", e); + return StatusCode::INTERNAL_SERVER_ERROR.into_response(); + } + }; - let cid = match logo_cid { + let cid_str = match logo_cid { Some(c) if !c.is_empty() => c, _ => return StatusCode::NOT_FOUND.into_response(), }; + let cid = crate::types::CidLink::new_unchecked(&cid_str); - let blob = match sqlx::query!( - "SELECT storage_key, mime_type FROM blobs WHERE cid = $1", - cid - ) - .fetch_optional(&state.db) - .await - { - Ok(Some(row)) => row, + let metadata = match state.blob_repo.get_blob_metadata(&cid).await { + Ok(Some(m)) => m, Ok(None) => return StatusCode::NOT_FOUND.into_response(), Err(e) => { error!("DB error fetching blob: {:?}", e); @@ -41,10 +32,10 @@ pub async fn get_logo(State(state): State) -> Response { } }; - match state.blob_store.get(&blob.storage_key).await { + match state.blob_store.get(&metadata.storage_key).await { Ok(data) => Response::builder() .status(StatusCode::OK) - .header(header::CONTENT_TYPE, &blob.mime_type) + .header(header::CONTENT_TYPE, &metadata.mime_type) .header(header::CACHE_CONTROL, "public, max-age=3600") .body(Body::from(data)) .unwrap(), diff --git a/crates/tranquil-pds/src/api/server/meta.rs b/crates/tranquil-pds/src/api/server/meta.rs index b1b6f6b..2a97ea7 100644 --- a/crates/tranquil-pds/src/api/server/meta.rs +++ b/crates/tranquil-pds/src/api/server/meta.rs @@ -1,7 +1,6 @@ use crate::state::AppState; use axum::{Json, extract::State, http::StatusCode, response::IntoResponse}; use serde_json::json; -use tracing::error; fn get_available_comms_channels() -> Vec<&'static str> { let mut channels = vec!["email"]; @@ -64,11 +63,8 @@ pub async fn describe_server() -> impl IntoResponse { })) } pub async fn health(State(state): State) -> impl IntoResponse { - match sqlx::query!("SELECT 1 as one").fetch_one(&state.db).await { - Ok(_) => (StatusCode::OK, "OK"), - Err(e) => { - error!("Health check failed: {:?}", e); - (StatusCode::SERVICE_UNAVAILABLE, "Service Unavailable") - } + match state.infra_repo.health_check().await { + Ok(true) => (StatusCode::OK, "OK"), + _ => (StatusCode::SERVICE_UNAVAILABLE, "Service Unavailable"), } } diff --git a/crates/tranquil-pds/src/api/server/migration.rs b/crates/tranquil-pds/src/api/server/migration.rs index 72e72bd..c4ebf49 100644 --- a/crates/tranquil-pds/src/api/server/migration.rs +++ b/crates/tranquil-pds/src/api/server/migration.rs @@ -6,7 +6,6 @@ use axum::{ http::StatusCode, response::{IntoResponse, Response}, }; -use chrono::Utc; use serde::{Deserialize, Serialize}; use serde_json::json; @@ -51,7 +50,8 @@ pub async fn update_did_document( std::env::var("PDS_HOSTNAME").unwrap_or_else(|_| "localhost".to_string()) ); let auth_user = match crate::auth::validate_token_with_dpop( - &state.db, + state.user_repo.as_ref(), + state.oauth_repo.as_ref(), &extracted.token, extracted.is_dpop, dpop_proof, @@ -73,14 +73,8 @@ pub async fn update_did_document( .into_response(); } - let user = match sqlx::query!( - "SELECT id, handle, deactivated_at FROM users WHERE did = $1", - &auth_user.did - ) - .fetch_optional(&state.db) - .await - { - Ok(Some(row)) => row, + let user = match state.user_repo.get_user_for_did_doc(&auth_user.did).await { + Ok(Some(u)) => u, Ok(None) => return ApiError::AccountNotFound.into_response(), Err(e) => { tracing::error!("DB error getting user: {:?}", e); @@ -137,48 +131,28 @@ pub async fn update_did_document( let also_known_as: Option> = input.also_known_as.clone(); - let now = Utc::now(); - - let upsert_result = sqlx::query!( - r#" - INSERT INTO did_web_overrides (user_id, verification_methods, also_known_as, updated_at) - VALUES ($1, COALESCE($2, '[]'::jsonb), COALESCE($3, '{}'::text[]), $4) - ON CONFLICT (user_id) DO UPDATE SET - verification_methods = CASE WHEN $2 IS NOT NULL THEN $2 ELSE did_web_overrides.verification_methods END, - also_known_as = CASE WHEN $3 IS NOT NULL THEN $3 ELSE did_web_overrides.also_known_as END, - updated_at = $4 - "#, - user.id, - verification_methods_json, - also_known_as.as_deref(), - now - ) - .execute(&state.db) - .await; - - if let Err(e) = upsert_result { + if let Err(e) = state + .user_repo + .upsert_did_web_overrides(user.id, verification_methods_json, also_known_as) + .await + { tracing::error!("DB error upserting did_web_overrides: {:?}", e); return ApiError::InternalError(None).into_response(); } if let Some(ref endpoint) = input.service_endpoint { let endpoint_clean = endpoint.trim().trim_end_matches('/'); - let update_result = sqlx::query!( - "UPDATE users SET migrated_to_pds = $1, migrated_at = $2 WHERE did = $3", - endpoint_clean, - now, - &auth_user.did - ) - .execute(&state.db) - .await; - - if let Err(e) = update_result { + if let Err(e) = state + .user_repo + .update_migrated_to_pds(&auth_user.did, endpoint_clean) + .await + { tracing::error!("DB error updating service endpoint: {:?}", e); return ApiError::InternalError(None).into_response(); } } - let did_doc = build_did_document(&state.db, &auth_user.did).await; + let did_doc = build_did_document(&state, &auth_user.did).await; tracing::info!("Updated DID document for {}", &auth_user.did); @@ -208,7 +182,8 @@ pub async fn get_did_document( std::env::var("PDS_HOSTNAME").unwrap_or_else(|_| "localhost".to_string()) ); let auth_user = match crate::auth::validate_token_with_dpop( - &state.db, + state.user_repo.as_ref(), + state.oauth_repo.as_ref(), &extracted.token, extracted.is_dpop, dpop_proof, @@ -230,21 +205,15 @@ pub async fn get_did_document( .into_response(); } - let did_doc = build_did_document(&state.db, &auth_user.did).await; + let did_doc = build_did_document(&state, &auth_user.did).await; (StatusCode::OK, Json(json!({ "didDocument": did_doc }))).into_response() } -async fn build_did_document(db: &sqlx::PgPool, did: &str) -> serde_json::Value { +async fn build_did_document(state: &AppState, did: &crate::types::Did) -> serde_json::Value { let hostname = std::env::var("PDS_HOSTNAME").unwrap_or_else(|_| "localhost".to_string()); - let user = match sqlx::query!( - "SELECT id, handle, migrated_to_pds FROM users WHERE did = $1", - did - ) - .fetch_optional(db) - .await - { + let user = match state.user_repo.get_user_for_did_doc_build(did).await { Ok(Some(row)) => row, _ => { return json!({ @@ -253,14 +222,12 @@ async fn build_did_document(db: &sqlx::PgPool, did: &str) -> serde_json::Value { } }; - let overrides = sqlx::query!( - "SELECT verification_methods, also_known_as FROM did_web_overrides WHERE user_id = $1", - user.id - ) - .fetch_optional(db) - .await - .ok() - .flatten(); + let overrides = state + .user_repo + .get_did_web_overrides(user.id) + .await + .ok() + .flatten(); let service_endpoint = user .migrated_to_pds @@ -299,20 +266,15 @@ async fn build_did_document(db: &sqlx::PgPool, did: &str) -> serde_json::Value { }); } - let key_row = sqlx::query!( - "SELECT key_bytes, encryption_version FROM user_keys WHERE user_id = $1", - user.id - ) - .fetch_optional(db) - .await; + let key_info = state.user_repo.get_user_key_by_id(user.id).await.ok().flatten(); - let public_key_multibase = match key_row { - Ok(Some(row)) => match crate::config::decrypt_key(&row.key_bytes, row.encryption_version) { + let public_key_multibase = match key_info { + Some(info) => match crate::config::decrypt_key(&info.key_bytes, info.encryption_version) { Ok(key_bytes) => crate::api::identity::did::get_public_key_multibase(&key_bytes) .unwrap_or_else(|_| "error".to_string()), Err(_) => "error".to_string(), }, - _ => "error".to_string(), + None => "error".to_string(), }; let also_known_as = if let Some(ref ovr) = overrides { diff --git a/crates/tranquil-pds/src/api/server/mod.rs b/crates/tranquil-pds/src/api/server/mod.rs index 5f9e81c..9f07a05 100644 --- a/crates/tranquil-pds/src/api/server/mod.rs +++ b/crates/tranquil-pds/src/api/server/mod.rs @@ -32,8 +32,8 @@ pub use passkey_account::{ request_passkey_recovery, start_passkey_registration_for_setup, }; pub use passkeys::{ - delete_passkey, finish_passkey_registration, has_passkeys_for_user, has_passkeys_for_user_db, - list_passkeys, start_passkey_registration, update_passkey, + delete_passkey, finish_passkey_registration, has_passkeys_for_user, list_passkeys, + start_passkey_registration, update_passkey, }; pub use password::{ change_password, get_password_status, remove_password, request_password_reset, reset_password, @@ -53,7 +53,7 @@ pub use session::{ pub use signing_key::reserve_signing_key; pub use totp::{ create_totp_secret, disable_totp, enable_totp, get_totp_status, has_totp_enabled, - has_totp_enabled_db, regenerate_backup_codes, verify_totp_or_backup_for_user, + regenerate_backup_codes, verify_totp_or_backup_for_user, }; pub use trusted_devices::{ extend_device_trust, is_device_trusted, list_trusted_devices, revoke_trusted_device, diff --git a/crates/tranquil-pds/src/api/server/passkey_account.rs b/crates/tranquil-pds/src/api/server/passkey_account.rs index d75f1d5..6d321d2 100644 --- a/crates/tranquil-pds/src/api/server/passkey_account.rs +++ b/crates/tranquil-pds/src/api/server/passkey_account.rs @@ -149,7 +149,8 @@ pub async fn create_passkey_account( .unwrap_or(false); let hostname = std::env::var("PDS_HOSTNAME").unwrap_or_else(|_| "localhost".to_string()); - let pds_suffix = format!(".{}", hostname); + let hostname_for_handles = hostname.split(':').next().unwrap_or(&hostname); + let pds_suffix = format!(".{}", hostname_for_handles); let handle = if !input.handle.contains('.') || input.handle.ends_with(&pds_suffix) { let handle_to_validate = if input.handle.ends_with(&pds_suffix) { @@ -161,7 +162,7 @@ pub async fn create_passkey_account( &input.handle }; match crate::api::validation::validate_short_handle(handle_to_validate) { - Ok(h) => format!("{}.{}", h, hostname), + Ok(h) => format!("{}.{}", h, hostname_for_handles), Err(_) => { return ApiError::InvalidHandle(None).into_response(); } @@ -182,17 +183,9 @@ pub async fn create_passkey_account( } if let Some(ref code) = input.invite_code { - let valid = sqlx::query_scalar!( - "SELECT available_uses > 0 AND NOT disabled FROM invite_codes WHERE code = $1", - code - ) - .fetch_optional(&state.db) - .await - .ok() - .flatten() - .unwrap_or(Some(false)); + let valid = state.infra_repo.is_invite_code_valid(code).await.unwrap_or(false); - if valid != Some(true) { + if !valid { return ApiError::InvalidInviteCode.into_response(); } } else { @@ -233,21 +226,8 @@ pub async fn create_passkey_account( let (secret_key_bytes, reserved_key_id): (Vec, Option) = if let Some(signing_key_did) = &input.signing_key { - let reserved = sqlx::query!( - r#" - SELECT id, private_key_bytes - FROM reserved_signing_keys - WHERE public_key_did_key = $1 - AND used_at IS NULL - AND expires_at > NOW() - FOR UPDATE - "#, - signing_key_did - ) - .fetch_optional(&state.db) - .await; - match reserved { - Ok(Some(row)) => (row.private_key_bytes, Some(row.id)), + match state.infra_repo.get_reserved_signing_key(signing_key_did).await { + Ok(Some(reserved)) => (reserved.private_key_bytes, Some(reserved.id)), Ok(None) => { return ApiError::InvalidSigningKey.into_response(); } @@ -271,7 +251,7 @@ pub async fn create_passkey_account( let did = match did_type { "web" => { - let subdomain_host = format!("{}.{}", input.handle, hostname); + let subdomain_host = format!("{}.{}", input.handle, hostname_for_handles); let encoded_subdomain = subdomain_host.replace(':', "%3A"); let self_hosted_did = format!("did:web:{}", encoded_subdomain); info!(did = %self_hosted_did, "Creating self-hosted did:web passkey account"); @@ -391,85 +371,12 @@ pub async fn create_passkey_account( }; let setup_expires_at = Utc::now() + Duration::hours(1); - let mut tx = match state.db.begin().await { - Ok(tx) => tx, - Err(e) => { - error!("Error starting transaction: {:?}", e); - return ApiError::InternalError(None).into_response(); - } - }; - - let is_first_user = sqlx::query_scalar!("SELECT COUNT(*) as count FROM users") - .fetch_one(&mut *tx) - .await - .map(|c| c.unwrap_or(0) == 0) - .unwrap_or(false); - let deactivated_at: Option> = if is_byod_did_web { Some(Utc::now()) } else { None }; - let user_insert: Result<(Uuid,), _> = sqlx::query_as( - r#"INSERT INTO users ( - handle, email, did, password_hash, password_required, - preferred_comms_channel, - discord_id, telegram_username, signal_number, - recovery_token, recovery_token_expires_at, - is_admin, deactivated_at - ) VALUES ($1, $2, $3, NULL, FALSE, $4::comms_channel, $5, $6, $7, $8, $9, $10, $11) RETURNING id"#, - ) - .bind(&handle) - .bind(&email) - .bind(&did) - .bind(verification_channel) - .bind( - input - .discord_id - .as_deref() - .map(|s| s.trim()) - .filter(|s| !s.is_empty()), - ) - .bind( - input - .telegram_username - .as_deref() - .map(|s| s.trim()) - .filter(|s| !s.is_empty()), - ) - .bind( - input - .signal_number - .as_deref() - .map(|s| s.trim()) - .filter(|s| !s.is_empty()), - ) - .bind(&setup_token_hash) - .bind(setup_expires_at) - .bind(is_first_user) - .bind(deactivated_at) - .fetch_one(&mut *tx) - .await; - - let user_id = match user_insert { - Ok((id,)) => id, - Err(e) => { - if let Some(db_err) = e.as_database_error() - && db_err.code().as_deref() == Some("23505") - { - let constraint = db_err.constraint().unwrap_or(""); - if constraint.contains("handle") { - return ApiError::HandleNotAvailable(None).into_response(); - } else if constraint.contains("email") { - return ApiError::EmailTaken.into_response(); - } - } - error!("Error inserting user: {:?}", e); - return ApiError::InternalError(None).into_response(); - } - }; - let encrypted_key_bytes = match crate::config::encrypt_key(&secret_key_bytes) { Ok(bytes) => bytes, Err(e) => { @@ -478,31 +385,6 @@ pub async fn create_passkey_account( } }; - if let Err(e) = sqlx::query!( - "INSERT INTO user_keys (user_id, key_bytes, encryption_version, encrypted_at) VALUES ($1, $2, $3, NOW())", - user_id, - &encrypted_key_bytes[..], - crate::config::ENCRYPTION_VERSION - ) - .execute(&mut *tx) - .await - { - error!("Error inserting user key: {:?}", e); - return ApiError::InternalError(None).into_response(); - } - - if let Some(key_id) = reserved_key_id - && let Err(e) = sqlx::query!( - "UPDATE reserved_signing_keys SET used_at = NOW() WHERE id = $1", - key_id - ) - .execute(&mut *tx) - .await - { - error!("Error marking reserved key as used: {:?}", e); - return ApiError::InternalError(None).into_response(); - } - let mst = Mst::new(Arc::new(state.block_store.clone())); let mst_root = match mst.persist().await { Ok(c) => c, @@ -528,80 +410,61 @@ pub async fn create_passkey_account( return ApiError::InternalError(None).into_response(); } }; - let commit_cid_str = commit_cid.to_string(); - let rev_str = rev.as_ref().to_string(); - if let Err(e) = sqlx::query!( - "INSERT INTO repos (user_id, repo_root_cid, repo_rev) VALUES ($1, $2, $3)", - user_id, - commit_cid_str, - rev_str - ) - .execute(&mut *tx) - .await - { - error!("Error inserting repo: {:?}", e); - return ApiError::InternalError(None).into_response(); - } let genesis_block_cids = vec![mst_root.to_bytes(), commit_cid.to_bytes()]; - if let Err(e) = sqlx::query!( - r#" - INSERT INTO user_blocks (user_id, block_cid) - SELECT $1, block_cid FROM UNNEST($2::bytea[]) AS t(block_cid) - ON CONFLICT (user_id, block_cid) DO NOTHING - "#, - user_id, - &genesis_block_cids - ) - .execute(&mut *tx) - .await - { - error!("Error inserting user_blocks: {:?}", e); - return ApiError::InternalError(None).into_response(); - } - - if let Some(ref code) = input.invite_code { - let _ = sqlx::query!( - "UPDATE invite_codes SET available_uses = available_uses - 1 WHERE code = $1", - code - ) - .execute(&mut *tx) - .await; - - let _ = sqlx::query!( - "INSERT INTO invite_code_uses (code, used_by_user) VALUES ($1, $2)", - code, - user_id - ) - .execute(&mut *tx) - .await; - } - if std::env::var("PDS_AGE_ASSURANCE_OVERRIDE").is_ok() { - let birthdate_pref = json!({ + let birthdate_pref = std::env::var("PDS_AGE_ASSURANCE_OVERRIDE").ok().map(|_| { + json!({ "$type": "app.bsky.actor.defs#personalDetailsPref", "birthDate": "1998-05-06T00:00:00.000Z" - }); - if let Err(e) = sqlx::query!( - "INSERT INTO account_preferences (user_id, name, value_json) VALUES ($1, $2, $3) - ON CONFLICT (user_id, name) DO NOTHING", - user_id, - "app.bsky.actor.defs#personalDetailsPref", - birthdate_pref - ) - .execute(&mut *tx) - .await - { - warn!("Failed to set default birthdate preference: {:?}", e); - } - } + }) + }); + + let preferred_comms_channel = match verification_channel { + "email" => tranquil_db_traits::CommsChannel::Email, + "discord" => tranquil_db_traits::CommsChannel::Discord, + "telegram" => tranquil_db_traits::CommsChannel::Telegram, + "signal" => tranquil_db_traits::CommsChannel::Signal, + _ => tranquil_db_traits::CommsChannel::Email, + }; - if let Err(e) = tx.commit().await { - error!("Error committing transaction: {:?}", e); - return ApiError::InternalError(None).into_response(); - } + let handle_typed = Handle::new_unchecked(&handle); + let create_input = tranquil_db_traits::CreatePasskeyAccountInput { + handle: handle_typed.clone(), + email: email.clone().unwrap_or_default(), + did: did_typed.clone(), + preferred_comms_channel, + discord_id: input.discord_id.as_deref().map(|s| s.trim()).filter(|s| !s.is_empty()).map(String::from), + telegram_username: input.telegram_username.as_deref().map(|s| s.trim()).filter(|s| !s.is_empty()).map(String::from), + signal_number: input.signal_number.as_deref().map(|s| s.trim()).filter(|s| !s.is_empty()).map(String::from), + setup_token_hash, + setup_expires_at, + deactivated_at, + encrypted_key_bytes, + encryption_version: crate::config::ENCRYPTION_VERSION, + reserved_key_id, + commit_cid: commit_cid.to_string(), + repo_rev: rev.as_ref().to_string(), + genesis_block_cids, + invite_code: input.invite_code.clone(), + birthdate_pref, + }; + + let create_result = match state.user_repo.create_passkey_account(&create_input).await { + Ok(r) => r, + Err(tranquil_db_traits::CreateAccountError::HandleTaken) => { + return ApiError::HandleNotAvailable(None).into_response(); + } + Err(tranquil_db_traits::CreateAccountError::EmailTaken) => { + return ApiError::EmailTaken.into_response(); + } + Err(e) => { + error!("Error creating passkey account: {:?}", e); + return ApiError::InternalError(None).into_response(); + } + }; + let user_id = create_result.user_id; if !is_byod_did_web { - let handle_typed = Handle::new_unchecked(&handle); if let Err(e) = crate::api::repo::record::sequence_identity_event( &state, &did_typed, @@ -642,13 +505,13 @@ pub async fn create_passkey_account( ); let formatted_token = crate::auth::verification_token::format_token_for_display(&verification_token); - if let Err(e) = crate::comms::enqueue_signup_verification( - &state.db, + if let Err(e) = crate::comms::comms_repo::enqueue_signup_verification( + state.infra_repo.as_ref(), user_id, verification_channel, &verification_recipient, &formatted_token, - None, + &hostname, ) .await { @@ -662,21 +525,19 @@ pub async fn create_passkey_account( Ok(token_meta) => { let refresh_jti = uuid::Uuid::new_v4().to_string(); let refresh_expires = chrono::Utc::now() + chrono::Duration::hours(24); - let no_scope: Option = None; - if let Err(e) = sqlx::query!( - "INSERT INTO session_tokens (did, access_jti, refresh_jti, access_expires_at, refresh_expires_at, legacy_login, mfa_verified, scope) VALUES ($1, $2, $3, $4, $5, $6, $7, $8)", - did, - token_meta.jti, + let session_data = tranquil_db::SessionTokenCreate { + did: did_typed.clone(), + access_jti: token_meta.jti.clone(), refresh_jti, - token_meta.expires_at, - refresh_expires, - false, - false, - no_scope - ) - .execute(&state.db) - .await - { + access_expires_at: token_meta.expires_at, + refresh_expires_at: refresh_expires, + legacy_login: false, + mfa_verified: false, + scope: None, + controller_did: None, + app_password_name: None, + }; + if let Err(e) = state.session_repo.create_session(&session_data).await { warn!(did = %did, "Failed to insert migration session: {:?}", e); } info!(did = %did, "Generated migration access token for BYOD passkey account"); @@ -723,15 +584,7 @@ pub async fn complete_passkey_setup( State(state): State, Json(input): Json, ) -> Response { - let user = sqlx::query!( - r#"SELECT id, handle, recovery_token, recovery_token_expires_at, password_required - FROM users WHERE did = $1"#, - input.did.as_str() - ) - .fetch_optional(&state.db) - .await; - - let user = match user { + let user = match state.user_repo.get_user_for_passkey_setup(&input.did).await { Ok(Some(u)) => u, Ok(None) => { return ApiError::AccountNotFound.into_response(); @@ -772,17 +625,26 @@ pub async fn complete_passkey_setup( } }; - let reg_state = - match crate::auth::webauthn::load_registration_state(&state.db, &input.did).await { - Ok(Some(s)) => s, - Ok(None) => { - return ApiError::NoChallengeInProgress.into_response(); - } + let reg_state = match state + .user_repo + .load_webauthn_challenge(&input.did, "registration") + .await + { + Ok(Some(json)) => match serde_json::from_str(&json) { + Ok(s) => s, Err(e) => { - error!("Error loading registration state: {:?}", e); + error!("Error deserializing registration state: {:?}", e); return ApiError::InternalError(None).into_response(); } - }; + }, + Ok(None) => { + return ApiError::NoChallengeInProgress.into_response(); + } + Err(e) => { + error!("Error loading registration state: {:?}", e); + return ApiError::InternalError(None).into_response(); + } + }; let credential: webauthn_rs::prelude::RegisterPublicKeyCredential = match serde_json::from_value(input.passkey_credential) { @@ -801,13 +663,23 @@ pub async fn complete_passkey_setup( } }; - if let Err(e) = crate::auth::webauthn::save_passkey( - &state.db, - &input.did, - &security_key, - input.passkey_friendly_name.as_deref(), - ) - .await + let credential_id = security_key.cred_id().to_vec(); + let public_key = match serde_json::to_vec(&security_key) { + Ok(pk) => pk, + Err(e) => { + error!("Error serializing security key: {:?}", e); + return ApiError::InternalError(None).into_response(); + } + }; + if let Err(e) = state + .user_repo + .save_passkey( + &input.did, + &credential_id, + &public_key, + input.passkey_friendly_name.as_deref(), + ) + .await { error!("Error saving passkey: {:?}", e); return ApiError::InternalError(None).into_response(); @@ -823,44 +695,21 @@ pub async fn complete_passkey_setup( } }; - let mut tx = match state.db.begin().await { - Ok(tx) => tx, - Err(e) => { - error!("Failed to begin transaction: {:?}", e); - return ApiError::InternalError(None).into_response(); - } + let setup_input = tranquil_db_traits::CompletePasskeySetupInput { + user_id: user.id, + did: input.did.clone(), + app_password_name: app_password_name.clone(), + app_password_hash: password_hash, }; - - if let Err(e) = sqlx::query!( - "INSERT INTO app_passwords (user_id, name, password_hash, privileged) VALUES ($1, $2, $3, FALSE)", - user.id, - app_password_name, - password_hash - ) - .execute(&mut *tx) - .await - { - error!("Error creating app password: {:?}", e); - return ApiError::InternalError(None).into_response(); - } - - if let Err(e) = sqlx::query!( - "UPDATE users SET recovery_token = NULL, recovery_token_expires_at = NULL WHERE did = $1", - input.did.as_str() - ) - .execute(&mut *tx) - .await - { - error!("Error clearing setup token: {:?}", e); + if let Err(e) = state.user_repo.complete_passkey_setup(&setup_input).await { + error!("Error completing passkey setup: {:?}", e); return ApiError::InternalError(None).into_response(); } - if let Err(e) = tx.commit().await { - error!("Failed to commit setup transaction: {:?}", e); - return ApiError::InternalError(None).into_response(); - } - - let _ = crate::auth::webauthn::delete_registration_state(&state.db, &input.did).await; + let _ = state + .user_repo + .delete_webauthn_challenge(&input.did, "registration") + .await; info!(did = %input.did, "Passkey-only account setup completed"); @@ -877,15 +726,7 @@ pub async fn start_passkey_registration_for_setup( State(state): State, Json(input): Json, ) -> Response { - let user = sqlx::query!( - r#"SELECT handle, recovery_token, recovery_token_expires_at, password_required - FROM users WHERE did = $1"#, - input.did.as_str() - ) - .fetch_optional(&state.db) - .await; - - let user = match user { + let user = match state.user_repo.get_user_for_passkey_setup(&input.did).await { Ok(Some(u)) => u, Ok(None) => { return ApiError::AccountNotFound.into_response(); @@ -926,7 +767,9 @@ pub async fn start_passkey_registration_for_setup( } }; - let existing_passkeys = crate::auth::webauthn::get_passkeys_for_user(&state.db, &input.did) + let existing_passkeys = state + .user_repo + .get_passkeys_for_user(&input.did) .await .unwrap_or_default(); @@ -950,8 +793,17 @@ pub async fn start_passkey_registration_for_setup( } }; - if let Err(e) = - crate::auth::webauthn::save_registration_state(&state.db, &input.did, ®_state).await + let state_json = match serde_json::to_string(®_state) { + Ok(json) => json, + Err(e) => { + error!("Failed to serialize registration state: {:?}", e); + return ApiError::InternalError(None).into_response(); + } + }; + if let Err(e) = state + .user_repo + .save_webauthn_challenge(&input.did, "registration", &state_json) + .await { error!("Failed to save registration state: {:?}", e); return ApiError::InternalError(None).into_response(); @@ -990,23 +842,16 @@ pub async fn request_passkey_recovery( } let pds_hostname = std::env::var("PDS_HOSTNAME").unwrap_or_else(|_| "localhost".to_string()); + let hostname_for_handles = pds_hostname.split(':').next().unwrap_or(&pds_hostname); let identifier = input.email.trim().to_lowercase(); let identifier = identifier.strip_prefix('@').unwrap_or(&identifier); let normalized_handle = if identifier.contains('@') || identifier.contains('.') { identifier.to_string() } else { - format!("{}.{}", identifier, pds_hostname) + format!("{}.{}", identifier, hostname_for_handles) }; - let user = sqlx::query!( - "SELECT id, did, handle, password_required FROM users WHERE LOWER(email) = $1 OR handle = $2", - identifier, - normalized_handle - ) - .fetch_optional(&state.db) - .await; - - let user = match user { + let user = match state.user_repo.get_user_for_passkey_recovery(identifier, &normalized_handle).await { Ok(Some(u)) if !u.password_required => u, _ => { return SuccessResponse::ok().into_response(); @@ -1022,15 +867,7 @@ pub async fn request_passkey_recovery( }; let expires_at = Utc::now() + Duration::hours(1); - if let Err(e) = sqlx::query!( - "UPDATE users SET recovery_token = $1, recovery_token_expires_at = $2 WHERE did = $3", - recovery_token_hash, - expires_at, - &user.did - ) - .execute(&state.db) - .await - { + if let Err(e) = state.user_repo.set_recovery_token(&user.did, &recovery_token_hash, expires_at).await { error!("Error updating recovery token: {:?}", e); return ApiError::InternalError(None).into_response(); } @@ -1043,8 +880,14 @@ pub async fn request_passkey_recovery( urlencoding::encode(&recovery_token) ); - let _ = - crate::comms::enqueue_passkey_recovery(&state.db, user.id, &recovery_url, &hostname).await; + let _ = crate::comms::comms_repo::enqueue_passkey_recovery( + state.user_repo.as_ref(), + state.infra_repo.as_ref(), + user.id, + &recovery_url, + &hostname, + ) + .await; info!(did = %user.did, "Passkey recovery requested"); SuccessResponse::ok().into_response() @@ -1066,14 +909,7 @@ pub async fn recover_passkey_account( return ApiError::InvalidRequest(e.to_string()).into_response(); } - let user = sqlx::query!( - "SELECT id, did, recovery_token, recovery_token_expires_at FROM users WHERE did = $1", - input.did.as_str() - ) - .fetch_optional(&state.db) - .await; - - let user = match user { + let user = match state.user_repo.get_user_for_recovery(&input.did).await { Ok(Some(u)) => u, _ => { return ApiError::InvalidRecoveryLink.into_response(); @@ -1104,44 +940,20 @@ pub async fn recover_passkey_account( } }; - let mut tx = match state.db.begin().await { - Ok(tx) => tx, - Err(e) => { - error!("Failed to begin transaction: {:?}", e); - return ApiError::InternalError(None).into_response(); - } - }; - - if let Err(e) = sqlx::query!( - "UPDATE users SET password_hash = $1, password_required = TRUE, recovery_token = NULL, recovery_token_expires_at = NULL WHERE did = $2", + let recover_input = tranquil_db_traits::RecoverPasskeyAccountInput { + did: input.did.clone(), password_hash, - input.did.as_str() - ) - .execute(&mut *tx) - .await - { - error!("Error updating password: {:?}", e); - return ApiError::InternalError(None).into_response(); - } - - let deleted = sqlx::query!("DELETE FROM passkeys WHERE did = $1", input.did.as_str()) - .execute(&mut *tx) - .await; - let passkeys_deleted = match deleted { - Ok(result) => result.rows_affected(), + }; + let result = match state.user_repo.recover_passkey_account(&recover_input).await { + Ok(r) => r, Err(e) => { - error!(did = %input.did, "Failed to delete passkeys during recovery: {:?}", e); + error!("Error recovering passkey account: {:?}", e); return ApiError::InternalError(None).into_response(); } }; - if let Err(e) = tx.commit().await { - error!("Failed to commit recovery transaction: {:?}", e); - return ApiError::InternalError(None).into_response(); - } - - if passkeys_deleted > 0 { - info!(did = %input.did, count = passkeys_deleted, "Deleted lost passkeys during account recovery"); + if result.passkeys_deleted > 0 { + info!(did = %input.did, count = result.passkeys_deleted, "Deleted lost passkeys during account recovery"); } info!(did = %input.did, "Passkey-only account recovered with temporary password"); SuccessResponse::ok().into_response() diff --git a/crates/tranquil-pds/src/api/server/passkeys.rs b/crates/tranquil-pds/src/api/server/passkeys.rs index dd10d80..78b6b14 100644 --- a/crates/tranquil-pds/src/api/server/passkeys.rs +++ b/crates/tranquil-pds/src/api/server/passkeys.rs @@ -1,11 +1,7 @@ use crate::api::EmptyResponse; use crate::api::error::ApiError; use crate::auth::BearerAuth; -use crate::auth::webauthn::{ - self, WebAuthnConfig, delete_passkey as db_delete_passkey, delete_registration_state, - get_passkeys_for_user, load_registration_state, save_passkey, save_registration_state, - update_passkey_name as db_update_passkey_name, -}; +use crate::auth::webauthn::WebAuthnConfig; use crate::state::AppState; use axum::{ Json, @@ -46,12 +42,8 @@ pub async fn start_passkey_registration( Err(e) => return e.into_response(), }; - let user = sqlx::query!("SELECT handle FROM users WHERE did = $1", &*auth.0.did) - .fetch_optional(&state.db) - .await; - - let handle = match user { - Ok(Some(row)) => row.handle, + let handle = match state.user_repo.get_handle_by_did(&auth.0.did).await { + Ok(Some(h)) => h, Ok(None) => { return ApiError::AccountNotFound.into_response(); } @@ -61,7 +53,7 @@ pub async fn start_passkey_registration( } }; - let existing_passkeys = match get_passkeys_for_user(&state.db, &auth.0.did).await { + let existing_passkeys = match state.user_repo.get_passkeys_for_user(&auth.0.did).await { Ok(passkeys) => passkeys, Err(e) => { error!("DB error fetching existing passkeys: {:?}", e); @@ -90,7 +82,19 @@ pub async fn start_passkey_registration( } }; - if let Err(e) = save_registration_state(&state.db, &auth.0.did, ®_state).await { + let state_json = match serde_json::to_string(®_state) { + Ok(s) => s, + Err(e) => { + error!("Failed to serialize registration state: {:?}", e); + return ApiError::InternalError(None).into_response(); + } + }; + + if let Err(e) = state + .user_repo + .save_webauthn_challenge(&auth.0.did, "registration", &state_json) + .await + { error!("Failed to save registration state: {:?}", e); return ApiError::InternalError(None).into_response(); } @@ -126,8 +130,12 @@ pub async fn finish_passkey_registration( Err(e) => return e.into_response(), }; - let reg_state = match load_registration_state(&state.db, &auth.0.did).await { - Ok(Some(state)) => state, + let reg_state_json = match state + .user_repo + .load_webauthn_challenge(&auth.0.did, "registration") + .await + { + Ok(Some(json)) => json, Ok(None) => { return ApiError::NoRegistrationInProgress.into_response(); } @@ -137,6 +145,14 @@ pub async fn finish_passkey_registration( } }; + let reg_state: SecurityKeyRegistration = match serde_json::from_str(®_state_json) { + Ok(s) => s, + Err(e) => { + error!("Failed to deserialize registration state: {:?}", e); + return ApiError::InternalError(None).into_response(); + } + }; + let credential: RegisterPublicKeyCredential = match serde_json::from_value(input.credential) { Ok(c) => c, Err(e) => { @@ -153,13 +169,23 @@ pub async fn finish_passkey_registration( } }; - let passkey_id = match save_passkey( - &state.db, - &auth.0.did, - &passkey, - input.friendly_name.as_deref(), - ) - .await + let public_key = match serde_json::to_vec(&passkey) { + Ok(pk) => pk, + Err(e) => { + error!("Failed to serialize passkey: {:?}", e); + return ApiError::InternalError(None).into_response(); + } + }; + + let passkey_id = match state + .user_repo + .save_passkey( + &auth.0.did, + passkey.cred_id(), + &public_key, + input.friendly_name.as_deref(), + ) + .await { Ok(id) => id, Err(e) => { @@ -168,7 +194,11 @@ pub async fn finish_passkey_registration( } }; - if let Err(e) = delete_registration_state(&state.db, &auth.0.did).await { + if let Err(e) = state + .user_repo + .delete_webauthn_challenge(&auth.0.did, "registration") + .await + { warn!("Failed to delete registration state: {:?}", e); } @@ -203,7 +233,7 @@ pub struct ListPasskeysResponse { } pub async fn list_passkeys(State(state): State, auth: BearerAuth) -> Response { - let passkeys = match get_passkeys_for_user(&state.db, &auth.0.did).await { + let passkeys = match state.user_repo.get_passkeys_for_user(&auth.0.did).await { Ok(pks) => pks, Err(e) => { error!("DB error fetching passkeys: {:?}", e); @@ -239,13 +269,13 @@ pub async fn delete_passkey( auth: BearerAuth, Json(input): Json, ) -> Response { - if !crate::api::server::reauth::check_legacy_session_mfa(&state.db, &auth.0.did).await { - return crate::api::server::reauth::legacy_mfa_required_response(&state.db, &auth.0.did) + if !crate::api::server::reauth::check_legacy_session_mfa(&*state.session_repo, &auth.0.did).await { + return crate::api::server::reauth::legacy_mfa_required_response(&*state.user_repo, &*state.session_repo, &auth.0.did) .await; } - if crate::api::server::reauth::check_reauth_required(&state.db, &auth.0.did).await { - return crate::api::server::reauth::reauth_required_response(&state.db, &auth.0.did).await; + if crate::api::server::reauth::check_reauth_required(&*state.session_repo, &auth.0.did).await { + return crate::api::server::reauth::reauth_required_response(&*state.user_repo, &*state.session_repo, &auth.0.did).await; } let id: uuid::Uuid = match input.id.parse() { @@ -255,7 +285,7 @@ pub async fn delete_passkey( } }; - match db_delete_passkey(&state.db, id, &auth.0.did).await { + match state.user_repo.delete_passkey(id, &auth.0.did).await { Ok(true) => { info!(did = %auth.0.did, passkey_id = %id, "Passkey deleted"); EmptyResponse::ok().into_response() @@ -287,7 +317,11 @@ pub async fn update_passkey( } }; - match db_update_passkey_name(&state.db, id, &auth.0.did, &input.friendly_name).await { + match state + .user_repo + .update_passkey_name(id, &auth.0.did, &input.friendly_name) + .await + { Ok(true) => { info!(did = %auth.0.did, passkey_id = %id, "Passkey renamed"); EmptyResponse::ok().into_response() @@ -300,10 +334,6 @@ pub async fn update_passkey( } } -pub async fn has_passkeys_for_user(state: &AppState, did: &str) -> bool { - has_passkeys_for_user_db(&state.db, did).await -} - -pub async fn has_passkeys_for_user_db(db: &sqlx::PgPool, did: &str) -> bool { - webauthn::has_passkeys(db, did).await.unwrap_or(false) +pub async fn has_passkeys_for_user(state: &AppState, did: &crate::types::Did) -> bool { + state.user_repo.has_passkeys(did).await.unwrap_or(false) } diff --git a/crates/tranquil-pds/src/api/server/password.rs b/crates/tranquil-pds/src/api/server/password.rs index 4e49b03..5b7dcea 100644 --- a/crates/tranquil-pds/src/api/server/password.rs +++ b/crates/tranquil-pds/src/api/server/password.rs @@ -14,7 +14,6 @@ use bcrypt::{DEFAULT_COST, hash, verify}; use chrono::{Duration, Utc}; use serde::Deserialize; use tracing::{error, info, warn}; -use uuid::Uuid; fn generate_reset_code() -> String { crate::util::generate_token_code() @@ -58,22 +57,20 @@ pub async fn request_password_reset( return ApiError::InvalidRequest("email or handle is required".into()).into_response(); } let pds_hostname = std::env::var("PDS_HOSTNAME").unwrap_or_else(|_| "localhost".to_string()); + let hostname_for_handles = pds_hostname.split(':').next().unwrap_or(&pds_hostname); let normalized = identifier.to_lowercase(); let normalized = normalized.strip_prefix('@').unwrap_or(&normalized); let normalized_handle = if normalized.contains('@') || normalized.contains('.') { normalized.to_string() } else { - format!("{}.{}", normalized, pds_hostname) + format!("{}.{}", normalized, hostname_for_handles) }; - let user = sqlx::query!( - "SELECT id FROM users WHERE LOWER(email) = $1 OR handle = $2", - normalized, - normalized_handle - ) - .fetch_optional(&state.db) - .await; - let user_id = match user { - Ok(Some(row)) => row.id, + let user_id = match state + .user_repo + .get_id_by_email_or_handle(&normalized, &normalized_handle) + .await + { + Ok(Some(id)) => id, Ok(None) => { info!("Password reset requested for unknown identifier"); return EmptyResponse::ok().into_response(); @@ -85,20 +82,23 @@ pub async fn request_password_reset( }; let code = generate_reset_code(); let expires_at = Utc::now() + Duration::minutes(10); - let update = sqlx::query!( - "UPDATE users SET password_reset_code = $1, password_reset_code_expires_at = $2 WHERE id = $3", - code, - expires_at, - user_id - ) - .execute(&state.db) - .await; - if let Err(e) = update { + if let Err(e) = state + .user_repo + .set_password_reset_code(user_id, &code, expires_at) + .await + { error!("DB error setting reset code: {:?}", e); return ApiError::InternalError(None).into_response(); } let hostname = std::env::var("PDS_HOSTNAME").unwrap_or_else(|_| "localhost".to_string()); - if let Err(e) = crate::comms::enqueue_password_reset(&state.db, user_id, &code, &hostname).await + if let Err(e) = crate::comms::comms_repo::enqueue_password_reset( + state.user_repo.as_ref(), + state.infra_repo.as_ref(), + user_id, + &code, + &hostname, + ) + .await { warn!("Failed to enqueue password reset notification: {:?}", e); } @@ -136,17 +136,8 @@ pub async fn reset_password( if let Err(e) = validate_password(password) { return ApiError::InvalidRequest(e.to_string()).into_response(); } - let user = sqlx::query!( - "SELECT id, password_reset_code, password_reset_code_expires_at FROM users WHERE password_reset_code = $1", - token - ) - .fetch_optional(&state.db) - .await; - let (user_id, expires_at) = match user { - Ok(Some(row)) => { - let expires = row.password_reset_code_expires_at; - (row.id, expires) - } + let user = match state.user_repo.get_user_by_reset_code(token).await { + Ok(Some(u)) => u, Ok(None) => { return ApiError::InvalidToken(None).into_response(); } @@ -155,21 +146,15 @@ pub async fn reset_password( return ApiError::InternalError(None).into_response(); } }; - if let Some(exp) = expires_at { - if Utc::now() > exp { - if let Err(e) = sqlx::query!( - "UPDATE users SET password_reset_code = NULL, password_reset_code_expires_at = NULL WHERE id = $1", - user_id - ) - .execute(&state.db) - .await - { - error!("Failed to clear expired reset code: {:?}", e); - } - return ApiError::ExpiredToken(None).into_response(); - } - } else { + let user_id = user.id; + let Some(exp) = user.expires_at else { return ApiError::InvalidToken(None).into_response(); + }; + if Utc::now() > exp { + if let Err(e) = state.user_repo.clear_password_reset_code(user_id).await { + error!("Failed to clear expired reset code: {:?}", e); + } + return ApiError::ExpiredToken(None).into_response(); } let password_clone = password.to_string(); let password_hash = @@ -184,63 +169,19 @@ pub async fn reset_password( return ApiError::InternalError(None).into_response(); } }; - let mut tx = match state.db.begin().await { - Ok(tx) => tx, - Err(e) => { - error!("Failed to begin transaction: {:?}", e); - return ApiError::InternalError(None).into_response(); - } - }; - if let Err(e) = sqlx::query!( - "UPDATE users SET password_hash = $1, password_reset_code = NULL, password_reset_code_expires_at = NULL, password_required = TRUE WHERE id = $2", - password_hash, - user_id - ) - .execute(&mut *tx) - .await - { - error!("DB error updating password: {:?}", e); - return ApiError::InternalError(None).into_response(); - } - let user_did = match sqlx::query_scalar!("SELECT did FROM users WHERE id = $1", user_id) - .fetch_one(&mut *tx) + let result = match state + .user_repo + .reset_password_with_sessions(user_id, &password_hash) .await { - Ok(did) => did, + Ok(r) => r, Err(e) => { - error!("Failed to get DID for user {}: {:?}", user_id, e); + error!("Failed to reset password: {:?}", e); return ApiError::InternalError(None).into_response(); } }; - let session_jtis: Vec = match sqlx::query_scalar!( - "SELECT access_jti FROM session_tokens WHERE did = $1", - user_did - ) - .fetch_all(&mut *tx) - .await - { - Ok(jtis) => jtis, - Err(e) => { - error!("Failed to fetch session JTIs: {:?}", e); - vec![] - } - }; - if let Err(e) = sqlx::query!("DELETE FROM session_tokens WHERE did = $1", user_did) - .execute(&mut *tx) - .await - { - error!( - "Failed to invalidate sessions after password reset: {:?}", - e - ); - return ApiError::InternalError(None).into_response(); - } - if let Err(e) = tx.commit().await { - error!("Failed to commit password reset transaction: {:?}", e); - return ApiError::InternalError(None).into_response(); - } - futures::future::join_all(session_jtis.into_iter().map(|jti| { - let cache_key = format!("auth:session:{}:{}", user_did, jti); + futures::future::join_all(result.session_jtis.iter().map(|jti| { + let cache_key = format!("auth:session:{}:{}", result.did, jti); let cache = state.cache.clone(); async move { if let Err(e) = cache.delete(&cache_key).await { @@ -268,8 +209,8 @@ pub async fn change_password( auth: BearerAuth, Json(input): Json, ) -> Response { - if !crate::api::server::reauth::check_legacy_session_mfa(&state.db, &auth.0.did).await { - return crate::api::server::reauth::legacy_mfa_required_response(&state.db, &auth.0.did) + if !crate::api::server::reauth::check_legacy_session_mfa(&*state.session_repo, &auth.0.did).await { + return crate::api::server::reauth::legacy_mfa_required_response(&*state.user_repo, &*state.session_repo, &auth.0.did) .await; } @@ -284,13 +225,8 @@ pub async fn change_password( if let Err(e) = validate_password(new_password) { return ApiError::InvalidRequest(e.to_string()).into_response(); } - let user = - sqlx::query_as::<_, (Uuid, String)>("SELECT id, password_hash FROM users WHERE did = $1") - .bind(&auth.0.did) - .fetch_optional(&state.db) - .await; - let (user_id, password_hash) = match user { - Ok(Some(row)) => row, + let user = match state.user_repo.get_id_and_password_hash_by_did(&auth.0.did).await { + Ok(Some(u)) => u, Ok(None) => { return ApiError::AccountNotFound.into_response(); } @@ -299,6 +235,7 @@ pub async fn change_password( return ApiError::InternalError(None).into_response(); } }; + let (user_id, password_hash) = (user.id, user.password_hash); let valid = match verify(current_password, &password_hash) { Ok(v) => v, Err(e) => { @@ -322,12 +259,7 @@ pub async fn change_password( return ApiError::InternalError(None).into_response(); } }; - if let Err(e) = sqlx::query("UPDATE users SET password_hash = $1 WHERE id = $2") - .bind(&new_hash) - .bind(user_id) - .execute(&state.db) - .await - { + if let Err(e) = state.user_repo.update_password_hash(user_id, &new_hash).await { error!("DB error updating password: {:?}", e); return ApiError::InternalError(None).into_response(); } @@ -336,17 +268,8 @@ pub async fn change_password( } pub async fn get_password_status(State(state): State, auth: BearerAuth) -> Response { - let user = sqlx::query!( - "SELECT password_hash IS NOT NULL as has_password FROM users WHERE did = $1", - &auth.0.did - ) - .fetch_optional(&state.db) - .await; - - match user { - Ok(Some(row)) => { - HasPasswordResponse::response(row.has_password.unwrap_or(false)).into_response() - } + match state.user_repo.has_password_by_did(&auth.0.did).await { + Ok(Some(has)) => HasPasswordResponse::response(has).into_response(), Ok(None) => ApiError::AccountNotFound.into_response(), Err(e) => { error!("DB error: {:?}", e); @@ -356,23 +279,22 @@ pub async fn get_password_status(State(state): State, auth: BearerAuth } pub async fn remove_password(State(state): State, auth: BearerAuth) -> Response { - if !crate::api::server::reauth::check_legacy_session_mfa(&state.db, &auth.0.did).await { - return crate::api::server::reauth::legacy_mfa_required_response(&state.db, &auth.0.did) + if !crate::api::server::reauth::check_legacy_session_mfa(&*state.session_repo, &auth.0.did).await { + return crate::api::server::reauth::legacy_mfa_required_response(&*state.user_repo, &*state.session_repo, &auth.0.did) .await; } if crate::api::server::reauth::check_reauth_required_cached( - &state.db, + &*state.session_repo, &state.cache, &auth.0.did, ) .await { - return crate::api::server::reauth::reauth_required_response(&state.db, &auth.0.did).await; + return crate::api::server::reauth::reauth_required_response(&*state.user_repo, &*state.session_repo, &auth.0.did).await; } - let has_passkeys = - crate::api::server::passkeys::has_passkeys_for_user_db(&state.db, &auth.0.did).await; + let has_passkeys = state.user_repo.has_passkeys(&auth.0.did).await.unwrap_or(false); if !has_passkeys { return ApiError::InvalidRequest( "You must have at least one passkey registered before removing your password".into(), @@ -380,14 +302,7 @@ pub async fn remove_password(State(state): State, auth: BearerAuth) -> .into_response(); } - let user = sqlx::query!( - "SELECT id, password_hash FROM users WHERE did = $1", - &auth.0.did - ) - .fetch_optional(&state.db) - .await; - - let user = match user { + let user = match state.user_repo.get_password_info_by_did(&auth.0.did).await { Ok(Some(u)) => u, Ok(None) => { return ApiError::AccountNotFound.into_response(); @@ -402,13 +317,7 @@ pub async fn remove_password(State(state): State, auth: BearerAuth) -> return ApiError::InvalidRequest("Account already has no password".into()).into_response(); } - if let Err(e) = sqlx::query!( - "UPDATE users SET password_hash = NULL, password_required = FALSE WHERE id = $1", - user.id - ) - .execute(&state.db) - .await - { + if let Err(e) = state.user_repo.remove_user_password(user.id).await { error!("DB error removing password: {:?}", e); return ApiError::InternalError(None).into_response(); } @@ -429,13 +338,13 @@ pub async fn set_password( Json(input): Json, ) -> Response { if crate::api::server::reauth::check_reauth_required_cached( - &state.db, + &*state.session_repo, &state.cache, &auth.0.did, ) .await { - return crate::api::server::reauth::reauth_required_response(&state.db, &auth.0.did).await; + return crate::api::server::reauth::reauth_required_response(&*state.user_repo, &*state.session_repo, &auth.0.did).await; } let new_password = &input.new_password; @@ -446,14 +355,7 @@ pub async fn set_password( return ApiError::InvalidRequest(e.to_string()).into_response(); } - let user = sqlx::query!( - "SELECT id, password_hash FROM users WHERE did = $1", - &auth.0.did - ) - .fetch_optional(&state.db) - .await; - - let user = match user { + let user = match state.user_repo.get_password_info_by_did(&auth.0.did).await { Ok(Some(u)) => u, Ok(None) => { return ApiError::AccountNotFound.into_response(); @@ -485,14 +387,7 @@ pub async fn set_password( } }; - if let Err(e) = sqlx::query!( - "UPDATE users SET password_hash = $1, password_required = TRUE WHERE id = $2", - new_hash, - user.id - ) - .execute(&state.db) - .await - { + if let Err(e) = state.user_repo.set_new_user_password(user.id, &new_hash).await { error!("DB error setting password: {:?}", e); return ApiError::InternalError(None).into_response(); } diff --git a/crates/tranquil-pds/src/api/server/reauth.rs b/crates/tranquil-pds/src/api/server/reauth.rs index b254588..d429ad6 100644 --- a/crates/tranquil-pds/src/api/server/reauth.rs +++ b/crates/tranquil-pds/src/api/server/reauth.rs @@ -7,8 +7,8 @@ use axum::{ }; use chrono::{DateTime, Utc}; use serde::{Deserialize, Serialize}; -use sqlx::PgPool; use tracing::{error, info, warn}; +use tranquil_db_traits::{SessionRepository, UserRepository}; use crate::auth::BearerAuth; use crate::state::{AppState, RateLimitKind}; @@ -25,16 +25,8 @@ pub struct ReauthStatusResponse { } pub async fn get_reauth_status(State(state): State, auth: BearerAuth) -> Response { - let session = sqlx::query!( - "SELECT last_reauth_at FROM session_tokens WHERE did = $1 ORDER BY created_at DESC LIMIT 1", - &auth.0.did - ) - .fetch_optional(&state.db) - .await; - - let last_reauth_at = match session { - Ok(Some(row)) => row.last_reauth_at, - Ok(None) => None, + let last_reauth_at = match state.session_repo.get_last_reauth_at(&auth.0.did).await { + Ok(t) => t, Err(e) => { error!("DB error: {:?}", e); return ApiError::InternalError(None).into_response(); @@ -42,7 +34,8 @@ pub async fn get_reauth_status(State(state): State, auth: BearerAuth) }; let reauth_required = is_reauth_required(last_reauth_at); - let available_methods = get_available_reauth_methods(&state.db, &auth.0.did).await; + let available_methods = + get_available_reauth_methods(&*state.user_repo, &*state.session_repo, &auth.0.did).await; Json(ReauthStatusResponse { last_reauth_at, @@ -69,15 +62,8 @@ pub async fn reauth_password( auth: BearerAuth, Json(input): Json, ) -> Response { - let user = sqlx::query!( - "SELECT password_hash FROM users WHERE did = $1", - &*&auth.0.did - ) - .fetch_optional(&state.db) - .await; - - let password_hash = match user { - Ok(Some(row)) => row.password_hash, + let password_hash = match state.user_repo.get_password_hash_by_did(&auth.0.did).await { + Ok(Some(hash)) => hash, Ok(None) => { return ApiError::AccountNotFound.into_response(); } @@ -87,25 +73,18 @@ pub async fn reauth_password( } }; - let password_valid = password_hash - .as_ref() - .map(|h| bcrypt::verify(&input.password, h).unwrap_or(false)) - .unwrap_or(false); + let password_valid = bcrypt::verify(&input.password, &password_hash).unwrap_or(false); if !password_valid { - let app_passwords = sqlx::query!( - "SELECT ap.password_hash FROM app_passwords ap - JOIN users u ON ap.user_id = u.id - WHERE u.did = $1", - &auth.0.did - ) - .fetch_all(&state.db) - .await - .unwrap_or_default(); + let app_password_hashes = state + .session_repo + .get_app_password_hashes_by_did(&auth.0.did) + .await + .unwrap_or_default(); - let app_password_valid = app_passwords + let app_password_valid = app_password_hashes .iter() - .any(|ap| bcrypt::verify(&input.password, &ap.password_hash).unwrap_or(false)); + .any(|h| bcrypt::verify(&input.password, h).unwrap_or(false)); if !app_password_valid { warn!(did = %&auth.0.did, "Re-auth failed: invalid password"); @@ -113,7 +92,7 @@ pub async fn reauth_password( } } - match update_last_reauth_cached(&state.db, &state.cache, &auth.0.did).await { + match update_last_reauth_cached(&*state.session_repo, &state.cache, &auth.0.did).await { Ok(reauthed_at) => { info!(did = %&auth.0.did, "Re-auth successful via password"); Json(ReauthResponse { reauthed_at }).into_response() @@ -156,7 +135,7 @@ pub async fn reauth_totp( return ApiError::InvalidCode(Some("Invalid TOTP or backup code".into())).into_response(); } - match update_last_reauth_cached(&state.db, &state.cache, &auth.0.did).await { + match update_last_reauth_cached(&*state.session_repo, &state.cache, &auth.0.did).await { Ok(reauthed_at) => { info!(did = %&auth.0.did, "Re-auth successful via TOTP"); Json(ReauthResponse { reauthed_at }).into_response() @@ -177,14 +156,13 @@ pub struct PasskeyReauthStartResponse { pub async fn reauth_passkey_start(State(state): State, auth: BearerAuth) -> Response { let pds_hostname = std::env::var("PDS_HOSTNAME").unwrap_or_else(|_| "localhost".to_string()); - let stored_passkeys = - match crate::auth::webauthn::get_passkeys_for_user(&state.db, &auth.0.did).await { - Ok(pks) => pks, - Err(e) => { - error!("Failed to get passkeys: {:?}", e); - return ApiError::InternalError(None).into_response(); - } - }; + let stored_passkeys = match state.user_repo.get_passkeys_for_user(&auth.0.did).await { + Ok(pks) => pks, + Err(e) => { + error!("Failed to get passkeys: {:?}", e); + return ApiError::InternalError(None).into_response(); + } + }; if stored_passkeys.is_empty() { return ApiError::NoPasskeys.into_response(); @@ -192,7 +170,7 @@ pub async fn reauth_passkey_start(State(state): State, auth: BearerAut let passkeys: Vec = stored_passkeys .iter() - .filter_map(|sp| sp.to_security_key().ok()) + .filter_map(|sp| serde_json::from_slice(&sp.public_key).ok()) .collect(); if passkeys.is_empty() { @@ -215,8 +193,18 @@ pub async fn reauth_passkey_start(State(state): State, auth: BearerAut } }; - if let Err(e) = - crate::auth::webauthn::save_authentication_state(&state.db, &auth.0.did, &auth_state).await + let state_json = match serde_json::to_string(&auth_state) { + Ok(s) => s, + Err(e) => { + error!("Failed to serialize authentication state: {:?}", e); + return ApiError::InternalError(None).into_response(); + } + }; + + if let Err(e) = state + .user_repo + .save_webauthn_challenge(&auth.0.did, "authentication", &state_json) + .await { error!("Failed to save authentication state: {:?}", e); return ApiError::InternalError(None).into_response(); @@ -239,14 +227,26 @@ pub async fn reauth_passkey_finish( ) -> Response { let pds_hostname = std::env::var("PDS_HOSTNAME").unwrap_or_else(|_| "localhost".to_string()); - let auth_state = - match crate::auth::webauthn::load_authentication_state(&state.db, &auth.0.did).await { - Ok(Some(s)) => s, - Ok(None) => { - return ApiError::NoChallengeInProgress.into_response(); - } + let auth_state_json = match state + .user_repo + .load_webauthn_challenge(&auth.0.did, "authentication") + .await + { + Ok(Some(json)) => json, + Ok(None) => { + return ApiError::NoChallengeInProgress.into_response(); + } + Err(e) => { + error!("Failed to load authentication state: {:?}", e); + return ApiError::InternalError(None).into_response(); + } + }; + + let auth_state: webauthn_rs::prelude::SecurityKeyAuthentication = + match serde_json::from_str(&auth_state_json) { + Ok(s) => s, Err(e) => { - error!("Failed to load authentication state: {:?}", e); + error!("Failed to deserialize authentication state: {:?}", e); return ApiError::InternalError(None).into_response(); } }; @@ -278,17 +278,17 @@ pub async fn reauth_passkey_finish( }; let cred_id_bytes = auth_result.cred_id().as_ref(); - match crate::auth::webauthn::update_passkey_counter( - &state.db, - cred_id_bytes, - auth_result.counter(), - ) - .await + match state + .user_repo + .update_passkey_counter(cred_id_bytes, auth_result.counter() as i32) + .await { Ok(false) => { warn!(did = %&auth.0.did, "Passkey counter anomaly detected - possible cloned key"); - let _ = - crate::auth::webauthn::delete_authentication_state(&state.db, &auth.0.did).await; + let _ = state + .user_repo + .delete_webauthn_challenge(&auth.0.did, "authentication") + .await; return ApiError::PasskeyCounterAnomaly.into_response(); } Err(e) => { @@ -297,9 +297,12 @@ pub async fn reauth_passkey_finish( Ok(true) => {} } - let _ = crate::auth::webauthn::delete_authentication_state(&state.db, &auth.0.did).await; + let _ = state + .user_repo + .delete_webauthn_challenge(&auth.0.did, "authentication") + .await; - match update_last_reauth_cached(&state.db, &state.cache, &auth.0.did).await { + match update_last_reauth_cached(&*state.session_repo, &state.cache, &auth.0.did).await { Ok(reauthed_at) => { info!(did = %&auth.0.did, "Re-auth successful via passkey"); Json(ReauthResponse { reauthed_at }).into_response() @@ -312,18 +315,11 @@ pub async fn reauth_passkey_finish( } pub async fn update_last_reauth_cached( - db: &PgPool, + session_repo: &dyn SessionRepository, cache: &std::sync::Arc, - did: &str, -) -> Result, sqlx::Error> { - let now = Utc::now(); - sqlx::query!( - "UPDATE session_tokens SET last_reauth_at = $1, mfa_verified = TRUE WHERE did = $2", - now, - did - ) - .execute(db) - .await?; + did: &crate::types::Did, +) -> Result, tranquil_db_traits::DbError> { + let now = session_repo.update_last_reauth(did).await?; let cache_key = format!("reauth:{}", did); let _ = cache .set( @@ -345,29 +341,30 @@ fn is_reauth_required(last_reauth_at: Option>) -> bool { } } -async fn get_available_reauth_methods(db: &PgPool, did: &str) -> Vec { +async fn get_available_reauth_methods( + user_repo: &dyn UserRepository, + _session_repo: &dyn SessionRepository, + did: &crate::types::Did, +) -> Vec { let mut methods = Vec::new(); - let has_password = sqlx::query_scalar!( - "SELECT password_hash IS NOT NULL as has_pw FROM users WHERE did = $1", - did - ) - .fetch_optional(db) - .await - .ok() - .flatten() - .unwrap_or(Some(false)); + let has_password = user_repo + .get_password_hash_by_did(did) + .await + .ok() + .flatten() + .is_some(); - if has_password == Some(true) { + if has_password { methods.push("password".to_string()); } - let has_totp = crate::api::server::totp::has_totp_enabled_db(db, did).await; + let has_totp = user_repo.has_totp_enabled(did).await.unwrap_or(false); if has_totp { methods.push("totp".to_string()); } - let has_passkeys = crate::api::server::passkeys::has_passkeys_for_user_db(db, did).await; + let has_passkeys = user_repo.has_passkeys(did).await.unwrap_or(false); if has_passkeys { methods.push("passkey".to_string()); } @@ -375,24 +372,17 @@ async fn get_available_reauth_methods(db: &PgPool, did: &str) -> Vec { methods } -pub async fn check_reauth_required(db: &PgPool, did: &str) -> bool { - let session = sqlx::query!( - "SELECT last_reauth_at FROM session_tokens WHERE did = $1 ORDER BY created_at DESC LIMIT 1", - did - ) - .fetch_optional(db) - .await; - - match session { - Ok(Some(row)) => is_reauth_required(row.last_reauth_at), +pub async fn check_reauth_required(session_repo: &dyn SessionRepository, did: &crate::types::Did) -> bool { + match session_repo.get_last_reauth_at(did).await { + Ok(last_reauth_at) => is_reauth_required(last_reauth_at), _ => true, } } pub async fn check_reauth_required_cached( - db: &PgPool, + session_repo: &dyn SessionRepository, cache: &std::sync::Arc, - did: &str, + did: &crate::types::Did, ) -> bool { let cache_key = format!("reauth:{}", did); if let Some(timestamp_str) = cache.get(&cache_key).await @@ -406,15 +396,8 @@ pub async fn check_reauth_required_cached( } } } - let session = sqlx::query!( - "SELECT last_reauth_at FROM session_tokens WHERE did = $1 ORDER BY created_at DESC LIMIT 1", - did - ) - .fetch_optional(db) - .await; - - match session { - Ok(Some(row)) => is_reauth_required(row.last_reauth_at), + match session_repo.get_last_reauth_at(did).await { + Ok(last_reauth_at) => is_reauth_required(last_reauth_at), _ => true, } } @@ -427,8 +410,12 @@ pub struct ReauthRequiredError { pub reauth_methods: Vec, } -pub async fn reauth_required_response(db: &PgPool, did: &str) -> Response { - let methods = get_available_reauth_methods(db, did).await; +pub async fn reauth_required_response( + user_repo: &dyn UserRepository, + session_repo: &dyn SessionRepository, + did: &crate::types::Did, +) -> Response { + let methods = get_available_reauth_methods(user_repo, session_repo, did).await; ( StatusCode::UNAUTHORIZED, Json(ReauthRequiredError { @@ -440,23 +427,16 @@ pub async fn reauth_required_response(db: &PgPool, did: &str) -> Response { .into_response() } -pub async fn check_legacy_session_mfa(db: &PgPool, did: &str) -> bool { - let session = sqlx::query!( - "SELECT legacy_login, mfa_verified, last_reauth_at FROM session_tokens WHERE did = $1 ORDER BY created_at DESC LIMIT 1", - did - ) - .fetch_optional(db) - .await; - - match session { - Ok(Some(row)) => { - if !row.legacy_login { +pub async fn check_legacy_session_mfa(session_repo: &dyn SessionRepository, did: &crate::types::Did) -> bool { + match session_repo.get_session_mfa_status(did).await { + Ok(Some(status)) => { + if !status.legacy_login { return true; } - if row.mfa_verified { + if status.mfa_verified { return true; } - if let Some(last_reauth) = row.last_reauth_at { + if let Some(last_reauth) = status.last_reauth_at { let elapsed = chrono::Utc::now().signed_duration_since(last_reauth); if elapsed.num_seconds() <= REAUTH_WINDOW_SECONDS { return true; @@ -468,18 +448,19 @@ pub async fn check_legacy_session_mfa(db: &PgPool, did: &str) -> bool { } } -pub async fn update_mfa_verified(db: &PgPool, did: &str) -> Result<(), sqlx::Error> { - sqlx::query!( - "UPDATE session_tokens SET mfa_verified = TRUE, last_reauth_at = NOW() WHERE did = $1", - did - ) - .execute(db) - .await?; - Ok(()) +pub async fn update_mfa_verified( + session_repo: &dyn SessionRepository, + did: &crate::types::Did, +) -> Result<(), tranquil_db_traits::DbError> { + session_repo.update_mfa_verified(did).await } -pub async fn legacy_mfa_required_response(db: &PgPool, did: &str) -> Response { - let methods = get_available_reauth_methods(db, did).await; +pub async fn legacy_mfa_required_response( + user_repo: &dyn UserRepository, + session_repo: &dyn SessionRepository, + did: &crate::types::Did, +) -> Response { + let methods = get_available_reauth_methods(user_repo, session_repo, did).await; ( StatusCode::FORBIDDEN, Json(MfaVerificationRequiredError { diff --git a/crates/tranquil-pds/src/api/server/service_auth.rs b/crates/tranquil-pds/src/api/server/service_auth.rs index ee06e81..55a0fb7 100644 --- a/crates/tranquil-pds/src/api/server/service_auth.rs +++ b/crates/tranquil-pds/src/api/server/service_auth.rs @@ -81,7 +81,7 @@ pub async fn get_service_auth( let auth_user = if is_dpop { match crate::oauth::verify::verify_oauth_access_token( - &state.db, + state.oauth_repo.as_ref(), &token, dpop_proof, "GET", @@ -119,7 +119,7 @@ pub async fn get_service_auth( } } } else { - match crate::auth::validate_bearer_token_for_service_auth(&state.db, &token).await { + match crate::auth::validate_bearer_token_for_service_auth(state.user_repo.as_ref(), &token).await { Ok(user) => user, Err(e) => { warn!(error = ?e, "getServiceAuth auth validation failed"); @@ -137,28 +137,27 @@ pub async fn get_service_auth( Some(kb) => kb.clone(), None => { warn!(did = %&auth_user.did, "getServiceAuth: OAuth token has no key_bytes, fetching from DB"); - match sqlx::query_as::<_, (Vec, Option)>( - "SELECT k.key_bytes, k.encryption_version - FROM users u - JOIN user_keys k ON u.id = k.user_id - WHERE u.did = $1", - ) - .bind(&auth_user.did) - .fetch_optional(&state.db) - .await - { - Ok(Some((key_bytes_enc, encryption_version))) => { - match crate::config::decrypt_key(&key_bytes_enc, encryption_version) { - Ok(key) => key, - Err(e) => { - error!(error = ?e, "Failed to decrypt user key for service auth"); - return ApiError::AuthenticationFailed(Some( - "Failed to get signing key".into(), - )) - .into_response(); + match state.user_repo.get_user_info_by_did(&auth_user.did).await { + Ok(Some(info)) => match info.key_bytes { + Some(key_bytes_enc) => { + match crate::config::decrypt_key(&key_bytes_enc, info.encryption_version) { + Ok(key) => key, + Err(e) => { + error!(error = ?e, "Failed to decrypt user key for service auth"); + return ApiError::AuthenticationFailed(Some( + "Failed to get signing key".into(), + )) + .into_response(); + } } } - } + None => { + return ApiError::AuthenticationFailed(Some( + "User has no signing key".into(), + )) + .into_response(); + } + }, Ok(None) => { return ApiError::AuthenticationFailed(Some("User has no signing key".into())) .into_response(); @@ -196,17 +195,13 @@ pub async fn get_service_auth( } } - let user_status = sqlx::query!( - "SELECT takedown_ref FROM users WHERE did = $1", - &auth_user.did - ) - .fetch_optional(&state.db) - .await; - - let is_takendown = match user_status { - Ok(Some(row)) => row.takedown_ref.is_some(), - _ => false, - }; + let is_takendown = state + .user_repo + .get_status_by_did(&auth_user.did) + .await + .ok() + .flatten() + .is_some_and(|s| s.takedown_ref.is_some()); if is_takendown && lxm != Some("com.atproto.server.createAccount") { return ApiError::InvalidToken(Some("Bad token scope".into())).into_response(); diff --git a/crates/tranquil-pds/src/api/server/session.rs b/crates/tranquil-pds/src/api/server/session.rs index 6301306..4829aa6 100644 --- a/crates/tranquil-pds/src/api/server/session.rs +++ b/crates/tranquil-pds/src/api/server/session.rs @@ -13,6 +13,7 @@ use bcrypt::verify; use serde::{Deserialize, Serialize}; use serde_json::json; use tracing::{error, info, warn}; +use tranquil_types::TokenId; fn extract_client_ip(headers: &HeaderMap) -> String { if let Some(forwarded) = headers.get("x-forwarded-for") @@ -90,26 +91,16 @@ pub async fn create_session( return ApiError::RateLimitExceeded(None).into_response(); } let pds_hostname = std::env::var("PDS_HOSTNAME").unwrap_or_else(|_| "localhost".to_string()); - let normalized_identifier = normalize_handle(&input.identifier, &pds_hostname); + let hostname_for_handles = pds_hostname.split(':').next().unwrap_or(&pds_hostname); + let normalized_identifier = normalize_handle(&input.identifier, hostname_for_handles); info!( "Normalized identifier: {} -> {}", input.identifier, normalized_identifier ); - let row = match sqlx::query!( - r#"SELECT - u.id, u.did, u.handle, u.password_hash, u.email, u.deactivated_at, u.takedown_ref, - u.email_verified, u.discord_verified, u.telegram_verified, u.signal_verified, - u.allow_legacy_login, u.migrated_to_pds, - u.preferred_comms_channel as "preferred_comms_channel: crate::comms::CommsChannel", - k.key_bytes, k.encryption_version, - (SELECT verified FROM user_totp WHERE did = u.did) as totp_enabled - FROM users u - JOIN user_keys k ON u.id = k.user_id - WHERE u.handle = $1 OR u.email = $1 OR u.did = $1"#, - normalized_identifier - ) - .fetch_optional(&state.db) - .await + let row = match state + .user_repo + .get_login_full_by_identifier(&normalized_identifier) + .await { Ok(Some(row)) => row, Ok(None) => { @@ -141,13 +132,11 @@ pub async fn create_session( { (true, None, None, None) } else { - let app_passwords = sqlx::query!( - "SELECT name, password_hash, scopes, created_by_controller_did FROM app_passwords WHERE user_id = $1 ORDER BY created_at DESC LIMIT 20", - row.id - ) - .fetch_all(&state.db) - .await - .unwrap_or_default(); + let app_passwords = state + .session_repo + .get_app_passwords_for_login(row.id) + .await + .unwrap_or_default(); let matched = app_passwords .iter() .find(|app| verify(&input.password, &app.password_hash).unwrap_or(false)); @@ -178,7 +167,9 @@ pub async fn create_session( } let is_verified = row.email_verified || row.discord_verified || row.telegram_verified || row.signal_verified; - let is_delegated = crate::delegation::is_delegated_account(&state.db, &row.did) + let is_delegated = state + .delegation_repo + .is_delegated_account(&row.did) .await .unwrap_or(false); if !is_verified && !is_delegated { @@ -193,7 +184,7 @@ pub async fn create_session( ) .into_response(); } - let has_totp = row.totp_enabled.unwrap_or(false); + let has_totp = row.totp_enabled; let is_legacy_login = has_totp; if has_totp && !row.allow_legacy_login { warn!("Legacy login blocked for TOTP-enabled account: {}", row.did); @@ -229,21 +220,20 @@ pub async fn create_session( }; let did_for_doc = row.did.clone(); let did_resolver = state.did_resolver.clone(); + let session_data = tranquil_db_traits::SessionTokenCreate { + did: row.did.clone(), + access_jti: access_meta.jti.clone(), + refresh_jti: refresh_meta.jti.clone(), + access_expires_at: access_meta.expires_at, + refresh_expires_at: refresh_meta.expires_at, + legacy_login: is_legacy_login, + mfa_verified: false, + scope: app_password_scopes.clone(), + controller_did: app_password_controller.clone(), + app_password_name: app_password_name.clone(), + }; let (insert_result, did_doc) = tokio::join!( - sqlx::query!( - "INSERT INTO session_tokens (did, access_jti, refresh_jti, access_expires_at, refresh_expires_at, legacy_login, mfa_verified, scope, controller_did, app_password_name) VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10)", - row.did, - access_meta.jti, - refresh_meta.jti, - access_meta.expires_at, - refresh_meta.expires_at, - is_legacy_login, - false, - app_password_scopes, - app_password_controller, - app_password_name - ) - .execute(&state.db), + state.session_repo.create_session(&session_data), did_resolver.resolve_did_document(&did_for_doc) ); if let Err(e) = insert_result { @@ -257,8 +247,9 @@ pub async fn create_session( "Legacy login on TOTP-enabled account - sending notification" ); let hostname = std::env::var("PDS_HOSTNAME").unwrap_or_else(|_| "localhost".to_string()); - if let Err(e) = crate::comms::queue_legacy_login_notification( - &state.db, + if let Err(e) = crate::comms::comms_repo::enqueue_legacy_login( + state.user_repo.as_ref(), + state.infra_repo.as_ref(), row.id, &hostname, &client_ip, @@ -296,24 +287,16 @@ pub async fn get_session( let did_for_doc = auth_user.did.clone(); let did_resolver = state.did_resolver.clone(); let (db_result, did_doc) = tokio::join!( - sqlx::query!( - r#"SELECT - handle, email, email_verified, is_admin, deactivated_at, takedown_ref, preferred_locale, - preferred_comms_channel as "preferred_channel: crate::comms::CommsChannel", - discord_verified, telegram_verified, signal_verified, migrated_to_pds, migrated_at - FROM users WHERE did = $1"#, - &auth_user.did - ) - .fetch_optional(&state.db), + state.user_repo.get_session_info_by_did(&auth_user.did), did_resolver.resolve_did_document(&did_for_doc) ); match db_result { Ok(Some(row)) => { - let (preferred_channel, preferred_channel_verified) = match row.preferred_channel { - crate::comms::CommsChannel::Email => ("email", row.email_verified), - crate::comms::CommsChannel::Discord => ("discord", row.discord_verified), - crate::comms::CommsChannel::Telegram => ("telegram", row.telegram_verified), - crate::comms::CommsChannel::Signal => ("signal", row.signal_verified), + let (preferred_channel, preferred_channel_verified) = match row.preferred_comms_channel { + tranquil_db_traits::CommsChannel::Email => ("email", row.email_verified), + tranquil_db_traits::CommsChannel::Discord => ("discord", row.discord_verified), + tranquil_db_traits::CommsChannel::Telegram => ("telegram", row.telegram_verified), + tranquil_db_traits::CommsChannel::Signal => ("signal", row.signal_verified), }; let pds_hostname = std::env::var("PDS_HOSTNAME").unwrap_or_else(|_| "localhost".to_string()); @@ -379,11 +362,8 @@ pub async fn delete_session( Err(_) => return ApiError::AuthenticationFailed(None).into_response(), }; let did = crate::auth::get_did_from_token(&extracted.token).ok(); - match sqlx::query!("DELETE FROM session_tokens WHERE access_jti = $1", jti) - .execute(&state.db) - .await - { - Ok(res) if res.rows_affected() > 0 => { + match state.session_repo.delete_session_by_access_jti(&jti).await { + Ok(rows) if rows > 0 => { if let Some(did) = did { let session_cache_key = format!("auth:session:{}:{}", did, jti); let _ = state.cache.delete(&session_cache_key).await; @@ -391,10 +371,7 @@ pub async fn delete_session( EmptyResponse::ok().into_response() } Ok(_) => ApiError::AuthenticationFailed(None).into_response(), - Err(e) => { - error!("Database error in delete_session: {:?}", e); - ApiError::AuthenticationFailed(None).into_response() - } + Err(_) => ApiError::AuthenticationFailed(None).into_response(), } } @@ -424,44 +401,21 @@ pub async fn refresh_session( .into_response(); } }; - let mut tx = match state.db.begin().await { - Ok(tx) => tx, - Err(e) => { - error!("Failed to begin transaction: {:?}", e); - return ApiError::InternalError(None).into_response(); - } - }; - if let Ok(Some(session_id)) = sqlx::query_scalar!( - "SELECT session_id FROM used_refresh_tokens WHERE refresh_jti = $1 FOR UPDATE", - refresh_jti - ) - .fetch_optional(&mut *tx) - .await + if let Ok(Some(_)) = state + .session_repo + .check_refresh_token_used(&refresh_jti) + .await { - warn!( - "Refresh token reuse detected! Revoking token family for session_id: {}", - session_id - ); - let _ = sqlx::query!("DELETE FROM session_tokens WHERE id = $1", session_id) - .execute(&mut *tx) - .await; - let _ = tx.commit().await; + warn!("Refresh token reuse detected for jti: {}", refresh_jti); return ApiError::AuthenticationFailed(Some( "Refresh token has been revoked due to suspected compromise".into(), )) .into_response(); } - let session_row = match sqlx::query!( - r#"SELECT st.id, st.did, st.scope, st.controller_did, k.key_bytes, k.encryption_version - FROM session_tokens st - JOIN users u ON st.did = u.did - JOIN user_keys k ON u.id = k.user_id - WHERE st.refresh_jti = $1 AND st.refresh_expires_at > NOW() - FOR UPDATE OF st"#, - refresh_jti - ) - .fetch_optional(&mut *tx) - .await + let session_row = match state + .session_repo + .get_session_for_refresh(&refresh_jti) + .await { Ok(Some(row)) => row, Ok(None) => { @@ -474,7 +428,7 @@ pub async fn refresh_session( } }; let key_bytes = - match crate::config::decrypt_key(&session_row.key_bytes, session_row.encryption_version) { + match crate::config::decrypt_key(&session_row.key_bytes, Some(session_row.encryption_version)) { Ok(k) => k, Err(e) => { error!("Failed to decrypt user key: {:?}", e); @@ -506,67 +460,52 @@ pub async fn refresh_session( return ApiError::InternalError(None).into_response(); } }; - match sqlx::query!( - "INSERT INTO used_refresh_tokens (refresh_jti, session_id) VALUES ($1, $2) ON CONFLICT (refresh_jti) DO NOTHING", - refresh_jti, - session_row.id - ) - .execute(&mut *tx) - .await + let refresh_data = tranquil_db_traits::SessionRefreshData { + old_refresh_jti: refresh_jti.clone(), + session_id: session_row.id, + new_access_jti: new_access_meta.jti.clone(), + new_refresh_jti: new_refresh_meta.jti.clone(), + new_access_expires_at: new_access_meta.expires_at, + new_refresh_expires_at: new_refresh_meta.expires_at, + }; + match state + .session_repo + .refresh_session_atomic(&refresh_data) + .await { - Ok(result) if result.rows_affected() == 0 => { - warn!("Concurrent refresh token reuse detected for session_id: {}", session_row.id); - let _ = sqlx::query!("DELETE FROM session_tokens WHERE id = $1", session_row.id) - .execute(&mut *tx) - .await; - let _ = tx.commit().await; - return ApiError::AuthenticationFailed(Some("Refresh token has been revoked due to suspected compromise".into())).into_response(); + Ok(tranquil_db_traits::RefreshSessionResult::Success) => {} + Ok(tranquil_db_traits::RefreshSessionResult::TokenAlreadyUsed) => { + warn!("Refresh token reuse detected during atomic operation"); + return ApiError::AuthenticationFailed(Some( + "Refresh token has been revoked due to suspected compromise".into(), + )) + .into_response(); + } + Ok(tranquil_db_traits::RefreshSessionResult::ConcurrentRefresh) => { + warn!("Concurrent refresh detected for session_id: {}", session_row.id); + return ApiError::AuthenticationFailed(Some( + "Refresh token has been revoked due to suspected compromise".into(), + )) + .into_response(); } Err(e) => { - error!("Failed to record used refresh token: {:?}", e); + error!("Database error during session refresh: {:?}", e); return ApiError::InternalError(None).into_response(); } - Ok(_) => {} - } - if let Err(e) = sqlx::query!( - "UPDATE session_tokens SET access_jti = $1, refresh_jti = $2, access_expires_at = $3, refresh_expires_at = $4, updated_at = NOW() WHERE id = $5", - new_access_meta.jti, - new_refresh_meta.jti, - new_access_meta.expires_at, - new_refresh_meta.expires_at, - session_row.id - ) - .execute(&mut *tx) - .await - { - error!("Database error updating session: {:?}", e); - return ApiError::InternalError(None).into_response(); - } - if let Err(e) = tx.commit().await { - error!("Failed to commit transaction: {:?}", e); - return ApiError::InternalError(None).into_response(); } let did_for_doc = session_row.did.clone(); let did_resolver = state.did_resolver.clone(); let (db_result, did_doc) = tokio::join!( - sqlx::query!( - r#"SELECT - handle, email, email_verified, is_admin, preferred_locale, deactivated_at, takedown_ref, - preferred_comms_channel as "preferred_channel: crate::comms::CommsChannel", - discord_verified, telegram_verified, signal_verified - FROM users WHERE did = $1"#, - session_row.did - ) - .fetch_optional(&state.db), + state.user_repo.get_session_info_by_did(&session_row.did), did_resolver.resolve_did_document(&did_for_doc) ); match db_result { Ok(Some(u)) => { - let (preferred_channel, preferred_channel_verified) = match u.preferred_channel { - crate::comms::CommsChannel::Email => ("email", u.email_verified), - crate::comms::CommsChannel::Discord => ("discord", u.discord_verified), - crate::comms::CommsChannel::Telegram => ("telegram", u.telegram_verified), - crate::comms::CommsChannel::Signal => ("signal", u.signal_verified), + let (preferred_channel, preferred_channel_verified) = match u.preferred_comms_channel { + tranquil_db_traits::CommsChannel::Email => ("email", u.email_verified), + tranquil_db_traits::CommsChannel::Discord => ("discord", u.discord_verified), + tranquil_db_traits::CommsChannel::Telegram => ("telegram", u.telegram_verified), + tranquil_db_traits::CommsChannel::Signal => ("signal", u.signal_verified), }; let pds_hostname = std::env::var("PDS_HOSTNAME").unwrap_or_else(|_| "localhost".to_string()); @@ -630,19 +569,10 @@ pub async fn confirm_signup( Json(input): Json, ) -> Response { info!("confirm_signup called for DID: {}", input.did); - let row = match sqlx::query!( - r#"SELECT - u.id, u.did, u.handle, u.email, - u.preferred_comms_channel as "channel: crate::comms::CommsChannel", - u.discord_id, u.telegram_username, u.signal_number, - k.key_bytes, k.encryption_version - FROM users u - JOIN user_keys k ON u.id = k.user_id - WHERE u.did = $1"#, - input.did.as_str() - ) - .fetch_optional(&state.db) - .await + let row = match state + .user_repo + .get_confirm_signup_by_did(&input.did) + .await { Ok(Some(row)) => row, Ok(None) => { @@ -657,15 +587,15 @@ pub async fn confirm_signup( }; let (channel_str, identifier) = match row.channel { - crate::comms::CommsChannel::Email => ("email", row.email.clone().unwrap_or_default()), - crate::comms::CommsChannel::Discord => { + tranquil_db_traits::CommsChannel::Email => ("email", row.email.clone().unwrap_or_default()), + tranquil_db_traits::CommsChannel::Discord => { ("discord", row.discord_id.clone().unwrap_or_default()) } - crate::comms::CommsChannel::Telegram => ( + tranquil_db_traits::CommsChannel::Telegram => ( "telegram", row.telegram_username.clone().unwrap_or_default(), ), - crate::comms::CommsChannel::Signal => { + tranquil_db_traits::CommsChannel::Signal => { ("signal", row.signal_number.clone().unwrap_or_default()) } }; @@ -721,64 +651,49 @@ pub async fn confirm_signup( } }; - let mut tx = match state.db.begin().await { - Ok(tx) => tx, - Err(e) => { - error!("Failed to begin transaction: {:?}", e); - return ApiError::InternalError(None).into_response(); - } - }; - - let verified_column = match row.channel { - crate::comms::CommsChannel::Email => "email_verified", - crate::comms::CommsChannel::Discord => "discord_verified", - crate::comms::CommsChannel::Telegram => "telegram_verified", - crate::comms::CommsChannel::Signal => "signal_verified", - }; - let update_query = format!("UPDATE users SET {} = TRUE WHERE did = $1", verified_column); - if let Err(e) = sqlx::query(&update_query) - .bind(input.did.as_str()) - .execute(&mut *tx) + if let Err(e) = state + .user_repo + .set_channel_verified(&input.did, row.channel.clone()) .await { error!("Failed to update verification status: {:?}", e); return ApiError::InternalError(None).into_response(); } - let no_scope: Option = None; - if let Err(e) = sqlx::query!( - "INSERT INTO session_tokens (did, access_jti, refresh_jti, access_expires_at, refresh_expires_at, legacy_login, mfa_verified, scope) VALUES ($1, $2, $3, $4, $5, $6, $7, $8)", - row.did, - access_meta.jti, - refresh_meta.jti, - access_meta.expires_at, - refresh_meta.expires_at, - false, - false, - no_scope - ) - .execute(&mut *tx) - .await - { + let session_data = tranquil_db_traits::SessionTokenCreate { + did: row.did.clone(), + access_jti: access_meta.jti.clone(), + refresh_jti: refresh_meta.jti.clone(), + access_expires_at: access_meta.expires_at, + refresh_expires_at: refresh_meta.expires_at, + legacy_login: false, + mfa_verified: false, + scope: None, + controller_did: None, + app_password_name: None, + }; + if let Err(e) = state.session_repo.create_session(&session_data).await { error!("Failed to insert session: {:?}", e); return ApiError::InternalError(None).into_response(); } - if let Err(e) = tx.commit().await { - error!("Failed to commit transaction: {:?}", e); - return ApiError::InternalError(None).into_response(); - } - let hostname = std::env::var("PDS_HOSTNAME").unwrap_or_else(|_| "localhost".to_string()); - if let Err(e) = crate::comms::enqueue_welcome(&state.db, row.id, &hostname).await { + if let Err(e) = crate::comms::comms_repo::enqueue_welcome( + state.user_repo.as_ref(), + state.infra_repo.as_ref(), + row.id, + &hostname, + ) + .await + { warn!("Failed to enqueue welcome notification: {:?}", e); } - let email_verified = matches!(row.channel, crate::comms::CommsChannel::Email); + let email_verified = matches!(row.channel, tranquil_db_traits::CommsChannel::Email); let preferred_channel = match row.channel { - crate::comms::CommsChannel::Email => "email", - crate::comms::CommsChannel::Discord => "discord", - crate::comms::CommsChannel::Telegram => "telegram", - crate::comms::CommsChannel::Signal => "signal", + tranquil_db_traits::CommsChannel::Email => "email", + tranquil_db_traits::CommsChannel::Discord => "discord", + tranquil_db_traits::CommsChannel::Telegram => "telegram", + tranquil_db_traits::CommsChannel::Signal => "signal", }; Json(ConfirmSignupOutput { access_jwt: access_meta.token, @@ -804,18 +719,10 @@ pub async fn resend_verification( Json(input): Json, ) -> Response { info!("resend_verification called for DID: {}", input.did); - let row = match sqlx::query!( - r#"SELECT - id, handle, email, - preferred_comms_channel as "channel: crate::comms::CommsChannel", - discord_id, telegram_username, signal_number, - email_verified, discord_verified, telegram_verified, signal_verified - FROM users - WHERE did = $1"#, - input.did.as_str() - ) - .fetch_optional(&state.db) - .await + let row = match state + .user_repo + .get_resend_verification_by_did(&input.did) + .await { Ok(Some(row)) => row, Ok(None) => { @@ -833,15 +740,15 @@ pub async fn resend_verification( } let (channel_str, recipient) = match row.channel { - crate::comms::CommsChannel::Email => ("email", row.email.clone().unwrap_or_default()), - crate::comms::CommsChannel::Discord => { + tranquil_db_traits::CommsChannel::Email => ("email", row.email.clone().unwrap_or_default()), + tranquil_db_traits::CommsChannel::Discord => { ("discord", row.discord_id.clone().unwrap_or_default()) } - crate::comms::CommsChannel::Telegram => ( + tranquil_db_traits::CommsChannel::Telegram => ( "telegram", row.telegram_username.clone().unwrap_or_default(), ), - crate::comms::CommsChannel::Signal => { + tranquil_db_traits::CommsChannel::Signal => { ("signal", row.signal_number.clone().unwrap_or_default()) } }; @@ -851,13 +758,14 @@ pub async fn resend_verification( let formatted_token = crate::auth::verification_token::format_token_for_display(&verification_token); - if let Err(e) = crate::comms::enqueue_signup_verification( - &state.db, + let hostname = std::env::var("PDS_HOSTNAME").unwrap_or_else(|_| "localhost".to_string()); + if let Err(e) = crate::comms::comms_repo::enqueue_signup_verification( + state.infra_repo.as_ref(), row.id, channel_str, &recipient, &formatted_token, - None, + &hostname, ) .await { @@ -894,26 +802,7 @@ pub async fn list_sessions( .and_then(|v| v.strip_prefix("Bearer ")) .and_then(|token| crate::auth::get_jti_from_token(token).ok()); - let jwt_rows = match sqlx::query_as::< - _, - ( - i32, - String, - chrono::DateTime, - chrono::DateTime, - ), - >( - r#" - SELECT id, access_jti, created_at, refresh_expires_at - FROM session_tokens - WHERE did = $1 AND refresh_expires_at > NOW() - ORDER BY created_at DESC - "#, - ) - .bind(&auth.0.did) - .fetch_all(&state.db) - .await - { + let jwt_rows = match state.session_repo.list_sessions_by_did(&auth.0.did).await { Ok(rows) => rows, Err(e) => { error!("DB error fetching JWT sessions: {:?}", e); @@ -921,27 +810,7 @@ pub async fn list_sessions( } }; - let oauth_rows = match sqlx::query_as::< - _, - ( - i32, - String, - chrono::DateTime, - chrono::DateTime, - String, - ), - >( - r#" - SELECT id, token_id, created_at, expires_at, client_id - FROM oauth_token - WHERE did = $1 AND expires_at > NOW() - ORDER BY created_at DESC - "#, - ) - .bind(&auth.0.did) - .fetch_all(&state.db) - .await - { + let oauth_rows = match state.oauth_repo.list_sessions_by_did(&auth.0.did).await { Ok(rows) => rows, Err(e) => { error!("DB error fetching OAuth sessions: {:?}", e); @@ -949,33 +818,28 @@ pub async fn list_sessions( } }; - let jwt_sessions = jwt_rows - .into_iter() - .map(|(id, access_jti, created_at, expires_at)| SessionInfo { - id: format!("jwt:{}", id), - session_type: "legacy".to_string(), - client_name: None, - created_at: created_at.to_rfc3339(), - expires_at: expires_at.to_rfc3339(), - is_current: current_jti.as_ref() == Some(&access_jti), - }); + let jwt_sessions = jwt_rows.into_iter().map(|row| SessionInfo { + id: format!("jwt:{}", row.id), + session_type: "legacy".to_string(), + client_name: None, + created_at: row.created_at.to_rfc3339(), + expires_at: row.refresh_expires_at.to_rfc3339(), + is_current: current_jti.as_ref() == Some(&row.access_jti), + }); let is_oauth = auth.0.is_oauth; - let oauth_sessions = - oauth_rows - .into_iter() - .map(|(id, token_id, created_at, expires_at, client_id)| { - let client_name = extract_client_name(&client_id); - let is_current_oauth = is_oauth && current_jti.as_ref() == Some(&token_id); - SessionInfo { - id: format!("oauth:{}", id), - session_type: "oauth".to_string(), - client_name: Some(client_name), - created_at: created_at.to_rfc3339(), - expires_at: expires_at.to_rfc3339(), - is_current: is_current_oauth, - } - }); + let oauth_sessions = oauth_rows.into_iter().map(|row| { + let client_name = extract_client_name(&row.client_id); + let is_current_oauth = is_oauth && current_jti.as_ref().map(|s| s.as_str()) == Some(row.token_id.as_str()); + SessionInfo { + id: format!("oauth:{}", row.id), + session_type: "oauth".to_string(), + client_name: Some(client_name), + created_at: row.created_at.to_rfc3339(), + expires_at: row.expires_at.to_rfc3339(), + is_current: is_current_oauth, + } + }); let mut sessions: Vec = jwt_sessions.chain(oauth_sessions).collect(); sessions.sort_by(|a, b| b.created_at.cmp(&a.created_at)); @@ -1008,15 +872,12 @@ pub async fn revoke_session( let Ok(session_id) = jwt_id.parse::() else { return ApiError::InvalidRequest("Invalid session ID".into()).into_response(); }; - let session = sqlx::query_as::<_, (String,)>( - "SELECT access_jti FROM session_tokens WHERE id = $1 AND did = $2", - ) - .bind(session_id) - .bind(&auth.0.did) - .fetch_optional(&state.db) - .await; - let access_jti = match session { - Ok(Some((jti,))) => jti, + let access_jti = match state + .session_repo + .get_session_access_jti_by_id(session_id, &auth.0.did) + .await + { + Ok(Some(jti)) => jti, Ok(None) => { return ApiError::SessionNotFound.into_response(); } @@ -1025,11 +886,7 @@ pub async fn revoke_session( return ApiError::InternalError(None).into_response(); } }; - if let Err(e) = sqlx::query("DELETE FROM session_tokens WHERE id = $1") - .bind(session_id) - .execute(&state.db) - .await - { + if let Err(e) = state.session_repo.delete_session_by_id(session_id).await { error!("DB error deleting session: {:?}", e); return ApiError::InternalError(None).into_response(); } @@ -1042,13 +899,12 @@ pub async fn revoke_session( let Ok(session_id) = oauth_id.parse::() else { return ApiError::InvalidRequest("Invalid session ID".into()).into_response(); }; - let result = sqlx::query("DELETE FROM oauth_token WHERE id = $1 AND did = $2") - .bind(session_id) - .bind(&auth.0.did) - .execute(&state.db) - .await; - match result { - Ok(r) if r.rows_affected() == 0 => { + match state + .oauth_repo + .delete_session_by_id(session_id, &auth.0.did) + .await + { + Ok(0) => { return ApiError::SessionNotFound.into_response(); } Err(e) => { @@ -1078,58 +934,35 @@ pub async fn revoke_all_sessions( return ApiError::InvalidToken(None).into_response(); }; - let mut tx = match state.db.begin().await { - Ok(tx) => tx, - Err(e) => { - error!("Failed to begin transaction: {:?}", e); - return ApiError::InternalError(None).into_response(); - } - }; - if auth.0.is_oauth { - if let Err(e) = sqlx::query("DELETE FROM session_tokens WHERE did = $1") - .bind(&auth.0.did) - .execute(&mut *tx) - .await - { + if let Err(e) = state.session_repo.delete_sessions_by_did(&auth.0.did).await { error!("DB error revoking JWT sessions: {:?}", e); return ApiError::InternalError(None).into_response(); } - if let Err(e) = sqlx::query("DELETE FROM oauth_token WHERE did = $1 AND token_id != $2") - .bind(&auth.0.did) - .bind(jti) - .execute(&mut *tx) + let jti_typed = TokenId::from(jti.clone()); + if let Err(e) = state + .oauth_repo + .delete_sessions_by_did_except(&auth.0.did, &jti_typed) .await { error!("DB error revoking OAuth sessions: {:?}", e); return ApiError::InternalError(None).into_response(); } } else { - if let Err(e) = - sqlx::query("DELETE FROM session_tokens WHERE did = $1 AND access_jti != $2") - .bind(&auth.0.did) - .bind(jti) - .execute(&mut *tx) - .await + if let Err(e) = state + .session_repo + .delete_sessions_by_did_except_jti(&auth.0.did, jti) + .await { error!("DB error revoking JWT sessions: {:?}", e); return ApiError::InternalError(None).into_response(); } - if let Err(e) = sqlx::query("DELETE FROM oauth_token WHERE did = $1") - .bind(&auth.0.did) - .execute(&mut *tx) - .await - { + if let Err(e) = state.oauth_repo.delete_sessions_by_did(&auth.0.did).await { error!("DB error revoking OAuth sessions: {:?}", e); return ApiError::InternalError(None).into_response(); } } - if let Err(e) = tx.commit().await { - error!("Failed to commit transaction: {:?}", e); - return ApiError::InternalError(None).into_response(); - } - info!(did = %&auth.0.did, "All other sessions revoked"); SuccessResponse::ok().into_response() } @@ -1145,21 +978,10 @@ pub async fn get_legacy_login_preference( State(state): State, auth: BearerAuth, ) -> Response { - let result = sqlx::query!( - r#"SELECT - u.allow_legacy_login, - (EXISTS(SELECT 1 FROM user_totp t WHERE t.did = u.did AND t.verified = TRUE) OR - EXISTS(SELECT 1 FROM passkeys p WHERE p.did = u.did)) as "has_mfa!" - FROM users u WHERE u.did = $1"#, - &auth.0.did - ) - .fetch_optional(&state.db) - .await; - - match result { - Ok(Some(row)) => Json(LegacyLoginPreferenceOutput { - allow_legacy_login: row.allow_legacy_login, - has_mfa: row.has_mfa, + match state.user_repo.get_legacy_login_pref(&auth.0.did).await { + Ok(Some(pref)) => Json(LegacyLoginPreferenceOutput { + allow_legacy_login: pref.allow_legacy_login, + has_mfa: pref.has_mfa, }) .into_response(), Ok(None) => ApiError::AccountNotFound.into_response(), @@ -1181,25 +1003,21 @@ pub async fn update_legacy_login_preference( auth: BearerAuth, Json(input): Json, ) -> Response { - if !crate::api::server::reauth::check_legacy_session_mfa(&state.db, &auth.0.did).await { - return crate::api::server::reauth::legacy_mfa_required_response(&state.db, &auth.0.did) + if !crate::api::server::reauth::check_legacy_session_mfa(&*state.session_repo, &auth.0.did).await { + return crate::api::server::reauth::legacy_mfa_required_response(&*state.user_repo, &*state.session_repo, &auth.0.did) .await; } - if crate::api::server::reauth::check_reauth_required(&state.db, &auth.0.did).await { - return crate::api::server::reauth::reauth_required_response(&state.db, &auth.0.did).await; + if crate::api::server::reauth::check_reauth_required(&*state.session_repo, &auth.0.did).await { + return crate::api::server::reauth::reauth_required_response(&*state.user_repo, &*state.session_repo, &auth.0.did).await; } - let result = sqlx::query!( - "UPDATE users SET allow_legacy_login = $1 WHERE did = $2 RETURNING did", - input.allow_legacy_login, - &auth.0.did - ) - .fetch_optional(&state.db) - .await; - - match result { - Ok(Some(_)) => { + match state + .user_repo + .update_legacy_login(&auth.0.did, input.allow_legacy_login) + .await + { + Ok(true) => { info!( did = %&auth.0.did, allow_legacy_login = input.allow_legacy_login, @@ -1210,7 +1028,7 @@ pub async fn update_legacy_login_preference( })) .into_response() } - Ok(None) => ApiError::AccountNotFound.into_response(), + Ok(false) => ApiError::AccountNotFound.into_response(), Err(e) => { error!("DB error: {:?}", e); ApiError::InternalError(None).into_response() @@ -1239,16 +1057,12 @@ pub async fn update_locale( .into_response(); } - let result = sqlx::query!( - "UPDATE users SET preferred_locale = $1 WHERE did = $2 RETURNING did", - input.preferred_locale, - &auth.0.did - ) - .fetch_optional(&state.db) - .await; - - match result { - Ok(Some(_)) => { + match state + .user_repo + .update_locale(&auth.0.did, &input.preferred_locale) + .await + { + Ok(true) => { info!( did = %&auth.0.did, locale = %input.preferred_locale, @@ -1259,7 +1073,7 @@ pub async fn update_locale( })) .into_response() } - Ok(None) => ApiError::AccountNotFound.into_response(), + Ok(false) => ApiError::AccountNotFound.into_response(), Err(e) => { error!("DB error updating locale: {:?}", e); ApiError::InternalError(None).into_response() diff --git a/crates/tranquil-pds/src/api/server/signing_key.rs b/crates/tranquil-pds/src/api/server/signing_key.rs index fc2eac2..cf47a4f 100644 --- a/crates/tranquil-pds/src/api/server/signing_key.rs +++ b/crates/tranquil-pds/src/api/server/signing_key.rs @@ -38,27 +38,30 @@ pub async fn reserve_signing_key( State(state): State, Json(input): Json, ) -> Response { + let did: Option = match input.did { + Some(ref d) => match d.parse() { + Ok(parsed) => Some(parsed), + Err(_) => return ApiError::InvalidDid("Invalid DID format".into()).into_response(), + }, + None => None, + }; let signing_key = SigningKey::random(&mut rand::thread_rng()); let private_key_bytes = signing_key.to_bytes(); let public_key_did_key = public_key_to_did_key(&signing_key); let expires_at = Utc::now() + Duration::hours(24); let private_bytes: &[u8] = &private_key_bytes; - let result = sqlx::query!( - r#" - INSERT INTO reserved_signing_keys (did, public_key_did_key, private_key_bytes, expires_at) - VALUES ($1, $2, $3, $4) - RETURNING id - "#, - input.did, - public_key_did_key, - private_bytes, - expires_at - ) - .fetch_one(&state.db) - .await; - match result { - Ok(row) => { - info!("Reserved signing key {} for did {:?}", row.id, input.did); + match state + .infra_repo + .reserve_signing_key( + did.as_ref(), + &public_key_did_key, + private_bytes, + expires_at, + ) + .await + { + Ok(key_id) => { + info!("Reserved signing key {} for did {:?}", key_id, input.did); ( StatusCode::OK, Json(ReserveSigningKeyOutput { diff --git a/crates/tranquil-pds/src/api/server/totp.rs b/crates/tranquil-pds/src/api/server/totp.rs index 74e6b4c..7c254fb 100644 --- a/crates/tranquil-pds/src/api/server/totp.rs +++ b/crates/tranquil-pds/src/api/server/totp.rs @@ -13,7 +13,6 @@ use axum::{ extract::State, response::{IntoResponse, Response}, }; -use chrono::Utc; use serde::{Deserialize, Serialize}; use tracing::{error, info, warn}; @@ -28,24 +27,18 @@ pub struct CreateTotpSecretResponse { } pub async fn create_totp_secret(State(state): State, auth: BearerAuth) -> Response { - let existing = sqlx::query_scalar!( - "SELECT verified FROM user_totp WHERE did = $1", - &*&auth.0.did - ) - .fetch_optional(&state.db) - .await; - - if let Ok(Some(true)) = existing { - return ApiError::TotpAlreadyEnabled.into_response(); + match state.user_repo.get_totp_record(&auth.0.did).await { + Ok(Some(record)) if record.verified => return ApiError::TotpAlreadyEnabled.into_response(), + Ok(_) => {} + Err(e) => { + error!("DB error checking TOTP: {:?}", e); + return ApiError::InternalError(None).into_response(); + } } let secret = generate_totp_secret(); - let handle = sqlx::query_scalar!("SELECT handle FROM users WHERE did = $1", &*&auth.0.did) - .fetch_optional(&state.db) - .await; - - let handle = match handle { + let handle = match state.user_repo.get_handle_by_did(&auth.0.did).await { Ok(Some(h)) => h, Ok(None) => return ApiError::AccountNotFound.into_response(), Err(e) => { @@ -74,25 +67,11 @@ pub async fn create_totp_secret(State(state): State, auth: BearerAuth) } }; - let result = sqlx::query!( - r#" - INSERT INTO user_totp (did, secret_encrypted, encryption_version, verified, created_at) - VALUES ($1, $2, $3, false, NOW()) - ON CONFLICT (did) DO UPDATE SET - secret_encrypted = $2, - encryption_version = $3, - verified = false, - created_at = NOW(), - last_used = NULL - "#, - &auth.0.did, - encrypted_secret, - ENCRYPTION_VERSION - ) - .execute(&state.db) - .await; - - if let Err(e) = result { + if let Err(e) = state + .user_repo + .upsert_totp_secret(&auth.0.did, &encrypted_secret, ENCRYPTION_VERSION) + .await + { error!("Failed to store TOTP secret: {:?}", e); return ApiError::InternalError(None).into_response(); } @@ -133,14 +112,7 @@ pub async fn enable_totp( return ApiError::RateLimitExceeded(None).into_response(); } - let totp_row = sqlx::query!( - "SELECT secret_encrypted, encryption_version, verified FROM user_totp WHERE did = $1", - &auth.0.did - ) - .fetch_optional(&state.db) - .await; - - let totp_row = match totp_row { + let totp_record = match state.user_repo.get_totp_record(&auth.0.did).await { Ok(Some(row)) => row, Ok(None) => return ApiError::TotpNotEnabled.into_response(), Err(e) => { @@ -149,18 +121,18 @@ pub async fn enable_totp( } }; - if totp_row.verified { + if totp_record.verified { return ApiError::TotpAlreadyEnabled.into_response(); } - let secret = match decrypt_totp_secret(&totp_row.secret_encrypted, totp_row.encryption_version) - { - Ok(s) => s, - Err(e) => { - error!("Failed to decrypt TOTP secret: {:?}", e); - return ApiError::InternalError(None).into_response(); - } - }; + let secret = + match decrypt_totp_secret(&totp_record.secret_encrypted, totp_record.encryption_version) { + Ok(s) => s, + Err(e) => { + error!("Failed to decrypt TOTP secret: {:?}", e); + return ApiError::InternalError(None).into_response(); + } + }; let code = input.code.trim(); if !verify_totp_code(&secret, code) { @@ -168,33 +140,6 @@ pub async fn enable_totp( } let backup_codes = generate_backup_codes(); - let mut tx = match state.db.begin().await { - Ok(tx) => tx, - Err(e) => { - error!("Failed to begin transaction: {:?}", e); - return ApiError::InternalError(None).into_response(); - } - }; - - if let Err(e) = sqlx::query!( - "UPDATE user_totp SET verified = true, last_used = NOW() WHERE did = $1", - &auth.0.did - ) - .execute(&mut *tx) - .await - { - error!("Failed to enable TOTP: {:?}", e); - return ApiError::InternalError(None).into_response(); - } - - if let Err(e) = sqlx::query!("DELETE FROM backup_codes WHERE did = $1", &*&auth.0.did) - .execute(&mut *tx) - .await - { - error!("Failed to clear old backup codes: {:?}", e); - return ApiError::InternalError(None).into_response(); - } - let backup_hashes: Result, _> = backup_codes.iter().map(|c| hash_backup_code(c)).collect(); let backup_hashes = match backup_hashes { @@ -205,23 +150,12 @@ pub async fn enable_totp( } }; - if let Err(e) = sqlx::query!( - r#" - INSERT INTO backup_codes (did, code_hash, created_at) - SELECT $1, hash, NOW() FROM UNNEST($2::text[]) AS t(hash) - "#, - &auth.0.did, - &backup_hashes[..] - ) - .execute(&mut *tx) - .await + if let Err(e) = state + .user_repo + .enable_totp_with_backup_codes(&auth.0.did, &backup_hashes) + .await { - error!("Failed to store backup codes: {:?}", e); - return ApiError::InternalError(None).into_response(); - } - - if let Err(e) = tx.commit().await { - error!("Failed to commit transaction: {:?}", e); + error!("Failed to enable TOTP: {:?}", e); return ApiError::InternalError(None).into_response(); } @@ -241,9 +175,14 @@ pub async fn disable_totp( auth: BearerAuth, Json(input): Json, ) -> Response { - if !crate::api::server::reauth::check_legacy_session_mfa(&state.db, &auth.0.did).await { - return crate::api::server::reauth::legacy_mfa_required_response(&state.db, &auth.0.did) - .await; + if !crate::api::server::reauth::check_legacy_session_mfa(&*state.session_repo, &auth.0.did).await + { + return crate::api::server::reauth::legacy_mfa_required_response( + &*state.user_repo, + &*state.session_repo, + &auth.0.did, + ) + .await; } if !state @@ -254,15 +193,8 @@ pub async fn disable_totp( return ApiError::RateLimitExceeded(None).into_response(); } - let user = sqlx::query!( - "SELECT password_hash FROM users WHERE did = $1", - &*&auth.0.did - ) - .fetch_optional(&state.db) - .await; - - let password_hash = match user { - Ok(Some(row)) => row.password_hash, + let password_hash = match state.user_repo.get_password_hash_by_did(&auth.0.did).await { + Ok(Some(hash)) => hash, Ok(None) => return ApiError::AccountNotFound.into_response(), Err(e) => { error!("DB error fetching user: {:?}", e); @@ -270,22 +202,12 @@ pub async fn disable_totp( } }; - let password_valid = password_hash - .as_ref() - .map(|h| bcrypt::verify(&input.password, h).unwrap_or(false)) - .unwrap_or(false); + let password_valid = bcrypt::verify(&input.password, &password_hash).unwrap_or(false); if !password_valid { return ApiError::InvalidPassword("Password is incorrect".into()).into_response(); } - let totp_row = sqlx::query!( - "SELECT secret_encrypted, encryption_version, verified FROM user_totp WHERE did = $1", - &auth.0.did - ) - .fetch_optional(&state.db) - .await; - - let totp_row = match totp_row { + let totp_record = match state.user_repo.get_totp_record(&auth.0.did).await { Ok(Some(row)) if row.verified => row, Ok(Some(_)) | Ok(None) => return ApiError::TotpNotEnabled.into_response(), Err(e) => { @@ -298,14 +220,16 @@ pub async fn disable_totp( let code_valid = if is_backup_code_format(code) { verify_backup_code_for_user(&state, &auth.0.did, code).await } else { - let secret = - match decrypt_totp_secret(&totp_row.secret_encrypted, totp_row.encryption_version) { - Ok(s) => s, - Err(e) => { - error!("Failed to decrypt TOTP secret: {:?}", e); - return ApiError::InternalError(None).into_response(); - } - }; + let secret = match decrypt_totp_secret( + &totp_record.secret_encrypted, + totp_record.encryption_version, + ) { + Ok(s) => s, + Err(e) => { + error!("Failed to decrypt TOTP secret: {:?}", e); + return ApiError::InternalError(None).into_response(); + } + }; verify_totp_code(&secret, code) }; @@ -313,35 +237,15 @@ pub async fn disable_totp( return ApiError::InvalidCode(Some("Invalid verification code".into())).into_response(); } - let mut tx = match state.db.begin().await { - Ok(tx) => tx, - Err(e) => { - error!("Failed to begin transaction: {:?}", e); - return ApiError::InternalError(None).into_response(); - } - }; - - if let Err(e) = sqlx::query!("DELETE FROM user_totp WHERE did = $1", &*&auth.0.did) - .execute(&mut *tx) + if let Err(e) = state + .user_repo + .delete_totp_and_backup_codes(&auth.0.did) .await { error!("Failed to delete TOTP: {:?}", e); return ApiError::InternalError(None).into_response(); } - if let Err(e) = sqlx::query!("DELETE FROM backup_codes WHERE did = $1", &*&auth.0.did) - .execute(&mut *tx) - .await - { - error!("Failed to delete backup codes: {:?}", e); - return ApiError::InternalError(None).into_response(); - } - - if let Err(e) = tx.commit().await { - error!("Failed to commit transaction: {:?}", e); - return ApiError::InternalError(None).into_response(); - } - info!(did = %&auth.0.did, "TOTP disabled"); EmptyResponse::ok().into_response() @@ -356,14 +260,7 @@ pub struct GetTotpStatusResponse { } pub async fn get_totp_status(State(state): State, auth: BearerAuth) -> Response { - let totp_row = sqlx::query!( - "SELECT verified FROM user_totp WHERE did = $1", - &*&auth.0.did - ) - .fetch_optional(&state.db) - .await; - - let enabled = match totp_row { + let enabled = match state.user_repo.get_totp_record(&auth.0.did).await { Ok(Some(row)) => row.verified, Ok(None) => false, Err(e) => { @@ -372,14 +269,13 @@ pub async fn get_totp_status(State(state): State, auth: BearerAuth) -> } }; - let backup_count_row = sqlx::query!( - "SELECT COUNT(*) as count FROM backup_codes WHERE did = $1 AND used_at IS NULL", - &auth.0.did - ) - .fetch_one(&state.db) - .await; - - let backup_count = backup_count_row.map(|r| r.count.unwrap_or(0)).unwrap_or(0); + let backup_count = match state.user_repo.count_unused_backup_codes(&auth.0.did).await { + Ok(count) => count, + Err(e) => { + error!("DB error counting backup codes: {:?}", e); + return ApiError::InternalError(None).into_response(); + } + }; Json(GetTotpStatusResponse { enabled, @@ -414,15 +310,8 @@ pub async fn regenerate_backup_codes( return ApiError::RateLimitExceeded(None).into_response(); } - let user = sqlx::query!( - "SELECT password_hash FROM users WHERE did = $1", - &*&auth.0.did - ) - .fetch_optional(&state.db) - .await; - - let password_hash = match user { - Ok(Some(row)) => row.password_hash, + let password_hash = match state.user_repo.get_password_hash_by_did(&auth.0.did).await { + Ok(Some(hash)) => hash, Ok(None) => return ApiError::AccountNotFound.into_response(), Err(e) => { error!("DB error fetching user: {:?}", e); @@ -430,22 +319,12 @@ pub async fn regenerate_backup_codes( } }; - let password_valid = password_hash - .as_ref() - .map(|h| bcrypt::verify(&input.password, h).unwrap_or(false)) - .unwrap_or(false); + let password_valid = bcrypt::verify(&input.password, &password_hash).unwrap_or(false); if !password_valid { return ApiError::InvalidPassword("Password is incorrect".into()).into_response(); } - let totp_row = sqlx::query!( - "SELECT secret_encrypted, encryption_version, verified FROM user_totp WHERE did = $1", - &auth.0.did - ) - .fetch_optional(&state.db) - .await; - - let totp_row = match totp_row { + let totp_record = match state.user_repo.get_totp_record(&auth.0.did).await { Ok(Some(row)) if row.verified => row, Ok(Some(_)) | Ok(None) => return ApiError::TotpNotEnabled.into_response(), Err(e) => { @@ -454,14 +333,14 @@ pub async fn regenerate_backup_codes( } }; - let secret = match decrypt_totp_secret(&totp_row.secret_encrypted, totp_row.encryption_version) - { - Ok(s) => s, - Err(e) => { - error!("Failed to decrypt TOTP secret: {:?}", e); - return ApiError::InternalError(None).into_response(); - } - }; + let secret = + match decrypt_totp_secret(&totp_record.secret_encrypted, totp_record.encryption_version) { + Ok(s) => s, + Err(e) => { + error!("Failed to decrypt TOTP secret: {:?}", e); + return ApiError::InternalError(None).into_response(); + } + }; let code = input.code.trim(); if !verify_totp_code(&secret, code) { @@ -469,22 +348,6 @@ pub async fn regenerate_backup_codes( } let backup_codes = generate_backup_codes(); - let mut tx = match state.db.begin().await { - Ok(tx) => tx, - Err(e) => { - error!("Failed to begin transaction: {:?}", e); - return ApiError::InternalError(None).into_response(); - } - }; - - if let Err(e) = sqlx::query!("DELETE FROM backup_codes WHERE did = $1", &*&auth.0.did) - .execute(&mut *tx) - .await - { - error!("Failed to clear old backup codes: {:?}", e); - return ApiError::InternalError(None).into_response(); - } - let backup_hashes: Result, _> = backup_codes.iter().map(|c| hash_backup_code(c)).collect(); let backup_hashes = match backup_hashes { @@ -495,23 +358,12 @@ pub async fn regenerate_backup_codes( } }; - if let Err(e) = sqlx::query!( - r#" - INSERT INTO backup_codes (did, code_hash, created_at) - SELECT $1, hash, NOW() FROM UNNEST($2::text[]) AS t(hash) - "#, - &auth.0.did, - &backup_hashes[..] - ) - .execute(&mut *tx) - .await + if let Err(e) = state + .user_repo + .replace_backup_codes(&auth.0.did, &backup_hashes) + .await { - error!("Failed to store backup codes: {:?}", e); - return ApiError::InternalError(None).into_response(); - } - - if let Err(e) = tx.commit().await { - error!("Failed to commit transaction: {:?}", e); + error!("Failed to regenerate backup codes: {:?}", e); return ApiError::InternalError(None).into_response(); } @@ -520,17 +372,10 @@ pub async fn regenerate_backup_codes( Json(RegenerateBackupCodesResponse { backup_codes }).into_response() } -async fn verify_backup_code_for_user(state: &AppState, did: &str, code: &str) -> bool { +async fn verify_backup_code_for_user(state: &AppState, did: &crate::types::Did, code: &str) -> bool { let code = code.trim().to_uppercase(); - let backup_codes = sqlx::query!( - "SELECT id, code_hash FROM backup_codes WHERE did = $1 AND used_at IS NULL", - did - ) - .fetch_all(&state.db) - .await; - - let backup_codes = match backup_codes { + let backup_codes = match state.user_repo.get_unused_backup_codes(did).await { Ok(codes) => codes, Err(e) => { warn!("Failed to fetch backup codes: {:?}", e); @@ -544,62 +389,43 @@ async fn verify_backup_code_for_user(state: &AppState, did: &str, code: &str) -> match matched { Some(row) => { - let _ = sqlx::query!( - "UPDATE backup_codes SET used_at = $1 WHERE id = $2", - Utc::now(), - row.id - ) - .execute(&state.db) - .await; + let _ = state.user_repo.mark_backup_code_used(row.id).await; true } None => false, } } -pub async fn verify_totp_or_backup_for_user(state: &AppState, did: &str, code: &str) -> bool { +pub async fn verify_totp_or_backup_for_user(state: &AppState, did: &crate::types::Did, code: &str) -> bool { let code = code.trim(); if is_backup_code_format(code) { return verify_backup_code_for_user(state, did, code).await; } - let totp_row = sqlx::query!( - "SELECT secret_encrypted, encryption_version, verified FROM user_totp WHERE did = $1", - did - ) - .fetch_optional(&state.db) - .await; - - let totp_row = match totp_row { + let totp_record = match state.user_repo.get_totp_record(did).await { Ok(Some(row)) if row.verified => row, _ => return false, }; - let secret = match decrypt_totp_secret(&totp_row.secret_encrypted, totp_row.encryption_version) - { - Ok(s) => s, - Err(_) => return false, - }; + let secret = + match decrypt_totp_secret(&totp_record.secret_encrypted, totp_record.encryption_version) { + Ok(s) => s, + Err(_) => return false, + }; if verify_totp_code(&secret, code) { - let _ = sqlx::query!("UPDATE user_totp SET last_used = NOW() WHERE did = $1", did) - .execute(&state.db) - .await; + let _ = state.user_repo.update_totp_last_used(did).await; return true; } false } -pub async fn has_totp_enabled(state: &AppState, did: &str) -> bool { - has_totp_enabled_db(&state.db, did).await -} - -pub async fn has_totp_enabled_db(db: &sqlx::PgPool, did: &str) -> bool { - let result = sqlx::query_scalar!("SELECT verified FROM user_totp WHERE did = $1", did) - .fetch_optional(db) - .await; - - matches!(result, Ok(Some(true))) +pub async fn has_totp_enabled(state: &AppState, did: &crate::types::Did) -> bool { + state + .user_repo + .has_totp_enabled(did) + .await + .unwrap_or(false) } diff --git a/crates/tranquil-pds/src/api/server/trusted_devices.rs b/crates/tranquil-pds/src/api/server/trusted_devices.rs index 12f7855..5dc2c59 100644 --- a/crates/tranquil-pds/src/api/server/trusted_devices.rs +++ b/crates/tranquil-pds/src/api/server/trusted_devices.rs @@ -7,8 +7,9 @@ use axum::{ }; use chrono::{DateTime, Duration, Utc}; use serde::{Deserialize, Serialize}; -use sqlx::PgPool; use tracing::{error, info}; +use tranquil_db_traits::OAuthRepository; +use tranquil_types::DeviceId; use crate::auth::BearerAuth; use crate::state::AppState; @@ -71,18 +72,7 @@ pub struct ListTrustedDevicesResponse { } pub async fn list_trusted_devices(State(state): State, auth: BearerAuth) -> Response { - let devices = sqlx::query!( - r#"SELECT od.id, od.user_agent, od.friendly_name, od.trusted_at, od.trusted_until, od.last_seen_at - FROM oauth_device od - JOIN oauth_account_device oad ON od.id = oad.device_id - WHERE oad.did = $1 AND od.trusted_until IS NOT NULL AND od.trusted_until > NOW() - ORDER BY od.last_seen_at DESC"#, - &auth.0.did - ) - .fetch_all(&state.db) - .await; - - match devices { + match state.oauth_repo.list_trusted_devices(&auth.0.did).await { Ok(rows) => { let devices = rows .into_iter() @@ -120,19 +110,14 @@ pub async fn revoke_trusted_device( auth: BearerAuth, Json(input): Json, ) -> Response { - let device_exists = sqlx::query_scalar!( - r#"SELECT 1 as one FROM oauth_device od - JOIN oauth_account_device oad ON od.id = oad.device_id - WHERE oad.did = $1 AND od.id = $2"#, - &auth.0.did, - input.device_id - ) - .fetch_optional(&state.db) - .await; - - match device_exists { - Ok(Some(_)) => {} - Ok(None) => { + let device_id = DeviceId::from(input.device_id.clone()); + match state + .oauth_repo + .device_belongs_to_user(&device_id, &auth.0.did) + .await + { + Ok(true) => {} + Ok(false) => { return ApiError::DeviceNotFound.into_response(); } Err(e) => { @@ -141,15 +126,8 @@ pub async fn revoke_trusted_device( } } - let result = sqlx::query!( - "UPDATE oauth_device SET trusted_at = NULL, trusted_until = NULL WHERE id = $1", - input.device_id - ) - .execute(&state.db) - .await; - - match result { - Ok(_) => { + match state.oauth_repo.revoke_device_trust(&device_id).await { + Ok(()) => { info!(did = %&auth.0.did, device_id = %input.device_id, "Trusted device revoked"); SuccessResponse::ok().into_response() } @@ -172,19 +150,14 @@ pub async fn update_trusted_device( auth: BearerAuth, Json(input): Json, ) -> Response { - let device_exists = sqlx::query_scalar!( - r#"SELECT 1 as one FROM oauth_device od - JOIN oauth_account_device oad ON od.id = oad.device_id - WHERE oad.did = $1 AND od.id = $2"#, - &auth.0.did, - input.device_id - ) - .fetch_optional(&state.db) - .await; - - match device_exists { - Ok(Some(_)) => {} - Ok(None) => { + let device_id = DeviceId::from(input.device_id.clone()); + match state + .oauth_repo + .device_belongs_to_user(&device_id, &auth.0.did) + .await + { + Ok(true) => {} + Ok(false) => { return ApiError::DeviceNotFound.into_response(); } Err(e) => { @@ -193,16 +166,12 @@ pub async fn update_trusted_device( } } - let result = sqlx::query!( - "UPDATE oauth_device SET friendly_name = $1 WHERE id = $2", - input.friendly_name, - input.device_id - ) - .execute(&state.db) - .await; - - match result { - Ok(_) => { + match state + .oauth_repo + .update_device_friendly_name(&device_id, input.friendly_name.as_deref()) + .await + { + Ok(()) => { info!(did = %auth.0.did, device_id = %input.device_id, "Trusted device updated"); SuccessResponse::ok().into_response() } @@ -213,55 +182,43 @@ pub async fn update_trusted_device( } } -pub async fn get_device_trust_state(db: &PgPool, device_id: &str, did: &str) -> DeviceTrustState { - let result = sqlx::query!( - r#"SELECT trusted_at, trusted_until FROM oauth_device od - JOIN oauth_account_device oad ON od.id = oad.device_id - WHERE od.id = $1 AND oad.did = $2"#, - device_id, - did - ) - .fetch_optional(db) - .await; - - match result { - Ok(Some(row)) => DeviceTrustState::from_timestamps(row.trusted_at, row.trusted_until), +pub async fn get_device_trust_state( + oauth_repo: &dyn OAuthRepository, + device_id: &str, + did: &tranquil_types::Did, +) -> DeviceTrustState { + let device_id_typed = DeviceId::from(device_id.to_string()); + match oauth_repo.get_device_trust_info(&device_id_typed, did).await { + Ok(Some(info)) => DeviceTrustState::from_timestamps(info.trusted_at, info.trusted_until), _ => DeviceTrustState::Untrusted, } } -pub async fn is_device_trusted(db: &PgPool, device_id: &str, did: &str) -> bool { - get_device_trust_state(db, device_id, did) +pub async fn is_device_trusted( + oauth_repo: &dyn OAuthRepository, + device_id: &str, + did: &tranquil_types::Did, +) -> bool { + get_device_trust_state(oauth_repo, device_id, did) .await .is_trusted() } -pub async fn trust_device(db: &PgPool, device_id: &str) -> Result<(), sqlx::Error> { +pub async fn trust_device( + oauth_repo: &dyn OAuthRepository, + device_id: &str, +) -> Result<(), tranquil_db_traits::DbError> { let now = Utc::now(); let trusted_until = now + Duration::days(TRUST_DURATION_DAYS); - - sqlx::query!( - "UPDATE oauth_device SET trusted_at = $1, trusted_until = $2 WHERE id = $3", - now, - trusted_until, - device_id - ) - .execute(db) - .await?; - - Ok(()) + let device_id_typed = DeviceId::from(device_id.to_string()); + oauth_repo.trust_device(&device_id_typed, now, trusted_until).await } -pub async fn extend_device_trust(db: &PgPool, device_id: &str) -> Result<(), sqlx::Error> { +pub async fn extend_device_trust( + oauth_repo: &dyn OAuthRepository, + device_id: &str, +) -> Result<(), tranquil_db_traits::DbError> { let trusted_until = Utc::now() + Duration::days(TRUST_DURATION_DAYS); - - sqlx::query!( - "UPDATE oauth_device SET trusted_until = $1 WHERE id = $2 AND trusted_until IS NOT NULL", - trusted_until, - device_id - ) - .execute(db) - .await?; - - Ok(()) + let device_id_typed = DeviceId::from(device_id.to_string()); + oauth_repo.extend_device_trust(&device_id_typed, trusted_until).await } diff --git a/crates/tranquil-pds/src/api/server/verify_email.rs b/crates/tranquil-pds/src/api/server/verify_email.rs index f21e2dd..f98296b 100644 --- a/crates/tranquil-pds/src/api/server/verify_email.rs +++ b/crates/tranquil-pds/src/api/server/verify_email.rs @@ -55,22 +55,15 @@ pub async fn resend_migration_verification( ) -> Result, ApiError> { let email = input.email.trim().to_lowercase(); - let user = sqlx::query!( - "SELECT id, did, email, email_verified, handle FROM users WHERE LOWER(email) = $1", - email - ) - .fetch_optional(&state.db) - .await - .map_err(|e| { - warn!(error = %e, "Database error during resend verification"); - ApiError::InternalError(None) - })?; - - let user = match user { - Some(u) => u, - None => { + let user = match state.user_repo.get_by_email(&email).await { + Ok(Some(u)) => u, + Ok(None) => { return Ok(Json(ResendMigrationVerificationOutput { sent: true })); } + Err(e) => { + warn!(error = ?e, "Database error during resend verification"); + return Err(ApiError::InternalError(None)); + } }; if user.email_verified { @@ -81,8 +74,9 @@ pub async fn resend_migration_verification( let token = crate::auth::verification_token::generate_migration_token(&user.did, &email); let formatted_token = crate::auth::verification_token::format_token_for_display(&token); - if let Err(e) = crate::comms::enqueue_migration_verification( - &state.db, + if let Err(e) = crate::comms::comms_repo::enqueue_migration_verification( + state.user_repo.as_ref(), + state.infra_repo.as_ref(), user.id, &email, &formatted_token, @@ -90,7 +84,7 @@ pub async fn resend_migration_verification( ) .await { - warn!(error = %e, "Failed to enqueue migration verification email"); + warn!(error = ?e, "Failed to enqueue migration verification email"); } info!(did = %user.did, "Resent migration verification email"); diff --git a/crates/tranquil-pds/src/api/server/verify_token.rs b/crates/tranquil-pds/src/api/server/verify_token.rs index 60de1e5..0154b6c 100644 --- a/crates/tranquil-pds/src/api/server/verify_token.rs +++ b/crates/tranquil-pds/src/api/server/verify_token.rs @@ -74,34 +74,30 @@ async fn handle_migration_verification( return Err(ApiError::InvalidChannel); } - let user = sqlx::query!( - "SELECT id, email, email_verified FROM users WHERE did = $1", - did - ) - .fetch_optional(&state.db) - .await - .map_err(|e| { - warn!(error = %e, "Database error during migration verification"); - ApiError::InternalError(None) - })?; - - let user = user.ok_or(ApiError::AccountNotFound)?; + let did_typed: Did = did.parse().map_err(|_| ApiError::InvalidDid("Invalid DID format".into()))?; + let user = state + .user_repo + .get_verification_info(&did_typed) + .await + .map_err(|e| { + warn!(error = ?e, "Database error during migration verification"); + ApiError::InternalError(None) + })? + .ok_or(ApiError::AccountNotFound)?; if user.email.as_ref().map(|e| e.to_lowercase()) != Some(identifier.to_string()) { return Err(ApiError::IdentifierMismatch); } if !user.email_verified { - sqlx::query!( - "UPDATE users SET email_verified = true WHERE id = $1", - user.id - ) - .execute(&state.db) - .await - .map_err(|e| { - warn!(error = %e, "Failed to update email_verified status"); - ApiError::InternalError(None) - })?; + state + .user_repo + .set_email_verified_flag(user.id) + .await + .map_err(|e| { + warn!(error = ?e, "Failed to update email_verified status"); + ApiError::InternalError(None) + })?; } info!(did = %did, "Migration email verified successfully"); @@ -120,49 +116,63 @@ async fn handle_channel_update( channel: &str, identifier: &str, ) -> Result, ApiError> { - let user_id = sqlx::query_scalar!("SELECT id FROM users WHERE did = $1", did) - .fetch_one(&state.db) + let did_typed: Did = did.parse().map_err(|_| ApiError::InvalidDid("Invalid DID format".into()))?; + let user_id = state + .user_repo + .get_id_by_did(&did_typed) .await - .map_err(|_| ApiError::InternalError(None))?; + .map_err(|_| ApiError::InternalError(None))? + .ok_or(ApiError::AccountNotFound)?; - let update_result = match channel { - "email" => sqlx::query!( - "UPDATE users SET email = $1, email_verified = TRUE, updated_at = NOW() WHERE id = $2", - identifier, - user_id - ).execute(&state.db).await, - "discord" => sqlx::query!( - "UPDATE users SET discord_id = $1, discord_verified = TRUE, updated_at = NOW() WHERE id = $2", - identifier, - user_id - ).execute(&state.db).await, - "telegram" => sqlx::query!( - "UPDATE users SET telegram_username = $1, telegram_verified = TRUE, updated_at = NOW() WHERE id = $2", - identifier, - user_id - ).execute(&state.db).await, - "signal" => sqlx::query!( - "UPDATE users SET signal_number = $1, signal_verified = TRUE, updated_at = NOW() WHERE id = $2", - identifier, - user_id - ).execute(&state.db).await, + match channel { + "email" => { + let success = state + .user_repo + .verify_email_channel(user_id, identifier) + .await + .map_err(|e| { + error!("Failed to update email channel: {:?}", e); + ApiError::InternalError(None) + })?; + if !success { + return Err(ApiError::EmailTaken); + } + } + "discord" => { + state + .user_repo + .verify_discord_channel(user_id, identifier) + .await + .map_err(|e| { + error!("Failed to update discord channel: {:?}", e); + ApiError::InternalError(None) + })?; + } + "telegram" => { + state + .user_repo + .verify_telegram_channel(user_id, identifier) + .await + .map_err(|e| { + error!("Failed to update telegram channel: {:?}", e); + ApiError::InternalError(None) + })?; + } + "signal" => { + state + .user_repo + .verify_signal_channel(user_id, identifier) + .await + .map_err(|e| { + error!("Failed to update signal channel: {:?}", e); + ApiError::InternalError(None) + })?; + } _ => { return Err(ApiError::InvalidChannel); } }; - if let Err(e) = update_result { - error!("Failed to update user channel: {:?}", e); - if channel == "email" - && e.as_database_error() - .map(|db| db.is_unique_violation()) - .unwrap_or(false) - { - return Err(ApiError::EmailTaken); - } - return Err(ApiError::InternalError(None)); - } - info!(did = %did, channel = %channel, "Channel verified successfully"); Ok(Json(VerifyTokenOutput { @@ -179,18 +189,16 @@ async fn handle_signup_verification( channel: &str, _identifier: &str, ) -> Result, ApiError> { - let user = sqlx::query!( - "SELECT id, handle, email, email_verified, discord_verified, telegram_verified, signal_verified FROM users WHERE did = $1", - did - ) - .fetch_optional(&state.db) - .await - .map_err(|e| { - warn!(error = %e, "Database error during signup verification"); - ApiError::InternalError(None) - })?; - - let user = user.ok_or(ApiError::AccountNotFound)?; + let did_typed: Did = did.parse().map_err(|_| ApiError::InvalidDid("Invalid DID format".into()))?; + let user = state + .user_repo + .get_verification_info(&did_typed) + .await + .map_err(|e| { + warn!(error = ?e, "Database error during signup verification"); + ApiError::InternalError(None) + })? + .ok_or(ApiError::AccountNotFound)?; let is_verified = user.email_verified || user.discord_verified @@ -206,49 +214,52 @@ async fn handle_signup_verification( })); } - let update_result = match channel { + match channel { "email" => { - sqlx::query!( - "UPDATE users SET email_verified = TRUE WHERE id = $1", - user.id - ) - .execute(&state.db) - .await + state + .user_repo + .set_email_verified_flag(user.id) + .await + .map_err(|e| { + warn!(error = ?e, "Failed to update email verified status"); + ApiError::InternalError(None) + })?; } "discord" => { - sqlx::query!( - "UPDATE users SET discord_verified = TRUE WHERE id = $1", - user.id - ) - .execute(&state.db) - .await + state + .user_repo + .set_discord_verified_flag(user.id) + .await + .map_err(|e| { + warn!(error = ?e, "Failed to update discord verified status"); + ApiError::InternalError(None) + })?; } "telegram" => { - sqlx::query!( - "UPDATE users SET telegram_verified = TRUE WHERE id = $1", - user.id - ) - .execute(&state.db) - .await + state + .user_repo + .set_telegram_verified_flag(user.id) + .await + .map_err(|e| { + warn!(error = ?e, "Failed to update telegram verified status"); + ApiError::InternalError(None) + })?; } "signal" => { - sqlx::query!( - "UPDATE users SET signal_verified = TRUE WHERE id = $1", - user.id - ) - .execute(&state.db) - .await + state + .user_repo + .set_signal_verified_flag(user.id) + .await + .map_err(|e| { + warn!(error = ?e, "Failed to update signal verified status"); + ApiError::InternalError(None) + })?; } _ => { return Err(ApiError::InvalidChannel); } }; - update_result.map_err(|e| { - warn!(error = %e, "Failed to update channel verified status"); - ApiError::InternalError(None) - })?; - info!(did = %did, channel = %channel, "Signup verified successfully"); Ok(Json(VerifyTokenOutput { diff --git a/crates/tranquil-pds/src/api/temp.rs b/crates/tranquil-pds/src/api/temp.rs index 065cb2f..a3e4a55 100644 --- a/crates/tranquil-pds/src/api/temp.rs +++ b/crates/tranquil-pds/src/api/temp.rs @@ -28,7 +28,8 @@ pub async fn check_signup_queue(State(state): State, headers: HeaderMa { let dpop_proof = headers.get("DPoP").and_then(|h| h.to_str().ok()); if let Ok(user) = validate_token_with_dpop( - &state.db, + state.user_repo.as_ref(), + state.oauth_repo.as_ref(), &extracted.token, extracted.is_dpop, dpop_proof, diff --git a/crates/tranquil-pds/src/auth/extractor.rs b/crates/tranquil-pds/src/auth/extractor.rs index 5d59e20..8cbd195 100644 --- a/crates/tranquil-pds/src/auth/extractor.rs +++ b/crates/tranquil-pds/src/auth/extractor.rs @@ -130,7 +130,8 @@ impl FromRequestParts for BearerAuth { let uri = build_full_url(&parts.uri.to_string()); match validate_token_with_dpop( - &state.db, + state.user_repo.as_ref(), + state.oauth_repo.as_ref(), &extracted.token, true, dpop_proof, @@ -148,8 +149,12 @@ impl FromRequestParts for BearerAuth { Err(_) => Err(AuthError::AuthenticationFailed), } } else { - match validate_bearer_token_cached(&state.db, state.cache.as_ref(), &extracted.token) - .await + match validate_bearer_token_cached( + state.user_repo.as_ref(), + state.cache.as_ref(), + &extracted.token, + ) + .await { Ok(user) => Ok(BearerAuth(user)), Err(TokenValidationError::AccountDeactivated) => Err(AuthError::AccountDeactivated), @@ -186,7 +191,8 @@ impl FromRequestParts for BearerAuthAllowDeactivated { let uri = build_full_url(&parts.uri.to_string()); match validate_token_with_dpop( - &state.db, + state.user_repo.as_ref(), + state.oauth_repo.as_ref(), &extracted.token, true, dpop_proof, @@ -204,7 +210,7 @@ impl FromRequestParts for BearerAuthAllowDeactivated { } } else { match validate_bearer_token_cached_allow_deactivated( - &state.db, + state.user_repo.as_ref(), state.cache.as_ref(), &extracted.token, ) @@ -244,7 +250,8 @@ impl FromRequestParts for BearerAuthAllowTakendown { let uri = build_full_url(&parts.uri.to_string()); match validate_token_with_dpop( - &state.db, + state.user_repo.as_ref(), + state.oauth_repo.as_ref(), &extracted.token, true, dpop_proof, @@ -261,7 +268,9 @@ impl FromRequestParts for BearerAuthAllowTakendown { Err(_) => Err(AuthError::AuthenticationFailed), } } else { - match validate_bearer_token_allow_takendown(&state.db, &extracted.token).await { + match validate_bearer_token_allow_takendown(state.user_repo.as_ref(), &extracted.token) + .await + { Ok(user) => Ok(BearerAuthAllowTakendown(user)), Err(TokenValidationError::AccountDeactivated) => Err(AuthError::AccountDeactivated), Err(TokenValidationError::TokenExpired) => Err(AuthError::TokenExpired), @@ -296,7 +305,8 @@ impl FromRequestParts for BearerAuthAdmin { let uri = build_full_url(&parts.uri.to_string()); match validate_token_with_dpop( - &state.db, + state.user_repo.as_ref(), + state.oauth_repo.as_ref(), &extracted.token, true, dpop_proof, @@ -320,8 +330,12 @@ impl FromRequestParts for BearerAuthAdmin { Err(_) => return Err(AuthError::AuthenticationFailed), } } else { - match validate_bearer_token_cached(&state.db, state.cache.as_ref(), &extracted.token) - .await + match validate_bearer_token_cached( + state.user_repo.as_ref(), + state.cache.as_ref(), + &extracted.token, + ) + .await { Ok(user) => user, Err(TokenValidationError::AccountDeactivated) => { diff --git a/crates/tranquil-pds/src/auth/mod.rs b/crates/tranquil-pds/src/auth/mod.rs index baa2b33..7366c39 100644 --- a/crates/tranquil-pds/src/auth/mod.rs +++ b/crates/tranquil-pds/src/auth/mod.rs @@ -1,5 +1,4 @@ use serde::{Deserialize, Serialize}; -use sqlx::PgPool; use std::fmt; use std::time::Duration; @@ -7,6 +6,8 @@ use crate::AccountStatus; use crate::cache::Cache; use crate::oauth::scopes::ScopePermissions; use crate::types::Did; +use tranquil_db::UserRepository; +use tranquil_db_traits::OAuthRepository; pub mod extractor; pub mod scope_check; @@ -62,6 +63,7 @@ pub enum TokenValidationError { AuthenticationFailed, TokenExpired, OAuthTokenExpired, + InvalidToken, } impl fmt::Display for TokenValidationError { @@ -72,6 +74,7 @@ impl fmt::Display for TokenValidationError { Self::KeyDecryptionFailed => write!(f, "KeyDecryptionFailed"), Self::AuthenticationFailed => write!(f, "AuthenticationFailed"), Self::TokenExpired | Self::OAuthTokenExpired => write!(f, "ExpiredToken"), + Self::InvalidToken => write!(f, "InvalidToken"), } } } @@ -105,51 +108,51 @@ impl AuthenticatedUser { } pub async fn validate_bearer_token( - db: &PgPool, + user_repo: &dyn UserRepository, token: &str, ) -> Result { - validate_bearer_token_with_options_internal(db, None, token, false, false).await + validate_bearer_token_with_options_internal(user_repo, None, token, false, false).await } pub async fn validate_bearer_token_allow_deactivated( - db: &PgPool, + user_repo: &dyn UserRepository, token: &str, ) -> Result { - validate_bearer_token_with_options_internal(db, None, token, true, false).await + validate_bearer_token_with_options_internal(user_repo, None, token, true, false).await } pub async fn validate_bearer_token_cached( - db: &PgPool, + user_repo: &dyn UserRepository, cache: &dyn Cache, token: &str, ) -> Result { - validate_bearer_token_with_options_internal(db, Some(cache), token, false, false).await + validate_bearer_token_with_options_internal(user_repo, Some(cache), token, false, false).await } pub async fn validate_bearer_token_cached_allow_deactivated( - db: &PgPool, + user_repo: &dyn UserRepository, cache: &dyn Cache, token: &str, ) -> Result { - validate_bearer_token_with_options_internal(db, Some(cache), token, true, false).await + validate_bearer_token_with_options_internal(user_repo, Some(cache), token, true, false).await } pub async fn validate_bearer_token_for_service_auth( - db: &PgPool, + user_repo: &dyn UserRepository, token: &str, ) -> Result { - validate_bearer_token_with_options_internal(db, None, token, true, true).await + validate_bearer_token_with_options_internal(user_repo, None, token, true, true).await } pub async fn validate_bearer_token_allow_takendown( - db: &PgPool, + user_repo: &dyn UserRepository, token: &str, ) -> Result { - validate_bearer_token_with_options_internal(db, None, token, false, true).await + validate_bearer_token_with_options_internal(user_repo, None, token, false, true).await } async fn validate_bearer_token_with_options_internal( - db: &PgPool, + user_repo: &dyn UserRepository, cache: Option<&dyn Cache>, token: &str, allow_deactivated: bool, @@ -157,8 +160,12 @@ async fn validate_bearer_token_with_options_internal( ) -> Result { let did_from_token = get_did_from_token(token).ok(); - if let Some(ref did) = did_from_token { - let key_cache_key = format!("auth:key:{}", did); + if let Some(ref did_str) = did_from_token { + let did: tranquil_types::Did = match did_str.parse() { + Ok(d) => d, + Err(_) => return Err(TokenValidationError::InvalidToken), + }; + let key_cache_key = format!("auth:key:{}", did_str); let mut cached_key: Option> = None; if let Some(c) = cache { @@ -172,7 +179,7 @@ async fn validate_bearer_token_with_options_internal( let (decrypted_key, deactivated_at, takedown_ref, is_admin) = if let Some(key) = cached_key { - let status_cache_key = format!("auth:status:{}", did); + let status_cache_key = format!("auth:status:{}", did_str); let cached_status: Option = if let Some(c) = cache { c.get(&status_cache_key) .await @@ -197,14 +204,7 @@ async fn validate_bearer_token_with_options_internal( status.is_admin, ) } else { - let user_status = sqlx::query!( - "SELECT deactivated_at, takedown_ref, is_admin FROM users WHERE did = $1", - did - ) - .fetch_optional(db) - .await - .ok() - .flatten(); + let user_status = user_repo.get_status_by_did(&did).await.ok().flatten(); match user_status { Some(status) => { @@ -234,18 +234,7 @@ async fn validate_bearer_token_with_options_internal( None => (None, None, None, false), } } - } else if let Some(user) = sqlx::query!( - "SELECT k.key_bytes, k.encryption_version, u.deactivated_at, u.takedown_ref, u.is_admin - FROM users u - JOIN user_keys k ON u.id = k.user_id - WHERE u.did = $1", - did - ) - .fetch_optional(db) - .await - .ok() - .flatten() - { + } else if let Some(user) = user_repo.get_with_key_by_did(&did).await.ok().flatten() { let key = crate::config::decrypt_key(&user.key_bytes, user.encryption_version) .map_err(|_| TokenValidationError::KeyDecryptionFailed)?; @@ -310,18 +299,14 @@ async fn validate_bearer_token_with_options_internal( } if !session_valid { - let session_row = sqlx::query!( - "SELECT access_expires_at FROM session_tokens WHERE did = $1 AND access_jti = $2", - did, - jti - ) - .fetch_optional(db) - .await - .ok() - .flatten(); - - if let Some(row) = session_row { - if row.access_expires_at > chrono::Utc::now() { + let session_expiry = user_repo + .get_session_access_expiry(&did, jti) + .await + .ok() + .flatten(); + + if let Some(expires_at) = session_expiry { + if expires_at > chrono::Utc::now() { session_valid = true; if let Some(c) = cache { let _ = c @@ -347,7 +332,7 @@ async fn validate_bearer_token_with_options_internal( let status = AccountStatus::from_db_fields(takedown_ref.as_deref(), deactivated_at); return Ok(AuthenticatedUser { - did: Did::new_unchecked(did.clone()), + did: did.clone(), key_bytes: Some(decrypted_key), is_oauth: false, is_admin, @@ -366,19 +351,11 @@ async fn validate_bearer_token_with_options_internal( } if let Ok(oauth_info) = crate::oauth::verify::extract_oauth_token_info(token) - && let Some(oauth_token) = sqlx::query!( - r#"SELECT t.did, t.expires_at, u.deactivated_at, u.takedown_ref, u.is_admin, - k.key_bytes as "key_bytes?", k.encryption_version as "encryption_version?" - FROM oauth_token t - JOIN users u ON t.did = u.did - LEFT JOIN user_keys k ON u.id = k.user_id - WHERE t.token_id = $1"#, - oauth_info.token_id - ) - .fetch_optional(db) - .await - .ok() - .flatten() + && let Some(oauth_token) = user_repo + .get_oauth_token_with_user(&oauth_info.token_id) + .await + .ok() + .flatten() { let status = AccountStatus::from_db_fields( oauth_token.takedown_ref.as_deref(), @@ -428,7 +405,8 @@ pub async fn invalidate_auth_cache(cache: &dyn Cache, did: &str) { #[allow(clippy::too_many_arguments)] pub async fn validate_token_with_dpop( - db: &PgPool, + user_repo: &dyn UserRepository, + oauth_repo: &dyn OAuthRepository, token: &str, is_dpop_token: bool, dpop_proof: Option<&str>, @@ -439,15 +417,15 @@ pub async fn validate_token_with_dpop( ) -> Result { if !is_dpop_token { if allow_takendown { - return validate_bearer_token_allow_takendown(db, token).await; + return validate_bearer_token_allow_takendown(user_repo, token).await; } else if allow_deactivated { - return validate_bearer_token_allow_deactivated(db, token).await; + return validate_bearer_token_allow_deactivated(user_repo, token).await; } else { - return validate_bearer_token(db, token).await; + return validate_bearer_token(user_repo, token).await; } } match crate::oauth::verify::verify_oauth_access_token( - db, + oauth_repo, token, dpop_proof, http_method, @@ -456,18 +434,12 @@ pub async fn validate_token_with_dpop( .await { Ok(result) => { - let user_info = sqlx::query!( - r#"SELECT u.deactivated_at, u.takedown_ref, u.is_admin, - k.key_bytes as "key_bytes?", k.encryption_version as "encryption_version?" - FROM users u - LEFT JOIN user_keys k ON u.id = k.user_id - WHERE u.did = $1"#, - result.did - ) - .fetch_optional(db) - .await - .ok() - .flatten(); + let result_did: Did = result.did.parse().map_err(|_| TokenValidationError::InvalidToken)?; + let user_info = user_repo + .get_user_info_by_did(&result_did) + .await + .ok() + .flatten(); let Some(user_info) = user_info else { return Err(TokenValidationError::AuthenticationFailed); }; diff --git a/crates/tranquil-pds/src/auth/webauthn.rs b/crates/tranquil-pds/src/auth/webauthn.rs index dbfab11..5e884da 100644 --- a/crates/tranquil-pds/src/auth/webauthn.rs +++ b/crates/tranquil-pds/src/auth/webauthn.rs @@ -1,6 +1,3 @@ -use base64::{Engine, engine::general_purpose::URL_SAFE_NO_PAD}; -use chrono::{Duration, Utc}; -use sqlx::{PgPool, Row}; use uuid::Uuid; use webauthn_rs::prelude::*; @@ -80,322 +77,3 @@ impl WebAuthnConfig { .map_err(|e| format!("Failed to finish authentication: {}", e)) } } - -pub async fn save_registration_state( - pool: &PgPool, - did: &str, - state: &SecurityKeyRegistration, -) -> Result { - let id = Uuid::new_v4(); - let state_json = serde_json::to_string(state) - .map_err(|e| sqlx::Error::Protocol(format!("Failed to serialize state: {}", e)))?; - let challenge = id.as_bytes().to_vec(); - let expires_at = Utc::now() + Duration::minutes(5); - - sqlx::query!( - r#" - INSERT INTO webauthn_challenges (id, did, challenge, challenge_type, state_json, expires_at) - VALUES ($1, $2, $3, 'registration', $4, $5) - "#, - id, - did, - challenge, - state_json, - expires_at, - ) - .execute(pool) - .await?; - - Ok(id) -} - -pub async fn load_registration_state( - pool: &PgPool, - did: &str, -) -> Result, sqlx::Error> { - let row = sqlx::query!( - r#" - SELECT state_json FROM webauthn_challenges - WHERE did = $1 AND challenge_type = 'registration' AND expires_at > NOW() - ORDER BY created_at DESC - LIMIT 1 - "#, - did, - ) - .fetch_optional(pool) - .await?; - - match row { - Some(r) => { - let state: SecurityKeyRegistration = - serde_json::from_str(&r.state_json).map_err(|e| { - sqlx::Error::Protocol(format!("Failed to deserialize state: {}", e)) - })?; - Ok(Some(state)) - } - None => Ok(None), - } -} - -pub async fn delete_registration_state(pool: &PgPool, did: &str) -> Result<(), sqlx::Error> { - sqlx::query!( - "DELETE FROM webauthn_challenges WHERE did = $1 AND challenge_type = 'registration'", - did, - ) - .execute(pool) - .await?; - Ok(()) -} - -pub async fn save_authentication_state( - pool: &PgPool, - did: &str, - state: &SecurityKeyAuthentication, -) -> Result { - let id = Uuid::new_v4(); - let state_json = serde_json::to_string(state) - .map_err(|e| sqlx::Error::Protocol(format!("Failed to serialize state: {}", e)))?; - let challenge = id.as_bytes().to_vec(); - let expires_at = Utc::now() + Duration::minutes(5); - - sqlx::query!( - r#" - INSERT INTO webauthn_challenges (id, did, challenge, challenge_type, state_json, expires_at) - VALUES ($1, $2, $3, 'authentication', $4, $5) - "#, - id, - did, - challenge, - state_json, - expires_at, - ) - .execute(pool) - .await?; - - Ok(id) -} - -pub async fn load_authentication_state( - pool: &PgPool, - did: &str, -) -> Result, sqlx::Error> { - let row = sqlx::query!( - r#" - SELECT state_json FROM webauthn_challenges - WHERE did = $1 AND challenge_type = 'authentication' AND expires_at > NOW() - ORDER BY created_at DESC - LIMIT 1 - "#, - did, - ) - .fetch_optional(pool) - .await?; - - match row { - Some(r) => { - let state: SecurityKeyAuthentication = - serde_json::from_str(&r.state_json).map_err(|e| { - sqlx::Error::Protocol(format!("Failed to deserialize state: {}", e)) - })?; - Ok(Some(state)) - } - None => Ok(None), - } -} - -pub async fn delete_authentication_state(pool: &PgPool, did: &str) -> Result<(), sqlx::Error> { - sqlx::query!( - "DELETE FROM webauthn_challenges WHERE did = $1 AND challenge_type = 'authentication'", - did, - ) - .execute(pool) - .await?; - Ok(()) -} - -pub async fn cleanup_expired_challenges(pool: &PgPool) -> Result { - let result = sqlx::query!("DELETE FROM webauthn_challenges WHERE expires_at < NOW()") - .execute(pool) - .await?; - Ok(result.rows_affected()) -} - -#[derive(Debug, Clone)] -pub struct StoredPasskey { - pub id: Uuid, - pub did: String, - pub credential_id: Vec, - pub public_key: Vec, - pub sign_count: i32, - pub created_at: chrono::DateTime, - pub last_used: Option>, - pub friendly_name: Option, - pub aaguid: Option>, - pub transports: Option>, -} - -impl StoredPasskey { - pub fn to_security_key(&self) -> Result { - serde_json::from_slice(&self.public_key) - .map_err(|e| format!("Failed to deserialize security key: {}", e)) - } - - pub fn credential_id_base64(&self) -> String { - URL_SAFE_NO_PAD.encode(&self.credential_id) - } -} - -pub async fn save_passkey( - pool: &PgPool, - did: &str, - security_key: &SecurityKey, - friendly_name: Option<&str>, -) -> Result { - let id = Uuid::new_v4(); - let credential_id = security_key.cred_id().to_vec(); - let public_key = serde_json::to_vec(security_key) - .map_err(|e| sqlx::Error::Protocol(format!("Failed to serialize security key: {}", e)))?; - let aaguid: Option> = None; - - sqlx::query!( - r#" - INSERT INTO passkeys (id, did, credential_id, public_key, sign_count, friendly_name, aaguid) - VALUES ($1, $2, $3, $4, 0, $5, $6) - "#, - id, - did, - credential_id, - public_key, - friendly_name, - aaguid, - ) - .execute(pool) - .await?; - - Ok(id) -} - -pub async fn get_passkeys_for_user( - pool: &PgPool, - did: &str, -) -> Result, sqlx::Error> { - let rows = sqlx::query!( - r#" - SELECT id, did, credential_id, public_key, sign_count, created_at, last_used, friendly_name, aaguid, transports - FROM passkeys - WHERE did = $1 - ORDER BY created_at DESC - "#, - did, - ) - .fetch_all(pool) - .await?; - - Ok(rows - .into_iter() - .map(|r| StoredPasskey { - id: r.id, - did: r.did, - credential_id: r.credential_id, - public_key: r.public_key, - sign_count: r.sign_count, - created_at: r.created_at, - last_used: r.last_used, - friendly_name: r.friendly_name, - aaguid: r.aaguid, - transports: r.transports, - }) - .collect()) -} - -pub async fn get_passkey_by_credential_id( - pool: &PgPool, - credential_id: &[u8], -) -> Result, sqlx::Error> { - let row = sqlx::query!( - r#" - SELECT id, did, credential_id, public_key, sign_count, created_at, last_used, friendly_name, aaguid, transports - FROM passkeys - WHERE credential_id = $1 - "#, - credential_id, - ) - .fetch_optional(pool) - .await?; - - Ok(row.map(|r| StoredPasskey { - id: r.id, - did: r.did, - credential_id: r.credential_id, - public_key: r.public_key, - sign_count: r.sign_count, - created_at: r.created_at, - last_used: r.last_used, - friendly_name: r.friendly_name, - aaguid: r.aaguid, - transports: r.transports, - })) -} - -pub async fn update_passkey_counter( - pool: &PgPool, - credential_id: &[u8], - new_counter: u32, -) -> Result { - let stored = get_passkey_by_credential_id(pool, credential_id).await?; - let Some(stored) = stored else { - return Err(sqlx::Error::RowNotFound); - }; - - if new_counter > 0 && new_counter <= stored.sign_count as u32 { - tracing::warn!( - credential_id = ?credential_id, - stored_counter = stored.sign_count, - new_counter = new_counter, - "Passkey counter did not increment - possible cloned key!" - ); - return Ok(false); - } - - sqlx::query!( - "UPDATE passkeys SET sign_count = $1, last_used = NOW() WHERE credential_id = $2", - new_counter as i32, - credential_id, - ) - .execute(pool) - .await?; - Ok(true) -} - -pub async fn delete_passkey(pool: &PgPool, id: Uuid, did: &str) -> Result { - let result = sqlx::query("DELETE FROM passkeys WHERE id = $1 AND did = $2") - .bind(id) - .bind(did) - .execute(pool) - .await?; - Ok(result.rows_affected() > 0) -} - -pub async fn update_passkey_name( - pool: &PgPool, - id: Uuid, - did: &str, - name: &str, -) -> Result { - let result = sqlx::query("UPDATE passkeys SET friendly_name = $1 WHERE id = $2 AND did = $3") - .bind(name) - .bind(id) - .bind(did) - .execute(pool) - .await?; - Ok(result.rows_affected() > 0) -} - -pub async fn has_passkeys(pool: &PgPool, did: &str) -> Result { - let row = sqlx::query("SELECT COUNT(*) as count FROM passkeys WHERE did = $1") - .bind(did) - .fetch_one(pool) - .await?; - let count: i64 = row.get("count"); - Ok(count > 0) -} diff --git a/crates/tranquil-pds/src/comms/mod.rs b/crates/tranquil-pds/src/comms/mod.rs index 8632fea..71be281 100644 --- a/crates/tranquil-pds/src/comms/mod.rs +++ b/crates/tranquil-pds/src/comms/mod.rs @@ -7,9 +7,4 @@ pub use tranquil_comms::{ sanitize_header_value, validate_locale, }; -pub use service::{ - CommsService, channel_display_name, enqueue_2fa_code, enqueue_account_deletion, enqueue_comms, - enqueue_email_update, enqueue_email_update_token, enqueue_migration_verification, - enqueue_passkey_recovery, enqueue_password_reset, enqueue_plc_operation, - enqueue_signup_verification, enqueue_welcome, queue_legacy_login_notification, -}; +pub use service::{CommsService, channel_display_name, repo as comms_repo}; diff --git a/crates/tranquil-pds/src/comms/service.rs b/crates/tranquil-pds/src/comms/service.rs index 2513868..49d54e6 100644 --- a/crates/tranquil-pds/src/comms/service.rs +++ b/crates/tranquil-pds/src/comms/service.rs @@ -3,25 +3,25 @@ use std::sync::Arc; use std::time::Duration; use chrono::Utc; -use sqlx::PgPool; use tokio::sync::watch; use tokio::time::interval; use tracing::{debug, error, info, warn}; use tranquil_comms::{ - CommsChannel, CommsSender, CommsStatus, CommsType, NewComms, QueuedComms, SendError, - format_message, get_strings, + CommsChannel, CommsSender, CommsStatus, CommsType, NewComms, SendError, format_message, + get_strings, }; +use tranquil_db_traits::{InfraRepository, QueuedComms, UserRepository}; use uuid::Uuid; pub struct CommsService { - db: PgPool, + infra_repo: Arc, senders: HashMap>, poll_interval: Duration, batch_size: i64, } impl CommsService { - pub fn new(db: PgPool) -> Self { + pub fn new(infra_repo: Arc) -> Self { let poll_interval_ms: u64 = std::env::var("NOTIFICATION_POLL_INTERVAL_MS") .ok() .and_then(|v| v.parse().ok()) @@ -31,7 +31,7 @@ impl CommsService { .and_then(|v| v.parse().ok()) .unwrap_or(100); Self { - db, + infra_repo, senders: HashMap::new(), poll_interval: Duration::from_millis(poll_interval_ms), batch_size, @@ -53,24 +53,39 @@ impl CommsService { self } - pub async fn enqueue(&self, item: NewComms) -> Result { - let id = sqlx::query_scalar!( - r#" - INSERT INTO comms_queue - (user_id, channel, comms_type, recipient, subject, body, metadata) - VALUES ($1, $2, $3, $4, $5, $6, $7) - RETURNING id - "#, - item.user_id, - item.channel as CommsChannel, - item.comms_type as CommsType, - item.recipient, - item.subject, - item.body, - item.metadata - ) - .fetch_one(&self.db) - .await?; + pub async fn enqueue(&self, item: NewComms) -> Result { + let channel = match item.channel { + CommsChannel::Email => tranquil_db_traits::CommsChannel::Email, + CommsChannel::Discord => tranquil_db_traits::CommsChannel::Discord, + CommsChannel::Telegram => tranquil_db_traits::CommsChannel::Telegram, + CommsChannel::Signal => tranquil_db_traits::CommsChannel::Signal, + }; + let comms_type = match item.comms_type { + CommsType::Welcome => tranquil_db_traits::CommsType::Welcome, + CommsType::EmailVerification => tranquil_db_traits::CommsType::EmailVerification, + CommsType::PasswordReset => tranquil_db_traits::CommsType::PasswordReset, + CommsType::EmailUpdate => tranquil_db_traits::CommsType::EmailUpdate, + CommsType::AccountDeletion => tranquil_db_traits::CommsType::AccountDeletion, + CommsType::AdminEmail => tranquil_db_traits::CommsType::AdminEmail, + CommsType::PlcOperation => tranquil_db_traits::CommsType::PlcOperation, + CommsType::TwoFactorCode => tranquil_db_traits::CommsType::TwoFactorCode, + CommsType::PasskeyRecovery => tranquil_db_traits::CommsType::PasskeyRecovery, + CommsType::LegacyLoginAlert => tranquil_db_traits::CommsType::LegacyLoginAlert, + CommsType::MigrationVerification => tranquil_db_traits::CommsType::MigrationVerification, + CommsType::ChannelVerification => tranquil_db_traits::CommsType::ChannelVerification, + }; + let id = self + .infra_repo + .enqueue_comms( + Some(item.user_id), + channel, + comms_type, + &item.recipient, + item.subject.as_deref(), + &item.body, + item.metadata, + ) + .await?; debug!(comms_id = %id, "Comms enqueued"); Ok(id) } @@ -109,7 +124,7 @@ impl CommsService { } } - async fn process_batch(&self) -> Result<(), sqlx::Error> { + async fn process_batch(&self) -> Result<(), tranquil_db_traits::DbError> { let items = self.fetch_pending().await?; if items.is_empty() { return Ok(()); @@ -119,43 +134,57 @@ impl CommsService { Ok(()) } - async fn fetch_pending(&self) -> Result, sqlx::Error> { + async fn fetch_pending(&self) -> Result, tranquil_db_traits::DbError> { let now = Utc::now(); - sqlx::query_as!( - QueuedComms, - r#" - UPDATE comms_queue - SET status = 'processing', updated_at = NOW() - WHERE id IN ( - SELECT id FROM comms_queue - WHERE status = 'pending' - AND scheduled_for <= $1 - AND attempts < max_attempts - ORDER BY scheduled_for ASC - LIMIT $2 - FOR UPDATE SKIP LOCKED - ) - RETURNING - id, user_id, - channel as "channel: CommsChannel", - comms_type as "comms_type: CommsType", - status as "status: CommsStatus", - recipient, subject, body, metadata, - attempts, max_attempts, last_error, - created_at, updated_at, scheduled_for, processed_at - "#, - now, - self.batch_size - ) - .fetch_all(&self.db) - .await + self.infra_repo.fetch_pending_comms(now, self.batch_size).await } async fn process_item(&self, item: QueuedComms) { let comms_id = item.id; - let channel = item.channel; + let channel = match item.channel { + tranquil_db_traits::CommsChannel::Email => CommsChannel::Email, + tranquil_db_traits::CommsChannel::Discord => CommsChannel::Discord, + tranquil_db_traits::CommsChannel::Telegram => CommsChannel::Telegram, + tranquil_db_traits::CommsChannel::Signal => CommsChannel::Signal, + }; + let comms_item = tranquil_comms::QueuedComms { + id: item.id, + user_id: item.user_id, + channel, + comms_type: match item.comms_type { + tranquil_db_traits::CommsType::Welcome => CommsType::Welcome, + tranquil_db_traits::CommsType::EmailVerification => CommsType::EmailVerification, + tranquil_db_traits::CommsType::PasswordReset => CommsType::PasswordReset, + tranquil_db_traits::CommsType::EmailUpdate => CommsType::EmailUpdate, + tranquil_db_traits::CommsType::AccountDeletion => CommsType::AccountDeletion, + tranquil_db_traits::CommsType::AdminEmail => CommsType::AdminEmail, + tranquil_db_traits::CommsType::PlcOperation => CommsType::PlcOperation, + tranquil_db_traits::CommsType::TwoFactorCode => CommsType::TwoFactorCode, + tranquil_db_traits::CommsType::PasskeyRecovery => CommsType::PasskeyRecovery, + tranquil_db_traits::CommsType::LegacyLoginAlert => CommsType::LegacyLoginAlert, + tranquil_db_traits::CommsType::MigrationVerification => CommsType::MigrationVerification, + tranquil_db_traits::CommsType::ChannelVerification => CommsType::ChannelVerification, + }, + status: match item.status { + tranquil_db_traits::CommsStatus::Pending => CommsStatus::Pending, + tranquil_db_traits::CommsStatus::Processing => CommsStatus::Processing, + tranquil_db_traits::CommsStatus::Sent => CommsStatus::Sent, + tranquil_db_traits::CommsStatus::Failed => CommsStatus::Failed, + }, + recipient: item.recipient, + subject: item.subject, + body: item.body, + metadata: item.metadata, + attempts: item.attempts, + max_attempts: item.max_attempts, + last_error: item.last_error, + created_at: item.created_at, + updated_at: item.updated_at, + scheduled_for: item.scheduled_for, + processed_at: item.processed_at, + }; let result = match self.senders.get(&channel) { - Some(sender) => sender.send(&item).await, + Some(sender) => sender.send(&comms_item).await, None => { warn!( comms_id = %comms_id, @@ -194,479 +223,390 @@ impl CommsService { } } - async fn mark_sent(&self, id: Uuid) -> Result<(), sqlx::Error> { - sqlx::query!( - r#" - UPDATE comms_queue - SET status = 'sent', processed_at = NOW(), updated_at = NOW() - WHERE id = $1 - "#, - id - ) - .execute(&self.db) - .await?; - Ok(()) + async fn mark_sent(&self, id: Uuid) -> Result<(), tranquil_db_traits::DbError> { + self.infra_repo.mark_comms_sent(id).await } - async fn mark_failed(&self, id: Uuid, error: &str) -> Result<(), sqlx::Error> { - sqlx::query!( - r#" - UPDATE comms_queue - SET - status = CASE - WHEN attempts + 1 >= max_attempts THEN 'failed'::comms_status - ELSE 'pending'::comms_status - END, - attempts = attempts + 1, - last_error = $2, - updated_at = NOW(), - scheduled_for = NOW() + (INTERVAL '1 minute' * (attempts + 1)) - WHERE id = $1 - "#, - id, - error - ) - .execute(&self.db) - .await?; - Ok(()) + async fn mark_failed(&self, id: Uuid, error: &str) -> Result<(), tranquil_db_traits::DbError> { + self.infra_repo.mark_comms_failed(id, error).await } } -pub async fn enqueue_comms(db: &PgPool, item: NewComms) -> Result { - sqlx::query_scalar!( - r#" - INSERT INTO comms_queue - (user_id, channel, comms_type, recipient, subject, body, metadata) - VALUES ($1, $2, $3, $4, $5, $6, $7) - RETURNING id - "#, - item.user_id, - item.channel as CommsChannel, - item.comms_type as CommsType, - item.recipient, - item.subject, - item.body, - item.metadata - ) - .fetch_one(db) - .await +pub fn channel_display_name(channel: CommsChannel) -> &'static str { + match channel { + CommsChannel::Email => "email", + CommsChannel::Discord => "Discord", + CommsChannel::Telegram => "Telegram", + CommsChannel::Signal => "Signal", + } } -pub struct UserCommsPrefs { - pub channel: CommsChannel, - pub email: Option, - pub handle: crate::types::Handle, - pub locale: String, +fn channel_from_str(s: &str) -> tranquil_db_traits::CommsChannel { + match s { + "discord" => tranquil_db_traits::CommsChannel::Discord, + "telegram" => tranquil_db_traits::CommsChannel::Telegram, + "signal" => tranquil_db_traits::CommsChannel::Signal, + _ => tranquil_db_traits::CommsChannel::Email, + } } -pub async fn get_user_comms_prefs( - db: &PgPool, - user_id: Uuid, -) -> Result { - let row = sqlx::query!( - r#" - SELECT - email, - handle, - preferred_comms_channel as "channel: CommsChannel", - preferred_locale - FROM users - WHERE id = $1 - "#, - user_id - ) - .fetch_one(db) - .await?; - Ok(UserCommsPrefs { - channel: row.channel, - email: row.email, - handle: row.handle.into(), - locale: row.preferred_locale.unwrap_or_else(|| "en".to_string()), - }) -} -pub async fn enqueue_welcome( - db: &PgPool, - user_id: Uuid, - hostname: &str, -) -> Result { - let prefs = get_user_comms_prefs(db, user_id).await?; - let strings = get_strings(&prefs.locale); - let body = format_message( - strings.welcome_body, - &[("hostname", hostname), ("handle", &prefs.handle)], - ); - let subject = format_message(strings.welcome_subject, &[("hostname", hostname)]); - enqueue_comms( - db, - NewComms::new( - user_id, - prefs.channel, +pub mod repo { + use super::*; + use tranquil_db_traits::DbError; + + pub async fn enqueue_welcome( + user_repo: &dyn UserRepository, + infra_repo: &dyn InfraRepository, + user_id: Uuid, + hostname: &str, + ) -> Result { + let prefs = user_repo.get_comms_prefs(user_id).await?.ok_or(DbError::NotFound)?; + let strings = get_strings(prefs.preferred_locale.as_deref().unwrap_or("en")); + let body = format_message( + strings.welcome_body, + &[("hostname", hostname), ("handle", &prefs.handle)], + ); + let subject = format_message(strings.welcome_subject, &[("hostname", hostname)]); + let channel = channel_from_str(&prefs.preferred_channel); + infra_repo.enqueue_comms( + Some(user_id), + channel, CommsType::Welcome, - prefs.email.unwrap_or_default(), - Some(subject), - body, - ), - ) - .await -} + &prefs.email.unwrap_or_default(), + Some(&subject), + &body, + None, + ).await + } -pub async fn enqueue_password_reset( - db: &PgPool, - user_id: Uuid, - code: &str, - hostname: &str, -) -> Result { - let prefs = get_user_comms_prefs(db, user_id).await?; - let strings = get_strings(&prefs.locale); - let body = format_message( - strings.password_reset_body, - &[("handle", &prefs.handle), ("code", code)], - ); - let subject = format_message(strings.password_reset_subject, &[("hostname", hostname)]); - enqueue_comms( - db, - NewComms::new( - user_id, - prefs.channel, + pub async fn enqueue_password_reset( + user_repo: &dyn UserRepository, + infra_repo: &dyn InfraRepository, + user_id: Uuid, + code: &str, + hostname: &str, + ) -> Result { + let prefs = user_repo.get_comms_prefs(user_id).await?.ok_or(DbError::NotFound)?; + let strings = get_strings(prefs.preferred_locale.as_deref().unwrap_or("en")); + let body = format_message( + strings.password_reset_body, + &[("handle", &prefs.handle), ("code", code)], + ); + let subject = format_message(strings.password_reset_subject, &[("hostname", hostname)]); + let channel = channel_from_str(&prefs.preferred_channel); + infra_repo.enqueue_comms( + Some(user_id), + channel, CommsType::PasswordReset, - prefs.email.unwrap_or_default(), - Some(subject), - body, - ), - ) - .await -} + &prefs.email.unwrap_or_default(), + Some(&subject), + &body, + None, + ).await + } -pub async fn enqueue_email_update( - db: &PgPool, - user_id: Uuid, - new_email: &str, - handle: &str, - code: &str, - hostname: &str, -) -> Result { - let prefs = get_user_comms_prefs(db, user_id).await?; - let strings = get_strings(&prefs.locale); - let encoded_email = urlencoding::encode(new_email); - let encoded_token = urlencoding::encode(code); - let verify_page = format!("https://{}/app/verify", hostname); - let verify_link = format!( - "https://{}/app/verify?token={}&identifier={}", - hostname, encoded_token, encoded_email - ); - let body = format_message( - strings.email_update_body, - &[ - ("handle", handle), - ("code", code), - ("verify_page", &verify_page), - ("verify_link", &verify_link), - ], - ); - let subject = format_message(strings.email_update_subject, &[("hostname", hostname)]); - enqueue_comms( - db, - NewComms::email( - user_id, + pub async fn enqueue_email_update( + infra_repo: &dyn InfraRepository, + user_id: Uuid, + new_email: &str, + handle: &str, + code: &str, + hostname: &str, + ) -> Result { + let strings = get_strings("en"); + let encoded_email = urlencoding::encode(new_email); + let encoded_token = urlencoding::encode(code); + let verify_page = format!("https://{}/app/verify", hostname); + let verify_link = format!( + "https://{}/app/verify?token={}&identifier={}", + hostname, encoded_token, encoded_email + ); + let body = format_message( + strings.email_update_body, + &[ + ("handle", handle), + ("code", code), + ("verify_page", &verify_page), + ("verify_link", &verify_link), + ], + ); + let subject = format_message(strings.email_update_subject, &[("hostname", hostname)]); + infra_repo.enqueue_comms( + Some(user_id), + tranquil_db_traits::CommsChannel::Email, CommsType::EmailUpdate, - new_email.to_string(), - subject, - body, - ), - ) - .await -} + new_email, + Some(&subject), + &body, + None, + ).await + } -pub async fn enqueue_email_update_token( - db: &PgPool, - user_id: Uuid, - code: &str, - hostname: &str, -) -> Result { - let prefs = get_user_comms_prefs(db, user_id).await?; - let strings = get_strings(&prefs.locale); - let current_email = prefs.email.unwrap_or_default(); - let verify_page = format!("https://{}/app/verify?type=email-update", hostname); - let verify_link = format!( - "https://{}/app/verify?type=email-update&token={}", - hostname, - urlencoding::encode(code) - ); - let body = format_message( - strings.email_update_body, - &[ - ("handle", &prefs.handle), - ("code", code), - ("verify_page", &verify_page), - ("verify_link", &verify_link), - ], - ); - let subject = format_message(strings.email_update_subject, &[("hostname", hostname)]); - enqueue_comms( - db, - NewComms::email( - user_id, + pub async fn enqueue_email_update_token( + user_repo: &dyn UserRepository, + infra_repo: &dyn InfraRepository, + user_id: Uuid, + code: &str, + hostname: &str, + ) -> Result { + let prefs = user_repo.get_comms_prefs(user_id).await?.ok_or(DbError::NotFound)?; + let strings = get_strings(prefs.preferred_locale.as_deref().unwrap_or("en")); + let current_email = prefs.email.unwrap_or_default(); + let verify_page = format!("https://{}/app/verify?type=email-update", hostname); + let verify_link = format!( + "https://{}/app/verify?type=email-update&token={}", + hostname, + urlencoding::encode(code) + ); + let body = format_message( + strings.email_update_body, + &[ + ("handle", &prefs.handle), + ("code", code), + ("verify_page", &verify_page), + ("verify_link", &verify_link), + ], + ); + let subject = format_message(strings.email_update_subject, &[("hostname", hostname)]); + infra_repo.enqueue_comms( + Some(user_id), + tranquil_db_traits::CommsChannel::Email, CommsType::EmailUpdate, - current_email, - subject, - body, - ), - ) - .await -} + ¤t_email, + Some(&subject), + &body, + None, + ).await + } -pub async fn enqueue_account_deletion( - db: &PgPool, - user_id: Uuid, - code: &str, - hostname: &str, -) -> Result { - let prefs = get_user_comms_prefs(db, user_id).await?; - let strings = get_strings(&prefs.locale); - let body = format_message( - strings.account_deletion_body, - &[("handle", &prefs.handle), ("code", code)], - ); - let subject = format_message(strings.account_deletion_subject, &[("hostname", hostname)]); - enqueue_comms( - db, - NewComms::new( - user_id, - prefs.channel, + pub async fn enqueue_account_deletion( + user_repo: &dyn UserRepository, + infra_repo: &dyn InfraRepository, + user_id: Uuid, + code: &str, + hostname: &str, + ) -> Result { + let prefs = user_repo.get_comms_prefs(user_id).await?.ok_or(DbError::NotFound)?; + let strings = get_strings(prefs.preferred_locale.as_deref().unwrap_or("en")); + let body = format_message( + strings.account_deletion_body, + &[("handle", &prefs.handle), ("code", code)], + ); + let subject = format_message(strings.account_deletion_subject, &[("hostname", hostname)]); + let channel = channel_from_str(&prefs.preferred_channel); + infra_repo.enqueue_comms( + Some(user_id), + channel, CommsType::AccountDeletion, - prefs.email.unwrap_or_default(), - Some(subject), - body, - ), - ) - .await -} + &prefs.email.unwrap_or_default(), + Some(&subject), + &body, + None, + ).await + } -pub async fn enqueue_plc_operation( - db: &PgPool, - user_id: Uuid, - token: &str, - hostname: &str, -) -> Result { - let prefs = get_user_comms_prefs(db, user_id).await?; - let strings = get_strings(&prefs.locale); - let body = format_message( - strings.plc_operation_body, - &[("handle", &prefs.handle), ("token", token)], - ); - let subject = format_message(strings.plc_operation_subject, &[("hostname", hostname)]); - enqueue_comms( - db, - NewComms::new( - user_id, - prefs.channel, + pub async fn enqueue_plc_operation( + user_repo: &dyn UserRepository, + infra_repo: &dyn InfraRepository, + user_id: Uuid, + token: &str, + hostname: &str, + ) -> Result { + let prefs = user_repo.get_comms_prefs(user_id).await?.ok_or(DbError::NotFound)?; + let strings = get_strings(prefs.preferred_locale.as_deref().unwrap_or("en")); + let body = format_message( + strings.plc_operation_body, + &[("handle", &prefs.handle), ("token", token)], + ); + let subject = format_message(strings.plc_operation_subject, &[("hostname", hostname)]); + let channel = channel_from_str(&prefs.preferred_channel); + infra_repo.enqueue_comms( + Some(user_id), + channel, CommsType::PlcOperation, - prefs.email.unwrap_or_default(), - Some(subject), - body, - ), - ) - .await -} - -pub async fn enqueue_2fa_code( - db: &PgPool, - user_id: Uuid, - code: &str, - hostname: &str, -) -> Result { - let prefs = get_user_comms_prefs(db, user_id).await?; - let strings = get_strings(&prefs.locale); - let body = format_message( - strings.two_factor_code_body, - &[("handle", &prefs.handle), ("code", code)], - ); - let subject = format_message(strings.two_factor_code_subject, &[("hostname", hostname)]); - enqueue_comms( - db, - NewComms::new( - user_id, - prefs.channel, - CommsType::TwoFactorCode, - prefs.email.unwrap_or_default(), - Some(subject), - body, - ), - ) - .await -} + &prefs.email.unwrap_or_default(), + Some(&subject), + &body, + None, + ).await + } -pub async fn enqueue_passkey_recovery( - db: &PgPool, - user_id: Uuid, - recovery_url: &str, - hostname: &str, -) -> Result { - let prefs = get_user_comms_prefs(db, user_id).await?; - let strings = get_strings(&prefs.locale); - let body = format_message( - strings.passkey_recovery_body, - &[("handle", &prefs.handle), ("url", recovery_url)], - ); - let subject = format_message(strings.passkey_recovery_subject, &[("hostname", hostname)]); - enqueue_comms( - db, - NewComms::new( - user_id, - prefs.channel, + pub async fn enqueue_passkey_recovery( + user_repo: &dyn UserRepository, + infra_repo: &dyn InfraRepository, + user_id: Uuid, + recovery_url: &str, + hostname: &str, + ) -> Result { + let prefs = user_repo.get_comms_prefs(user_id).await?.ok_or(DbError::NotFound)?; + let strings = get_strings(prefs.preferred_locale.as_deref().unwrap_or("en")); + let body = format_message( + strings.passkey_recovery_body, + &[("handle", &prefs.handle), ("url", recovery_url)], + ); + let subject = format_message(strings.passkey_recovery_subject, &[("hostname", hostname)]); + let channel = channel_from_str(&prefs.preferred_channel); + infra_repo.enqueue_comms( + Some(user_id), + channel, CommsType::PasskeyRecovery, - prefs.email.unwrap_or_default(), - Some(subject), - body, - ), - ) - .await -} + &prefs.email.unwrap_or_default(), + Some(&subject), + &body, + None, + ).await + } -pub fn channel_display_name(channel: CommsChannel) -> &'static str { - match channel { - CommsChannel::Email => "email", - CommsChannel::Discord => "Discord", - CommsChannel::Telegram => "Telegram", - CommsChannel::Signal => "Signal", + pub async fn enqueue_migration_verification( + user_repo: &dyn UserRepository, + infra_repo: &dyn InfraRepository, + user_id: Uuid, + email: &str, + token: &str, + hostname: &str, + ) -> Result { + let prefs = user_repo.get_comms_prefs(user_id).await?.ok_or(DbError::NotFound)?; + let strings = get_strings(prefs.preferred_locale.as_deref().unwrap_or("en")); + let encoded_email = urlencoding::encode(email); + let encoded_token = urlencoding::encode(token); + let verify_page = format!("https://{}/app/verify", hostname); + let verify_link = format!( + "https://{}/app/verify?token={}&identifier={}", + hostname, encoded_token, encoded_email + ); + let body = format_message( + strings.migration_verification_body, + &[ + ("code", token), + ("hostname", hostname), + ("verify_page", &verify_page), + ("verify_link", &verify_link), + ], + ); + let subject = format_message( + strings.migration_verification_subject, + &[("hostname", hostname)], + ); + infra_repo.enqueue_comms( + Some(user_id), + tranquil_db_traits::CommsChannel::Email, + CommsType::MigrationVerification, + email, + Some(&subject), + &body, + None, + ).await } -} -pub async fn enqueue_signup_verification( - db: &PgPool, - user_id: Uuid, - channel: &str, - recipient: &str, - code: &str, - locale: Option<&str>, -) -> Result { - let hostname = std::env::var("PDS_HOSTNAME").unwrap_or_else(|_| "localhost".to_string()); - let comms_channel = match channel { - "email" => CommsChannel::Email, - "discord" => CommsChannel::Discord, - "telegram" => CommsChannel::Telegram, - "signal" => CommsChannel::Signal, - _ => CommsChannel::Email, - }; - let strings = get_strings(locale.unwrap_or("en")); - let (verify_page, verify_link) = if comms_channel == CommsChannel::Email { - let encoded_email = urlencoding::encode(recipient); - let encoded_token = urlencoding::encode(code); - ( - format!("https://{}/app/verify", hostname), - format!( - "https://{}/app/verify?token={}&identifier={}", - hostname, encoded_token, encoded_email - ), - ) - } else { - (String::new(), String::new()) - }; - let body = format_message( - strings.signup_verification_body, - &[ - ("code", code), - ("hostname", &hostname), - ("verify_page", &verify_page), - ("verify_link", &verify_link), - ], - ); - let subject = match comms_channel { - CommsChannel::Email => Some(format_message( - strings.signup_verification_subject, - &[("hostname", &hostname)], - )), - _ => None, - }; - enqueue_comms( - db, - NewComms::new( - user_id, + pub async fn enqueue_signup_verification( + infra_repo: &dyn InfraRepository, + user_id: Uuid, + channel: &str, + recipient: &str, + code: &str, + hostname: &str, + ) -> Result { + let comms_channel = channel_from_str(channel); + let strings = get_strings("en"); + let (verify_page, verify_link) = match comms_channel { + tranquil_db_traits::CommsChannel::Email => { + let encoded_email = urlencoding::encode(recipient); + let encoded_token = urlencoding::encode(code); + ( + format!("https://{}/app/verify", hostname), + format!( + "https://{}/app/verify?token={}&identifier={}", + hostname, encoded_token, encoded_email + ), + ) + } + _ => (String::new(), String::new()), + }; + let body = format_message( + strings.signup_verification_body, + &[ + ("code", code), + ("hostname", hostname), + ("verify_page", &verify_page), + ("verify_link", &verify_link), + ], + ); + let subject = match comms_channel { + tranquil_db_traits::CommsChannel::Email => Some(format_message( + strings.signup_verification_subject, + &[("hostname", hostname)], + )), + _ => None, + }; + infra_repo.enqueue_comms( + Some(user_id), comms_channel, CommsType::EmailVerification, - recipient.to_string(), - subject, - body, - ), - ) - .await -} + recipient, + subject.as_deref(), + &body, + None, + ).await + } -pub async fn enqueue_migration_verification( - db: &PgPool, - user_id: Uuid, - email: &str, - token: &str, - hostname: &str, -) -> Result { - let prefs = get_user_comms_prefs(db, user_id).await?; - let strings = get_strings(&prefs.locale); - let encoded_email = urlencoding::encode(email); - let encoded_token = urlencoding::encode(token); - let verify_page = format!("https://{}/app/verify", hostname); - let verify_link = format!( - "https://{}/app/verify?token={}&identifier={}", - hostname, encoded_token, encoded_email - ); - let body = format_message( - strings.migration_verification_body, - &[ - ("code", token), - ("hostname", hostname), - ("verify_page", &verify_page), - ("verify_link", &verify_link), - ], - ); - let subject = format_message( - strings.migration_verification_subject, - &[("hostname", hostname)], - ); - enqueue_comms( - db, - NewComms::email( - user_id, - CommsType::MigrationVerification, - email.to_string(), - subject, - body, - ), - ) - .await -} + pub async fn enqueue_2fa_code( + user_repo: &dyn UserRepository, + infra_repo: &dyn InfraRepository, + user_id: Uuid, + code: &str, + hostname: &str, + ) -> Result { + let prefs = user_repo.get_comms_prefs(user_id).await?.ok_or(DbError::NotFound)?; + let strings = get_strings(prefs.preferred_locale.as_deref().unwrap_or("en")); + let body = format_message( + strings.two_factor_code_body, + &[("handle", &prefs.handle), ("code", code)], + ); + let subject = format_message(strings.two_factor_code_subject, &[("hostname", hostname)]); + let channel = channel_from_str(&prefs.preferred_channel); + infra_repo.enqueue_comms( + Some(user_id), + channel, + CommsType::TwoFactorCode, + &prefs.email.unwrap_or_default(), + Some(&subject), + &body, + None, + ).await + } -pub async fn queue_legacy_login_notification( - db: &PgPool, - user_id: Uuid, - hostname: &str, - client_ip: &str, - channel: CommsChannel, -) -> Result { - let prefs = get_user_comms_prefs(db, user_id).await?; - let strings = get_strings(&prefs.locale); - let timestamp = chrono::Utc::now() - .format("%Y-%m-%d %H:%M:%S UTC") - .to_string(); - let body = format_message( - strings.legacy_login_body, - &[ - ("handle", &prefs.handle), - ("timestamp", ×tamp), - ("ip", client_ip), - ("hostname", hostname), - ], - ); - let subject = format_message(strings.legacy_login_subject, &[("hostname", hostname)]); - enqueue_comms( - db, - NewComms::new( - user_id, + pub async fn enqueue_legacy_login( + user_repo: &dyn UserRepository, + infra_repo: &dyn InfraRepository, + user_id: Uuid, + hostname: &str, + client_ip: &str, + channel: tranquil_db_traits::CommsChannel, + ) -> Result { + let prefs = user_repo.get_comms_prefs(user_id).await?.ok_or(DbError::NotFound)?; + let strings = get_strings(prefs.preferred_locale.as_deref().unwrap_or("en")); + let timestamp = chrono::Utc::now() + .format("%Y-%m-%d %H:%M:%S UTC") + .to_string(); + let body = format_message( + strings.legacy_login_body, + &[ + ("handle", &prefs.handle), + ("timestamp", ×tamp), + ("ip", client_ip), + ("hostname", hostname), + ], + ); + let subject = format_message(strings.legacy_login_subject, &[("hostname", hostname)]); + infra_repo.enqueue_comms( + Some(user_id), channel, CommsType::LegacyLoginAlert, - prefs.email.unwrap_or_default(), - Some(subject), - body, - ), - ) - .await + &prefs.email.unwrap_or_default(), + Some(&subject), + &body, + None, + ).await + } } diff --git a/crates/tranquil-pds/src/delegation/audit.rs b/crates/tranquil-pds/src/delegation/audit.rs deleted file mode 100644 index 92fd826..0000000 --- a/crates/tranquil-pds/src/delegation/audit.rs +++ /dev/null @@ -1,143 +0,0 @@ -use chrono::{DateTime, Utc}; -use serde::{Deserialize, Serialize}; -use sqlx::PgPool; -use uuid::Uuid; - -#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, sqlx::Type)] -#[sqlx(type_name = "delegation_action_type", rename_all = "snake_case")] -pub enum DelegationActionType { - GrantCreated, - GrantRevoked, - ScopesModified, - TokenIssued, - RepoWrite, - BlobUpload, - AccountAction, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct AuditLogEntry { - pub id: Uuid, - pub delegated_did: String, - pub actor_did: String, - pub controller_did: Option, - pub action_type: DelegationActionType, - pub action_details: Option, - pub ip_address: Option, - pub user_agent: Option, - pub created_at: DateTime, -} - -#[allow(clippy::too_many_arguments)] -pub async fn log_delegation_action( - pool: &PgPool, - delegated_did: &str, - actor_did: &str, - controller_did: Option<&str>, - action_type: DelegationActionType, - action_details: Option, - ip_address: Option<&str>, - user_agent: Option<&str>, -) -> Result { - let id = sqlx::query_scalar!( - r#" - INSERT INTO delegation_audit_log - (delegated_did, actor_did, controller_did, action_type, action_details, ip_address, user_agent) - VALUES ($1, $2, $3, $4, $5, $6, $7) - RETURNING id - "#, - delegated_did, - actor_did, - controller_did, - action_type as DelegationActionType, - action_details, - ip_address, - user_agent - ) - .fetch_one(pool) - .await?; - - Ok(id) -} - -pub async fn get_audit_log_for_account( - pool: &PgPool, - delegated_did: &str, - limit: i64, - offset: i64, -) -> Result, sqlx::Error> { - let entries = sqlx::query_as!( - AuditLogEntry, - r#" - SELECT - id, - delegated_did, - actor_did, - controller_did, - action_type as "action_type: DelegationActionType", - action_details, - ip_address, - user_agent, - created_at - FROM delegation_audit_log - WHERE delegated_did = $1 - ORDER BY created_at DESC - LIMIT $2 OFFSET $3 - "#, - delegated_did, - limit, - offset - ) - .fetch_all(pool) - .await?; - - Ok(entries) -} - -pub async fn get_audit_log_by_controller( - pool: &PgPool, - controller_did: &str, - limit: i64, - offset: i64, -) -> Result, sqlx::Error> { - let entries = sqlx::query_as!( - AuditLogEntry, - r#" - SELECT - id, - delegated_did, - actor_did, - controller_did, - action_type as "action_type: DelegationActionType", - action_details, - ip_address, - user_agent, - created_at - FROM delegation_audit_log - WHERE controller_did = $1 - ORDER BY created_at DESC - LIMIT $2 OFFSET $3 - "#, - controller_did, - limit, - offset - ) - .fetch_all(pool) - .await?; - - Ok(entries) -} - -pub async fn count_audit_log_entries( - pool: &PgPool, - delegated_did: &str, -) -> Result { - let count = sqlx::query_scalar!( - r#"SELECT COUNT(*) as "count!" FROM delegation_audit_log WHERE delegated_did = $1"#, - delegated_did - ) - .fetch_one(pool) - .await?; - - Ok(count) -} diff --git a/crates/tranquil-pds/src/delegation/db.rs b/crates/tranquil-pds/src/delegation/db.rs deleted file mode 100644 index b519a58..0000000 --- a/crates/tranquil-pds/src/delegation/db.rs +++ /dev/null @@ -1,268 +0,0 @@ -use crate::types::Handle; -use chrono::{DateTime, Utc}; -use serde::{Deserialize, Serialize}; -use sqlx::PgPool; -use uuid::Uuid; - -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct DelegationGrant { - pub id: Uuid, - pub delegated_did: String, - pub controller_did: String, - pub granted_scopes: String, - pub granted_at: DateTime, - pub granted_by: String, - pub revoked_at: Option>, - pub revoked_by: Option, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct DelegatedAccountInfo { - pub did: String, - pub handle: Handle, - pub granted_scopes: String, - pub granted_at: DateTime, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ControllerInfo { - pub did: String, - pub handle: Handle, - pub granted_scopes: String, - pub granted_at: DateTime, - pub is_active: bool, -} - -pub async fn is_delegated_account(pool: &PgPool, did: &str) -> Result { - let result = sqlx::query_scalar!( - r#"SELECT account_type::text = 'delegated' as "is_delegated!" FROM users WHERE did = $1"#, - did - ) - .fetch_optional(pool) - .await?; - - Ok(result.unwrap_or(false)) -} - -pub async fn create_delegation( - pool: &PgPool, - delegated_did: &str, - controller_did: &str, - granted_scopes: &str, - granted_by: &str, -) -> Result { - let id = sqlx::query_scalar!( - r#" - INSERT INTO account_delegations (delegated_did, controller_did, granted_scopes, granted_by) - VALUES ($1, $2, $3, $4) - RETURNING id - "#, - delegated_did, - controller_did, - granted_scopes, - granted_by - ) - .fetch_one(pool) - .await?; - - Ok(id) -} - -pub async fn revoke_delegation( - pool: &PgPool, - delegated_did: &str, - controller_did: &str, - revoked_by: &str, -) -> Result { - let result = sqlx::query!( - r#" - UPDATE account_delegations - SET revoked_at = NOW(), revoked_by = $1 - WHERE delegated_did = $2 AND controller_did = $3 AND revoked_at IS NULL - "#, - revoked_by, - delegated_did, - controller_did - ) - .execute(pool) - .await?; - - Ok(result.rows_affected() > 0) -} - -pub async fn update_delegation_scopes( - pool: &PgPool, - delegated_did: &str, - controller_did: &str, - new_scopes: &str, -) -> Result { - let result = sqlx::query!( - r#" - UPDATE account_delegations - SET granted_scopes = $1 - WHERE delegated_did = $2 AND controller_did = $3 AND revoked_at IS NULL - "#, - new_scopes, - delegated_did, - controller_did - ) - .execute(pool) - .await?; - - Ok(result.rows_affected() > 0) -} - -pub async fn get_delegation( - pool: &PgPool, - delegated_did: &str, - controller_did: &str, -) -> Result, sqlx::Error> { - let grant = sqlx::query_as!( - DelegationGrant, - r#" - SELECT id, delegated_did, controller_did, granted_scopes, - granted_at, granted_by, revoked_at, revoked_by - FROM account_delegations - WHERE delegated_did = $1 AND controller_did = $2 AND revoked_at IS NULL - "#, - delegated_did, - controller_did - ) - .fetch_optional(pool) - .await?; - - Ok(grant) -} - -pub async fn get_delegations_for_account( - pool: &PgPool, - delegated_did: &str, -) -> Result, sqlx::Error> { - let controllers = sqlx::query_as!( - ControllerInfo, - r#" - SELECT - u.did, - u.handle, - d.granted_scopes, - d.granted_at, - (u.deactivated_at IS NULL AND u.takedown_ref IS NULL) as "is_active!" - FROM account_delegations d - JOIN users u ON u.did = d.controller_did - WHERE d.delegated_did = $1 AND d.revoked_at IS NULL - ORDER BY d.granted_at DESC - "#, - delegated_did - ) - .fetch_all(pool) - .await?; - - Ok(controllers) -} - -pub async fn get_accounts_controlled_by( - pool: &PgPool, - controller_did: &str, -) -> Result, sqlx::Error> { - let accounts = sqlx::query_as!( - DelegatedAccountInfo, - r#" - SELECT - u.did, - u.handle, - d.granted_scopes, - d.granted_at - FROM account_delegations d - JOIN users u ON u.did = d.delegated_did - WHERE d.controller_did = $1 - AND d.revoked_at IS NULL - AND u.deactivated_at IS NULL - AND u.takedown_ref IS NULL - ORDER BY d.granted_at DESC - "#, - controller_did - ) - .fetch_all(pool) - .await?; - - Ok(accounts) -} - -pub async fn get_active_controllers_for_account( - pool: &PgPool, - delegated_did: &str, -) -> Result, sqlx::Error> { - let controllers = sqlx::query_as!( - ControllerInfo, - r#" - SELECT - u.did, - u.handle, - d.granted_scopes, - d.granted_at, - true as "is_active!" - FROM account_delegations d - JOIN users u ON u.did = d.controller_did - WHERE d.delegated_did = $1 - AND d.revoked_at IS NULL - AND u.deactivated_at IS NULL - AND u.takedown_ref IS NULL - ORDER BY d.granted_at DESC - "#, - delegated_did - ) - .fetch_all(pool) - .await?; - - Ok(controllers) -} - -pub async fn count_active_controllers( - pool: &PgPool, - delegated_did: &str, -) -> Result { - let count = sqlx::query_scalar!( - r#" - SELECT COUNT(*) as "count!" - FROM account_delegations d - JOIN users u ON u.did = d.controller_did - WHERE d.delegated_did = $1 - AND d.revoked_at IS NULL - AND u.deactivated_at IS NULL - AND u.takedown_ref IS NULL - "#, - delegated_did - ) - .fetch_one(pool) - .await?; - - Ok(count) -} - -pub async fn has_any_controllers(pool: &PgPool, did: &str) -> Result { - let exists = sqlx::query_scalar!( - r#"SELECT EXISTS( - SELECT 1 FROM account_delegations - WHERE delegated_did = $1 AND revoked_at IS NULL - ) as "exists!""#, - did - ) - .fetch_one(pool) - .await?; - - Ok(exists) -} - -pub async fn controls_any_accounts(pool: &PgPool, did: &str) -> Result { - let exists = sqlx::query_scalar!( - r#"SELECT EXISTS( - SELECT 1 FROM account_delegations - WHERE controller_did = $1 AND revoked_at IS NULL - ) as "exists!""#, - did - ) - .fetch_one(pool) - .await?; - - Ok(exists) -} diff --git a/crates/tranquil-pds/src/delegation/mod.rs b/crates/tranquil-pds/src/delegation/mod.rs index ab2f890..0020d7f 100644 --- a/crates/tranquil-pds/src/delegation/mod.rs +++ b/crates/tranquil-pds/src/delegation/mod.rs @@ -1,11 +1,4 @@ -pub mod audit; -pub mod db; pub mod scopes; -pub use audit::{DelegationActionType, log_delegation_action}; -pub use db::{ - DelegationGrant, controls_any_accounts, create_delegation, get_accounts_controlled_by, - get_delegation, get_delegations_for_account, has_any_controllers, is_delegated_account, - revoke_delegation, update_delegation_scopes, -}; pub use scopes::{SCOPE_PRESETS, ScopePreset, intersect_scopes}; +pub use tranquil_db_traits::DelegationActionType; diff --git a/crates/tranquil-pds/src/main.rs b/crates/tranquil-pds/src/main.rs index 066de79..da0a9f4 100644 --- a/crates/tranquil-pds/src/main.rs +++ b/crates/tranquil-pds/src/main.rs @@ -32,18 +32,18 @@ async fn run() -> Result<(), Box> { let (shutdown_tx, shutdown_rx) = watch::channel(false); - let backfill_db = state.db.clone(); + let backfill_repo_repo = state.repo_repo.clone(); let backfill_block_store = state.block_store.clone(); tokio::spawn(async move { tokio::join!( - backfill_genesis_commit_blocks(&backfill_db, backfill_block_store.clone()), - backfill_repo_rev(&backfill_db, backfill_block_store.clone()), - backfill_user_blocks(&backfill_db, backfill_block_store.clone()), - backfill_record_blobs(&backfill_db, backfill_block_store), + backfill_genesis_commit_blocks(backfill_repo_repo.clone(), backfill_block_store.clone()), + backfill_repo_rev(backfill_repo_repo.clone(), backfill_block_store.clone()), + backfill_user_blocks(backfill_repo_repo.clone(), backfill_block_store.clone()), + backfill_record_blobs(backfill_repo_repo, backfill_block_store), ); }); - let mut comms_service = CommsService::new(state.db.clone()); + let mut comms_service = CommsService::new(state.infra_repo.clone()); if let Some(email_sender) = EmailSender::from_env() { info!("Email comms enabled"); @@ -88,7 +88,8 @@ async fn run() -> Result<(), Box> { let backup_handle = if let Some(backup_storage) = state.backup_storage.clone() { info!("Backup service enabled"); Some(tokio::spawn(start_backup_tasks( - state.db.clone(), + state.repo_repo.clone(), + state.backup_repo.clone(), state.block_store.clone(), backup_storage, shutdown_rx.clone(), @@ -99,7 +100,8 @@ async fn run() -> Result<(), Box> { }; let scheduled_handle = tokio::spawn(start_scheduled_tasks( - state.db.clone(), + state.user_repo.clone(), + state.blob_repo.clone(), state.blob_store.clone(), shutdown_rx, )); diff --git a/crates/tranquil-pds/src/oauth/db/client.rs b/crates/tranquil-pds/src/oauth/db/client.rs deleted file mode 100644 index 5f07126..0000000 --- a/crates/tranquil-pds/src/oauth/db/client.rs +++ /dev/null @@ -1,46 +0,0 @@ -use super::super::{AuthorizedClientData, OAuthError}; -use super::helpers::{from_json, to_json}; -use sqlx::PgPool; - -pub async fn upsert_authorized_client( - pool: &PgPool, - did: &str, - client_id: &str, - data: &AuthorizedClientData, -) -> Result<(), OAuthError> { - let data_json = to_json(data)?; - sqlx::query!( - r#" - INSERT INTO oauth_authorized_client (did, client_id, created_at, updated_at, data) - VALUES ($1, $2, NOW(), NOW(), $3) - ON CONFLICT (did, client_id) DO UPDATE SET updated_at = NOW(), data = $3 - "#, - did, - client_id, - data_json - ) - .execute(pool) - .await?; - Ok(()) -} - -pub async fn get_authorized_client( - pool: &PgPool, - did: &str, - client_id: &str, -) -> Result, OAuthError> { - let row = sqlx::query_scalar!( - r#" - SELECT data FROM oauth_authorized_client - WHERE did = $1 AND client_id = $2 - "#, - did, - client_id - ) - .fetch_optional(pool) - .await?; - match row { - Some(v) => Ok(Some(from_json(v)?)), - None => Ok(None), - } -} diff --git a/crates/tranquil-pds/src/oauth/db/device.rs b/crates/tranquil-pds/src/oauth/db/device.rs deleted file mode 100644 index 0022644..0000000 --- a/crates/tranquil-pds/src/oauth/db/device.rs +++ /dev/null @@ -1,148 +0,0 @@ -use super::super::{DeviceData, OAuthError}; -use crate::types::Handle; -use chrono::{DateTime, Utc}; -use sqlx::PgPool; - -pub struct DeviceAccountRow { - pub did: String, - pub handle: Handle, - pub email: Option, - pub last_used_at: DateTime, -} - -pub async fn create_device( - pool: &PgPool, - device_id: &str, - data: &DeviceData, -) -> Result<(), OAuthError> { - sqlx::query!( - r#" - INSERT INTO oauth_device (id, session_id, user_agent, ip_address, last_seen_at) - VALUES ($1, $2, $3, $4, $5) - "#, - device_id, - data.session_id, - data.user_agent, - data.ip_address, - data.last_seen_at, - ) - .execute(pool) - .await?; - Ok(()) -} - -pub async fn get_device(pool: &PgPool, device_id: &str) -> Result, OAuthError> { - let row = sqlx::query!( - r#" - SELECT session_id, user_agent, ip_address, last_seen_at - FROM oauth_device - WHERE id = $1 - "#, - device_id - ) - .fetch_optional(pool) - .await?; - Ok(row.map(|r| DeviceData { - session_id: r.session_id, - user_agent: r.user_agent, - ip_address: r.ip_address, - last_seen_at: r.last_seen_at, - })) -} - -pub async fn update_device_last_seen(pool: &PgPool, device_id: &str) -> Result<(), OAuthError> { - sqlx::query!( - r#" - UPDATE oauth_device - SET last_seen_at = NOW() - WHERE id = $1 - "#, - device_id - ) - .execute(pool) - .await?; - Ok(()) -} - -pub async fn delete_device(pool: &PgPool, device_id: &str) -> Result<(), OAuthError> { - sqlx::query!( - r#" - DELETE FROM oauth_device WHERE id = $1 - "#, - device_id - ) - .execute(pool) - .await?; - Ok(()) -} - -pub async fn upsert_account_device( - pool: &PgPool, - did: &str, - device_id: &str, -) -> Result<(), OAuthError> { - sqlx::query!( - r#" - INSERT INTO oauth_account_device (did, device_id, created_at, updated_at) - VALUES ($1, $2, NOW(), NOW()) - ON CONFLICT (did, device_id) DO UPDATE SET updated_at = NOW() - "#, - did, - device_id - ) - .execute(pool) - .await?; - Ok(()) -} - -pub async fn get_device_accounts( - pool: &PgPool, - device_id: &str, -) -> Result, OAuthError> { - let rows = sqlx::query!( - r#" - SELECT u.did, u.handle, u.email, ad.updated_at as last_used_at - FROM oauth_account_device ad - JOIN users u ON u.did = ad.did - WHERE ad.device_id = $1 - AND u.deactivated_at IS NULL - AND u.takedown_ref IS NULL - ORDER BY ad.updated_at DESC - "#, - device_id - ) - .fetch_all(pool) - .await?; - Ok(rows - .into_iter() - .map(|r| DeviceAccountRow { - did: r.did, - handle: r.handle.into(), - email: r.email, - last_used_at: r.last_used_at, - }) - .collect()) -} - -pub async fn verify_account_on_device( - pool: &PgPool, - device_id: &str, - did: &str, -) -> Result { - let row = sqlx::query!( - r#" - SELECT 1 as exists - FROM oauth_account_device ad - JOIN users u ON u.did = ad.did - WHERE ad.device_id = $1 - AND ad.did = $2 - AND u.deactivated_at IS NULL - AND u.takedown_ref IS NULL - "#, - device_id, - did - ) - .fetch_optional(pool) - .await?; - Ok(row.is_some()) -} diff --git a/crates/tranquil-pds/src/oauth/db/dpop.rs b/crates/tranquil-pds/src/oauth/db/dpop.rs deleted file mode 100644 index b471e82..0000000 --- a/crates/tranquil-pds/src/oauth/db/dpop.rs +++ /dev/null @@ -1,32 +0,0 @@ -use super::super::OAuthError; -use sqlx::PgPool; - -pub async fn check_and_record_dpop_jti(pool: &PgPool, jti: &str) -> Result { - let result = sqlx::query!( - r#" - INSERT INTO oauth_dpop_jti (jti) - VALUES ($1) - ON CONFLICT (jti) DO NOTHING - "#, - jti - ) - .execute(pool) - .await?; - Ok(result.rows_affected() > 0) -} - -pub async fn cleanup_expired_dpop_jtis( - pool: &PgPool, - max_age_secs: i64, -) -> Result { - let result = sqlx::query!( - r#" - DELETE FROM oauth_dpop_jti - WHERE created_at < NOW() - INTERVAL '1 second' * $1 - "#, - max_age_secs as f64 - ) - .execute(pool) - .await?; - Ok(result.rows_affected()) -} diff --git a/crates/tranquil-pds/src/oauth/db/helpers.rs b/crates/tranquil-pds/src/oauth/db/helpers.rs deleted file mode 100644 index 9e40cf6..0000000 --- a/crates/tranquil-pds/src/oauth/db/helpers.rs +++ /dev/null @@ -1,16 +0,0 @@ -use super::super::OAuthError; -use serde::{Serialize, de::DeserializeOwned}; - -pub fn to_json(value: &T) -> Result { - serde_json::to_value(value).map_err(|e| { - tracing::error!("JSON serialization error: {}", e); - OAuthError::ServerError("Internal serialization error".to_string()) - }) -} - -pub fn from_json(value: serde_json::Value) -> Result { - serde_json::from_value(value).map_err(|e| { - tracing::error!("JSON deserialization error: {}", e); - OAuthError::ServerError("Internal data corruption".to_string()) - }) -} diff --git a/crates/tranquil-pds/src/oauth/db/mod.rs b/crates/tranquil-pds/src/oauth/db/mod.rs index ab6247d..b222984 100644 --- a/crates/tranquil-pds/src/oauth/db/mod.rs +++ b/crates/tranquil-pds/src/oauth/db/mod.rs @@ -1,37 +1,8 @@ -mod client; -mod device; -mod dpop; -mod helpers; -mod request; mod scope_preference; mod token; mod two_factor; -pub use client::{get_authorized_client, upsert_authorized_client}; -pub use device::{ - DeviceAccountRow, create_device, delete_device, get_device, get_device_accounts, - update_device_last_seen, upsert_account_device, verify_account_on_device, -}; -pub use dpop::{check_and_record_dpop_jti, cleanup_expired_dpop_jtis}; -pub use request::{ - consume_authorization_request_by_code, create_authorization_request, - delete_authorization_request, delete_expired_authorization_requests, get_authorization_request, - get_authorization_request_with_state, mark_request_authenticated, set_authorization_did, - set_controller_did, set_request_did, update_authorization_request, update_request_scope, -}; -pub use scope_preference::{ - ScopePreference, delete_scope_preferences, get_scope_preferences, should_show_consent, - upsert_scope_preferences, -}; -pub use token::{ - RefreshTokenLookup, check_refresh_token_used, count_tokens_for_user, create_token, - delete_oldest_tokens_for_user, delete_token, delete_token_family, enforce_token_limit_for_user, - get_token_by_id, get_token_by_previous_refresh_token, get_token_by_refresh_token, - list_tokens_for_user, lookup_refresh_token, revoke_tokens_for_client, - revoke_tokens_for_controller, rotate_token, -}; -pub use two_factor::{ - TwoFactorChallenge, check_user_2fa_enabled, cleanup_expired_2fa_challenges, - create_2fa_challenge, delete_2fa_challenge, delete_2fa_challenge_by_request_uri, - generate_2fa_code, get_2fa_challenge, increment_2fa_attempts, -}; +pub use scope_preference::{ScopePreference, should_show_consent}; +pub use token::{RefreshTokenLookup, enforce_token_limit_for_user, lookup_refresh_token}; +pub use tranquil_db_traits::{DeviceAccountRow, TwoFactorChallenge}; +pub use two_factor::generate_2fa_code; diff --git a/crates/tranquil-pds/src/oauth/db/request.rs b/crates/tranquil-pds/src/oauth/db/request.rs deleted file mode 100644 index 86291b2..0000000 --- a/crates/tranquil-pds/src/oauth/db/request.rs +++ /dev/null @@ -1,265 +0,0 @@ -use super::super::{ - AuthFlowState, AuthorizationRequestParameters, ClientAuth, OAuthError, RequestData, -}; -use super::helpers::{from_json, to_json}; -use sqlx::PgPool; - -pub async fn get_authorization_request_with_state( - pool: &PgPool, - request_id: &str, -) -> Result, OAuthError> { - match get_authorization_request(pool, request_id).await? { - Some(data) => { - let state = AuthFlowState::from_request_data(&data); - Ok(Some((data, state))) - } - None => Ok(None), - } -} - -pub async fn create_authorization_request( - pool: &PgPool, - request_id: &str, - data: &RequestData, -) -> Result<(), OAuthError> { - let client_auth_json = match &data.client_auth { - Some(ca) => Some(to_json(ca)?), - None => None, - }; - let parameters_json = to_json(&data.parameters)?; - sqlx::query!( - r#" - INSERT INTO oauth_authorization_request - (id, did, device_id, client_id, client_auth, parameters, expires_at, code) - VALUES ($1, $2, $3, $4, $5, $6, $7, $8) - "#, - request_id, - data.did, - data.device_id, - data.client_id, - client_auth_json, - parameters_json, - data.expires_at, - data.code, - ) - .execute(pool) - .await?; - Ok(()) -} - -pub async fn get_authorization_request( - pool: &PgPool, - request_id: &str, -) -> Result, OAuthError> { - let row = sqlx::query!( - r#" - SELECT did, device_id, client_id, client_auth, parameters, expires_at, code, controller_did - FROM oauth_authorization_request - WHERE id = $1 - "#, - request_id - ) - .fetch_optional(pool) - .await?; - match row { - Some(r) => { - let client_auth: Option = match r.client_auth { - Some(v) => Some(from_json(v)?), - None => None, - }; - let parameters: AuthorizationRequestParameters = from_json(r.parameters)?; - Ok(Some(RequestData { - client_id: r.client_id, - client_auth, - parameters, - expires_at: r.expires_at, - did: r.did, - device_id: r.device_id, - code: r.code, - controller_did: r.controller_did, - })) - } - None => Ok(None), - } -} - -pub async fn set_authorization_did( - pool: &PgPool, - request_id: &str, - did: &str, - device_id: Option<&str>, -) -> Result<(), OAuthError> { - sqlx::query!( - r#" - UPDATE oauth_authorization_request - SET did = $2, device_id = $3 - WHERE id = $1 - "#, - request_id, - did, - device_id - ) - .execute(pool) - .await?; - Ok(()) -} - -pub async fn update_authorization_request( - pool: &PgPool, - request_id: &str, - did: &str, - device_id: Option<&str>, - code: &str, -) -> Result<(), OAuthError> { - sqlx::query!( - r#" - UPDATE oauth_authorization_request - SET did = $2, device_id = $3, code = $4 - WHERE id = $1 - "#, - request_id, - did, - device_id, - code - ) - .execute(pool) - .await?; - Ok(()) -} - -pub async fn consume_authorization_request_by_code( - pool: &PgPool, - code: &str, -) -> Result, OAuthError> { - let row = sqlx::query!( - r#" - DELETE FROM oauth_authorization_request - WHERE code = $1 - RETURNING did, device_id, client_id, client_auth, parameters, expires_at, code, controller_did - "#, - code - ) - .fetch_optional(pool) - .await?; - match row { - Some(r) => { - let client_auth: Option = match r.client_auth { - Some(v) => Some(from_json(v)?), - None => None, - }; - let parameters: AuthorizationRequestParameters = from_json(r.parameters)?; - Ok(Some(RequestData { - client_id: r.client_id, - client_auth, - parameters, - expires_at: r.expires_at, - did: r.did, - device_id: r.device_id, - code: r.code, - controller_did: r.controller_did, - })) - } - None => Ok(None), - } -} - -pub async fn delete_authorization_request( - pool: &PgPool, - request_id: &str, -) -> Result<(), OAuthError> { - sqlx::query!( - r#" - DELETE FROM oauth_authorization_request WHERE id = $1 - "#, - request_id - ) - .execute(pool) - .await?; - Ok(()) -} - -pub async fn delete_expired_authorization_requests(pool: &PgPool) -> Result { - let result = sqlx::query!( - r#" - DELETE FROM oauth_authorization_request - WHERE expires_at < NOW() - "# - ) - .execute(pool) - .await?; - Ok(result.rows_affected()) -} - -pub async fn mark_request_authenticated( - pool: &PgPool, - request_id: &str, - did: &str, - device_id: Option<&str>, -) -> Result<(), OAuthError> { - sqlx::query!( - r#" - UPDATE oauth_authorization_request - SET did = $2, device_id = $3 - WHERE id = $1 - "#, - request_id, - did, - device_id - ) - .execute(pool) - .await?; - Ok(()) -} - -pub async fn update_request_scope( - pool: &PgPool, - request_id: &str, - scope: &str, -) -> Result<(), OAuthError> { - sqlx::query!( - r#" - UPDATE oauth_authorization_request - SET parameters = jsonb_set(parameters, '{scope}', to_jsonb($2::text)) - WHERE id = $1 - "#, - request_id, - scope - ) - .execute(pool) - .await?; - Ok(()) -} - -pub async fn set_controller_did( - pool: &PgPool, - request_id: &str, - controller_did: &str, -) -> Result<(), OAuthError> { - sqlx::query!( - r#" - UPDATE oauth_authorization_request - SET controller_did = $2 - WHERE id = $1 - "#, - request_id, - controller_did - ) - .execute(pool) - .await?; - Ok(()) -} - -pub async fn set_request_did(pool: &PgPool, request_id: &str, did: &str) -> Result<(), OAuthError> { - sqlx::query!( - r#" - UPDATE oauth_authorization_request - SET did = $2 - WHERE id = $1 - "#, - request_id, - did - ) - .execute(pool) - .await?; - Ok(()) -} diff --git a/crates/tranquil-pds/src/oauth/db/scope_preference.rs b/crates/tranquil-pds/src/oauth/db/scope_preference.rs index f0b9d2f..ebc85a3 100644 --- a/crates/tranquil-pds/src/oauth/db/scope_preference.rs +++ b/crates/tranquil-pds/src/oauth/db/scope_preference.rs @@ -1,73 +1,23 @@ use super::super::OAuthError; -use serde::{Deserialize, Serialize}; -use sqlx::PgPool; +use tranquil_db_traits::OAuthRepository; +use tranquil_types::{ClientId, Did}; -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ScopePreference { - pub scope: String, - pub granted: bool, -} - -pub async fn get_scope_preferences( - pool: &PgPool, - did: &str, - client_id: &str, -) -> Result, OAuthError> { - let rows = sqlx::query!( - r#" - SELECT scope, granted FROM oauth_scope_preference - WHERE did = $1 AND client_id = $2 - "#, - did, - client_id - ) - .fetch_all(pool) - .await?; - - Ok(rows - .into_iter() - .map(|r| ScopePreference { - scope: r.scope, - granted: r.granted, - }) - .collect()) -} - -pub async fn upsert_scope_preferences( - pool: &PgPool, - did: &str, - client_id: &str, - prefs: &[ScopePreference], -) -> Result<(), OAuthError> { - for pref in prefs { - sqlx::query!( - r#" - INSERT INTO oauth_scope_preference (did, client_id, scope, granted, created_at, updated_at) - VALUES ($1, $2, $3, $4, NOW(), NOW()) - ON CONFLICT (did, client_id, scope) DO UPDATE SET granted = $4, updated_at = NOW() - "#, - did, - client_id, - pref.scope, - pref.granted - ) - .execute(pool) - .await?; - } - Ok(()) -} +pub use tranquil_db_traits::ScopePreference; pub async fn should_show_consent( - pool: &PgPool, - did: &str, - client_id: &str, + oauth_repo: &dyn OAuthRepository, + did: &Did, + client_id: &ClientId, requested_scopes: &[String], ) -> Result { if requested_scopes.is_empty() { return Ok(false); } - let stored_prefs = get_scope_preferences(pool, did, client_id).await?; + let stored_prefs = oauth_repo + .get_scope_preferences(did, client_id) + .await + .map_err(crate::oauth::db_err_to_oauth)?; if stored_prefs.is_empty() { return Ok(true); } @@ -79,21 +29,3 @@ pub async fn should_show_consent( .iter() .any(|scope| !stored_scopes.contains(scope.as_str()))) } - -pub async fn delete_scope_preferences( - pool: &PgPool, - did: &str, - client_id: &str, -) -> Result<(), OAuthError> { - sqlx::query!( - r#" - DELETE FROM oauth_scope_preference - WHERE did = $1 AND client_id = $2 - "#, - did, - client_id - ) - .execute(pool) - .await?; - Ok(()) -} diff --git a/crates/tranquil-pds/src/oauth/db/token.rs b/crates/tranquil-pds/src/oauth/db/token.rs index b626644..941068b 100644 --- a/crates/tranquil-pds/src/oauth/db/token.rs +++ b/crates/tranquil-pds/src/oauth/db/token.rs @@ -1,51 +1,22 @@ -use super::super::{OAuthError, RefreshTokenState, TokenData}; -use super::helpers::{from_json, to_json}; -use chrono::{DateTime, Utc}; -use sqlx::PgPool; - -pub enum RefreshTokenLookup { - Valid { - db_id: i32, - token_data: TokenData, - }, - InGracePeriod { - db_id: i32, - token_data: TokenData, - rotated_at: DateTime, - }, - Used { - original_token_id: i32, - }, - Expired { - db_id: i32, - }, - NotFound, -} - -impl RefreshTokenLookup { - pub fn state(&self) -> RefreshTokenState { - match self { - RefreshTokenLookup::Valid { .. } => RefreshTokenState::Valid, - RefreshTokenLookup::InGracePeriod { rotated_at, .. } => { - RefreshTokenState::InGracePeriod { - rotated_at: *rotated_at, - } - } - RefreshTokenLookup::Used { .. } => RefreshTokenState::Used { at: Utc::now() }, - RefreshTokenLookup::Expired { .. } => RefreshTokenState::Expired, - RefreshTokenLookup::NotFound => RefreshTokenState::Revoked, - } - } -} +use super::super::OAuthError; +use tranquil_db_traits::OAuthRepository; +use tranquil_types::{Did, RefreshToken}; +pub use tranquil_db_traits::RefreshTokenLookup; pub async fn lookup_refresh_token( - pool: &PgPool, - refresh_token: &str, + oauth_repo: &dyn OAuthRepository, + refresh_token: &RefreshToken, ) -> Result { - if let Some(token_id) = check_refresh_token_used(pool, refresh_token).await? { - if let Some((db_id, token_data)) = - get_token_by_previous_refresh_token(pool, refresh_token).await? - { + let token_id = oauth_repo + .check_refresh_token_used(refresh_token) + .await + .map_err(crate::oauth::db_err_to_oauth)?; + if let Some(token_id) = token_id { + let prev_token = oauth_repo + .get_token_by_previous_refresh_token(refresh_token) + .await + .map_err(crate::oauth::db_err_to_oauth)?; + if let Some((db_id, token_data)) = prev_token { let rotated_at = token_data.updated_at; return Ok(RefreshTokenLookup::InGracePeriod { db_id, @@ -58,9 +29,13 @@ pub async fn lookup_refresh_token( }); } - match get_token_by_refresh_token(pool, refresh_token).await? { + let token = oauth_repo + .get_token_by_refresh_token(refresh_token) + .await + .map_err(crate::oauth::db_err_to_oauth)?; + match token { Some((db_id, token_data)) => { - if token_data.expires_at < Utc::now() { + if token_data.expires_at < chrono::Utc::now() { Ok(RefreshTokenLookup::Expired { db_id }) } else { Ok(RefreshTokenLookup::Valid { db_id, token_data }) @@ -70,345 +45,22 @@ pub async fn lookup_refresh_token( } } -pub async fn create_token(pool: &PgPool, data: &TokenData) -> Result { - let client_auth_json = to_json(&data.client_auth)?; - let parameters_json = to_json(&data.parameters)?; - let row = sqlx::query!( - r#" - INSERT INTO oauth_token - (did, token_id, created_at, updated_at, expires_at, client_id, client_auth, - device_id, parameters, details, code, current_refresh_token, scope, controller_did) - VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12, $13, $14) - RETURNING id - "#, - data.did, - data.token_id, - data.created_at, - data.updated_at, - data.expires_at, - data.client_id, - client_auth_json, - data.device_id, - parameters_json, - data.details, - data.code, - data.current_refresh_token, - data.scope, - data.controller_did, - ) - .fetch_one(pool) - .await?; - Ok(row.id) -} - -pub async fn get_token_by_id( - pool: &PgPool, - token_id: &str, -) -> Result, OAuthError> { - let row = sqlx::query!( - r#" - SELECT did, token_id, created_at, updated_at, expires_at, client_id, client_auth, - device_id, parameters, details, code, current_refresh_token, scope, controller_did - FROM oauth_token - WHERE token_id = $1 - "#, - token_id - ) - .fetch_optional(pool) - .await?; - match row { - Some(r) => Ok(Some(TokenData { - did: r.did, - token_id: r.token_id, - created_at: r.created_at, - updated_at: r.updated_at, - expires_at: r.expires_at, - client_id: r.client_id, - client_auth: from_json(r.client_auth)?, - device_id: r.device_id, - parameters: from_json(r.parameters)?, - details: r.details, - code: r.code, - current_refresh_token: r.current_refresh_token, - scope: r.scope, - controller_did: r.controller_did, - })), - None => Ok(None), - } -} - -pub async fn get_token_by_refresh_token( - pool: &PgPool, - refresh_token: &str, -) -> Result, OAuthError> { - let row = sqlx::query!( - r#" - SELECT id, did, token_id, created_at, updated_at, expires_at, client_id, client_auth, - device_id, parameters, details, code, current_refresh_token, scope, controller_did - FROM oauth_token - WHERE current_refresh_token = $1 - "#, - refresh_token - ) - .fetch_optional(pool) - .await?; - match row { - Some(r) => Ok(Some(( - r.id, - TokenData { - did: r.did, - token_id: r.token_id, - created_at: r.created_at, - updated_at: r.updated_at, - expires_at: r.expires_at, - client_id: r.client_id, - client_auth: from_json(r.client_auth)?, - device_id: r.device_id, - parameters: from_json(r.parameters)?, - details: r.details, - code: r.code, - current_refresh_token: r.current_refresh_token, - scope: r.scope, - controller_did: r.controller_did, - }, - ))), - None => Ok(None), - } -} - -pub async fn rotate_token( - pool: &PgPool, - old_db_id: i32, - new_refresh_token: &str, - new_expires_at: DateTime, -) -> Result<(), OAuthError> { - let mut tx = pool.begin().await?; - let old_refresh = sqlx::query_scalar!( - r#" - SELECT current_refresh_token FROM oauth_token WHERE id = $1 - "#, - old_db_id - ) - .fetch_one(&mut *tx) - .await?; - if let Some(ref old_rt) = old_refresh { - sqlx::query!( - r#" - INSERT INTO oauth_used_refresh_token (refresh_token, token_id) - VALUES ($1, $2) - "#, - old_rt, - old_db_id - ) - .execute(&mut *tx) - .await?; - } - sqlx::query!( - r#" - UPDATE oauth_token - SET current_refresh_token = $2, expires_at = $3, updated_at = NOW(), - previous_refresh_token = $4, rotated_at = NOW() - WHERE id = $1 - "#, - old_db_id, - new_refresh_token, - new_expires_at, - old_refresh - ) - .execute(&mut *tx) - .await?; - tx.commit().await?; - Ok(()) -} - -pub async fn check_refresh_token_used( - pool: &PgPool, - refresh_token: &str, -) -> Result, OAuthError> { - let row = sqlx::query_scalar!( - r#" - SELECT token_id FROM oauth_used_refresh_token WHERE refresh_token = $1 - "#, - refresh_token - ) - .fetch_optional(pool) - .await?; - Ok(row) -} - -const REFRESH_GRACE_PERIOD_SECS: i64 = 60; - -pub async fn get_token_by_previous_refresh_token( - pool: &PgPool, - refresh_token: &str, -) -> Result, OAuthError> { - let grace_cutoff = Utc::now() - chrono::Duration::seconds(REFRESH_GRACE_PERIOD_SECS); - let row = sqlx::query!( - r#" - SELECT id, did, token_id, created_at, updated_at, expires_at, client_id, client_auth, - device_id, parameters, details, code, current_refresh_token, scope, controller_did - FROM oauth_token - WHERE previous_refresh_token = $1 AND rotated_at > $2 - "#, - refresh_token, - grace_cutoff - ) - .fetch_optional(pool) - .await?; - match row { - Some(r) => Ok(Some(( - r.id, - TokenData { - did: r.did, - token_id: r.token_id, - created_at: r.created_at, - updated_at: r.updated_at, - expires_at: r.expires_at, - client_id: r.client_id, - client_auth: from_json(r.client_auth)?, - device_id: r.device_id, - parameters: from_json(r.parameters)?, - details: r.details, - code: r.code, - current_refresh_token: r.current_refresh_token, - scope: r.scope, - controller_did: r.controller_did, - }, - ))), - None => Ok(None), - } -} - -pub async fn delete_token(pool: &PgPool, token_id: &str) -> Result<(), OAuthError> { - sqlx::query!( - r#" - DELETE FROM oauth_token WHERE token_id = $1 - "#, - token_id - ) - .execute(pool) - .await?; - Ok(()) -} - -pub async fn delete_token_family(pool: &PgPool, db_id: i32) -> Result<(), OAuthError> { - sqlx::query!( - r#" - DELETE FROM oauth_token WHERE id = $1 - "#, - db_id - ) - .execute(pool) - .await?; - Ok(()) -} - -pub async fn list_tokens_for_user(pool: &PgPool, did: &str) -> Result, OAuthError> { - let rows = sqlx::query!( - r#" - SELECT did, token_id, created_at, updated_at, expires_at, client_id, client_auth, - device_id, parameters, details, code, current_refresh_token, scope, controller_did - FROM oauth_token - WHERE did = $1 - "#, - did - ) - .fetch_all(pool) - .await?; - rows.into_iter() - .map(|r| { - Ok(TokenData { - did: r.did, - token_id: r.token_id, - created_at: r.created_at, - updated_at: r.updated_at, - expires_at: r.expires_at, - client_id: r.client_id, - client_auth: from_json(r.client_auth)?, - device_id: r.device_id, - parameters: from_json(r.parameters)?, - details: r.details, - code: r.code, - current_refresh_token: r.current_refresh_token, - scope: r.scope, - controller_did: r.controller_did, - }) - }) - .collect() -} - -pub async fn count_tokens_for_user(pool: &PgPool, did: &str) -> Result { - let count = sqlx::query_scalar!( - r#" - SELECT COUNT(*) as "count!" FROM oauth_token WHERE did = $1 - "#, - did - ) - .fetch_one(pool) - .await?; - Ok(count) -} - -pub async fn delete_oldest_tokens_for_user( - pool: &PgPool, - did: &str, - keep_count: i64, -) -> Result { - let result = sqlx::query!( - r#" - DELETE FROM oauth_token - WHERE id IN ( - SELECT id FROM oauth_token - WHERE did = $1 - ORDER BY updated_at ASC - OFFSET $2 - ) - "#, - did, - keep_count - ) - .execute(pool) - .await?; - Ok(result.rows_affected()) -} - const MAX_TOKENS_PER_USER: i64 = 100; -pub async fn enforce_token_limit_for_user(pool: &PgPool, did: &str) -> Result<(), OAuthError> { - let count = count_tokens_for_user(pool, did).await?; +pub async fn enforce_token_limit_for_user( + oauth_repo: &dyn OAuthRepository, + did: &Did, +) -> Result<(), OAuthError> { + let count = oauth_repo + .count_tokens_for_user(did) + .await + .map_err(crate::oauth::db_err_to_oauth)?; if count > MAX_TOKENS_PER_USER { let to_keep = MAX_TOKENS_PER_USER - 1; - delete_oldest_tokens_for_user(pool, did, to_keep).await?; + oauth_repo + .delete_oldest_tokens_for_user(did, to_keep) + .await + .map_err(crate::oauth::db_err_to_oauth)?; } Ok(()) } - -pub async fn revoke_tokens_for_client( - pool: &PgPool, - did: &str, - client_id: &str, -) -> Result { - let result = sqlx::query!( - "DELETE FROM oauth_token WHERE did = $1 AND client_id = $2", - did, - client_id - ) - .execute(pool) - .await?; - Ok(result.rows_affected()) -} - -pub async fn revoke_tokens_for_controller( - pool: &PgPool, - delegated_did: &str, - controller_did: &str, -) -> Result { - let result = sqlx::query!( - "DELETE FROM oauth_token WHERE did = $1 AND controller_did = $2", - delegated_did, - controller_did - ) - .execute(pool) - .await?; - Ok(result.rows_affected()) -} diff --git a/crates/tranquil-pds/src/oauth/db/two_factor.rs b/crates/tranquil-pds/src/oauth/db/two_factor.rs index 0dd4aeb..ff617fc 100644 --- a/crates/tranquil-pds/src/oauth/db/two_factor.rs +++ b/crates/tranquil-pds/src/oauth/db/two_factor.rs @@ -1,144 +1,7 @@ -use super::super::OAuthError; -use chrono::{DateTime, Duration, Utc}; use rand::Rng; -use sqlx::PgPool; -use uuid::Uuid; - -pub struct TwoFactorChallenge { - pub id: Uuid, - pub did: String, - pub request_uri: String, - pub code: String, - pub attempts: i32, - pub created_at: DateTime, - pub expires_at: DateTime, -} pub fn generate_2fa_code() -> String { let mut rng = rand::thread_rng(); let code: u32 = rng.gen_range(0..1_000_000); format!("{:06}", code) } - -pub async fn create_2fa_challenge( - pool: &PgPool, - did: &str, - request_uri: &str, -) -> Result { - let code = generate_2fa_code(); - let expires_at = Utc::now() + Duration::minutes(10); - let row = sqlx::query!( - r#" - INSERT INTO oauth_2fa_challenge (did, request_uri, code, expires_at) - VALUES ($1, $2, $3, $4) - RETURNING id, did, request_uri, code, attempts, created_at, expires_at - "#, - did, - request_uri, - code, - expires_at, - ) - .fetch_one(pool) - .await?; - Ok(TwoFactorChallenge { - id: row.id, - did: row.did, - request_uri: row.request_uri, - code: row.code, - attempts: row.attempts, - created_at: row.created_at, - expires_at: row.expires_at, - }) -} - -pub async fn get_2fa_challenge( - pool: &PgPool, - request_uri: &str, -) -> Result, OAuthError> { - let row = sqlx::query!( - r#" - SELECT id, did, request_uri, code, attempts, created_at, expires_at - FROM oauth_2fa_challenge - WHERE request_uri = $1 - "#, - request_uri - ) - .fetch_optional(pool) - .await?; - Ok(row.map(|r| TwoFactorChallenge { - id: r.id, - did: r.did, - request_uri: r.request_uri, - code: r.code, - attempts: r.attempts, - created_at: r.created_at, - expires_at: r.expires_at, - })) -} - -pub async fn increment_2fa_attempts(pool: &PgPool, id: Uuid) -> Result { - let row = sqlx::query!( - r#" - UPDATE oauth_2fa_challenge - SET attempts = attempts + 1 - WHERE id = $1 - RETURNING attempts - "#, - id - ) - .fetch_one(pool) - .await?; - Ok(row.attempts) -} - -pub async fn delete_2fa_challenge(pool: &PgPool, id: Uuid) -> Result<(), OAuthError> { - sqlx::query!( - r#" - DELETE FROM oauth_2fa_challenge WHERE id = $1 - "#, - id - ) - .execute(pool) - .await?; - Ok(()) -} - -pub async fn delete_2fa_challenge_by_request_uri( - pool: &PgPool, - request_uri: &str, -) -> Result<(), OAuthError> { - sqlx::query!( - r#" - DELETE FROM oauth_2fa_challenge WHERE request_uri = $1 - "#, - request_uri - ) - .execute(pool) - .await?; - Ok(()) -} - -pub async fn cleanup_expired_2fa_challenges(pool: &PgPool) -> Result { - let result = sqlx::query!( - r#" - DELETE FROM oauth_2fa_challenge WHERE expires_at < NOW() - "# - ) - .execute(pool) - .await?; - Ok(result.rows_affected()) -} - -pub async fn check_user_2fa_enabled(pool: &PgPool, did: &str) -> Result { - let row = sqlx::query!( - r#" - SELECT two_factor_enabled - FROM users - WHERE did = $1 - "#, - did - ) - .fetch_optional(pool) - .await?; - Ok(row.map(|r| r.two_factor_enabled).unwrap_or(false)) -} diff --git a/crates/tranquil-pds/src/oauth/endpoints/authorize.rs b/crates/tranquil-pds/src/oauth/endpoints/authorize.rs index 394449d..f4b5863 100644 --- a/crates/tranquil-pds/src/oauth/endpoints/authorize.rs +++ b/crates/tranquil-pds/src/oauth/endpoints/authorize.rs @@ -1,9 +1,12 @@ -use crate::comms::{CommsChannel, channel_display_name, enqueue_2fa_code}; +use crate::comms::{channel_display_name, comms_repo::enqueue_2fa_code}; use crate::oauth::{ - AuthFlowState, ClientMetadataCache, Code, DeviceData, DeviceId, OAuthError, SessionId, db, + AuthFlowState, ClientMetadataCache, Code, DeviceData, DeviceId, OAuthError, SessionId, + db::should_show_consent, }; use crate::state::{AppState, RateLimitKind}; -use crate::types::{Handle, PlainPassword}; +use tranquil_db_traits::ScopePreference; +use crate::types::{Did, Handle, PlainPassword}; +use tranquil_types::{AuthorizationCode, ClientId, DeviceId as DeviceIdType, RequestId}; use axum::{ Json, extract::{Query, State}, @@ -203,7 +206,8 @@ pub async fn authorize_get( ); } }; - let request_data = match db::get_authorization_request(&state.db, &request_uri).await { + let request_id = RequestId::from(request_uri.clone()); + let request_data = match state.oauth_repo.get_authorization_request(&request_id).await { Ok(Some(data)) => data, Ok(None) => { if wants_json(&headers) { @@ -235,7 +239,7 @@ pub async fn authorize_get( } }; if request_data.expires_at < Utc::now() { - let _ = db::delete_authorization_request(&state.db, &request_uri).await; + let _ = state.oauth_repo.delete_authorization_request(&request_id).await; if wants_json(&headers) { return ( StatusCode::BAD_REQUEST, @@ -273,25 +277,22 @@ pub async fn authorize_get( tracing::info!(login_hint = %login_hint, "Checking login_hint for delegation"); let pds_hostname = std::env::var("PDS_HOSTNAME").unwrap_or_else(|_| "localhost".to_string()); + let hostname_for_handles = pds_hostname.split(':').next().unwrap_or(&pds_hostname); let normalized = if login_hint.contains('@') || login_hint.starts_with("did:") { login_hint.clone() } else if !login_hint.contains('.') { - format!("{}.{}", login_hint.to_lowercase(), pds_hostname) + format!("{}.{}", login_hint.to_lowercase(), hostname_for_handles) } else { login_hint.to_lowercase() }; tracing::info!(normalized = %normalized, "Normalized login_hint"); - match sqlx::query!( - "SELECT did, password_hash FROM users WHERE handle = $1 OR email = $1", - normalized - ) - .fetch_optional(&state.db) - .await - { + match state.user_repo.get_login_check_by_handle_or_email(&normalized).await { Ok(Some(user)) => { tracing::info!(did = %user.did, has_password = user.password_hash.is_some(), "Found user for login_hint"); - let is_delegated = crate::delegation::is_delegated_account(&state.db, &user.did) + let is_delegated = state + .delegation_repo + .is_delegated_account(&user.did) .await .unwrap_or(false); let has_password = user.password_hash.is_some(); @@ -319,7 +320,7 @@ pub async fn authorize_get( if !force_new_account && let Some(device_id) = extract_device_cookie(&headers) - && let Ok(accounts) = db::get_device_accounts(&state.db, &device_id).await + && let Ok(accounts) = state.oauth_repo.get_device_accounts(&DeviceIdType::from(device_id.clone())).await && !accounts.is_empty() { return redirect_see_other(&format!( @@ -340,11 +341,13 @@ pub async fn authorize_get_json( let request_uri = query .request_uri .ok_or_else(|| OAuthError::InvalidRequest("request_uri is required".to_string()))?; - let request_data = db::get_authorization_request(&state.db, &request_uri) - .await? + let request_id_json = RequestId::from(request_uri.clone()); + let request_data = state.oauth_repo.get_authorization_request(&request_id_json) + .await + .map_err(crate::oauth::db_err_to_oauth)? .ok_or_else(|| OAuthError::InvalidRequest("Invalid or expired request_uri".to_string()))?; if request_data.expires_at < Utc::now() { - db::delete_authorization_request(&state.db, &request_uri).await?; + let _ = state.oauth_repo.delete_authorization_request(&request_id_json).await; return Err(OAuthError::InvalidRequest( "request_uri has expired".to_string(), )); @@ -417,7 +420,8 @@ pub async fn authorize_accounts( .into_response(); } }; - let accounts = match db::get_device_accounts(&state.db, &device_id).await { + let device_id_typed = DeviceIdType::from(device_id.clone()); + let accounts = match state.oauth_repo.get_device_accounts(&device_id_typed).await { Ok(accts) => accts, Err(_) => { return Json(AccountsResponse { @@ -430,7 +434,7 @@ pub async fn authorize_accounts( let account_infos: Vec = accounts .into_iter() .map(|row| AccountInfo { - did: row.did, + did: row.did.to_string(), handle: row.handle, email: row.email.map(|e| mask_email(&e)), }) @@ -469,7 +473,8 @@ pub async fn authorize_post( "Too many login attempts. Please try again later.", ); } - let request_data = match db::get_authorization_request(&state.db, &form.request_uri).await { + let form_request_id = RequestId::from(form.request_uri.clone()); + let request_data = match state.oauth_repo.get_authorization_request(&form_request_id).await { Ok(Some(data)) => data, Ok(None) => { if json_response { @@ -502,7 +507,7 @@ pub async fn authorize_post( } }; if request_data.expires_at < Utc::now() { - let _ = db::delete_authorization_request(&state.db, &form.request_uri).await; + let _ = state.oauth_repo.delete_authorization_request(&form_request_id).await; if json_response { return ( axum::http::StatusCode::BAD_REQUEST, @@ -536,6 +541,7 @@ pub async fn authorize_post( )) }; let pds_hostname = std::env::var("PDS_HOSTNAME").unwrap_or_else(|_| "localhost".to_string()); + let hostname_for_handles = pds_hostname.split(':').next().unwrap_or(&pds_hostname); let normalized_username = form.username.trim(); let normalized_username = normalized_username .strip_prefix('@') @@ -543,7 +549,7 @@ pub async fn authorize_post( let normalized_username = if normalized_username.contains('@') { normalized_username.to_string() } else if !normalized_username.contains('.') { - format!("{}.{}", normalized_username, pds_hostname) + format!("{}.{}", normalized_username, hostname_for_handles) } else { normalized_username.to_string() }; @@ -553,21 +559,7 @@ pub async fn authorize_post( pds_hostname = %pds_hostname, "Normalized username for lookup" ); - let user = match sqlx::query!( - r#" - SELECT id, did, email, password_hash, password_required, two_factor_enabled, - preferred_comms_channel as "preferred_comms_channel: CommsChannel", - deactivated_at, takedown_ref, - email_verified, discord_verified, telegram_verified, signal_verified, - account_type::text as "account_type!" - FROM users - WHERE handle = $1 OR email = $1 - "#, - normalized_username - ) - .fetch_optional(&state.db) - .await - { + let user = match state.user_repo.get_login_info_by_handle_or_email(&normalized_username).await { Ok(Some(u)) => u, Ok(None) => { let _ = bcrypt::verify( @@ -596,7 +588,7 @@ pub async fn authorize_post( } if user.account_type == "delegated" { - if db::set_authorization_did(&state.db, &form.request_uri, &user.did, None) + if state.oauth_repo.set_authorization_did(&form_request_id, &user.did, None) .await .is_err() { @@ -622,7 +614,7 @@ pub async fn authorize_post( } if !user.password_required { - if db::set_authorization_did(&state.db, &form.request_uri, &user.did, None) + if state.oauth_repo.set_authorization_did(&form_request_id, &user.did, None) .await .is_err() { @@ -661,17 +653,17 @@ pub async fn authorize_post( if has_totp { let device_cookie = extract_device_cookie(&headers); let device_is_trusted = if let Some(ref dev_id) = device_cookie { - crate::api::server::is_device_trusted(&state.db, dev_id, &user.did).await + crate::api::server::is_device_trusted(state.oauth_repo.as_ref(), dev_id, &user.did).await } else { false }; if device_is_trusted { if let Some(ref dev_id) = device_cookie { - let _ = crate::api::server::extend_device_trust(&state.db, dev_id).await; + let _ = crate::api::server::extend_device_trust(state.oauth_repo.as_ref(), dev_id).await; } } else { - if db::set_authorization_did(&state.db, &form.request_uri, &user.did, None) + if state.oauth_repo.set_authorization_did(&form_request_id, &user.did, None) .await .is_err() { @@ -690,13 +682,13 @@ pub async fn authorize_post( } } if user.two_factor_enabled { - let _ = db::delete_2fa_challenge_by_request_uri(&state.db, &form.request_uri).await; - match db::create_2fa_challenge(&state.db, &user.did, &form.request_uri).await { + let _ = state.oauth_repo.delete_2fa_challenge_by_request_uri(&form_request_id).await; + match state.oauth_repo.create_2fa_challenge(&user.did, &form_request_id).await { Ok(challenge) => { let hostname = std::env::var("PDS_HOSTNAME").unwrap_or_else(|_| "localhost".to_string()); if let Err(e) = - enqueue_2fa_code(&state.db, user.id, &challenge.code, &hostname).await + enqueue_2fa_code(state.user_repo.as_ref(), state.infra_repo.as_ref(), user.id, &challenge.code, &hostname).await { tracing::warn!( did = %user.did, @@ -736,7 +728,8 @@ pub async fn authorize_post( ip_address: extract_client_ip(&headers), last_seen_at: Utc::now(), }; - if db::create_device(&state.db, &new_id.0, &device_data) + let new_device_id_typed = DeviceIdType::from(new_id.0.clone()); + if state.oauth_repo.create_device(&new_device_id_typed, &device_data) .await .is_ok() { @@ -745,16 +738,15 @@ pub async fn authorize_post( } new_id.0 }; - let _ = db::upsert_account_device(&state.db, &user.did, &final_device_id).await; + let final_device_typed = DeviceIdType::from(final_device_id.clone()); + let _ = state.oauth_repo.upsert_account_device(&user.did, &final_device_typed).await; } - if db::set_authorization_did( - &state.db, - &form.request_uri, - &user.did, - device_id.as_deref(), - ) - .await - .is_err() + let set_auth_device_id = device_id.as_ref().map(|d| DeviceIdType::from(d.clone())); + if state + .oauth_repo + .set_authorization_did(&form_request_id, &user.did, set_auth_device_id.as_ref()) + .await + .is_err() { return show_login_error("An error occurred. Please try again.", json_response); } @@ -767,10 +759,11 @@ pub async fn authorize_post( .split_whitespace() .map(|s| s.to_string()) .collect(); - let needs_consent = db::should_show_consent( - &state.db, + let client_id_typed = ClientId::from(request_data.parameters.client_id.clone()); + let needs_consent = should_show_consent( + state.oauth_repo.as_ref(), &user.did, - &request_data.parameters.client_id, + &client_id_typed, &requested_scopes, ) .await @@ -801,12 +794,13 @@ pub async fn authorize_post( return redirect_see_other(&consent_url); } let code = Code::generate(); - if db::update_authorization_request( - &state.db, - &form.request_uri, + let auth_post_device_id = device_id.as_ref().map(|d| DeviceIdType::from(d.clone())); + let auth_post_code = AuthorizationCode::from(code.0.clone()); + if state.oauth_repo.update_authorization_request( + &form_request_id, &user.did, - device_id.as_deref(), - &code.0, + auth_post_device_id.as_ref(), + &auth_post_code, ) .await .is_err() @@ -864,7 +858,8 @@ pub async fn authorize_select( ) .into_response() }; - let request_data = match db::get_authorization_request(&state.db, &form.request_uri).await { + let select_request_id = RequestId::from(form.request_uri.clone()); + let request_data = match state.oauth_repo.get_authorization_request(&select_request_id).await { Ok(Some(data)) => data, Ok(None) => { return json_error( @@ -882,7 +877,7 @@ pub async fn authorize_select( } }; if request_data.expires_at < Utc::now() { - let _ = db::delete_authorization_request(&state.db, &form.request_uri).await; + let _ = state.oauth_repo.delete_authorization_request(&select_request_id).await; return json_error( StatusCode::BAD_REQUEST, "invalid_request", @@ -899,7 +894,18 @@ pub async fn authorize_select( ); } }; - let account_valid = match db::verify_account_on_device(&state.db, &device_id, &form.did).await { + let did: Did = match form.did.parse() { + Ok(d) => d, + Err(_) => { + return json_error( + StatusCode::BAD_REQUEST, + "invalid_request", + "Invalid DID format.", + ); + } + }; + let verify_device_id = DeviceIdType::from(device_id.clone()); + let account_valid = match state.oauth_repo.verify_account_on_device(&verify_device_id, &did).await { Ok(valid) => valid, Err(_) => { return json_error( @@ -916,19 +922,7 @@ pub async fn authorize_select( "This account is not available on this device. Please sign in.", ); } - let user = match sqlx::query!( - r#" - SELECT id, two_factor_enabled, - preferred_comms_channel as "preferred_comms_channel: CommsChannel", - email_verified, discord_verified, telegram_verified, signal_verified - FROM users - WHERE did = $1 - "#, - form.did - ) - .fetch_optional(&state.db) - .await - { + let user = match state.user_repo.get_2fa_status_by_did(&did).await { Ok(Some(u)) => u, Ok(None) => { return json_error( @@ -956,9 +950,10 @@ pub async fn authorize_select( "Please verify your account before logging in.", ); } - let has_totp = crate::api::server::has_totp_enabled(&state, &form.did).await; + let has_totp = crate::api::server::has_totp_enabled(&state, &did).await; + let select_early_device_typed = DeviceIdType::from(device_id.clone()); if has_totp { - if db::set_authorization_did(&state.db, &form.request_uri, &form.did, Some(&device_id)) + if state.oauth_repo.set_authorization_did(&select_request_id, &did, Some(&select_early_device_typed)) .await .is_err() { @@ -974,13 +969,13 @@ pub async fn authorize_select( .into_response(); } if user.two_factor_enabled { - let _ = db::delete_2fa_challenge_by_request_uri(&state.db, &form.request_uri).await; - match db::create_2fa_challenge(&state.db, &form.did, &form.request_uri).await { + let _ = state.oauth_repo.delete_2fa_challenge_by_request_uri(&select_request_id).await; + match state.oauth_repo.create_2fa_challenge(&did, &select_request_id).await { Ok(challenge) => { let hostname = std::env::var("PDS_HOSTNAME").unwrap_or_else(|_| "localhost".to_string()); if let Err(e) = - enqueue_2fa_code(&state.db, user.id, &challenge.code, &hostname).await + enqueue_2fa_code(state.user_repo.as_ref(), state.infra_repo.as_ref(), user.id, &challenge.code, &hostname).await { tracing::warn!( did = %form.did, @@ -1004,14 +999,15 @@ pub async fn authorize_select( } } } - let _ = db::upsert_account_device(&state.db, &form.did, &device_id).await; + let select_device_typed = DeviceIdType::from(device_id.clone()); + let _ = state.oauth_repo.upsert_account_device(&did, &select_device_typed).await; let code = Code::generate(); - if db::update_authorization_request( - &state.db, - &form.request_uri, - &form.did, - Some(&device_id), - &code.0, + let select_code = AuthorizationCode::from(code.0.clone()); + if state.oauth_repo.update_authorization_request( + &select_request_id, + &did, + Some(&select_device_typed), + &select_code, ) .await .is_err() @@ -1124,7 +1120,8 @@ pub async fn authorize_deny( State(state): State, Json(form): Json, ) -> Response { - let request_data = match db::get_authorization_request(&state.db, &form.request_uri).await { + let deny_request_id = RequestId::from(form.request_uri.clone()); + let request_data = match state.oauth_repo.get_authorization_request(&deny_request_id).await { Ok(Some(data)) => data, Ok(None) => { return ( @@ -1147,7 +1144,7 @@ pub async fn authorize_deny( .into_response(); } }; - let _ = db::delete_authorization_request(&state.db, &form.request_uri).await; + let _ = state.oauth_repo.delete_authorization_request(&deny_request_id).await; let redirect_uri = &request_data.parameters.redirect_uri; let mut redirect_url = redirect_uri.to_string(); let separator = if redirect_url.contains('?') { '&' } else { '?' }; @@ -1188,7 +1185,8 @@ pub async fn authorize_2fa_get( State(state): State, Query(query): Query, ) -> Response { - let challenge = match db::get_2fa_challenge(&state.db, &query.request_uri).await { + let twofa_request_id = RequestId::from(query.request_uri.clone()); + let challenge = match state.oauth_repo.get_2fa_challenge(&twofa_request_id).await { Ok(Some(c)) => c, Ok(None) => { return redirect_to_frontend_error( @@ -1204,13 +1202,13 @@ pub async fn authorize_2fa_get( } }; if challenge.expires_at < Utc::now() { - let _ = db::delete_2fa_challenge(&state.db, challenge.id).await; + let _ = state.oauth_repo.delete_2fa_challenge(challenge.id).await; return redirect_to_frontend_error( "invalid_request", "2FA code has expired. Please start over.", ); } - let _request_data = match db::get_authorization_request(&state.db, &query.request_uri).await { + let _request_data = match state.oauth_repo.get_authorization_request(&twofa_request_id).await { Ok(Some(d)) => d, Ok(None) => { return redirect_to_frontend_error( @@ -1279,9 +1277,10 @@ pub async fn consent_get( State(state): State, Query(query): Query, ) -> Response { - let (request_data, flow_state) = - match db::get_authorization_request_with_state(&state.db, &query.request_uri).await { - Ok(Some(result)) => result, + let consent_request_id = RequestId::from(query.request_uri.clone()); + let request_data = + match state.oauth_repo.get_authorization_request(&consent_request_id).await { + Ok(Some(data)) => data, Ok(None) => { return json_error( StatusCode::BAD_REQUEST, @@ -1297,15 +1296,26 @@ pub async fn consent_get( ); } }; + let flow_state = AuthFlowState::from_request_data(&request_data); if let Some(err_response) = validate_auth_flow_state(&flow_state, true) { if flow_state.is_expired() { - let _ = db::delete_authorization_request(&state.db, &query.request_uri).await; + let _ = state.oauth_repo.delete_authorization_request(&consent_request_id).await; } return err_response; } - let did = flow_state.did().unwrap().to_string(); + let did_str = flow_state.did().unwrap().to_string(); + let did: Did = match did_str.parse() { + Ok(d) => d, + Err(_) => { + return json_error( + StatusCode::BAD_REQUEST, + "invalid_request", + "Invalid DID format in request.", + ); + } + }; let client_cache = ClientMetadataCache::new(3600); let client_metadata = client_cache .get(&request_data.parameters.client_id) @@ -1318,8 +1328,11 @@ pub async fn consent_get( .filter(|s| !s.trim().is_empty()) .unwrap_or("atproto"); - let delegation_grant = if let Some(ref ctrl_did) = request_data.controller_did { - crate::delegation::get_delegation(&state.db, &did, ctrl_did) + let controller_did_parsed: Option = request_data.controller_did.as_ref().and_then(|s| s.parse().ok()); + let delegation_grant = if let Some(ref ctrl_did) = controller_did_parsed { + state + .delegation_repo + .get_delegation(&did, ctrl_did) .await .ok() .flatten() @@ -1328,14 +1341,15 @@ pub async fn consent_get( }; let effective_scope_str = if let Some(ref grant) = delegation_grant { - crate::delegation::scopes::intersect_scopes(requested_scope_str, &grant.granted_scopes) + crate::delegation::intersect_scopes(requested_scope_str, &grant.granted_scopes) } else { requested_scope_str.to_string() }; let requested_scopes: Vec<&str> = effective_scope_str.split_whitespace().collect(); + let consent_client_id = ClientId::from(request_data.parameters.client_id.clone()); let preferences = - db::get_scope_preferences(&state.db, &did, &request_data.parameters.client_id) + state.oauth_repo.get_scope_preferences(&did, &consent_client_id) .await .unwrap_or_default(); let pref_map: std::collections::HashMap<_, _> = preferences @@ -1344,10 +1358,10 @@ pub async fn consent_get( .collect(); let requested_scope_strings: Vec = requested_scopes.iter().map(|s| s.to_string()).collect(); - let show_consent = db::should_show_consent( - &state.db, + let show_consent = should_show_consent( + state.oauth_repo.as_ref(), &did, - &request_data.parameters.client_id, + &consent_client_id, &requested_scope_strings, ) .await @@ -1389,19 +1403,18 @@ pub async fn consent_get( } }) .collect(); - let (is_delegation, controller_did, controller_handle, delegation_level) = - if let Some(ref ctrl_did) = request_data.controller_did { - let ctrl_handle = - sqlx::query_scalar!("SELECT handle FROM users WHERE did = $1", ctrl_did) - .fetch_optional(&state.db) - .await - .ok() - .flatten(); + let (is_delegation, controller_did_resp, controller_handle, delegation_level) = + if let Some(ref ctrl_did) = controller_did_parsed { + let ctrl_handle = state + .user_repo + .get_handle_by_did(ctrl_did) + .await + .ok() + .flatten() + .map(|h| h.to_string()); let level = if let Some(ref grant) = delegation_grant { - let preset = crate::delegation::SCOPE_PRESETS - .iter() - .find(|p| p.scopes == grant.granted_scopes); + let preset = crate::delegation::SCOPE_PRESETS.iter().find(|p| p.scopes == grant.granted_scopes); preset .map(|p| p.label.to_string()) .unwrap_or_else(|| "Custom".to_string()) @@ -1409,7 +1422,7 @@ pub async fn consent_get( "Unknown".to_string() }; - (Some(true), Some(ctrl_did.clone()), ctrl_handle, Some(level)) + (Some(true), Some(ctrl_did.to_string()), ctrl_handle, Some(level)) } else { (None, None, None, None) }; @@ -1422,9 +1435,9 @@ pub async fn consent_get( logo_uri: client_metadata.as_ref().and_then(|m| m.logo_uri.clone()), scopes, show_consent, - did, + did: did_str, is_delegation, - controller_did, + controller_did: controller_did_resp, controller_handle, delegation_level, }) @@ -1440,9 +1453,10 @@ pub async fn consent_post( form.approved_scopes, form.remember ); - let (request_data, flow_state) = - match db::get_authorization_request_with_state(&state.db, &form.request_uri).await { - Ok(Some(result)) => result, + let consent_post_request_id = RequestId::from(form.request_uri.clone()); + let request_data = + match state.oauth_repo.get_authorization_request(&consent_post_request_id).await { + Ok(Some(data)) => data, Ok(None) => { return json_error( StatusCode::BAD_REQUEST, @@ -1458,9 +1472,10 @@ pub async fn consent_post( ); } }; + let flow_state = AuthFlowState::from_request_data(&request_data); if flow_state.is_expired() { - let _ = db::delete_authorization_request(&state.db, &form.request_uri).await; + let _ = state.oauth_repo.delete_authorization_request(&consent_post_request_id).await; return json_error( StatusCode::BAD_REQUEST, "invalid_request", @@ -1471,24 +1486,40 @@ pub async fn consent_post( return json_error(StatusCode::FORBIDDEN, "access_denied", "Not authenticated"); } - let did = flow_state.did().unwrap().to_string(); + let did_str = flow_state.did().unwrap().to_string(); + let did: Did = match did_str.parse() { + Ok(d) => d, + Err(_) => { + return json_error( + StatusCode::BAD_REQUEST, + "invalid_request", + "Invalid DID format", + ); + } + }; let original_scope_str = request_data .parameters .scope .as_deref() .unwrap_or("atproto"); - let delegation_grant = if let Some(ref ctrl_did) = request_data.controller_did { - crate::delegation::get_delegation(&state.db, &did, ctrl_did) + let controller_did_parsed: Option = request_data + .controller_did + .as_ref() + .and_then(|s| s.parse().ok()); + + let delegation_grant = match controller_did_parsed.as_ref() { + Some(ctrl_did) => state + .delegation_repo + .get_delegation(&did, ctrl_did) .await .ok() - .flatten() - } else { - None + .flatten(), + None => None, }; let effective_scope_str = if let Some(ref grant) = delegation_grant { - crate::delegation::scopes::intersect_scopes(original_scope_str, &grant.granted_scopes) + crate::delegation::intersect_scopes(original_scope_str, &grant.granted_scopes) } else { original_scope_str.to_string() }; @@ -1537,33 +1568,34 @@ pub async fn consent_post( ); } if form.remember { - let preferences: Vec = requested_scopes + let preferences: Vec = requested_scopes .iter() - .map(|s| db::ScopePreference { + .map(|s| ScopePreference { scope: s.to_string(), granted: form.approved_scopes.contains(&s.to_string()), }) .collect(); - let _ = db::upsert_scope_preferences( - &state.db, + let consent_post_client_id = ClientId::from(request_data.parameters.client_id.clone()); + let _ = state.oauth_repo.upsert_scope_preferences( &did, - &request_data.parameters.client_id, + &consent_post_client_id, &preferences, ) .await; } if let Err(e) = - db::update_request_scope(&state.db, &form.request_uri, &approved_scope_str).await + state.oauth_repo.update_request_scope(&consent_post_request_id, &approved_scope_str).await { tracing::warn!("Failed to update request scope: {:?}", e); } let code = Code::generate(); - if db::update_authorization_request( - &state.db, - &form.request_uri, + let consent_post_device_id = request_data.device_id.as_ref().map(|d| DeviceIdType::from(d.clone())); + let consent_post_code = AuthorizationCode::from(code.0.clone()); + if state.oauth_repo.update_authorization_request( + &consent_post_request_id, &did, - request_data.device_id.as_deref(), - &code.0, + consent_post_device_id.as_ref(), + &consent_post_code, ) .await .is_err() @@ -1616,7 +1648,8 @@ pub async fn authorize_2fa_post( "Too many attempts. Please try again later.", ); } - let request_data = match db::get_authorization_request(&state.db, &form.request_uri).await { + let twofa_post_request_id = RequestId::from(form.request_uri.clone()); + let request_data = match state.oauth_repo.get_authorization_request(&twofa_post_request_id).await { Ok(Some(d)) => d, Ok(None) => { return json_error( @@ -1634,20 +1667,20 @@ pub async fn authorize_2fa_post( } }; if request_data.expires_at < Utc::now() { - let _ = db::delete_authorization_request(&state.db, &form.request_uri).await; + let _ = state.oauth_repo.delete_authorization_request(&twofa_post_request_id).await; return json_error( StatusCode::BAD_REQUEST, "invalid_request", "Authorization request has expired.", ); } - let challenge = db::get_2fa_challenge(&state.db, &form.request_uri) + let challenge = state.oauth_repo.get_2fa_challenge(&twofa_post_request_id) .await .ok() .flatten(); if let Some(challenge) = challenge { if challenge.expires_at < Utc::now() { - let _ = db::delete_2fa_challenge(&state.db, challenge.id).await; + let _ = state.oauth_repo.delete_2fa_challenge( challenge.id).await; return json_error( StatusCode::BAD_REQUEST, "invalid_request", @@ -1655,7 +1688,7 @@ pub async fn authorize_2fa_post( ); } if challenge.attempts >= MAX_2FA_ATTEMPTS { - let _ = db::delete_2fa_challenge(&state.db, challenge.id).await; + let _ = state.oauth_repo.delete_2fa_challenge( challenge.id).await; return json_error( StatusCode::FORBIDDEN, "access_denied", @@ -1669,22 +1702,23 @@ pub async fn authorize_2fa_post( .ct_eq(challenge.code.as_bytes()) .into(); if !code_valid { - let _ = db::increment_2fa_attempts(&state.db, challenge.id).await; + let _ = state.oauth_repo.increment_2fa_attempts(challenge.id).await; return json_error( StatusCode::FORBIDDEN, "invalid_code", "Invalid verification code. Please try again.", ); } - let _ = db::delete_2fa_challenge(&state.db, challenge.id).await; + let _ = state.oauth_repo.delete_2fa_challenge(challenge.id).await; let code = Code::generate(); let device_id = extract_device_cookie(&headers); - if db::update_authorization_request( - &state.db, - &form.request_uri, + let twofa_totp_device_id = device_id.as_ref().map(|d| DeviceIdType::from(d.clone())); + let twofa_totp_code = AuthorizationCode::from(code.0.clone()); + if state.oauth_repo.update_authorization_request( + &twofa_post_request_id, &challenge.did, - device_id.as_deref(), - &code.0, + twofa_totp_device_id.as_ref(), + &twofa_totp_code, ) .await .is_err() @@ -1706,7 +1740,7 @@ pub async fn authorize_2fa_post( })) .into_response(); } - let did = match &request_data.did { + let did_str = match &request_data.did { Some(d) => d.clone(), None => { return json_error( @@ -1716,6 +1750,12 @@ pub async fn authorize_2fa_post( ); } }; + let did: tranquil_types::Did = match did_str.parse() { + Ok(d) => d, + Err(_) => { + return json_error(StatusCode::BAD_REQUEST, "invalid_request", "Invalid DID format."); + } + }; if !crate::api::server::has_totp_enabled(&state, &did).await { return json_error( StatusCode::BAD_REQUEST, @@ -1747,7 +1787,7 @@ pub async fn authorize_2fa_post( if form.trust_device && let Some(ref dev_id) = device_id { - let _ = crate::api::server::trust_device(&state.db, dev_id).await; + let _ = crate::api::server::trust_device(state.oauth_repo.as_ref(), dev_id).await; } let requested_scope_str = request_data .parameters @@ -1758,10 +1798,11 @@ pub async fn authorize_2fa_post( .split_whitespace() .map(|s| s.to_string()) .collect(); - let needs_consent = db::should_show_consent( - &state.db, + let twofa_post_client_id = ClientId::from(request_data.parameters.client_id.clone()); + let needs_consent = should_show_consent( + state.oauth_repo.as_ref(), &did, - &request_data.parameters.client_id, + &twofa_post_client_id, &requested_scopes, ) .await @@ -1774,12 +1815,13 @@ pub async fn authorize_2fa_post( return Json(serde_json::json!({"redirect_uri": consent_url})).into_response(); } let code = Code::generate(); - if db::update_authorization_request( - &state.db, - &form.request_uri, + let twofa_final_device_id = device_id.as_ref().map(|d| DeviceIdType::from(d.clone())); + let twofa_final_code = AuthorizationCode::from(code.0.clone()); + if state.oauth_repo.update_authorization_request( + &twofa_post_request_id, &did, - device_id.as_deref(), - &code.0, + twofa_final_device_id.as_ref(), + &twofa_final_code, ) .await .is_err() @@ -1819,24 +1861,23 @@ pub async fn check_user_has_passkeys( Query(query): Query, ) -> Response { let pds_hostname = std::env::var("PDS_HOSTNAME").unwrap_or_else(|_| "localhost".to_string()); + let hostname_for_handles = pds_hostname.split(':').next().unwrap_or(&pds_hostname); let normalized_identifier = query.identifier.trim(); let normalized_identifier = normalized_identifier .strip_prefix('@') .unwrap_or(normalized_identifier); let normalized_identifier = if let Some(bare_handle) = - normalized_identifier.strip_suffix(&format!(".{}", pds_hostname)) + normalized_identifier.strip_suffix(&format!(".{}", hostname_for_handles)) { bare_handle.to_string() } else { normalized_identifier.to_string() }; - let user = sqlx::query!( - "SELECT did FROM users WHERE handle = $1 OR email = $1", - normalized_identifier - ) - .fetch_optional(&state.db) - .await; + let user = state + .user_repo + .get_login_check_by_handle_or_email(&normalized_identifier) + .await; let has_passkeys = match user { Ok(Some(u)) => crate::api::server::has_passkeys_for_user(&state, &u.did).await, @@ -1862,22 +1903,21 @@ pub async fn check_user_security_status( Query(query): Query, ) -> Response { let pds_hostname = std::env::var("PDS_HOSTNAME").unwrap_or_else(|_| "localhost".to_string()); + let hostname_for_handles = pds_hostname.split(':').next().unwrap_or(&pds_hostname); let identifier = query.identifier.trim(); let identifier = identifier.strip_prefix('@').unwrap_or(identifier); let normalized_identifier = if identifier.contains('@') || identifier.starts_with("did:") { identifier.to_string() } else if !identifier.contains('.') { - format!("{}.{}", identifier.to_lowercase(), pds_hostname) + format!("{}.{}", identifier.to_lowercase(), hostname_for_handles) } else { identifier.to_lowercase() }; - let user = sqlx::query!( - "SELECT did, password_hash FROM users WHERE handle = $1 OR email = $1", - normalized_identifier - ) - .fetch_optional(&state.db) - .await; + let user = state + .user_repo + .get_login_check_by_handle_or_email(&normalized_identifier) + .await; let (has_passkeys, has_totp, has_password, is_delegated, did): ( bool, @@ -1890,10 +1930,12 @@ pub async fn check_user_security_status( let passkeys = crate::api::server::has_passkeys_for_user(&state, &u.did).await; let totp = crate::api::server::has_totp_enabled(&state, &u.did).await; let has_pw = u.password_hash.is_some(); - let has_controllers = crate::delegation::is_delegated_account(&state.db, &u.did) + let has_controllers = state + .delegation_repo + .is_delegated_account(&u.did) .await .unwrap_or(false); - (passkeys, totp, has_pw, has_controllers, Some(u.did)) + (passkeys, totp, has_pw, has_controllers, Some(u.did.to_string())) } _ => (false, false, false, false, None), }; @@ -1942,7 +1984,8 @@ pub async fn passkey_start( .into_response(); } - let request_data = match db::get_authorization_request(&state.db, &form.request_uri).await { + let passkey_start_request_id = RequestId::from(form.request_uri.clone()); + let request_data = match state.oauth_repo.get_authorization_request(&passkey_start_request_id).await { Ok(Some(data)) => data, Ok(None) => { return ( @@ -1967,7 +2010,7 @@ pub async fn passkey_start( }; if request_data.expires_at < Utc::now() { - let _ = db::delete_authorization_request(&state.db, &form.request_uri).await; + let _ = state.oauth_repo.delete_authorization_request(&passkey_start_request_id).await; return ( StatusCode::BAD_REQUEST, Json(serde_json::json!({ @@ -1979,6 +2022,7 @@ pub async fn passkey_start( } let pds_hostname = std::env::var("PDS_HOSTNAME").unwrap_or_else(|_| "localhost".to_string()); + let hostname_for_handles = pds_hostname.split(':').next().unwrap_or(&pds_hostname); let normalized_username = form.identifier.trim(); let normalized_username = normalized_username .strip_prefix('@') @@ -1986,23 +2030,12 @@ pub async fn passkey_start( let normalized_username = if normalized_username.contains('@') { normalized_username.to_string() } else if !normalized_username.contains('.') { - format!("{}.{}", normalized_username, pds_hostname) + format!("{}.{}", normalized_username, hostname_for_handles) } else { normalized_username.to_string() }; - let user = match sqlx::query!( - r#" - SELECT did, deactivated_at, takedown_ref, - email_verified, discord_verified, telegram_verified, signal_verified - FROM users - WHERE handle = $1 OR email = $1 - "#, - normalized_username - ) - .fetch_optional(&state.db) - .await - { + let user = match state.user_repo.get_login_info_by_handle_or_email(&normalized_username).await { Ok(Some(u)) => u, Ok(None) => { return ( @@ -2064,21 +2097,20 @@ pub async fn passkey_start( .into_response(); } - let stored_passkeys = - match crate::auth::webauthn::get_passkeys_for_user(&state.db, &user.did).await { - Ok(pks) => pks, - Err(e) => { - tracing::error!(error = %e, "Failed to get passkeys"); - return ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({ - "error": "server_error", - "error_description": "An error occurred." - })), - ) - .into_response(); - } - }; + let stored_passkeys = match state.user_repo.get_passkeys_for_user(&user.did).await { + Ok(pks) => pks, + Err(e) => { + tracing::error!(error = %e, "Failed to get passkeys"); + return ( + StatusCode::INTERNAL_SERVER_ERROR, + Json(serde_json::json!({ + "error": "server_error", + "error_description": "An error occurred." + })), + ) + .into_response(); + } + }; if stored_passkeys.is_empty() { return ( @@ -2093,7 +2125,7 @@ pub async fn passkey_start( let passkeys: Vec = stored_passkeys .iter() - .filter_map(|sp| sp.to_security_key().ok()) + .filter_map(|sp| serde_json::from_slice(&sp.public_key).ok()) .collect(); if passkeys.is_empty() { @@ -2137,8 +2169,25 @@ pub async fn passkey_start( } }; - if let Err(e) = - crate::auth::webauthn::save_authentication_state(&state.db, &user.did, &auth_state).await + let state_json = match serde_json::to_string(&auth_state) { + Ok(j) => j, + Err(e) => { + tracing::error!(error = %e, "Failed to serialize authentication state"); + return ( + StatusCode::INTERNAL_SERVER_ERROR, + Json(serde_json::json!({ + "error": "server_error", + "error_description": "An error occurred." + })), + ) + .into_response(); + } + }; + + if let Err(e) = state + .user_repo + .save_webauthn_challenge(&user.did, "authentication", &state_json) + .await { tracing::error!(error = %e, "Failed to save authentication state"); return ( @@ -2151,7 +2200,7 @@ pub async fn passkey_start( .into_response(); } - if db::set_authorization_did(&state.db, &form.request_uri, &user.did, None) + if state.oauth_repo.set_authorization_did(&passkey_start_request_id, &user.did, None) .await .is_err() { @@ -2181,7 +2230,8 @@ pub async fn passkey_finish( headers: HeaderMap, Json(form): Json, ) -> Response { - let request_data = match db::get_authorization_request(&state.db, &form.request_uri).await { + let passkey_finish_request_id = RequestId::from(form.request_uri.clone()); + let request_data = match state.oauth_repo.get_authorization_request(&passkey_finish_request_id).await { Ok(Some(data)) => data, Ok(None) => { return ( @@ -2206,7 +2256,7 @@ pub async fn passkey_finish( }; if request_data.expires_at < Utc::now() { - let _ = db::delete_authorization_request(&state.db, &form.request_uri).await; + let _ = state.oauth_repo.delete_authorization_request(&passkey_finish_request_id).await; return ( StatusCode::BAD_REQUEST, Json(serde_json::json!({ @@ -2217,7 +2267,7 @@ pub async fn passkey_finish( .into_response(); } - let did = match request_data.did { + let did_str = match request_data.did { Some(d) => d, None => { return ( @@ -2230,8 +2280,25 @@ pub async fn passkey_finish( .into_response(); } }; + let did: tranquil_types::Did = match did_str.parse() { + Ok(d) => d, + Err(_) => { + return ( + StatusCode::BAD_REQUEST, + Json(serde_json::json!({ + "error": "invalid_request", + "error_description": "Invalid DID format." + })), + ) + .into_response(); + } + }; - let auth_state = match crate::auth::webauthn::load_authentication_state(&state.db, &did).await { + let auth_state_json = match state + .user_repo + .load_webauthn_challenge(&did, "authentication") + .await + { Ok(Some(s)) => s, Ok(None) => { return ( @@ -2256,6 +2323,22 @@ pub async fn passkey_finish( } }; + let auth_state: webauthn_rs::prelude::SecurityKeyAuthentication = + match serde_json::from_str(&auth_state_json) { + Ok(s) => s, + Err(e) => { + tracing::error!(error = %e, "Failed to deserialize authentication state"); + return ( + StatusCode::INTERNAL_SERVER_ERROR, + Json(serde_json::json!({ + "error": "server_error", + "error_description": "An error occurred." + })), + ) + .into_response(); + } + }; + let credential: webauthn_rs::prelude::PublicKeyCredential = match serde_json::from_value(form.credential) { Ok(c) => c, @@ -2303,17 +2386,20 @@ pub async fn passkey_finish( } }; - if let Err(e) = crate::auth::webauthn::delete_authentication_state(&state.db, &did).await { + if let Err(e) = state + .user_repo + .delete_webauthn_challenge(&did, "authentication") + .await + { tracing::warn!(error = %e, "Failed to delete authentication state"); } if auth_result.needs_update() { - match crate::auth::webauthn::update_passkey_counter( - &state.db, - auth_result.cred_id(), - auth_result.counter(), - ) - .await + let cred_id_bytes = auth_result.cred_id().as_slice(); + match state + .user_repo + .update_passkey_counter(cred_id_bytes, auth_result.counter() as i32) + .await { Ok(false) => { tracing::warn!(did = %did, "Passkey counter anomaly detected - possible cloned key"); @@ -2343,23 +2429,18 @@ pub async fn passkey_finish( .into_response(); } - let user = sqlx::query!( - "SELECT two_factor_enabled, preferred_comms_channel as \"preferred_comms_channel: CommsChannel\", id FROM users WHERE did = $1", - did - ) - .fetch_optional(&state.db) - .await; + let user = state.user_repo.get_2fa_status_by_did(&did).await; if let Ok(Some(user)) = user && user.two_factor_enabled { - let _ = db::delete_2fa_challenge_by_request_uri(&state.db, &form.request_uri).await; - match db::create_2fa_challenge(&state.db, &did, &form.request_uri).await { + let _ = state.oauth_repo.delete_2fa_challenge_by_request_uri(&passkey_finish_request_id).await; + match state.oauth_repo.create_2fa_challenge(&did, &passkey_finish_request_id).await { Ok(challenge) => { let hostname = std::env::var("PDS_HOSTNAME").unwrap_or_else(|_| "localhost".to_string()); if let Err(e) = - enqueue_2fa_code(&state.db, user.id, &challenge.code, &hostname).await + enqueue_2fa_code(state.user_repo.as_ref(), state.infra_repo.as_ref(), user.id, &challenge.code, &hostname).await { tracing::warn!(did = %did, error = %e, "Failed to enqueue 2FA notification"); } @@ -2394,10 +2475,11 @@ pub async fn passkey_finish( .map(|s| s.to_string()) .collect(); - let needs_consent = db::should_show_consent( - &state.db, + let passkey_finish_client_id = ClientId::from(request_data.parameters.client_id.clone()); + let needs_consent = should_show_consent( + state.oauth_repo.as_ref(), &did, - &request_data.parameters.client_id, + &passkey_finish_client_id, &requested_scopes, ) .await @@ -2412,12 +2494,13 @@ pub async fn passkey_finish( } let code = Code::generate(); - if db::update_authorization_request( - &state.db, - &form.request_uri, + let passkey_final_device_id = device_id.as_ref().map(|d| DeviceIdType::from(d.clone())); + let passkey_final_code = AuthorizationCode::from(code.0.clone()); + if state.oauth_repo.update_authorization_request( + &passkey_finish_request_id, &did, - device_id.as_deref(), - &code.0, + passkey_final_device_id.as_ref(), + &passkey_final_code, ) .await .is_err() @@ -2463,7 +2546,8 @@ pub async fn authorize_passkey_start( ) -> Response { let pds_hostname = std::env::var("PDS_HOSTNAME").unwrap_or_else(|_| "localhost".to_string()); - let request_data = match db::get_authorization_request(&state.db, &query.request_uri).await { + let auth_passkey_start_request_id = RequestId::from(query.request_uri.clone()); + let request_data = match state.oauth_repo.get_authorization_request(&auth_passkey_start_request_id).await { Ok(Some(d)) => d, Ok(None) => { return ( @@ -2488,7 +2572,7 @@ pub async fn authorize_passkey_start( }; if request_data.expires_at < Utc::now() { - let _ = db::delete_authorization_request(&state.db, &query.request_uri).await; + let _ = state.oauth_repo.delete_authorization_request(&auth_passkey_start_request_id).await; return ( StatusCode::BAD_REQUEST, Json(serde_json::json!({ @@ -2499,7 +2583,7 @@ pub async fn authorize_passkey_start( .into_response(); } - let did = match &request_data.did { + let did_str = match &request_data.did { Some(d) => d.clone(), None => { return ( @@ -2513,16 +2597,29 @@ pub async fn authorize_passkey_start( } }; - let stored_passkeys = match crate::auth::webauthn::get_passkeys_for_user(&state.db, &did).await - { + let did: tranquil_types::Did = match did_str.parse() { + Ok(d) => d, + Err(_) => { + return ( + StatusCode::BAD_REQUEST, + Json(serde_json::json!({ + "error": "invalid_request", + "error_description": "Invalid DID format." + })), + ) + .into_response(); + } + }; + + let stored_passkeys = match state.user_repo.get_passkeys_for_user(&did).await { Ok(pks) => pks, Err(e) => { tracing::error!("Failed to get passkeys: {:?}", e); return ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": "server_error", "error_description": "An error occurred."})), - ) - .into_response(); + StatusCode::INTERNAL_SERVER_ERROR, + Json(serde_json::json!({"error": "server_error", "error_description": "An error occurred."})), + ) + .into_response(); } }; @@ -2539,7 +2636,7 @@ pub async fn authorize_passkey_start( let passkeys: Vec = stored_passkeys .iter() - .filter_map(|sp| sp.to_security_key().ok()) + .filter_map(|sp| serde_json::from_slice(&sp.public_key).ok()) .collect(); if passkeys.is_empty() { @@ -2574,8 +2671,22 @@ pub async fn authorize_passkey_start( } }; - if let Err(e) = - crate::auth::webauthn::save_authentication_state(&state.db, &did, &auth_state).await + let state_json = match serde_json::to_string(&auth_state) { + Ok(j) => j, + Err(e) => { + tracing::error!("Failed to serialize authentication state: {:?}", e); + return ( + StatusCode::INTERNAL_SERVER_ERROR, + Json(serde_json::json!({"error": "server_error", "error_description": "An error occurred."})), + ) + .into_response(); + } + }; + + if let Err(e) = state + .user_repo + .save_webauthn_challenge(&did, "authentication", &state_json) + .await { tracing::error!("Failed to save authentication state: {:?}", e); return ( @@ -2606,8 +2717,9 @@ pub async fn authorize_passkey_finish( Json(form): Json, ) -> Response { let pds_hostname = std::env::var("PDS_HOSTNAME").unwrap_or_else(|_| "localhost".to_string()); + let passkey_finish_request_id = RequestId::from(form.request_uri.clone()); - let request_data = match db::get_authorization_request(&state.db, &form.request_uri).await { + let request_data = match state.oauth_repo.get_authorization_request(&passkey_finish_request_id).await { Ok(Some(d)) => d, Ok(None) => { return ( @@ -2632,7 +2744,7 @@ pub async fn authorize_passkey_finish( }; if request_data.expires_at < Utc::now() { - let _ = db::delete_authorization_request(&state.db, &form.request_uri).await; + let _ = state.oauth_repo.delete_authorization_request(&passkey_finish_request_id).await; return ( StatusCode::BAD_REQUEST, Json(serde_json::json!({ @@ -2643,7 +2755,7 @@ pub async fn authorize_passkey_finish( .into_response(); } - let did = match &request_data.did { + let did_str = match &request_data.did { Some(d) => d.clone(), None => { return ( @@ -2657,7 +2769,25 @@ pub async fn authorize_passkey_finish( } }; - let auth_state = match crate::auth::webauthn::load_authentication_state(&state.db, &did).await { + let did: tranquil_types::Did = match did_str.parse() { + Ok(d) => d, + Err(_) => { + return ( + StatusCode::BAD_REQUEST, + Json(serde_json::json!({ + "error": "invalid_request", + "error_description": "Invalid DID format." + })), + ) + .into_response(); + } + }; + + let auth_state_json = match state + .user_repo + .load_webauthn_challenge(&did, "authentication") + .await + { Ok(Some(s)) => s, Ok(None) => { return ( @@ -2672,12 +2802,25 @@ pub async fn authorize_passkey_finish( Err(e) => { tracing::error!("Failed to load authentication state: {:?}", e); return ( + StatusCode::INTERNAL_SERVER_ERROR, + Json(serde_json::json!({"error": "server_error", "error_description": "An error occurred."})), + ) + .into_response(); + } + }; + + let auth_state: webauthn_rs::prelude::SecurityKeyAuthentication = + match serde_json::from_str(&auth_state_json) { + Ok(s) => s, + Err(e) => { + tracing::error!("Failed to deserialize authentication state: {:?}", e); + return ( StatusCode::INTERNAL_SERVER_ERROR, Json(serde_json::json!({"error": "server_error", "error_description": "An error occurred."})), ) .into_response(); - } - }; + } + }; let credential: webauthn_rs::prelude::PublicKeyCredential = match serde_json::from_value(form.credential.clone()) { @@ -2722,14 +2865,15 @@ pub async fn authorize_passkey_finish( } }; - let _ = crate::auth::webauthn::delete_authentication_state(&state.db, &did).await; + let _ = state + .user_repo + .delete_webauthn_challenge(&did, "authentication") + .await; - match crate::auth::webauthn::update_passkey_counter( - &state.db, - credential.id.as_ref(), - auth_result.counter(), - ) - .await + match state + .user_repo + .update_passkey_counter(credential.id.as_ref(), auth_result.counter() as i32) + .await { Ok(false) => { tracing::warn!(did = %did, "Passkey counter anomaly detected - possible cloned key"); @@ -2748,27 +2892,21 @@ pub async fn authorize_passkey_finish( Ok(true) => {} } - let has_totp = crate::api::server::has_totp_enabled_db(&state.db, &did).await; + let has_totp = state.user_repo.has_totp_enabled(&did).await.unwrap_or(false); if has_totp { let device_cookie = extract_device_cookie(&headers); let device_is_trusted = if let Some(ref dev_id) = device_cookie { - crate::api::server::is_device_trusted(&state.db, dev_id, &did).await + crate::api::server::is_device_trusted(state.oauth_repo.as_ref(), dev_id, &did).await } else { false }; if device_is_trusted { if let Some(ref dev_id) = device_cookie { - let _ = crate::api::server::extend_device_trust(&state.db, dev_id).await; + let _ = crate::api::server::extend_device_trust(state.oauth_repo.as_ref(), dev_id).await; } } else { - let user = match sqlx::query!( - r#"SELECT id, preferred_comms_channel as "preferred_comms_channel: CommsChannel" FROM users WHERE did = $1"#, - did - ) - .fetch_optional(&state.db) - .await - { + let user = match state.user_repo.get_2fa_status_by_did(&did).await { Ok(Some(u)) => u, _ => { return ( @@ -2779,11 +2917,11 @@ pub async fn authorize_passkey_finish( } }; - let _ = db::delete_2fa_challenge_by_request_uri(&state.db, &form.request_uri).await; - match db::create_2fa_challenge(&state.db, &did, &form.request_uri).await { + let _ = state.oauth_repo.delete_2fa_challenge_by_request_uri(&passkey_finish_request_id).await; + match state.oauth_repo.create_2fa_challenge(&did, &passkey_finish_request_id).await { Ok(challenge) => { if let Err(e) = - enqueue_2fa_code(&state.db, user.id, &challenge.code, &pds_hostname).await + enqueue_2fa_code(state.user_repo.as_ref(), state.infra_repo.as_ref(), user.id, &challenge.code, &pds_hostname).await { tracing::warn!(did = %did, error = %e, "Failed to enqueue 2FA notification"); } diff --git a/crates/tranquil-pds/src/oauth/endpoints/delegation.rs b/crates/tranquil-pds/src/oauth/endpoints/delegation.rs index b3c7949..59a67a7 100644 --- a/crates/tranquil-pds/src/oauth/endpoints/delegation.rs +++ b/crates/tranquil-pds/src/oauth/endpoints/delegation.rs @@ -1,5 +1,4 @@ -use crate::delegation; -use crate::oauth::db; +use crate::delegation::DelegationActionType; use crate::state::{AppState, RateLimitKind}; use crate::types::PlainPassword; use crate::util::extract_client_ip; @@ -10,6 +9,7 @@ use axum::{ response::{IntoResponse, Response}, }; use serde::{Deserialize, Serialize}; +use tranquil_types::{Did, RequestId}; #[derive(Debug, Deserialize)] pub struct DelegationAuthSubmit { @@ -54,7 +54,12 @@ pub async fn delegation_auth( .into_response(); } - let request = match db::get_authorization_request(&state.db, &form.request_uri).await { + let request_id = RequestId::from(form.request_uri.clone()); + let request = match state + .oauth_repo + .get_authorization_request(&request_id) + .await + { Ok(Some(r)) => r, Ok(None) => { return Json(DelegationAuthResponse { @@ -76,7 +81,7 @@ pub async fn delegation_auth( } }; - let delegated_did = match form.delegated_did.as_ref().or(request.did.as_ref()) { + let delegated_did_str = match form.delegated_did.as_ref().or(request.did.as_ref()) { Some(did) => did.clone(), None => { return Json(DelegationAuthResponse { @@ -89,48 +94,68 @@ pub async fn delegation_auth( } }; - if db::set_request_did(&state.db, &form.request_uri, &delegated_did) + let delegated_did: Did = match delegated_did_str.parse() { + Ok(d) => d, + Err(_) => { + return Json(DelegationAuthResponse { + success: false, + needs_totp: None, + redirect_uri: None, + error: Some("Invalid delegated DID".to_string()), + }) + .into_response(); + } + }; + + let controller_did: Did = match form.controller_did.parse() { + Ok(d) => d, + Err(_) => { + return Json(DelegationAuthResponse { + success: false, + needs_totp: None, + redirect_uri: None, + error: Some("Invalid controller DID".to_string()), + }) + .into_response(); + } + }; + + if state + .oauth_repo + .set_request_did(&request_id, &delegated_did) .await .is_err() { tracing::warn!("Failed to set delegated DID on authorization request"); } - let grant = - match delegation::get_delegation(&state.db, &delegated_did, &form.controller_did).await { - Ok(Some(g)) => g, - Ok(None) => { - return Json(DelegationAuthResponse { - success: false, - needs_totp: None, - redirect_uri: None, - error: Some("No delegation grant found for this controller".to_string()), - }) - .into_response(); - } - Err(_) => { - return Json(DelegationAuthResponse { - success: false, - needs_totp: None, - redirect_uri: None, - error: Some("Server error".to_string()), - }) - .into_response(); - } - }; - - let controller = match sqlx::query!( - r#" - SELECT id, did, password_hash, deactivated_at, takedown_ref, - email_verified, discord_verified, telegram_verified, signal_verified - FROM users - WHERE did = $1 - "#, - form.controller_did - ) - .fetch_optional(&state.db) - .await + let grant = match state + .delegation_repo + .get_delegation(&delegated_did, &controller_did) + .await { + Ok(Some(g)) => g, + Ok(None) => { + return Json(DelegationAuthResponse { + success: false, + needs_totp: None, + redirect_uri: None, + error: Some("No delegation grant found for this controller".to_string()), + }) + .into_response(); + } + Err(_) => { + return Json(DelegationAuthResponse { + success: false, + needs_totp: None, + redirect_uri: None, + error: Some("Server error".to_string()), + }) + .into_response(); + } + }; + + let controller = match state.user_repo.get_auth_info_by_did(&controller_did).await { Ok(Some(u)) => u, Ok(None) => { return Json(DelegationAuthResponse { @@ -188,7 +213,9 @@ pub async fn delegation_auth( .into_response(); } - if db::set_controller_did(&state.db, &form.request_uri, &form.controller_did) + if state + .oauth_repo + .set_controller_did(&request_id, &controller_did) .await .is_err() { @@ -201,7 +228,7 @@ pub async fn delegation_auth( .into_response(); } - let has_totp = crate::api::server::has_totp_enabled(&state, &form.controller_did).await; + let has_totp = crate::api::server::has_totp_enabled(&state, &controller_did).await; if has_totp { return Json(DelegationAuthResponse { success: true, @@ -221,20 +248,21 @@ pub async fn delegation_auth( .and_then(|v| v.to_str().ok()) .map(|s| s.to_string()); - let _ = delegation::log_delegation_action( - &state.db, - &delegated_did, - &form.controller_did, - Some(&form.controller_did), - delegation::DelegationActionType::TokenIssued, - Some(serde_json::json!({ - "client_id": request.client_id, - "granted_scopes": grant.granted_scopes - })), - Some(&ip), - user_agent.as_deref(), - ) - .await; + let _ = state + .delegation_repo + .log_delegation_action( + &delegated_did, + &controller_did, + Some(&controller_did), + DelegationActionType::TokenIssued, + Some(serde_json::json!({ + "client_id": request.client_id, + "granted_scopes": grant.granted_scopes + })), + Some(&ip), + user_agent.as_deref(), + ) + .await; Json(DelegationAuthResponse { success: true, @@ -276,7 +304,12 @@ pub async fn delegation_totp_verify( .into_response(); } - let request = match db::get_authorization_request(&state.db, &form.request_uri).await { + let totp_request_id = RequestId::from(form.request_uri.clone()); + let request = match state + .oauth_repo + .get_authorization_request(&totp_request_id) + .await + { Ok(Some(r)) => r, Ok(None) => { return Json(DelegationAuthResponse { @@ -298,7 +331,7 @@ pub async fn delegation_totp_verify( } }; - let controller_did = match &request.controller_did { + let controller_did_str = match &request.controller_did { Some(did) => did.clone(), None => { return Json(DelegationAuthResponse { @@ -311,7 +344,20 @@ pub async fn delegation_totp_verify( } }; - let delegated_did = match &request.did { + let controller_did: Did = match controller_did_str.parse() { + Ok(d) => d, + Err(_) => { + return Json(DelegationAuthResponse { + success: false, + needs_totp: None, + redirect_uri: None, + error: Some("Invalid controller DID".to_string()), + }) + .into_response(); + } + }; + + let delegated_did_str = match &request.did { Some(did) => did.clone(), None => { return Json(DelegationAuthResponse { @@ -324,7 +370,24 @@ pub async fn delegation_totp_verify( } }; - let grant = match delegation::get_delegation(&state.db, &delegated_did, &controller_did).await { + let delegated_did: Did = match delegated_did_str.parse() { + Ok(d) => d, + Err(_) => { + return Json(DelegationAuthResponse { + success: false, + needs_totp: None, + redirect_uri: None, + error: Some("Invalid delegated DID".to_string()), + }) + .into_response(); + } + }; + + let grant = match state + .delegation_repo + .get_delegation(&delegated_did, &controller_did) + .await + { Ok(Some(g)) => g, _ => { return Json(DelegationAuthResponse { @@ -356,20 +419,21 @@ pub async fn delegation_totp_verify( .and_then(|v| v.to_str().ok()) .map(|s| s.to_string()); - let _ = delegation::log_delegation_action( - &state.db, - &delegated_did, - &controller_did, - Some(&controller_did), - delegation::DelegationActionType::TokenIssued, - Some(serde_json::json!({ - "client_id": request.client_id, - "granted_scopes": grant.granted_scopes - })), - Some(&ip), - user_agent.as_deref(), - ) - .await; + let _ = state + .delegation_repo + .log_delegation_action( + &delegated_did, + &controller_did, + Some(&controller_did), + DelegationActionType::TokenIssued, + Some(serde_json::json!({ + "client_id": request.client_id, + "granted_scopes": grant.granted_scopes + })), + Some(&ip), + user_agent.as_deref(), + ) + .await; Json(DelegationAuthResponse { success: true, diff --git a/crates/tranquil-pds/src/oauth/endpoints/par.rs b/crates/tranquil-pds/src/oauth/endpoints/par.rs index f3472f2..ace9459 100644 --- a/crates/tranquil-pds/src/oauth/endpoints/par.rs +++ b/crates/tranquil-pds/src/oauth/endpoints/par.rs @@ -1,9 +1,10 @@ use crate::oauth::{ AuthorizationRequestParameters, ClientAuth, ClientMetadataCache, OAuthError, RequestData, - RequestId, db, + RequestId, scopes::{ParsedScope, parse_scope}, }; use crate::state::{AppState, RateLimitKind}; +use tranquil_types::RequestId as RequestIdType; use axum::body::Bytes; use axum::{Json, extract::State, http::HeaderMap}; use chrono::{Duration, Utc}; @@ -131,11 +132,16 @@ pub async fn pushed_authorization_request( code: None, controller_did: None, }; - db::create_authorization_request(&state.db, &request_id.0, &request_data).await?; + let request_id_typed = RequestIdType::from(request_id.0.clone()); + state + .oauth_repo + .create_authorization_request(&request_id_typed, &request_data) + .await + .map_err(crate::oauth::db_err_to_oauth)?; tokio::spawn({ - let pool = state.db.clone(); + let oauth_repo = state.oauth_repo.clone(); async move { - if let Err(e) = db::delete_expired_authorization_requests(&pool).await { + if let Err(e) = oauth_repo.delete_expired_authorization_requests().await { tracing::warn!("Failed to cleanup expired authorization requests: {:?}", e); } } diff --git a/crates/tranquil-pds/src/oauth/endpoints/token/grants.rs b/crates/tranquil-pds/src/oauth/endpoints/token/grants.rs index 28ab02a..4e73c12 100644 --- a/crates/tranquil-pds/src/oauth/endpoints/token/grants.rs +++ b/crates/tranquil-pds/src/oauth/endpoints/token/grants.rs @@ -1,15 +1,17 @@ use super::helpers::{create_access_token_with_delegation, verify_pkce}; use super::types::{TokenGrant, TokenResponse, ValidatedTokenRequest}; use crate::config::AuthConfig; -use crate::delegation; +use crate::delegation::intersect_scopes; use crate::oauth::{ AuthFlowState, ClientAuth, ClientMetadataCache, DPoPVerifier, OAuthError, RefreshToken, TokenData, TokenId, - db::{self, RefreshTokenLookup}, + db::{lookup_refresh_token, enforce_token_limit_for_user}, scopes::expand_include_scopes, verify_client_auth, }; use crate::state::AppState; +use tranquil_db_traits::RefreshTokenLookup; +use tranquil_types::{AuthorizationCode, Did, RefreshToken as RefreshTokenType}; use axum::Json; use axum::http::HeaderMap; use chrono::{Duration, Utc}; @@ -41,8 +43,12 @@ pub async fn handle_authorization_code_grant( )); } }; - let auth_request = db::consume_authorization_request_by_code(&state.db, &code) - .await? + let auth_code = AuthorizationCode::from(code); + let auth_request = state + .oauth_repo + .consume_authorization_request_by_code(&auth_code) + .await + .map_err(crate::oauth::db_err_to_oauth)? .ok_or_else(|| OAuthError::InvalidGrant("Invalid or expired code".to_string()))?; let flow_state = AuthFlowState::from_request_data(&auth_request); @@ -100,7 +106,12 @@ pub async fn handle_authorization_code_grant( std::env::var("PDS_HOSTNAME").unwrap_or_else(|_| "localhost".to_string()); let token_endpoint = format!("https://{}/oauth/token", pds_hostname); let result = verifier.verify_proof(proof, "POST", &token_endpoint, None)?; - if !db::check_and_record_dpop_jti(&state.db, &result.jti).await? { + if !state + .oauth_repo + .check_and_record_dpop_jti(&result.jti) + .await + .map_err(crate::oauth::db_err_to_oauth)? + { return Err(OAuthError::InvalidDpopProof( "DPoP proof has already been used".to_string(), )); @@ -125,7 +136,15 @@ pub async fn handle_authorization_code_grant( let now = Utc::now(); let (raw_scope, controller_did) = if let Some(ref controller) = auth_request.controller_did { - let grant = delegation::get_delegation(&state.db, &did, controller) + let did_parsed: Did = did.parse().map_err(|_| { + OAuthError::InvalidRequest("Invalid DID format".to_string()) + })?; + let controller_parsed: Did = controller.parse().map_err(|_| { + OAuthError::InvalidRequest("Invalid controller DID format".to_string()) + })?; + let grant = state + .delegation_repo + .get_delegation(&did_parsed, &controller_parsed) .await .ok() .flatten(); @@ -135,7 +154,7 @@ pub async fn handle_authorization_code_grant( .scope .as_deref() .unwrap_or("atproto"); - let intersected = delegation::intersect_scopes(requested, &granted_scopes); + let intersected = intersect_scopes(requested, &granted_scopes); (Some(intersected), Some(controller.clone())) } else { (auth_request.parameters.scope.clone(), None) @@ -182,7 +201,11 @@ pub async fn handle_authorization_code_grant( scope: final_scope.clone(), controller_did: controller_did.clone(), }; - db::create_token(&state.db, &token_data).await?; + state + .oauth_repo + .create_token(&token_data) + .await + .map_err(crate::oauth::db_err_to_oauth)?; tracing::info!( did = %did, token_id = %token_id.0, @@ -190,11 +213,13 @@ pub async fn handle_authorization_code_grant( "Authorization code grant completed, token created" ); tokio::spawn({ - let pool = state.db.clone(); + let oauth_repo = state.oauth_repo.clone(); let did_clone = did.clone(); async move { - if let Err(e) = db::enforce_token_limit_for_user(&pool, &did_clone).await { - tracing::warn!("Failed to enforce token limit for user: {:?}", e); + if let Ok(did_typed) = did_clone.parse::() { + if let Err(e) = enforce_token_limit_for_user(oauth_repo.as_ref(), &did_typed).await { + tracing::warn!("Failed to enforce token limit for user: {:?}", e); + } } } }); @@ -236,7 +261,8 @@ pub async fn handle_refresh_token_grant( "Refresh token grant requested" ); - let lookup = db::lookup_refresh_token(&state.db, &refresh_token_str).await?; + let refresh_token_typed = RefreshTokenType::from(refresh_token_str.clone()); + let lookup = lookup_refresh_token(state.oauth_repo.as_ref(), &refresh_token_typed).await?; let token_state = lookup.state(); tracing::debug!(state = %token_state, "Refresh token state"); @@ -281,14 +307,22 @@ pub async fn handle_refresh_token_grant( refresh_token_prefix = %token_prefix, "Refresh token reuse detected, revoking token family" ); - db::delete_token_family(&state.db, original_token_id).await?; + state + .oauth_repo + .delete_token_family(original_token_id) + .await + .map_err(crate::oauth::db_err_to_oauth)?; return Err(OAuthError::InvalidGrant( "Refresh token reuse detected, token family revoked".to_string(), )); } RefreshTokenLookup::Expired { db_id } => { tracing::warn!(refresh_token_prefix = %token_prefix, "Refresh token has expired"); - db::delete_token_family(&state.db, db_id).await?; + state + .oauth_repo + .delete_token_family(db_id) + .await + .map_err(crate::oauth::db_err_to_oauth)?; return Err(OAuthError::InvalidGrant( "Refresh token has expired".to_string(), )); @@ -307,7 +341,12 @@ pub async fn handle_refresh_token_grant( std::env::var("PDS_HOSTNAME").unwrap_or_else(|_| "localhost".to_string()); let token_endpoint = format!("https://{}/oauth/token", pds_hostname); let result = verifier.verify_proof(proof, "POST", &token_endpoint, None)?; - if !db::check_and_record_dpop_jti(&state.db, &result.jti).await? { + if !state + .oauth_repo + .check_and_record_dpop_jti(&result.jti) + .await + .map_err(crate::oauth::db_err_to_oauth)? + { return Err(OAuthError::InvalidDpopProof( "DPoP proof has already been used".to_string(), )); @@ -334,7 +373,12 @@ pub async fn handle_refresh_token_grant( REFRESH_TOKEN_EXPIRY_DAYS_CONFIDENTIAL }; let new_expires_at = Utc::now() + Duration::days(refresh_expiry_days); - db::rotate_token(&state.db, db_id, &new_refresh_token.0, new_expires_at).await?; + let new_refresh_typed = RefreshTokenType::from(new_refresh_token.0.clone()); + state + .oauth_repo + .rotate_token(db_id, &new_refresh_typed, new_expires_at) + .await + .map_err(crate::oauth::db_err_to_oauth)?; tracing::info!( did = %token_data.did, new_expires_at = %new_expires_at, diff --git a/crates/tranquil-pds/src/oauth/endpoints/token/introspect.rs b/crates/tranquil-pds/src/oauth/endpoints/token/introspect.rs index 624d0dc..6f40473 100644 --- a/crates/tranquil-pds/src/oauth/endpoints/token/introspect.rs +++ b/crates/tranquil-pds/src/oauth/endpoints/token/introspect.rs @@ -1,11 +1,12 @@ use super::helpers::extract_token_claims; -use crate::oauth::{OAuthError, db}; +use crate::oauth::OAuthError; use crate::state::{AppState, RateLimitKind}; use axum::extract::State; use axum::http::{HeaderMap, StatusCode}; use axum::{Form, Json}; use chrono::Utc; use serde::{Deserialize, Serialize}; +use tranquil_types::{RefreshToken, TokenId}; #[derive(Debug, Deserialize)] pub struct RevokeRequest { @@ -28,10 +29,25 @@ pub async fn revoke_token( return Err(OAuthError::RateLimited); } if let Some(token) = &request.token { - if let Some((db_id, _)) = db::get_token_by_refresh_token(&state.db, token).await? { - db::delete_token_family(&state.db, db_id).await?; + let refresh_token = RefreshToken::from(token.clone()); + if let Some((db_id, _)) = state + .oauth_repo + .get_token_by_refresh_token(&refresh_token) + .await + .map_err(crate::oauth::db_err_to_oauth)? + { + state + .oauth_repo + .delete_token_family(db_id) + .await + .map_err(crate::oauth::db_err_to_oauth)?; } else { - db::delete_token(&state.db, token).await?; + let token_id = TokenId::from(token.clone()); + state + .oauth_repo + .delete_token(&token_id) + .await + .map_err(crate::oauth::db_err_to_oauth)?; } } Ok(StatusCode::OK) @@ -102,7 +118,8 @@ pub async fn introspect_token( Ok(info) => info, Err(_) => return Ok(Json(inactive_response)), }; - let token_data = match db::get_token_by_id(&state.db, &token_info.sid).await { + let token_id = TokenId::from(token_info.sid.clone()); + let token_data = match state.oauth_repo.get_token_by_id(&token_id).await { Ok(Some(data)) => data, _ => return Ok(Json(inactive_response)), }; diff --git a/crates/tranquil-pds/src/oauth/mod.rs b/crates/tranquil-pds/src/oauth/mod.rs index 1f5f162..d34cdbb 100644 --- a/crates/tranquil-pds/src/oauth/mod.rs +++ b/crates/tranquil-pds/src/oauth/mod.rs @@ -4,6 +4,11 @@ pub mod jwks; pub mod scopes; pub mod verify; +pub fn db_err_to_oauth(err: tranquil_db::DbError) -> OAuthError { + tracing::error!("Database error in OAuth flow: {}", err); + OAuthError::ServerError("An internal error occurred".to_string()) +} + pub use tranquil_oauth::{ AuthFlowState, AuthorizationRequestParameters, AuthorizationServerMetadata, AuthorizedClientData, ClientAuth, ClientMetadata, ClientMetadataCache, Code, DPoPClaims, diff --git a/crates/tranquil-pds/src/oauth/verify.rs b/crates/tranquil-pds/src/oauth/verify.rs index 8232d08..a8b41ee 100644 --- a/crates/tranquil-pds/src/oauth/verify.rs +++ b/crates/tranquil-pds/src/oauth/verify.rs @@ -8,10 +8,10 @@ use base64::{Engine as _, engine::general_purpose::URL_SAFE_NO_PAD}; use hmac::{Hmac, Mac}; use serde_json::json; use sha2::Sha256; -use sqlx::PgPool; use subtle::ConstantTimeEq; +use tranquil_db_traits::{OAuthRepository, UserRepository}; +use tranquil_types::TokenId; -use super::db; use super::scopes::ScopePermissions; use super::{DPoPVerifier, OAuthError}; use crate::config::AuthConfig; @@ -34,7 +34,7 @@ pub struct VerifyResult { } pub async fn verify_oauth_access_token( - pool: &PgPool, + oauth_repo: &dyn OAuthRepository, access_token: &str, dpop_proof: Option<&str>, http_method: &str, @@ -46,10 +46,13 @@ pub async fn verify_oauth_access_token( has_dpop_proof = dpop_proof.is_some(), "Verifying OAuth access token" ); - let token_data = db::get_token_by_id(pool, &token_info.token_id) - .await? + let token_id = TokenId::from(token_info.token_id.clone()); + let token_data = oauth_repo + .get_token_by_id(&token_id) + .await + .map_err(crate::oauth::db_err_to_oauth)? .ok_or_else(|| { - tracing::warn!(token_id = %token_info.token_id, "Token not found in database"); + tracing::warn!(token_id = %token_id, "Token not found in database"); OAuthError::InvalidToken("Token not found or revoked".to_string()) })?; let now = chrono::Utc::now(); @@ -73,7 +76,11 @@ pub async fn verify_oauth_access_token( tracing::warn!(error = ?e, http_method = %http_method, http_uri = %http_uri, "DPoP proof verification failed"); e })?; - if !db::check_and_record_dpop_jti(pool, &result.jti).await? { + if !oauth_repo + .check_and_record_dpop_jti(&result.jti) + .await + .map_err(crate::oauth::db_err_to_oauth)? + { return Err(OAuthError::InvalidDpopProof( "DPoP proof has already been used".to_string(), )); @@ -86,7 +93,7 @@ pub async fn verify_oauth_access_token( } Ok(VerifyResult { did: token_data.did, - token_id: token_info.token_id, + token_id: token_id.to_string(), client_id: token_data.client_id, scope: token_data.scope, }) @@ -271,7 +278,7 @@ impl FromRequestParts for OAuthUser { }); }; let dpop_proof = parts.headers.get("DPoP").and_then(|v| v.to_str().ok()); - if let Ok(result) = try_legacy_auth(&state.db, token).await { + if let Ok(result) = try_legacy_auth(state.user_repo.as_ref(), token).await { return Ok(OAuthUser { did: result.did, client_id: None, @@ -282,7 +289,14 @@ impl FromRequestParts for OAuthUser { } let http_method = parts.method.as_str(); let http_uri = crate::util::build_full_url(&parts.uri.to_string()); - match verify_oauth_access_token(&state.db, token, dpop_proof, http_method, &http_uri).await + match verify_oauth_access_token( + state.oauth_repo.as_ref(), + token, + dpop_proof, + http_method, + &http_uri, + ) + .await { Ok(result) => { let permissions = ScopePermissions::from_scope_string(result.scope.as_deref()); @@ -371,8 +385,8 @@ struct LegacyAuthResult { did: String, } -async fn try_legacy_auth(pool: &PgPool, token: &str) -> Result { - match crate::auth::validate_bearer_token(pool, token).await { +async fn try_legacy_auth(user_repo: &dyn UserRepository, token: &str) -> Result { + match crate::auth::validate_bearer_token(user_repo, token).await { Ok(user) if !user.is_oauth => Ok(LegacyAuthResult { did: user.did.to_string(), }), diff --git a/crates/tranquil-pds/src/scheduled.rs b/crates/tranquil-pds/src/scheduled.rs index a35d328..20b36f7 100644 --- a/crates/tranquil-pds/src/scheduled.rs +++ b/crates/tranquil-pds/src/scheduled.rs @@ -2,147 +2,26 @@ use cid::Cid; use ipld_core::ipld::Ipld; use jacquard_repo::commit::Commit; use jacquard_repo::storage::BlockStore; -use sqlx::PgPool; use std::str::FromStr; use std::sync::Arc; use std::time::Duration; use tokio::sync::watch; use tokio::time::interval; use tracing::{debug, error, info, warn}; +use tranquil_db_traits::{ + BackupRepository, BlobRepository, BrokenGenesisCommit, RepoRepository, UserRepository, +}; +use tranquil_types::{AtUri, CidLink, Did}; use crate::repo::PostgresBlockStore; use crate::storage::{BackupStorage, BlobStorage}; use crate::sync::car::encode_car_header; -async fn update_genesis_blocks_cids(db: &PgPool, blocks_cids: &[String], seq: i64) -> Result<(), sqlx::Error> { - sqlx::query!( - "UPDATE repo_seq SET blocks_cids = $1 WHERE seq = $2", - blocks_cids, - seq - ) - .execute(db) - .await?; - Ok(()) -} - -async fn update_repo_rev(db: &PgPool, rev: &str, user_id: uuid::Uuid) -> Result<(), sqlx::Error> { - sqlx::query!( - "UPDATE repos SET repo_rev = $1 WHERE user_id = $2", - rev, - user_id - ) - .execute(db) - .await?; - Ok(()) -} - -async fn insert_user_blocks(db: &PgPool, user_id: uuid::Uuid, block_cids: &[Vec]) -> Result<(), sqlx::Error> { - sqlx::query!( - r#" - INSERT INTO user_blocks (user_id, block_cid) - SELECT $1, block_cid FROM UNNEST($2::bytea[]) AS t(block_cid) - ON CONFLICT (user_id, block_cid) DO NOTHING - "#, - user_id, - block_cids - ) - .execute(db) - .await?; - Ok(()) -} - -async fn fetch_user_records(db: &PgPool, user_id: uuid::Uuid) -> Result, sqlx::Error> { - let rows = sqlx::query!( - "SELECT collection, rkey, record_cid FROM records WHERE repo_id = $1", - user_id - ) - .fetch_all(db) - .await?; - Ok(rows.into_iter().map(|r| (r.collection, r.rkey, r.record_cid)).collect()) -} - -async fn insert_record_blobs(db: &PgPool, user_id: uuid::Uuid, record_uris: &[String], blob_cids: &[String]) -> Result<(), sqlx::Error> { - sqlx::query!( - r#" - INSERT INTO record_blobs (repo_id, record_uri, blob_cid) - SELECT $1, record_uri, blob_cid - FROM UNNEST($2::text[], $3::text[]) AS t(record_uri, blob_cid) - ON CONFLICT (repo_id, record_uri, blob_cid) DO NOTHING - "#, - user_id, - record_uris, - blob_cids - ) - .execute(db) - .await?; - Ok(()) -} - -async fn delete_backup_record(db: &PgPool, id: uuid::Uuid) -> Result<(), sqlx::Error> { - sqlx::query!("DELETE FROM account_backups WHERE id = $1", id) - .execute(db) - .await?; - Ok(()) -} - -async fn fetch_old_backups( - db: &PgPool, - user_id: uuid::Uuid, - retention_count: i64, -) -> Result, sqlx::Error> { - let rows = sqlx::query!( - r#" - SELECT id, storage_key - FROM account_backups - WHERE user_id = $1 - ORDER BY created_at DESC - OFFSET $2 - "#, - user_id, - retention_count - ) - .fetch_all(db) - .await?; - Ok(rows.into_iter().map(|r| (r.id, r.storage_key)).collect()) -} - -async fn insert_backup_record( - db: &PgPool, - user_id: uuid::Uuid, - storage_key: &str, - repo_root_cid: &str, - repo_rev: &str, - block_count: i32, - size_bytes: i64, -) -> Result<(), sqlx::Error> { - sqlx::query!( - r#" - INSERT INTO account_backups (user_id, storage_key, repo_root_cid, repo_rev, block_count, size_bytes) - VALUES ($1, $2, $3, $4, $5, $6) - "#, - user_id, - storage_key, - repo_root_cid, - repo_rev, - block_count, - size_bytes - ) - .execute(db) - .await?; - Ok(()) -} - -struct GenesisCommitRow { - seq: i64, - did: String, - commit_cid: Option, -} - async fn process_genesis_commit( - db: &PgPool, + repo_repo: &dyn RepoRepository, block_store: &PostgresBlockStore, - row: GenesisCommitRow, -) -> Result<(String, i64), (i64, &'static str)> { + row: BrokenGenesisCommit, +) -> Result<(Did, i64), (i64, &'static str)> { let commit_cid_str = row.commit_cid.ok_or((row.seq, "missing commit_cid"))?; let commit_cid = Cid::from_str(&commit_cid_str).map_err(|_| (row.seq, "invalid CID"))?; let block = block_store @@ -152,28 +31,24 @@ async fn process_genesis_commit( .ok_or((row.seq, "block not found"))?; let commit = Commit::from_cbor(&block).map_err(|_| (row.seq, "failed to parse commit"))?; let blocks_cids = vec![commit.data.to_string(), commit_cid.to_string()]; - update_genesis_blocks_cids(db, &blocks_cids, row.seq) + repo_repo + .update_seq_blocks_cids(row.seq, &blocks_cids) .await .map_err(|_| (row.seq, "failed to update"))?; Ok((row.did, row.seq)) } -pub async fn backfill_genesis_commit_blocks(db: &PgPool, block_store: PostgresBlockStore) { - let broken_genesis_commits = match sqlx::query!( - r#" - SELECT seq, did, commit_cid - FROM repo_seq - WHERE event_type = 'commit' - AND prev_cid IS NULL - AND (blocks_cids IS NULL OR array_length(blocks_cids, 1) IS NULL OR array_length(blocks_cids, 1) = 0) - "# - ) - .fetch_all(db) - .await - { +pub async fn backfill_genesis_commit_blocks( + repo_repo: Arc, + block_store: PostgresBlockStore, +) { + let broken_genesis_commits = match repo_repo.get_broken_genesis_commits().await { Ok(rows) => rows, Err(e) => { - error!("Failed to query repo_seq for genesis commit backfill: {}", e); + error!( + "Failed to query repo_seq for genesis commit backfill: {:?}", + e + ); return; } }; @@ -189,15 +64,9 @@ pub async fn backfill_genesis_commit_blocks(db: &PgPool, block_store: PostgresBl ); let results = futures::future::join_all(broken_genesis_commits.into_iter().map(|row| { - process_genesis_commit( - db, - &block_store, - GenesisCommitRow { - seq: row.seq, - did: row.did, - commit_cid: row.commit_cid, - }, - ) + let repo_repo = repo_repo.clone(); + let block_store = block_store.clone(); + async move { process_genesis_commit(repo_repo.as_ref(), &block_store, row).await } })) .await; @@ -219,7 +88,7 @@ pub async fn backfill_genesis_commit_blocks(db: &PgPool, block_store: PostgresBl } async fn process_repo_rev( - db: &PgPool, + repo_repo: &dyn RepoRepository, block_store: &PostgresBlockStore, user_id: uuid::Uuid, repo_root_cid: String, @@ -233,24 +102,24 @@ async fn process_repo_rev( .ok_or(user_id)?; let commit = Commit::from_cbor(&block).map_err(|_| user_id)?; let rev = commit.rev().to_string(); - update_repo_rev(db, &rev, user_id) + repo_repo + .update_repo_rev(user_id, &rev) .await .map_err(|_| user_id)?; Ok(user_id) } -pub async fn backfill_repo_rev(db: &PgPool, block_store: PostgresBlockStore) { - let repos_missing_rev = - match sqlx::query!("SELECT user_id, repo_root_cid FROM repos WHERE repo_rev IS NULL") - .fetch_all(db) - .await - { - Ok(rows) => rows, - Err(e) => { - error!("Failed to query repos for backfill: {}", e); - return; - } - }; +pub async fn backfill_repo_rev( + repo_repo: Arc, + block_store: PostgresBlockStore, +) { + let repos_missing_rev = match repo_repo.get_repos_without_rev().await { + Ok(rows) => rows, + Err(e) => { + error!("Failed to query repos for backfill: {:?}", e); + return; + } + }; if repos_missing_rev.is_empty() { debug!("No repos need repo_rev backfill"); @@ -263,28 +132,32 @@ pub async fn backfill_repo_rev(db: &PgPool, block_store: PostgresBlockStore) { ); let results = futures::future::join_all(repos_missing_rev.into_iter().map(|repo| { - process_repo_rev(db, &block_store, repo.user_id, repo.repo_root_cid) + let repo_repo = repo_repo.clone(); + let block_store = block_store.clone(); + async move { + process_repo_rev(repo_repo.as_ref(), &block_store, repo.user_id, repo.repo_root_cid.to_string()) + .await + } })) .await; - let (success, failed) = results - .iter() - .fold((0, 0), |(s, f), r| match r { - Ok(_) => (s + 1, f), - Err(user_id) => { - warn!(user_id = %user_id, "Failed to update repo_rev"); - (s, f + 1) - } - }); + let (success, failed) = results.iter().fold((0, 0), |(s, f), r| match r { + Ok(_) => (s + 1, f), + Err(user_id) => { + warn!(user_id = %user_id, "Failed to update repo_rev"); + (s, f + 1) + } + }); info!(success, failed, "Completed repo_rev backfill"); } async fn process_user_blocks( - db: &PgPool, + repo_repo: &dyn RepoRepository, block_store: &PostgresBlockStore, user_id: uuid::Uuid, repo_root_cid: String, + repo_rev: Option, ) -> Result<(uuid::Uuid, usize), uuid::Uuid> { let root_cid = Cid::from_str(&repo_root_cid).map_err(|_| user_id)?; let block_cids = collect_current_repo_blocks(block_store, &root_cid) @@ -294,27 +167,22 @@ async fn process_user_blocks( return Err(user_id); } let count = block_cids.len(); - insert_user_blocks(db, user_id, &block_cids) + let rev = repo_rev.unwrap_or_else(|| "0".to_string()); + repo_repo + .insert_user_blocks(user_id, &block_cids, &rev) .await .map_err(|_| user_id)?; Ok((user_id, count)) } -pub async fn backfill_user_blocks(db: &PgPool, block_store: PostgresBlockStore) { - let users_without_blocks = match sqlx::query!( - r#" - SELECT u.id as user_id, r.repo_root_cid - FROM users u - JOIN repos r ON r.user_id = u.id - WHERE NOT EXISTS (SELECT 1 FROM user_blocks ub WHERE ub.user_id = u.id) - "# - ) - .fetch_all(db) - .await - { +pub async fn backfill_user_blocks( + repo_repo: Arc, + block_store: PostgresBlockStore, +) { + let users_without_blocks = match repo_repo.get_users_without_blocks().await { Ok(rows) => rows, Err(e) => { - error!("Failed to query users for user_blocks backfill: {}", e); + error!("Failed to query users for user_blocks backfill: {:?}", e); return; } }; @@ -330,7 +198,18 @@ pub async fn backfill_user_blocks(db: &PgPool, block_store: PostgresBlockStore) ); let results = futures::future::join_all(users_without_blocks.into_iter().map(|user| { - process_user_blocks(db, &block_store, user.user_id, user.repo_root_cid) + let repo_repo = repo_repo.clone(); + let block_store = block_store.clone(); + async move { + process_user_blocks( + repo_repo.as_ref(), + &block_store, + user.user_id, + user.repo_root_cid.to_string(), + user.repo_rev, + ) + .await + } })) .await; @@ -401,22 +280,23 @@ pub async fn collect_current_repo_blocks( } async fn process_record_blobs( - db: &PgPool, + repo_repo: &dyn RepoRepository, block_store: &PostgresBlockStore, user_id: uuid::Uuid, - did: String, -) -> Result<(uuid::Uuid, String, usize), (uuid::Uuid, &'static str)> { - let records = fetch_user_records(db, user_id) + did: Did, +) -> Result<(uuid::Uuid, Did, usize), (uuid::Uuid, &'static str)> { + let records = repo_repo + .get_all_records(user_id) .await .map_err(|_| (user_id, "failed to fetch records"))?; - let mut batch_record_uris: Vec = Vec::new(); - let mut batch_blob_cids: Vec = Vec::new(); + let mut batch_record_uris: Vec = Vec::new(); + let mut batch_blob_cids: Vec = Vec::new(); - futures::future::join_all(records.into_iter().map(|(collection, rkey, record_cid)| { + futures::future::join_all(records.into_iter().map(|record| { let did = did.clone(); async move { - let cid = Cid::from_str(&record_cid).ok()?; + let cid = Cid::from_str(&record.record_cid).ok()?; let block_bytes = block_store.get(&cid).await.ok()??; let record_ipld: Ipld = serde_ipld_dagcbor::from_slice(&block_bytes).ok()?; let blob_refs = crate::sync::import::find_blob_refs_ipld(&record_ipld, 0); @@ -424,8 +304,9 @@ async fn process_record_blobs( blob_refs .into_iter() .map(|blob_ref| { - let record_uri = format!("at://{}/{}/{}", did, collection, rkey); - (record_uri, blob_ref.cid) + let record_uri = + AtUri::from_parts(did.as_str(), record.collection.as_str(), record.rkey.as_str()); + (record_uri, CidLink::new_unchecked(blob_ref.cid)) }) .collect::>(), ) @@ -442,29 +323,22 @@ async fn process_record_blobs( let blob_refs_found = batch_record_uris.len(); if !batch_record_uris.is_empty() { - insert_record_blobs(db, user_id, &batch_record_uris, &batch_blob_cids) + repo_repo + .insert_record_blobs(user_id, &batch_record_uris, &batch_blob_cids) .await .map_err(|_| (user_id, "failed to insert"))?; } Ok((user_id, did, blob_refs_found)) } -pub async fn backfill_record_blobs(db: &PgPool, block_store: PostgresBlockStore) { - let users_needing_backfill = match sqlx::query!( - r#" - SELECT DISTINCT u.id as user_id, u.did - FROM users u - JOIN records r ON r.repo_id = u.id - WHERE NOT EXISTS (SELECT 1 FROM record_blobs rb WHERE rb.repo_id = u.id) - LIMIT 100 - "# - ) - .fetch_all(db) - .await - { +pub async fn backfill_record_blobs( + repo_repo: Arc, + block_store: PostgresBlockStore, +) { + let users_needing_backfill = match repo_repo.get_users_needing_record_blobs_backfill(100).await { Ok(rows) => rows, Err(e) => { - error!("Failed to query users for record_blobs backfill: {}", e); + error!("Failed to query users for record_blobs backfill: {:?}", e); return; } }; @@ -480,7 +354,11 @@ pub async fn backfill_record_blobs(db: &PgPool, block_store: PostgresBlockStore) ); let results = futures::future::join_all(users_needing_backfill.into_iter().map(|user| { - process_record_blobs(db, &block_store, user.user_id, user.did) + let repo_repo = repo_repo.clone(); + let block_store = block_store.clone(); + async move { + process_record_blobs(repo_repo.as_ref(), &block_store, user.user_id, user.did).await + } })) .await; @@ -501,7 +379,8 @@ pub async fn backfill_record_blobs(db: &PgPool, block_store: PostgresBlockStore) } pub async fn start_scheduled_tasks( - db: PgPool, + user_repo: Arc, + blob_repo: Arc, blob_store: Arc, mut shutdown_rx: watch::Receiver, ) { @@ -529,7 +408,11 @@ pub async fn start_scheduled_tasks( } } _ = ticker.tick() => { - if let Err(e) = process_scheduled_deletions(&db, blob_store.as_ref()).await { + if let Err(e) = process_scheduled_deletions( + user_repo.as_ref(), + blob_repo.as_ref(), + blob_store.as_ref(), + ).await { error!("Error processing scheduled deletions: {}", e); } } @@ -538,22 +421,14 @@ pub async fn start_scheduled_tasks( } async fn process_scheduled_deletions( - db: &PgPool, + user_repo: &dyn UserRepository, + blob_repo: &dyn BlobRepository, blob_store: &dyn BlobStorage, ) -> Result<(), String> { - let accounts_to_delete = sqlx::query!( - r#" - SELECT did, handle - FROM users - WHERE delete_after IS NOT NULL - AND delete_after < NOW() - AND deactivated_at IS NOT NULL - LIMIT 100 - "# - ) - .fetch_all(db) - .await - .map_err(|e| format!("DB error fetching accounts to delete: {}", e))?; + let accounts_to_delete = user_repo + .get_accounts_scheduled_for_deletion(100) + .await + .map_err(|e| format!("DB error fetching accounts to delete: {:?}", e))?; if accounts_to_delete.is_empty() { debug!("No accounts scheduled for deletion"); @@ -566,7 +441,8 @@ async fn process_scheduled_deletions( ); futures::future::join_all(accounts_to_delete.into_iter().map(|account| async move { - let result = delete_account_data(db, blob_store, &account.did, &account.handle).await; + let result = + delete_account_data(user_repo, blob_repo, blob_store, account.id, &account.did).await; (account.did, account.handle, result) })) .await @@ -580,23 +456,16 @@ async fn process_scheduled_deletions( } async fn delete_account_data( - db: &PgPool, + user_repo: &dyn UserRepository, + blob_repo: &dyn BlobRepository, blob_store: &dyn BlobStorage, - did: &str, - _handle: &str, + user_id: uuid::Uuid, + did: &Did, ) -> Result<(), String> { - let user_id: uuid::Uuid = sqlx::query_scalar!("SELECT id FROM users WHERE did = $1", did) - .fetch_one(db) + let blob_storage_keys = blob_repo + .get_blob_storage_keys_by_user(user_id) .await - .map_err(|e| format!("DB error fetching user: {}", e))?; - - let blob_storage_keys: Vec = sqlx::query_scalar!( - r#"SELECT storage_key as "storage_key!" FROM blobs WHERE created_by_user = $1"#, - user_id - ) - .fetch_all(db) - .await - .map_err(|e| format!("DB error fetching blob keys: {}", e))?; + .map_err(|e| format!("DB error fetching blob keys: {:?}", e))?; futures::future::join_all(blob_storage_keys.iter().map(|storage_key| async move { (storage_key, blob_store.delete(storage_key).await) @@ -608,50 +477,10 @@ async fn delete_account_data( warn!(storage_key = %key, error = %e, "Failed to delete blob from storage (continuing anyway)"); }); - let mut tx = db - .begin() - .await - .map_err(|e| format!("Failed to begin transaction: {}", e))?; - - sqlx::query!("DELETE FROM blobs WHERE created_by_user = $1", user_id) - .execute(&mut *tx) - .await - .map_err(|e| format!("Failed to delete blobs: {}", e))?; - - sqlx::query!("DELETE FROM users WHERE id = $1", user_id) - .execute(&mut *tx) - .await - .map_err(|e| format!("Failed to delete user: {}", e))?; - - let account_seq = sqlx::query_scalar!( - r#" - INSERT INTO repo_seq (did, event_type, active, status) - VALUES ($1, 'account', false, 'deleted') - RETURNING seq - "#, - did - ) - .fetch_one(&mut *tx) - .await - .map_err(|e| format!("Failed to sequence account deletion: {}", e))?; - - sqlx::query!( - "DELETE FROM repo_seq WHERE did = $1 AND seq != $2", - did, - account_seq - ) - .execute(&mut *tx) - .await - .map_err(|e| format!("Failed to cleanup sequences: {}", e))?; - - tx.commit() + let _account_seq = user_repo + .delete_account_with_firehose(user_id, did) .await - .map_err(|e| format!("Failed to commit transaction: {}", e))?; - - sqlx::query(&format!("NOTIFY repo_updates, '{}'", account_seq)) - .execute(db) - .await - .map_err(|e| format!("Failed to notify: {}", e))?; + .map_err(|e| format!("Failed to delete account: {:?}", e))?; info!( did = %did, @@ -663,7 +492,8 @@ async fn delete_account_data( } pub async fn start_backup_tasks( - db: PgPool, + repo_repo: Arc, + backup_repo: Arc, block_store: PostgresBlockStore, backup_storage: Arc, mut shutdown_rx: watch::Receiver, @@ -688,7 +518,12 @@ pub async fn start_backup_tasks( } } _ = ticker.tick() => { - if let Err(e) = process_scheduled_backups(&db, &block_store, &backup_storage).await { + if let Err(e) = process_scheduled_backups( + repo_repo.as_ref(), + backup_repo.as_ref(), + &block_store, + &backup_storage, + ).await { error!("Error processing scheduled backups: {}", e); } } @@ -711,7 +546,8 @@ enum BackupOutcome { } async fn process_single_backup( - db: &PgPool, + repo_repo: &dyn RepoRepository, + backup_repo: &dyn BackupRepository, block_store: &PostgresBlockStore, backup_storage: &BackupStorage, user_id: uuid::Uuid, @@ -729,7 +565,7 @@ async fn process_single_backup( Err(_) => return BackupOutcome::Skipped(did, "invalid repo_root_cid"), }; - let car_bytes = match generate_full_backup(db, block_store, user_id, &head_cid).await { + let car_bytes = match generate_full_backup(repo_repo, block_store, user_id, &head_cid).await { Ok(bytes) => bytes, Err(e) => return BackupOutcome::Failed(did, format!("CAR generation: {}", e)), }; @@ -742,16 +578,16 @@ async fn process_single_backup( Err(e) => return BackupOutcome::Failed(did, format!("S3 upload: {}", e)), }; - if let Err(e) = insert_backup_record( - db, - user_id, - &storage_key, - &repo_root_cid, - &repo_rev, - block_count, - size_bytes, - ) - .await + if let Err(e) = backup_repo + .insert_backup( + user_id, + &storage_key, + &repo_root_cid, + &repo_rev, + block_count, + size_bytes, + ) + .await { if let Err(rollback_err) = backup_storage.delete_backup(&storage_key).await { error!( @@ -761,7 +597,7 @@ async fn process_single_backup( "Failed to rollback orphaned backup from S3" ); } - return BackupOutcome::Failed(did, format!("DB insert: {}", e)); + return BackupOutcome::Failed(did, format!("DB insert: {:?}", e)); } BackupOutcome::Success(BackupResult { @@ -774,35 +610,18 @@ async fn process_single_backup( } async fn process_scheduled_backups( - db: &PgPool, + repo_repo: &dyn RepoRepository, + backup_repo: &dyn BackupRepository, block_store: &PostgresBlockStore, backup_storage: &BackupStorage, ) -> Result<(), String> { let backup_interval_secs = BackupStorage::interval_secs() as i64; let retention_count = BackupStorage::retention_count(); - let users_needing_backup = sqlx::query!( - r#" - SELECT u.id as user_id, u.did, r.repo_root_cid, r.repo_rev - FROM users u - JOIN repos r ON r.user_id = u.id - WHERE u.backup_enabled = true - AND u.deactivated_at IS NULL - AND ( - NOT EXISTS ( - SELECT 1 FROM account_backups ab WHERE ab.user_id = u.id - ) - OR ( - SELECT MAX(ab.created_at) FROM account_backups ab WHERE ab.user_id = u.id - ) < NOW() - make_interval(secs => $1) - ) - LIMIT 50 - "#, - backup_interval_secs as f64 - ) - .fetch_all(db) - .await - .map_err(|e| format!("DB error fetching users for backup: {}", e))?; + let users_needing_backup = backup_repo + .get_users_needing_backup(backup_interval_secs, 50) + .await + .map_err(|e| format!("DB error fetching users for backup: {:?}", e))?; if users_needing_backup.is_empty() { debug!("No accounts need backup"); @@ -816,12 +635,13 @@ async fn process_scheduled_backups( let results = futures::future::join_all(users_needing_backup.into_iter().map(|user| { process_single_backup( - db, + repo_repo, + backup_repo, block_store, backup_storage, - user.user_id, - user.did, - user.repo_root_cid, + user.id, + user.did.to_string(), + user.repo_root_cid.to_string(), user.repo_rev, ) })) @@ -838,7 +658,8 @@ async fn process_scheduled_backups( "Created backup" ); if let Err(e) = - cleanup_old_backups(db, backup_storage, result.user_id, retention_count).await + cleanup_old_backups(backup_repo, backup_storage, result.user_id, retention_count) + .await { warn!(did = %result.did, error = %e, "Failed to cleanup old backups"); } @@ -905,21 +726,19 @@ fn encode_car_block(cid: &Cid, block: &[u8]) -> Vec { } pub async fn generate_repo_car_from_user_blocks( - db: &PgPool, + repo_repo: &dyn tranquil_db_traits::RepoRepository, block_store: &PostgresBlockStore, user_id: uuid::Uuid, _head_cid: &Cid, ) -> Result, String> { use std::str::FromStr; - let repo_root_cid_str: String = sqlx::query_scalar!( - "SELECT repo_root_cid FROM repos WHERE user_id = $1", - user_id - ) - .fetch_optional(db) - .await - .map_err(|e| format!("Failed to fetch repo: {}", e))? - .ok_or_else(|| "Repository not found".to_string())?; + let repo_root_cid_str: String = repo_repo + .get_repo_root_cid_by_user_id(user_id) + .await + .map_err(|e| format!("Failed to fetch repo: {:?}", e))? + .ok_or_else(|| "Repository not found".to_string())? + .to_string(); let actual_head_cid = Cid::from_str(&repo_root_cid_str).map_err(|e| format!("Invalid repo_root_cid: {}", e))?; @@ -928,12 +747,12 @@ pub async fn generate_repo_car_from_user_blocks( } pub async fn generate_full_backup( - db: &PgPool, + repo_repo: &dyn tranquil_db_traits::RepoRepository, block_store: &PostgresBlockStore, user_id: uuid::Uuid, head_cid: &Cid, ) -> Result, String> { - generate_repo_car_from_user_blocks(db, block_store, user_id, head_cid).await + generate_repo_car_from_user_blocks(repo_repo, block_store, user_id, head_cid).await } pub fn count_car_blocks(car_bytes: &[u8]) -> i32 { @@ -977,24 +796,25 @@ fn read_varint(data: &[u8]) -> Option<(u64, usize)> { } async fn cleanup_old_backups( - db: &PgPool, + backup_repo: &dyn BackupRepository, backup_storage: &BackupStorage, user_id: uuid::Uuid, retention_count: u32, ) -> Result<(), String> { - let old_backups = fetch_old_backups(db, user_id, retention_count as i64) + let old_backups = backup_repo + .get_old_backups(user_id, retention_count as i64) .await - .map_err(|e| format!("DB error fetching old backups: {}", e))?; + .map_err(|e| format!("DB error fetching old backups: {:?}", e))?; - let results = futures::future::join_all(old_backups.into_iter().map(|(id, storage_key)| async move { - match backup_storage.delete_backup(&storage_key).await { - Ok(()) => match delete_backup_record(db, id).await { + let results = futures::future::join_all(old_backups.into_iter().map(|backup| async move { + match backup_storage.delete_backup(&backup.storage_key).await { + Ok(()) => match backup_repo.delete_backup(backup.id).await { Ok(()) => Ok(()), - Err(e) => Err(format!("DB delete failed for {}: {}", storage_key, e)), + Err(e) => Err(format!("DB delete failed for {}: {:?}", backup.storage_key, e)), }, Err(e) => { warn!( - storage_key = %storage_key, + storage_key = %backup.storage_key, error = %e, "Failed to delete old backup from storage, skipping DB cleanup to avoid orphan" ); diff --git a/crates/tranquil-pds/src/state.rs b/crates/tranquil-pds/src/state.rs index 103291b..dc17cd6 100644 --- a/crates/tranquil-pds/src/state.rs +++ b/crates/tranquil-pds/src/state.rs @@ -10,10 +10,25 @@ use sqlx::PgPool; use std::error::Error; use std::sync::Arc; use tokio::sync::broadcast; +use tranquil_db::{ + BacklinkRepository, BackupRepository, BlobRepository, DelegationRepository, InfraRepository, + OAuthRepository, PostgresRepositories, RepoEventNotifier, RepoRepository, SessionRepository, + UserRepository, +}; #[derive(Clone)] pub struct AppState { - pub db: PgPool, + pub repos: Arc, + pub user_repo: Arc, + pub oauth_repo: Arc, + pub session_repo: Arc, + pub delegation_repo: Arc, + pub repo_repo: Arc, + pub blob_repo: Arc, + pub infra_repo: Arc, + pub backup_repo: Arc, + pub backlink_repo: Arc, + pub event_notifier: Arc, pub block_store: PostgresBlockStore, pub blob_store: Arc, pub backup_storage: Option>, @@ -133,7 +148,8 @@ impl AppState { pub async fn from_db(db: PgPool) -> Self { AuthConfig::init(); - let block_store = PostgresBlockStore::new(db.clone()); + let repos = Arc::new(PostgresRepositories::new(db.clone())); + let block_store = PostgresBlockStore::new(db); let blob_store = S3BlobStorage::new().await; let backup_storage = BackupStorage::new().await.map(Arc::new); @@ -149,7 +165,17 @@ impl AppState { let did_resolver = Arc::new(DidResolver::new()); Self { - db, + user_repo: repos.user.clone(), + oauth_repo: repos.oauth.clone(), + session_repo: repos.session.clone(), + delegation_repo: repos.delegation.clone(), + repo_repo: repos.repo.clone(), + blob_repo: repos.blob.clone(), + infra_repo: repos.infra.clone(), + backup_repo: repos.backup.clone(), + backlink_repo: repos.backlink.clone(), + event_notifier: repos.event_notifier.clone(), + repos, block_store, blob_store: Arc::new(blob_store), backup_storage, diff --git a/crates/tranquil-pds/src/sync/blob.rs b/crates/tranquil-pds/src/sync/blob.rs index 2bb6e0d..52ff131 100644 --- a/crates/tranquil-pds/src/sync/blob.rs +++ b/crates/tranquil-pds/src/sync/blob.rs @@ -11,6 +11,7 @@ use axum::{ }; use serde::{Deserialize, Serialize}; use tracing::error; +use tranquil_types::{CidLink, Did}; #[derive(Deserialize)] pub struct GetBlobParams { @@ -22,36 +23,36 @@ pub async fn get_blob( State(state): State, Query(params): Query, ) -> Response { - let did = params.did.trim(); - let cid = params.cid.trim(); - if did.is_empty() { + let did_str = params.did.trim(); + let cid_str = params.cid.trim(); + if did_str.is_empty() { return ApiError::InvalidRequest("did is required".into()).into_response(); } - if cid.is_empty() { + if cid_str.is_empty() { return ApiError::InvalidRequest("cid is required".into()).into_response(); } + let did: Did = match did_str.parse() { + Ok(d) => d, + Err(_) => return ApiError::InvalidRequest("invalid did".into()).into_response(), + }; + let cid: CidLink = match cid_str.parse() { + Ok(c) => c, + Err(_) => return ApiError::InvalidRequest("invalid cid".into()).into_response(), + }; - let _account = match assert_repo_availability(&state.db, did, false).await { + let _account = match assert_repo_availability(state.repo_repo.as_ref(), &did, false).await { Ok(a) => a, Err(e) => return e.into_response(), }; - let blob_result = sqlx::query!( - "SELECT storage_key, mime_type, size_bytes FROM blobs WHERE cid = $1", - cid - ) - .fetch_optional(&state.db) - .await; + let blob_result = state.blob_repo.get_blob_metadata(&cid).await; match blob_result { - Ok(Some(row)) => { - let storage_key = &row.storage_key; - let mime_type = &row.mime_type; - let size_bytes = row.size_bytes; - match state.blob_store.get(storage_key).await { + Ok(Some(metadata)) => { + match state.blob_store.get(&metadata.storage_key).await { Ok(data) => Response::builder() .status(StatusCode::OK) - .header(header::CONTENT_TYPE, mime_type) - .header(header::CONTENT_LENGTH, size_bytes.to_string()) + .header(header::CONTENT_TYPE, &metadata.mime_type) + .header(header::CONTENT_LENGTH, metadata.size_bytes.to_string()) .header("x-content-type-options", "nosniff") .header("content-security-policy", "default-src 'none'; sandbox") .body(Body::from(data)) @@ -65,7 +66,7 @@ pub async fn get_blob( Ok(None) => ApiError::BlobNotFound(Some("Blob not found".into())).into_response(), Err(e) => { error!("DB error in get_blob: {:?}", e); - ApiError::InternalError(Some("Database error".into())).into_response() + ApiError::InternalError(Some(format!("Database error: {}", e))).into_response() } } } @@ -89,12 +90,16 @@ pub async fn list_blobs( State(state): State, Query(params): Query, ) -> Response { - let did = params.did.trim(); - if did.is_empty() { + let did_str = params.did.trim(); + if did_str.is_empty() { return ApiError::InvalidRequest("did is required".into()).into_response(); } + let did: Did = match did_str.parse() { + Ok(d) => d, + Err(_) => return ApiError::InvalidRequest("invalid did".into()).into_response(), + }; - let account = match assert_repo_availability(&state.db, did, false).await { + let account = match assert_repo_availability(state.repo_repo.as_ref(), &did, false).await { Ok(a) => a, Err(e) => return e.into_response(), }; @@ -103,40 +108,26 @@ pub async fn list_blobs( let cursor_cid = params.cursor.as_deref().unwrap_or(""); let user_id = account.user_id; - let cids_result: Result, sqlx::Error> = if let Some(since) = ¶ms.since { - sqlx::query_scalar!( - r#" - SELECT DISTINCT unnest(blobs) as "cid!" - FROM repo_seq - WHERE did = $1 AND rev > $2 AND blobs IS NOT NULL - "#, - did, - since - ) - .fetch_all(&state.db) - .await - .map(|mut cids| { - cids.sort(); - cids.into_iter() - .filter(|c| c.as_str() > cursor_cid) - .take((limit + 1) as usize) - .collect() - }) + let cids_result: Result, _> = if let Some(since) = ¶ms.since { + state + .blob_repo + .list_blobs_since_rev(&did, since) + .await + .map(|cids| { + let mut cid_strs: Vec = cids.into_iter().map(|c| c.to_string()).collect(); + cid_strs.sort(); + cid_strs + .into_iter() + .filter(|c| c.as_str() > cursor_cid) + .take((limit + 1) as usize) + .collect() + }) } else { - sqlx::query!( - r#" - SELECT cid FROM blobs - WHERE created_by_user = $1 AND cid > $2 - ORDER BY cid ASC - LIMIT $3 - "#, - user_id, - cursor_cid, - limit + 1 - ) - .fetch_all(&state.db) - .await - .map(|rows| rows.into_iter().map(|r| r.cid).collect()) + state + .blob_repo + .list_blobs_by_user(user_id, Some(cursor_cid), limit + 1) + .await + .map(|cids| cids.into_iter().map(|c| c.to_string()).collect()) }; match cids_result { Ok(cids) => { @@ -154,7 +145,7 @@ pub async fn list_blobs( } Err(e) => { error!("DB error in list_blobs: {:?}", e); - ApiError::InternalError(Some("Database error".into())).into_response() + ApiError::InternalError(Some(format!("Database error: {}", e))).into_response() } } } diff --git a/crates/tranquil-pds/src/sync/commit.rs b/crates/tranquil-pds/src/sync/commit.rs index a1f5f89..8dc4a6d 100644 --- a/crates/tranquil-pds/src/sync/commit.rs +++ b/crates/tranquil-pds/src/sync/commit.rs @@ -13,6 +13,7 @@ use jacquard_repo::storage::BlockStore; use serde::{Deserialize, Serialize}; use std::str::FromStr; use tracing::error; +use tranquil_types::Did; async fn get_rev_from_commit(state: &AppState, cid_str: &str) -> Option { let cid = Cid::from_str(cid_str).ok()?; @@ -36,12 +37,16 @@ pub async fn get_latest_commit( State(state): State, Query(params): Query, ) -> Response { - let did = params.did.trim(); - if did.is_empty() { + let did_str = params.did.trim(); + if did_str.is_empty() { return ApiError::InvalidRequest("did is required".into()).into_response(); } + let did: Did = match did_str.parse() { + Ok(d) => d, + Err(_) => return ApiError::InvalidRequest("invalid did".into()).into_response(), + }; - let account = match assert_repo_availability(&state.db, did, false).await { + let account = match assert_repo_availability(state.repo_repo.as_ref(), &did, false).await { Ok(a) => a, Err(e) => return e.into_response(), }; @@ -53,7 +58,7 @@ pub async fn get_latest_commit( let Some(rev) = get_rev_from_commit(&state, &repo_root_cid).await else { error!( "Failed to parse commit for DID {}: CID {}", - did, repo_root_cid + did_str, repo_root_cid ); return ApiError::InternalError(Some("Failed to read repo commit".into())).into_response(); }; @@ -97,27 +102,19 @@ pub async fn list_repos( Query(params): Query, ) -> Response { let limit = params.limit.unwrap_or(50).clamp(1, 1000); - let cursor_did = params.cursor.as_deref().unwrap_or(""); - let result = sqlx::query!( - r#" - SELECT u.did, u.deactivated_at, u.takedown_ref, r.repo_root_cid, r.repo_rev - FROM repos r - JOIN users u ON r.user_id = u.id - WHERE u.did > $1 - ORDER BY u.did ASC - LIMIT $2 - "#, - cursor_did, - limit + 1 - ) - .fetch_all(&state.db) - .await; + let cursor_did: Option = params + .cursor + .as_ref() + .and_then(|s| s.parse().ok()); + let cursor_ref = cursor_did.as_ref(); + let result = state.repo_repo.list_repos_paginated(cursor_ref, limit + 1).await; match result { Ok(rows) => { let has_more = rows.len() as i64 > limit; let mut repos: Vec = Vec::new(); for row in rows.iter().take(limit as usize) { - let rev = match get_rev_from_commit(&state, &row.repo_root_cid).await { + let cid_str = row.repo_root_cid.to_string(); + let rev = match get_rev_from_commit(&state, &cid_str).await { Some(r) => r, None => { if let Some(ref stored_rev) = row.repo_rev { @@ -140,8 +137,8 @@ pub async fn list_repos( AccountStatus::Active }; repos.push(RepoInfo { - did: row.did.clone(), - head: row.repo_root_cid.clone(), + did: row.did.to_string(), + head: cid_str, rev, active: status.is_active(), status: status.as_str().map(String::from), @@ -187,15 +184,19 @@ pub async fn get_repo_status( State(state): State, Query(params): Query, ) -> Response { - let did = params.did.trim(); - if did.is_empty() { + let did_str = params.did.trim(); + if did_str.is_empty() { return ApiError::InvalidRequest("did is required".into()).into_response(); } + let did: Did = match did_str.parse() { + Ok(d) => d, + Err(_) => return ApiError::InvalidRequest("invalid did".into()).into_response(), + }; - let account = match get_account_with_status(&state.db, did).await { + let account = match get_account_with_status(state.repo_repo.as_ref(), &did).await { Ok(Some(a)) => a, Ok(None) => { - return ApiError::RepoNotFound(Some(format!("Could not find repo for DID: {}", did))) + return ApiError::RepoNotFound(Some(format!("Could not find repo for DID: {}", did_str))) .into_response(); } Err(e) => { diff --git a/crates/tranquil-pds/src/sync/deprecated.rs b/crates/tranquil-pds/src/sync/deprecated.rs index 95fbc7f..1548600 100644 --- a/crates/tranquil-pds/src/sync/deprecated.rs +++ b/crates/tranquil-pds/src/sync/deprecated.rs @@ -2,6 +2,7 @@ use crate::api::error::ApiError; use crate::state::AppState; use crate::sync::car::encode_car_header; use crate::sync::util::assert_repo_availability; +use tranquil_types::Did; use axum::{ Json, extract::{Query, State}, @@ -27,7 +28,8 @@ async fn check_admin_or_self(state: &AppState, headers: &HeaderMap, did: &str) - let dpop_proof = headers.get("DPoP").and_then(|h| h.to_str().ok()); let http_uri = "/"; match crate::auth::validate_token_with_dpop( - &state.db, + state.user_repo.as_ref(), + state.oauth_repo.as_ref(), &extracted.token, extracted.is_dpop, dpop_proof, @@ -58,18 +60,22 @@ pub async fn get_head( headers: HeaderMap, Query(params): Query, ) -> Response { - let did = params.did.trim(); - if did.is_empty() { + let did_str = params.did.trim(); + if did_str.is_empty() { return ApiError::InvalidRequest("did is required".into()).into_response(); } - let is_admin_or_self = check_admin_or_self(&state, &headers, did).await; - let account = match assert_repo_availability(&state.db, did, is_admin_or_self).await { + let did: Did = match did_str.parse() { + Ok(d) => d, + Err(_) => return ApiError::InvalidRequest("invalid did".into()).into_response(), + }; + let is_admin_or_self = check_admin_or_self(&state, &headers, did_str).await; + let account = match assert_repo_availability(state.repo_repo.as_ref(), &did, is_admin_or_self).await { Ok(a) => a, Err(e) => return e.into_response(), }; match account.repo_root_cid { Some(root) => (StatusCode::OK, Json(GetHeadOutput { root })).into_response(), - None => ApiError::RepoNotFound(Some(format!("Could not find root for DID: {}", did))) + None => ApiError::RepoNotFound(Some(format!("Could not find root for DID: {}", did_str))) .into_response(), } } @@ -84,12 +90,16 @@ pub async fn get_checkout( headers: HeaderMap, Query(params): Query, ) -> Response { - let did = params.did.trim(); - if did.is_empty() { + let did_str = params.did.trim(); + if did_str.is_empty() { return ApiError::InvalidRequest("did is required".into()).into_response(); } - let is_admin_or_self = check_admin_or_self(&state, &headers, did).await; - let account = match assert_repo_availability(&state.db, did, is_admin_or_self).await { + let did: Did = match did_str.parse() { + Ok(d) => d, + Err(_) => return ApiError::InvalidRequest("invalid did".into()).into_response(), + }; + let is_admin_or_self = check_admin_or_self(&state, &headers, did_str).await; + let account = match assert_repo_availability(state.repo_repo.as_ref(), &did, is_admin_or_self).await { Ok(a) => a, Err(e) => return e.into_response(), }; diff --git a/crates/tranquil-pds/src/sync/firehose.rs b/crates/tranquil-pds/src/sync/firehose.rs index f87438d..4f92e39 100644 --- a/crates/tranquil-pds/src/sync/firehose.rs +++ b/crates/tranquil-pds/src/sync/firehose.rs @@ -1,21 +1 @@ -use chrono::{DateTime, Utc}; -use serde::{Deserialize, Serialize}; -use serde_json::Value; - -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct SequencedEvent { - pub seq: i64, - pub did: String, - pub created_at: DateTime, - pub event_type: String, - pub commit_cid: Option, - pub prev_cid: Option, - pub prev_data_cid: Option, - pub ops: Option, - pub blobs: Option>, - pub blocks_cids: Option>, - pub handle: Option, - pub active: Option, - pub status: Option, - pub rev: Option, -} +pub use tranquil_db_traits::SequencedEvent; diff --git a/crates/tranquil-pds/src/sync/frame.rs b/crates/tranquil-pds/src/sync/frame.rs index 992fa0c..b8570f2 100644 --- a/crates/tranquil-pds/src/sync/frame.rs +++ b/crates/tranquil-pds/src/sync/frame.rs @@ -198,14 +198,14 @@ impl TryFrom for CommitFrame { type Error = CommitFrameError; fn try_from(event: SequencedEvent) -> Result { - let commit_cid_str = event.commit_cid.ok_or_else(|| { + let commit_cid = event.commit_cid.ok_or_else(|| { CommitFrameError::InvalidCommitCid("Missing commit_cid in event".to_string()) })?; let builder = CommitFrameBuilder::new( event.seq, - event.did, - &commit_cid_str, - event.prev_cid.as_deref(), + event.did.to_string(), + commit_cid.as_str(), + event.prev_cid.as_ref().map(|c| c.as_str()), event.ops.unwrap_or_default(), event.blobs.unwrap_or_default(), event.created_at, diff --git a/crates/tranquil-pds/src/sync/import.rs b/crates/tranquil-pds/src/sync/import.rs index c4b122f..e44cb87 100644 --- a/crates/tranquil-pds/src/sync/import.rs +++ b/crates/tranquil-pds/src/sync/import.rs @@ -3,11 +3,12 @@ use cid::Cid; use ipld_core::ipld::Ipld; use iroh_car::CarReader; use serde_json::Value as JsonValue; -use sqlx::PgPool; use std::collections::HashMap; use std::io::Cursor; +use std::sync::Arc; use thiserror::Error; use tracing::debug; +use tranquil_db::{ImportBlock, ImportRecord, ImportRepoError, RepoRepository}; use uuid::Uuid; #[derive(Error, Debug)] @@ -21,7 +22,7 @@ pub enum ImportError { #[error("Invalid CBOR: {0}")] InvalidCbor(String), #[error("Database error: {0}")] - Database(#[from] sqlx::Error), + Database(String), #[error("Block store error: {0}")] BlockStore(String), #[error("Import size limit exceeded")] @@ -38,6 +39,16 @@ pub enum ImportError { DidMismatch { car_did: String, auth_did: String }, } +impl From for ImportError { + fn from(e: ImportRepoError) -> Self { + match e { + ImportRepoError::RepoNotFound => ImportError::RepoNotFound, + ImportRepoError::ConcurrentModification => ImportError::ConcurrentModification, + ImportRepoError::Database(msg) => ImportError::Database(msg), + } + } +} + #[derive(Debug, Clone)] pub struct BlobRef { pub cid: String, @@ -307,7 +318,7 @@ fn extract_commit_info(commit: &Ipld) -> Result<(Cid, CommitInfo), ImportError> } pub async fn apply_import( - db: &PgPool, + repo_repo: &Arc, user_id: Uuid, root: Cid, blocks: HashMap, @@ -329,62 +340,33 @@ pub async fn apply_import( records.len(), user_id ); - let mut tx = db.begin().await?; - let repo = sqlx::query!( - "SELECT repo_root_cid FROM repos WHERE user_id = $1 FOR UPDATE NOWAIT", - user_id - ) - .fetch_optional(&mut *tx) - .await - .map_err(|e| { - if let sqlx::Error::Database(ref db_err) = e - && db_err.code().as_deref() == Some("55P03") - { - return ImportError::ConcurrentModification; - } - ImportError::Database(e) - })?; - if repo.is_none() { - return Err(ImportError::RepoNotFound); - } - let block_chunks: Vec> = blocks + + let import_blocks: Vec = blocks .iter() - .collect::>() - .chunks(100) - .map(|c| c.to_vec()) + .map(|(cid, data)| ImportBlock { + cid_bytes: cid.to_bytes(), + data: data.to_vec(), + }) .collect(); - for chunk in block_chunks { - for (cid, data) in chunk { - let cid_bytes = cid.to_bytes(); - sqlx::query!( - "INSERT INTO blocks (cid, data) VALUES ($1, $2) ON CONFLICT (cid) DO NOTHING", - &cid_bytes, - data.as_ref() - ) - .execute(&mut *tx) - .await?; - } - } - sqlx::query!("DELETE FROM records WHERE repo_id = $1", user_id) - .execute(&mut *tx) - .await?; - for record in &records { - let record_cid_str = record.cid.to_string(); - sqlx::query!( - r#" - INSERT INTO records (repo_id, collection, rkey, record_cid) - VALUES ($1, $2, $3, $4) - ON CONFLICT (repo_id, collection, rkey) DO UPDATE SET record_cid = $4 - "#, - user_id, - record.collection, - record.rkey, - record_cid_str - ) - .execute(&mut *tx) + + let import_records: Vec = records + .iter() + .filter_map(|r| { + let collection = r.collection.parse().ok()?; + let rkey = r.rkey.parse().ok()?; + let record_cid = r.cid.to_string().parse().ok()?; + Some(ImportRecord { + collection, + rkey, + record_cid, + }) + }) + .collect(); + + repo_repo + .import_repo_data(user_id, &import_blocks, &import_records) .await?; - } - tx.commit().await?; + debug!( "Successfully imported {} blocks and {} records", blocks.len(), diff --git a/crates/tranquil-pds/src/sync/listener.rs b/crates/tranquil-pds/src/sync/listener.rs index 672830a..357abe9 100644 --- a/crates/tranquil-pds/src/sync/listener.rs +++ b/crates/tranquil-pds/src/sync/listener.rs @@ -1,17 +1,12 @@ use crate::state::AppState; use crate::sync::firehose::SequencedEvent; -use sqlx::postgres::PgListener; use std::sync::atomic::{AtomicI64, Ordering}; use tracing::{debug, error, info, warn}; static LAST_BROADCAST_SEQ: AtomicI64 = AtomicI64::new(0); pub async fn start_sequencer_listener(state: AppState) { - let initial_seq = sqlx::query_scalar!("SELECT COALESCE(MAX(seq), 0) as max FROM repo_seq") - .fetch_one(&state.db) - .await - .unwrap_or(Some(0)) - .unwrap_or(0); + let initial_seq = state.repo_repo.get_max_seq().await.unwrap_or(0); LAST_BROADCAST_SEQ.store(initial_seq, Ordering::SeqCst); info!(initial_seq = initial_seq, "Initialized sequencer listener"); tokio::spawn(async move { @@ -26,22 +21,18 @@ pub async fn start_sequencer_listener(state: AppState) { } async fn listen_loop(state: AppState) -> anyhow::Result<()> { - let mut listener = PgListener::connect_with(&state.db).await?; - listener.listen("repo_updates").await?; - info!("Connected to Postgres and listening for 'repo_updates'"); + let mut receiver = state + .event_notifier + .subscribe() + .await + .map_err(|e| anyhow::anyhow!("Failed to subscribe to events: {:?}", e))?; + info!("Connected to database and listening for repo updates"); let catchup_start = LAST_BROADCAST_SEQ.load(Ordering::SeqCst); - let events = sqlx::query_as!( - SequencedEvent, - r#" - SELECT seq, did, created_at, event_type, commit_cid, prev_cid, prev_data_cid, ops, blobs, blocks_cids, handle, active, status, rev - FROM repo_seq - WHERE seq > $1 - ORDER BY seq ASC - "#, - catchup_start - ) - .fetch_all(&state.db) - .await?; + let events = state + .repo_repo + .get_events_since_seq(catchup_start, None) + .await + .map_err(|e| anyhow::anyhow!("Failed to fetch catchup events: {:?}", e))?; if !events.is_empty() { info!( count = events.len(), @@ -50,24 +41,16 @@ async fn listen_loop(state: AppState) -> anyhow::Result<()> { ); events.into_iter().for_each(|event| { let seq = event.seq; - let _ = state.firehose_tx.send(event); + let firehose_event = to_firehose_event(event); + let _ = state.firehose_tx.send(firehose_event); LAST_BROADCAST_SEQ.store(seq, Ordering::SeqCst); }); } loop { - let notification = listener.recv().await?; - let payload = notification.payload(); - debug!(payload = %payload, "Received postgres notification"); - let seq_id: i64 = match payload.parse() { - Ok(id) => id, - Err(e) => { - warn!( - "Received invalid payload in repo_updates: '{}'. Error: {}", - payload, e - ); - continue; - } + let Some(seq_id) = receiver.recv().await else { + return Err(anyhow::anyhow!("Event receiver disconnected")); }; + debug!(seq = seq_id, "Received event notification"); let last_seq = LAST_BROADCAST_SEQ.load(Ordering::SeqCst); if seq_id <= last_seq { debug!( @@ -78,41 +61,26 @@ async fn listen_loop(state: AppState) -> anyhow::Result<()> { continue; } if seq_id > last_seq + 1 { - let gap_events = sqlx::query_as!( - SequencedEvent, - r#" - SELECT seq, did, created_at, event_type, commit_cid, prev_cid, prev_data_cid, ops, blobs, blocks_cids, handle, active, status, rev - FROM repo_seq - WHERE seq > $1 AND seq < $2 - ORDER BY seq ASC - "#, - last_seq, - seq_id - ) - .fetch_all(&state.db) - .await?; + let gap_events = state + .repo_repo + .get_events_in_seq_range(last_seq, seq_id) + .await + .unwrap_or_default(); if !gap_events.is_empty() { debug!(count = gap_events.len(), "Filling sequence gap"); gap_events.into_iter().for_each(|event| { let seq = event.seq; - let _ = state.firehose_tx.send(event); + let firehose_event = to_firehose_event(event); + let _ = state.firehose_tx.send(firehose_event); LAST_BROADCAST_SEQ.store(seq, Ordering::SeqCst); }); } } - let event = sqlx::query_as!( - SequencedEvent, - r#" - SELECT seq, did, created_at, event_type, commit_cid, prev_cid, prev_data_cid, ops, blobs, blocks_cids, handle, active, status, rev - FROM repo_seq - WHERE seq = $1 - "#, - seq_id - ) - .fetch_optional(&state.db) - .await?; + let event = state.repo_repo.get_event_by_seq(seq_id).await.ok().flatten(); if let Some(event) = event { - match state.firehose_tx.send(event) { + let seq = event.seq; + let firehose_event = to_firehose_event(event); + match state.firehose_tx.send(firehose_event) { Ok(receiver_count) => { debug!( seq = seq_id, @@ -124,7 +92,7 @@ async fn listen_loop(state: AppState) -> anyhow::Result<()> { warn!(seq = seq_id, error = %e, "Failed to broadcast event (no receivers?)"); } } - LAST_BROADCAST_SEQ.store(seq_id, Ordering::SeqCst); + LAST_BROADCAST_SEQ.store(seq, Ordering::SeqCst); } else { warn!( seq = seq_id, @@ -133,3 +101,22 @@ async fn listen_loop(state: AppState) -> anyhow::Result<()> { } } } + +fn to_firehose_event(event: tranquil_db_traits::SequencedEvent) -> SequencedEvent { + SequencedEvent { + seq: event.seq, + did: event.did, + created_at: event.created_at, + event_type: event.event_type, + commit_cid: event.commit_cid, + prev_cid: event.prev_cid, + prev_data_cid: event.prev_data_cid, + ops: event.ops, + blobs: event.blobs, + blocks_cids: event.blocks_cids, + handle: event.handle, + active: event.active, + status: event.status, + rev: event.rev, + } +} diff --git a/crates/tranquil-pds/src/sync/repo.rs b/crates/tranquil-pds/src/sync/repo.rs index 560de6a..bc8a114 100644 --- a/crates/tranquil-pds/src/sync/repo.rs +++ b/crates/tranquil-pds/src/sync/repo.rs @@ -14,6 +14,7 @@ use serde::Deserialize; use std::io::Write; use std::str::FromStr; use tracing::error; +use tranquil_types::Did; fn parse_get_blocks_query(query_string: &str) -> Result<(String, Vec), String> { let did = crate::util::parse_repeated_query_param(Some(query_string), "did") @@ -29,12 +30,16 @@ pub async fn get_blocks(State(state): State, RawQuery(query): RawQuery return ApiError::InvalidRequest("Missing query parameters".into()).into_response(); }; - let (did, cid_strings) = match parse_get_blocks_query(&query_string) { + let (did_str, cid_strings) = match parse_get_blocks_query(&query_string) { Ok(parsed) => parsed, Err(msg) => return ApiError::InvalidRequest(msg).into_response(), }; + let did: Did = match did_str.parse() { + Ok(d) => d, + Err(_) => return ApiError::InvalidRequest("invalid did".into()).into_response(), + }; - let _account = match assert_repo_availability(&state.db, &did, false).await { + let _account = match assert_repo_availability(state.repo_repo.as_ref(), &did, false).await { Ok(a) => a, Err(e) => return e.into_response(), }; @@ -119,7 +124,11 @@ pub async fn get_repo( State(state): State, Query(query): Query, ) -> Response { - let account = match assert_repo_availability(&state.db, &query.did, false).await { + let did: Did = match query.did.parse() { + Ok(d) => d, + Err(_) => return ApiError::InvalidRequest("invalid did".into()).into_response(), + }; + let account = match assert_repo_availability(state.repo_repo.as_ref(), &did, false).await { Ok(a) => a, Err(e) => return e.into_response(), }; @@ -133,11 +142,11 @@ pub async fn get_repo( }; if let Some(since) = &query.since { - return get_repo_since(&state, &query.did, &head_cid, since).await; + return get_repo_since(&state, &did, &head_cid, since).await; } let car_bytes = match generate_repo_car_from_user_blocks( - &state.db, + state.repo_repo.as_ref(), &state.block_store, account.user_id, &head_cid, @@ -159,48 +168,32 @@ pub async fn get_repo( .into_response() } -async fn get_repo_since(state: &AppState, did: &str, head_cid: &Cid, since: &str) -> Response { - let events = sqlx::query!( - r#" - SELECT blocks_cids, commit_cid - FROM repo_seq - WHERE did = $1 AND rev > $2 - ORDER BY seq DESC - "#, - did, - since - ) - .fetch_all(&state.db) - .await; +async fn get_repo_since(state: &AppState, did: &Did, head_cid: &Cid, since: &str) -> Response { + let user_id = match state.user_repo.get_id_by_did(did).await { + Ok(Some(id)) => id, + Ok(None) => { + return ApiError::RepoNotFound(Some(format!("Could not find repo for DID: {}", did))) + .into_response(); + } + Err(e) => { + error!("DB error looking up user: {:?}", e); + return ApiError::InternalError(Some("Database error".into())).into_response(); + } + }; - let events = match events { - Ok(e) => e, + let block_cid_bytes = match state.repo_repo.get_user_block_cids_since_rev(user_id, since).await + { + Ok(cids) => cids, Err(e) => { error!("DB error in get_repo_since: {:?}", e); return ApiError::InternalError(Some("Database error".into())).into_response(); } }; - let block_cids: Vec = events + let block_cids: Vec = block_cid_bytes .iter() - .flat_map(|event| { - let block_cids = event - .blocks_cids - .as_ref() - .map(|cids| cids.iter().filter_map(|s| Cid::from_str(s).ok()).collect()) - .unwrap_or_else(Vec::new); - let commit_cid = event - .commit_cid - .as_ref() - .and_then(|s| Cid::from_str(s).ok()); - block_cids.into_iter().chain(commit_cid) - }) - .fold(Vec::new(), |mut acc, cid| { - if !acc.contains(&cid) { - acc.push(cid); - } - acc - }); + .filter_map(|bytes| Cid::try_from(bytes.as_slice()).ok()) + .collect(); let mut car_bytes = match encode_car_header(head_cid) { Ok(h) => h, @@ -269,7 +262,11 @@ pub async fn get_record( use std::collections::BTreeMap; use std::sync::Arc; - let account = match assert_repo_availability(&state.db, &query.did, false).await { + let did: Did = match query.did.parse() { + Ok(d) => d, + Err(_) => return ApiError::InvalidRequest("invalid did".into()).into_response(), + }; + let account = match assert_repo_availability(state.repo_repo.as_ref(), &did, false).await { Ok(a) => a, Err(e) => return e.into_response(), }; diff --git a/crates/tranquil-pds/src/sync/subscribe_repos.rs b/crates/tranquil-pds/src/sync/subscribe_repos.rs index 0a75a36..7d84f6e 100644 --- a/crates/tranquil-pds/src/sync/subscribe_repos.rs +++ b/crates/tranquil-pds/src/sync/subscribe_repos.rs @@ -72,12 +72,7 @@ async fn handle_socket_inner( let mut last_seen: i64 = -1; if let Some(cursor) = params.cursor { - let current_seq = sqlx::query_scalar!("SELECT MAX(seq) FROM repo_seq") - .fetch_one(&state.db) - .await - .ok() - .flatten() - .unwrap_or(0); + let current_seq = state.repo_repo.get_max_seq().await.unwrap_or(0); if cursor > current_seq { if let Ok(error_bytes) = @@ -91,21 +86,12 @@ async fn handle_socket_inner( let backfill_time = chrono::Utc::now() - chrono::Duration::hours(get_backfill_hours()); - let first_event = sqlx::query_as!( - SequencedEvent, - r#" - SELECT seq, did, created_at, event_type, commit_cid, prev_cid, prev_data_cid, ops, blobs, blocks_cids, handle, active, status, rev - FROM repo_seq - WHERE seq > $1 - ORDER BY seq ASC - LIMIT 1 - "#, - cursor - ) - .fetch_optional(&state.db) - .await - .ok() - .flatten(); + let first_event = state + .repo_repo + .get_events_since_cursor(cursor, 1) + .await + .ok() + .and_then(|events| events.into_iter().next()); let mut current_cursor = cursor; @@ -119,14 +105,7 @@ async fn handle_socket_inner( let _ = socket.send(Message::Binary(info_bytes.into())).await; } - let earliest = sqlx::query_scalar!( - "SELECT MIN(seq) FROM repo_seq WHERE created_at >= $1", - backfill_time - ) - .fetch_one(&state.db) - .await - .ok() - .flatten(); + let earliest = state.repo_repo.get_min_seq_since(backfill_time).await.ok().flatten(); if let Some(earliest_seq) = earliest { current_cursor = earliest_seq - 1; @@ -136,20 +115,10 @@ async fn handle_socket_inner( last_seen = current_cursor; loop { - let events = sqlx::query_as!( - SequencedEvent, - r#" - SELECT seq, did, created_at, event_type, commit_cid, prev_cid, prev_data_cid, ops, blobs, blocks_cids, handle, active, status, rev - FROM repo_seq - WHERE seq > $1 - ORDER BY seq ASC - LIMIT $2 - "#, - current_cursor, - BACKFILL_BATCH_SIZE - ) - .fetch_all(&state.db) - .await; + let events = state + .repo_repo + .get_events_since_cursor(current_cursor, BACKFILL_BATCH_SIZE) + .await; match events { Ok(events) => { if events.is_empty() { @@ -186,25 +155,17 @@ async fn handle_socket_inner( } } Err(e) => { - error!("Failed to fetch backfill events: {}", e); + error!("Failed to fetch backfill events: {:?}", e); socket.close().await.ok(); return Err(()); } } } - let cutover_events = sqlx::query_as!( - SequencedEvent, - r#" - SELECT seq, did, created_at, event_type, commit_cid, prev_cid, prev_data_cid, ops, blobs, blocks_cids, handle, active, status, rev - FROM repo_seq - WHERE seq > $1 - ORDER BY seq ASC - "#, - last_seen - ) - .fetch_all(&state.db) - .await; + let cutover_events = state + .repo_repo + .get_events_since_seq(last_seen, None) + .await; if let Ok(events) = cutover_events && !events.is_empty() diff --git a/crates/tranquil-pds/src/sync/util.rs b/crates/tranquil-pds/src/sync/util.rs index fb931aa..cd0b7cb 100644 --- a/crates/tranquil-pds/src/sync/util.rs +++ b/crates/tranquil-pds/src/sync/util.rs @@ -12,7 +12,8 @@ use iroh_car::{CarHeader, CarWriter}; use jacquard_repo::commit::Commit; use jacquard_repo::storage::BlockStore; use serde::Serialize; -use sqlx::PgPool; +use tranquil_db_traits::RepoRepository; +use tranquil_types::Did; use std::collections::{BTreeMap, HashMap}; use std::io::Cursor; use std::str::FromStr; @@ -134,20 +135,10 @@ impl IntoResponse for RepoAvailabilityError { } pub async fn get_account_with_status( - db: &PgPool, - did: &str, -) -> Result, sqlx::Error> { - let row = sqlx::query!( - r#" - SELECT u.id, u.did, u.deactivated_at, u.takedown_ref, r.repo_root_cid - FROM users u - LEFT JOIN repos r ON r.user_id = u.id - WHERE u.did = $1 - "#, - did - ) - .fetch_optional(db) - .await?; + repo_repo: &dyn RepoRepository, + did: &Did, +) -> Result, tranquil_db_traits::DbError> { + let row = repo_repo.get_account_with_repo(did).await?; Ok(row.map(|r| { let status = if r.takedown_ref.is_some() { @@ -159,26 +150,27 @@ pub async fn get_account_with_status( }; RepoAccount { - did: r.did, - user_id: r.id, + did: r.did.to_string(), + user_id: r.user_id, status, - repo_root_cid: Some(r.repo_root_cid), + repo_root_cid: r.repo_root_cid.map(|c| c.to_string()), } })) } pub async fn assert_repo_availability( - db: &PgPool, - did: &str, + repo_repo: &dyn RepoRepository, + did: &Did, is_admin_or_self: bool, ) -> Result { - let account = get_account_with_status(db, did) + let account = get_account_with_status(repo_repo, did) .await .map_err(|e| RepoAvailabilityError::Internal(e.to_string()))?; + let did_str = did.to_string(); let account = match account { Some(a) => a, - None => return Err(RepoAvailabilityError::NotFound(did.to_string())), + None => return Err(RepoAvailabilityError::NotFound(did_str)), }; if is_admin_or_self { @@ -186,9 +178,9 @@ pub async fn assert_repo_availability( } match account.status { - AccountStatus::Takendown => return Err(RepoAvailabilityError::Takendown(did.to_string())), + AccountStatus::Takendown => return Err(RepoAvailabilityError::Takendown(did_str)), AccountStatus::Deactivated => { - return Err(RepoAvailabilityError::Deactivated(did.to_string())); + return Err(RepoAvailabilityError::Deactivated(did_str)); } _ => {} } @@ -239,8 +231,8 @@ fn format_atproto_time(dt: chrono::DateTime) -> String { fn format_identity_event(event: &SequencedEvent) -> Result, anyhow::Error> { let frame = IdentityFrame { - did: event.did.clone(), - handle: event.handle.clone(), + did: event.did.to_string(), + handle: event.handle.as_ref().map(|h| h.to_string()), seq: event.seq, time: format_atproto_time(event.created_at), }; @@ -256,7 +248,7 @@ fn format_identity_event(event: &SequencedEvent) -> Result, anyhow::Erro fn format_account_event(event: &SequencedEvent) -> Result, anyhow::Error> { let frame = AccountFrame { - did: event.did.clone(), + did: event.did.to_string(), active: event.active.unwrap_or(true), status: event.status.clone(), seq: event.seq, @@ -303,7 +295,7 @@ async fn format_sync_event( }; let car_bytes = write_car_blocks(commit_cid, Some(commit_bytes), BTreeMap::new()).await?; let frame = SyncFrame { - did: event.did.clone(), + did: event.did.to_string(), rev, blocks: car_bytes, seq: event.seq, @@ -330,18 +322,18 @@ pub async fn format_event_for_sending( _ => {} } let block_cids_str = event.blocks_cids.clone().unwrap_or_default(); - let prev_cid_str = event.prev_cid.clone(); - let prev_data_cid_str = event.prev_data_cid.clone(); + let prev_cid_link = event.prev_cid.clone(); + let prev_data_cid_link = event.prev_data_cid.clone(); let mut frame: CommitFrame = event .try_into() .map_err(|e| anyhow::anyhow!("Invalid event: {}", e))?; - if let Some(ref pdc) = prev_data_cid_str - && let Ok(cid) = Cid::from_str(pdc) + if let Some(ref pdc) = prev_data_cid_link + && let Ok(cid) = Cid::from_str(pdc.as_str()) { frame.prev_data = Some(cid); } let commit_cid = frame.commit; - let prev_cid = prev_cid_str.as_ref().and_then(|s| Cid::from_str(s).ok()); + let prev_cid = prev_cid_link.as_ref().and_then(|c| Cid::from_str(c.as_str()).ok()); let mut all_cids: Vec = block_cids_str .iter() .filter_map(|s| Cid::from_str(s).ok()) @@ -443,7 +435,7 @@ fn format_sync_event_with_prefetched( BTreeMap::new(), ))?; let frame = SyncFrame { - did: event.did.clone(), + did: event.did.to_string(), rev, blocks: car_bytes, seq: event.seq, @@ -470,18 +462,18 @@ pub async fn format_event_with_prefetched_blocks( _ => {} } let block_cids_str = event.blocks_cids.clone().unwrap_or_default(); - let prev_cid_str = event.prev_cid.clone(); - let prev_data_cid_str = event.prev_data_cid.clone(); + let prev_cid_link = event.prev_cid.clone(); + let prev_data_cid_link = event.prev_data_cid.clone(); let mut frame: CommitFrame = event .try_into() .map_err(|e| anyhow::anyhow!("Invalid event: {}", e))?; - if let Some(ref pdc) = prev_data_cid_str - && let Ok(cid) = Cid::from_str(pdc) + if let Some(ref pdc) = prev_data_cid_link + && let Ok(cid) = Cid::from_str(pdc.as_str()) { frame.prev_data = Some(cid); } let commit_cid = frame.commit; - let prev_cid = prev_cid_str.as_ref().and_then(|s| Cid::from_str(s).ok()); + let prev_cid = prev_cid_link.as_ref().and_then(|c| Cid::from_str(c.as_str()).ok()); let mut all_cids: Vec = block_cids_str .iter() .filter_map(|s| Cid::from_str(s).ok()) diff --git a/crates/tranquil-pds/src/util.rs b/crates/tranquil-pds/src/util.rs index 24fab14..781dce5 100644 --- a/crates/tranquil-pds/src/util.rs +++ b/crates/tranquil-pds/src/util.rs @@ -3,13 +3,9 @@ use cid::Cid; use ipld_core::ipld::Ipld; use rand::Rng; use serde_json::Value as JsonValue; -use sqlx::PgPool; use std::collections::BTreeMap; use std::str::FromStr; use std::sync::OnceLock; -use uuid::Uuid; - -use crate::types::{Did, Handle}; const BASE32_ALPHABET: &str = "abcdefghijklmnopqrstuvwxyz234567"; const DEFAULT_MAX_BLOB_SIZE: usize = 10 * 1024 * 1024 * 1024; @@ -43,66 +39,6 @@ pub fn generate_token_code_parts(parts: usize, part_len: usize) -> String { .join("-") } -#[derive(Debug)] -pub enum DbLookupError { - NotFound, - DatabaseError(sqlx::Error), -} - -impl From for DbLookupError { - fn from(e: sqlx::Error) -> Self { - DbLookupError::DatabaseError(e) - } -} - -pub async fn get_user_id_by_did(db: &PgPool, did: &str) -> Result { - sqlx::query_scalar!("SELECT id FROM users WHERE did = $1", did) - .fetch_optional(db) - .await? - .ok_or(DbLookupError::NotFound) -} - -pub struct UserInfo { - pub id: Uuid, - pub did: Did, - pub handle: Handle, -} - -pub async fn get_user_by_did(db: &PgPool, did: &str) -> Result { - sqlx::query_as!( - UserInfo, - "SELECT id, did, handle FROM users WHERE did = $1", - did - ) - .fetch_optional(db) - .await? - .ok_or(DbLookupError::NotFound) -} - -pub async fn get_user_by_identifier( - db: &PgPool, - identifier: &str, -) -> Result { - sqlx::query_as!( - UserInfo, - "SELECT id, did, handle FROM users WHERE did = $1 OR handle = $1", - identifier - ) - .fetch_optional(db) - .await? - .ok_or(DbLookupError::NotFound) -} - -pub async fn is_account_migrated(db: &PgPool, did: &str) -> Result { - let row = sqlx::query!( - r#"SELECT (migrated_to_pds IS NOT NULL AND deactivated_at IS NOT NULL) as "migrated!: bool" FROM users WHERE did = $1"#, - did - ) - .fetch_optional(db) - .await?; - Ok(row.map(|r| r.migrated).unwrap_or(false)) -} - pub fn parse_repeated_query_param(query: Option<&str>, key: &str) -> Vec { query .map(|q| { diff --git a/crates/tranquil-pds/tests/account_notifications.rs b/crates/tranquil-pds/tests/account_notifications.rs index 8fbdaec..d31807b 100644 --- a/crates/tranquil-pds/tests/account_notifications.rs +++ b/crates/tranquil-pds/tests/account_notifications.rs @@ -1,7 +1,7 @@ mod common; use common::{base_url, client, create_account_and_login, get_test_db_pool}; use serde_json::{Value, json}; -use tranquil_pds::comms::{CommsType, NewComms, enqueue_comms}; +use sqlx::Row; #[tokio::test] async fn test_get_notification_history() { @@ -10,20 +10,24 @@ async fn test_get_notification_history() { let pool = get_test_db_pool().await; let (token, did) = create_account_and_login(&client).await; - let user_id: uuid::Uuid = sqlx::query_scalar!("SELECT id FROM users WHERE did = $1", did) + let user_id: uuid::Uuid = sqlx::query_scalar("SELECT id FROM users WHERE did = $1") + .bind(&did) .fetch_one(pool) .await .expect("User not found"); for i in 0..3 { - let comms = NewComms::email( - user_id, - CommsType::Welcome, - "test@example.com".to_string(), - format!("Subject {}", i), - format!("Body {}", i), - ); - enqueue_comms(pool, comms).await.expect("Failed to enqueue"); + sqlx::query( + r#"INSERT INTO comms_queue (user_id, channel, comms_type, recipient, subject, body) + VALUES ($1, 'email', 'welcome', $2, $3, $4)"#, + ) + .bind(user_id) + .bind("test@example.com") + .bind(format!("Subject {}", i)) + .bind(format!("Body {}", i)) + .execute(pool) + .await + .expect("Failed to enqueue"); } let resp = client @@ -69,21 +73,22 @@ async fn test_verify_channel_discord() { ); let pool = get_test_db_pool().await; - let user_id: uuid::Uuid = sqlx::query_scalar!("SELECT id FROM users WHERE did = $1", did) + let user_id: uuid::Uuid = sqlx::query_scalar("SELECT id FROM users WHERE did = $1") + .bind(&did) .fetch_one(pool) .await .expect("User not found"); - let row = sqlx::query!( + let row = sqlx::query( "SELECT body, metadata FROM comms_queue WHERE user_id = $1 AND comms_type = 'channel_verification' ORDER BY created_at DESC LIMIT 1", - user_id ) + .bind(user_id) .fetch_one(pool) .await .expect("Verification code not found"); - let code = row - .metadata + let metadata: Option = row.get("metadata"); + let code = metadata .as_ref() .and_then(|m| m.get("code")) .and_then(|c| c.as_str()) @@ -203,15 +208,16 @@ async fn test_update_email_via_notification_prefs() { .contains(&json!("email")) ); - let user_id: uuid::Uuid = sqlx::query_scalar!("SELECT id FROM users WHERE did = $1", did) + let user_id: uuid::Uuid = sqlx::query_scalar("SELECT id FROM users WHERE did = $1") + .bind(&did) .fetch_one(pool) .await .expect("User not found"); - let body_text: String = sqlx::query_scalar!( + let body_text: String = sqlx::query_scalar( "SELECT body FROM comms_queue WHERE user_id = $1 AND comms_type = 'email_update' ORDER BY created_at DESC LIMIT 1", - user_id ) + .bind(user_id) .fetch_one(pool) .await .expect("Verification code not found"); diff --git a/crates/tranquil-pds/tests/common/mod.rs b/crates/tranquil-pds/tests/common/mod.rs index bb16086..2283d41 100644 --- a/crates/tranquil-pds/tests/common/mod.rs +++ b/crates/tranquil-pds/tests/common/mod.rs @@ -76,6 +76,16 @@ pub fn app_port() -> u16 { *APP_PORT.get().expect("APP_PORT not initialized") } +#[allow(dead_code)] +pub fn pds_hostname() -> String { + std::env::var("PDS_HOSTNAME").unwrap_or_else(|_| format!("pds.test:{}", app_port())) +} + +#[allow(dead_code)] +pub fn pds_endpoint() -> String { + format!("https://{}", pds_hostname()) +} + pub async fn base_url() -> &'static str { SERVER_URL.get_or_init(|| { let (tx, rx) = std::sync::mpsc::channel(); @@ -457,7 +467,7 @@ async fn spawn_app(database_url: String) -> String { let addr = listener.local_addr().unwrap(); APP_PORT.set(addr.port()).ok(); unsafe { - std::env::set_var("PDS_HOSTNAME", addr.to_string()); + std::env::set_var("PDS_HOSTNAME", format!("pds.test:{}", addr.port())); } let rate_limiters = RateLimiters::new() .with_login_limit(10000) @@ -474,7 +484,7 @@ async fn spawn_app(database_url: String) -> String { tokio::spawn(async move { axum::serve(listener, app).await.unwrap(); }); - format!("http://{}", addr) + format!("http://localhost:{}", addr.port()) } #[allow(dead_code)] diff --git a/crates/tranquil-pds/tests/did_web.rs b/crates/tranquil-pds/tests/did_web.rs index 4221631..074134c 100644 --- a/crates/tranquil-pds/tests/did_web.rs +++ b/crates/tranquil-pds/tests/did_web.rs @@ -94,17 +94,18 @@ async fn test_create_self_hosted_did_web() { #[tokio::test] async fn test_external_did_web_no_local_doc() { let client = client(); + let base = base_url().await; let mock_server = MockServer::start().await; let mock_uri = mock_server.uri(); let mock_addr = mock_uri.trim_start_matches("http://"); let did = format!("did:web:{}", mock_addr.replace(":", "%3A")); let handle = format!("xw{}", &uuid::Uuid::new_v4().simple().to_string()[..12]); - let pds_endpoint = base_url().await.replace("http://", "https://"); + let pds_endpoint = common::pds_endpoint(); let reserve_res = client .post(format!( "{}/xrpc/com.atproto.server.reserveSigningKey", - base_url().await + base )) .json(&json!({ "did": did })) .send() @@ -150,7 +151,7 @@ async fn test_external_did_web_no_local_doc() { let res = client .post(format!( "{}/xrpc/com.atproto.server.createAccount", - base_url().await + base )) .json(&payload) .send() @@ -161,7 +162,7 @@ async fn test_external_did_web_no_local_doc() { panic!("createAccount failed: {:?}", body); } let res = client - .get(format!("{}/u/{}/did.json", base_url().await, handle)) + .get(format!("{}/u/{}/did.json", base, handle)) .send() .await .expect("Failed to fetch DID doc"); @@ -383,6 +384,7 @@ fn create_service_jwt(signing_key: &SigningKey, did: &str, aud: &str) -> String #[tokio::test] async fn test_did_web_byod_flow() { let client = client(); + let base = base_url().await; let mock_server = MockServer::start().await; let mock_uri = mock_server.uri(); let mock_addr = mock_uri.trim_start_matches("http://"); @@ -393,8 +395,9 @@ async fn test_did_web_byod_flow() { unique_id ); let handle = format!("by{}", &uuid::Uuid::new_v4().simple().to_string()[..12]); - let pds_endpoint = base_url().await.replace("http://", "https://"); - let pds_did = format!("did:web:{}", pds_endpoint.trim_start_matches("https://")); + let pds_endpoint = common::pds_endpoint(); + let pds_hostname = common::pds_hostname(); + let pds_did = format!("did:web:{}", pds_hostname); let temp_key = SigningKey::random(&mut rand::thread_rng()); let public_key_multibase = signing_key_to_multibase(&temp_key); @@ -430,7 +433,7 @@ async fn test_did_web_byod_flow() { let res = client .post(format!( "{}/xrpc/com.atproto.server.createAccount", - base_url().await + base )) .header("Authorization", format!("Bearer {}", service_jwt)) .json(&payload) @@ -454,7 +457,7 @@ async fn test_did_web_byod_flow() { let res = client .get(format!( "{}/xrpc/com.atproto.server.checkAccountStatus", - base_url().await + base )) .bearer_auth(&access_jwt) .send() @@ -470,7 +473,7 @@ async fn test_did_web_byod_flow() { let res = client .get(format!( "{}/xrpc/com.atproto.identity.getRecommendedDidCredentials", - base_url().await + base )) .bearer_auth(&access_jwt) .send() @@ -493,7 +496,7 @@ async fn test_did_web_byod_flow() { let res = client .post(format!( "{}/xrpc/com.atproto.server.activateAccount", - base_url().await + base )) .bearer_auth(&access_jwt) .send() @@ -508,7 +511,7 @@ async fn test_did_web_byod_flow() { let res = client .get(format!( "{}/xrpc/com.atproto.server.checkAccountStatus", - base_url().await + base )) .bearer_auth(&access_jwt) .send() @@ -524,7 +527,7 @@ async fn test_did_web_byod_flow() { let res = client .post(format!( "{}/xrpc/com.atproto.repo.createRecord", - base_url().await + base )) .bearer_auth(&access_jwt) .json(&json!({ diff --git a/crates/tranquil-pds/tests/firehose_validation.rs b/crates/tranquil-pds/tests/firehose_validation.rs index 62da708..b8e059a 100644 --- a/crates/tranquil-pds/tests/firehose_validation.rs +++ b/crates/tranquil-pds/tests/firehose_validation.rs @@ -850,3 +850,134 @@ async fn test_firehose_outdated_cursor_info() { "Should have received commits even with outdated cursor" ); } + +#[tokio::test] +async fn test_firehose_car_contains_mst_blocks() { + let client = client(); + let (token, did) = create_account_and_login(&client).await; + + for i in 0..3 { + let post_payload = json!({ + "repo": did, + "collection": "app.bsky.feed.post", + "record": { + "$type": "app.bsky.feed.post", + "text": format!("Setup post {}", i), + "createdAt": chrono::Utc::now().to_rfc3339(), + } + }); + client + .post(format!( + "{}/xrpc/com.atproto.repo.createRecord", + base_url().await + )) + .bearer_auth(&token) + .json(&post_payload) + .send() + .await + .expect("Failed to create setup post"); + tokio::time::sleep(std::time::Duration::from_millis(50)).await; + } + + let url = format!( + "ws://127.0.0.1:{}/xrpc/com.atproto.sync.subscribeRepos", + app_port() + ); + let (mut ws_stream, _) = connect_async(&url).await.expect("Failed to connect"); + tokio::time::sleep(std::time::Duration::from_millis(100)).await; + + let post_payload = json!({ + "repo": did, + "collection": "app.bsky.feed.post", + "record": { + "$type": "app.bsky.feed.post", + "text": "Test post for MST block validation", + "createdAt": chrono::Utc::now().to_rfc3339(), + } + }); + let res = client + .post(format!( + "{}/xrpc/com.atproto.repo.createRecord", + base_url().await + )) + .bearer_auth(&token) + .json(&post_payload) + .send() + .await + .expect("Failed to create post"); + assert_eq!(res.status(), StatusCode::OK); + let create_result: Value = res.json().await.unwrap(); + let record_cid_str = create_result["cid"].as_str().unwrap(); + let expected_record_cid: Cid = record_cid_str.parse().unwrap(); + + let mut frame_opt: Option = None; + let timeout = tokio::time::timeout(std::time::Duration::from_secs(10), async { + loop { + let msg = ws_stream.next().await.unwrap().unwrap(); + let raw_bytes = match msg { + tungstenite::Message::Binary(bin) => bin, + _ => continue, + }; + if let Ok((_, f)) = parse_frame(&raw_bytes) + && f.repo == did + && f.ops.iter().any(|op| op.cid == Some(expected_record_cid)) + { + frame_opt = Some(f); + break; + } + } + }) + .await; + assert!(timeout.is_ok(), "Timed out waiting for firehose event"); + let frame = frame_opt.expect("No matching frame found"); + + let mut car_reader = CarReader::new(Cursor::new(&frame.blocks)).await.unwrap(); + + let mut block_count = 0; + let mut found_commit = false; + let mut found_record = false; + let mut mst_block_count = 0; + + while let Ok(Some((cid, data))) = car_reader.next_block().await { + block_count += 1; + + if cid == frame.commit { + found_commit = true; + continue; + } + + if cid == expected_record_cid { + found_record = true; + continue; + } + + if data.len() > 10 && data.len() < 5000 { + mst_block_count += 1; + } + } + + println!("CAR block analysis:"); + println!(" Total blocks: {}", block_count); + println!(" Found commit: {}", found_commit); + println!(" Found record: {}", found_record); + println!(" MST/other blocks: {}", mst_block_count); + + assert!(found_commit, "CAR must contain commit block"); + assert!(found_record, "CAR must contain record block"); + + assert!( + block_count >= 3, + "CAR should contain at least commit + record + MST node(s), got {} blocks. \ + This may indicate firehose is not including all relevant blocks.", + block_count + ); + + assert!( + mst_block_count >= 1, + "CAR should contain MST node blocks for repo validation, got {} MST blocks. \ + Firehose must include relevant MST blocks, not just new ones.", + mst_block_count + ); + + ws_stream.send(tungstenite::Message::Close(None)).await.ok(); +} diff --git a/crates/tranquil-pds/tests/identity.rs b/crates/tranquil-pds/tests/identity.rs index 6e008df..8822a11 100644 --- a/crates/tranquil-pds/tests/identity.rs +++ b/crates/tranquil-pds/tests/identity.rs @@ -48,11 +48,12 @@ async fn test_resolve_handle_success() { #[tokio::test] async fn test_resolve_handle_not_found() { let client = client(); - let params = [("handle", "nonexistent_handle_12345")]; + let _base = base_url().await; + let params = [("handle", "nonexistent.handle.test")]; let res = client .get(format!( "{}/xrpc/com.atproto.identity.resolveHandle", - base_url().await + _base )) .query(¶ms) .send() @@ -99,12 +100,13 @@ async fn test_create_did_web_account_and_resolve() { let mock_addr = mock_uri.trim_start_matches("http://"); let did = format!("did:web:{}", mock_addr.replace(":", "%3A")); let handle = format!("wu{}", &uuid::Uuid::new_v4().simple().to_string()[..12]); - let pds_endpoint = base_url().await.replace("http://", "https://"); + let base = base_url().await; + let pds_endpoint = common::pds_endpoint(); let reserve_res = client .post(format!( "{}/xrpc/com.atproto.server.reserveSigningKey", - base_url().await + base )) .json(&json!({ "did": did })) .send() @@ -149,7 +151,7 @@ async fn test_create_did_web_account_and_resolve() { let res = client .post(format!( "{}/xrpc/com.atproto.server.createAccount", - base_url().await + base )) .json(&payload) .send() @@ -169,7 +171,7 @@ async fn test_create_did_web_account_and_resolve() { .expect("createAccount response was not JSON"); assert_eq!(body["did"], did); let res = client - .get(format!("{}/u/{}/did.json", base_url().await, handle)) + .get(format!("{}/u/{}/did.json", base, handle)) .send() .await .expect("Failed to fetch DID doc"); @@ -217,18 +219,19 @@ async fn test_create_account_duplicate_handle() { #[tokio::test] async fn test_did_web_lifecycle() { let client = client(); + let base = base_url().await; let mock_server = MockServer::start().await; let mock_uri = mock_server.uri(); let mock_addr = mock_uri.trim_start_matches("http://"); let handle = format!("lc{}", &uuid::Uuid::new_v4().simple().to_string()[..12]); let did = format!("did:web:{}:u:{}", mock_addr.replace(":", "%3A"), handle); let email = format!("{}@test.com", handle); - let pds_endpoint = base_url().await.replace("http://", "https://"); + let pds_endpoint = common::pds_endpoint(); let reserve_res = client .post(format!( "{}/xrpc/com.atproto.server.reserveSigningKey", - base_url().await + base )) .json(&json!({ "did": did })) .send() @@ -273,7 +276,7 @@ async fn test_did_web_lifecycle() { let res = client .post(format!( "{}/xrpc/com.atproto.server.createAccount", - base_url().await + base )) .json(&create_payload) .send() diff --git a/crates/tranquil-pds/tests/lifecycle_record.rs b/crates/tranquil-pds/tests/lifecycle_record.rs index 65e93fa..847104f 100644 --- a/crates/tranquil-pds/tests/lifecycle_record.rs +++ b/crates/tranquil-pds/tests/lifecycle_record.rs @@ -132,8 +132,8 @@ async fn test_record_crud_lifecycle() { .expect("Failed to send stale update"); assert_eq!( stale_res.status(), - StatusCode::CONFLICT, - "Stale update should cause 409" + StatusCode::BAD_REQUEST, + "Stale update should cause 400 InvalidSwap" ); let good_update_payload = json!({ "repo": did, diff --git a/crates/tranquil-pds/tests/notifications.rs b/crates/tranquil-pds/tests/notifications.rs index 4837a31..7c7607f 100644 --- a/crates/tranquil-pds/tests/notifications.rs +++ b/crates/tranquil-pds/tests/notifications.rs @@ -1,113 +1,91 @@ mod common; -use tranquil_pds::comms::{ - CommsChannel, CommsStatus, CommsType, NewComms, enqueue_comms, enqueue_welcome, -}; +use sqlx::Row; +use tranquil_pds::comms::{CommsChannel, CommsStatus, CommsType}; #[tokio::test] async fn test_enqueue_comms() { let pool = common::get_test_db_pool().await; let (_, did) = common::create_account_and_login(&common::client()).await; - let user_id: uuid::Uuid = sqlx::query_scalar!("SELECT id FROM users WHERE did = $1", did) + let user_id: uuid::Uuid = sqlx::query_scalar("SELECT id FROM users WHERE did = $1") + .bind(&did) .fetch_one(pool) .await .expect("User not found"); - let item = NewComms::email( - user_id, - CommsType::Welcome, - "test@example.com".to_string(), - "Test Subject".to_string(), - "Test body".to_string(), - ); - let comms_id = enqueue_comms(pool, item) - .await - .expect("Failed to enqueue comms"); - let row = sqlx::query!( - r#" - SELECT - id, user_id, recipient, subject, body, - channel as "channel: CommsChannel", - comms_type as "comms_type: CommsType", - status as "status: CommsStatus" - FROM comms_queue - WHERE id = $1 - "#, - comms_id + let comms_id: uuid::Uuid = sqlx::query_scalar( + r#"INSERT INTO comms_queue (user_id, channel, comms_type, recipient, subject, body) + VALUES ($1, 'email', 'welcome', $2, $3, $4) + RETURNING id"#, ) + .bind(user_id) + .bind("test@example.com") + .bind("Test Subject") + .bind("Test body") .fetch_one(pool) .await - .expect("Comms not found"); - assert_eq!(row.user_id, user_id); - assert_eq!(row.recipient, "test@example.com"); - assert_eq!(row.subject.as_deref(), Some("Test Subject")); - assert_eq!(row.body, "Test body"); - assert_eq!(row.channel, CommsChannel::Email); - assert_eq!(row.comms_type, CommsType::Welcome); - assert_eq!(row.status, CommsStatus::Pending); -} - -#[tokio::test] -async fn test_enqueue_welcome() { - let pool = common::get_test_db_pool().await; - let (_, did) = common::create_account_and_login(&common::client()).await; - let user_row = sqlx::query!("SELECT id, email, handle FROM users WHERE did = $1", did) - .fetch_one(pool) - .await - .expect("User not found"); - let comms_id = enqueue_welcome(pool, user_row.id, "example.com") - .await - .expect("Failed to enqueue welcome comms"); - let row = sqlx::query!( + .expect("Failed to enqueue comms"); + let row = sqlx::query( r#" - SELECT - recipient, subject, body, - comms_type as "comms_type: CommsType" + SELECT id, user_id, recipient, subject, body, channel, comms_type, status FROM comms_queue WHERE id = $1 "#, - comms_id ) + .bind(comms_id) .fetch_one(pool) .await .expect("Comms not found"); - assert_eq!(Some(row.recipient), user_row.email); - assert_eq!(row.subject.as_deref(), Some("Welcome to example.com")); - assert!(row.body.contains(&format!("@{}", user_row.handle))); - assert_eq!(row.comms_type, CommsType::Welcome); + let row_user_id: uuid::Uuid = row.get("user_id"); + let row_recipient: String = row.get("recipient"); + let row_subject: Option = row.get("subject"); + let row_body: String = row.get("body"); + let row_channel: CommsChannel = row.get("channel"); + let row_comms_type: CommsType = row.get("comms_type"); + let row_status: CommsStatus = row.get("status"); + assert_eq!(row_user_id, user_id); + assert_eq!(row_recipient, "test@example.com"); + assert_eq!(row_subject.as_deref(), Some("Test Subject")); + assert_eq!(row_body, "Test body"); + assert_eq!(row_channel, CommsChannel::Email); + assert_eq!(row_comms_type, CommsType::Welcome); + assert_eq!(row_status, CommsStatus::Pending); } #[tokio::test] async fn test_comms_queue_status_index() { let pool = common::get_test_db_pool().await; let (_, did) = common::create_account_and_login(&common::client()).await; - let user_id: uuid::Uuid = sqlx::query_scalar!("SELECT id FROM users WHERE did = $1", did) + let user_id: uuid::Uuid = sqlx::query_scalar("SELECT id FROM users WHERE did = $1") + .bind(&did) .fetch_one(pool) .await .expect("User not found"); - let initial_count: i64 = sqlx::query_scalar!( + let initial_count: i64 = sqlx::query_scalar( "SELECT COUNT(*) FROM comms_queue WHERE status = 'pending' AND user_id = $1", - user_id ) + .bind(user_id) .fetch_one(pool) .await - .expect("Failed to count") - .unwrap_or(0); - for i in 0..5 { - let item = NewComms::email( - user_id, - CommsType::PasswordReset, - format!("test{}@example.com", i), - "Test".to_string(), - "Body".to_string(), - ); - enqueue_comms(pool, item).await.expect("Failed to enqueue"); - } - let final_count: i64 = sqlx::query_scalar!( + .expect("Failed to count"); + let inserts = (0..5).map(|i| { + sqlx::query( + r#"INSERT INTO comms_queue (user_id, channel, comms_type, recipient, subject, body) + VALUES ($1, 'email', 'password_reset', $2, $3, $4)"#, + ) + .bind(user_id) + .bind(format!("test{}@example.com", i)) + .bind("Test") + .bind("Body") + .execute(pool) + }); + futures::future::try_join_all(inserts) + .await + .expect("Failed to enqueue"); + let final_count: i64 = sqlx::query_scalar( "SELECT COUNT(*) FROM comms_queue WHERE status = 'pending' AND user_id = $1", - user_id ) + .bind(user_id) .fetch_one(pool) .await - .expect("Failed to count") - .unwrap_or(0); + .expect("Failed to count"); assert_eq!(final_count - initial_count, 5); } diff --git a/crates/tranquil-pds/tests/sync_conformance.rs b/crates/tranquil-pds/tests/sync_conformance.rs index a1164a5..c3368a2 100644 --- a/crates/tranquil-pds/tests/sync_conformance.rs +++ b/crates/tranquil-pds/tests/sync_conformance.rs @@ -352,6 +352,8 @@ async fn test_get_repo_since_returns_partial() { let initial_body: Value = initial_commit_res.json().await.unwrap(); let initial_rev = initial_body["rev"].as_str().unwrap(); + create_post(&client, &did, &jwt, "Test post for since param").await; + let full_repo_res = client .get(format!( "{}/xrpc/com.atproto.sync.getRepo", @@ -365,8 +367,6 @@ async fn test_get_repo_since_returns_partial() { let full_repo_bytes = full_repo_res.bytes().await.unwrap(); let full_repo_size = full_repo_bytes.len(); - create_post(&client, &did, &jwt, "Test post for since param").await; - let partial_repo_res = client .get(format!( "{}/xrpc/com.atproto.sync.getRepo", diff --git a/crates/tranquil-types/src/lib.rs b/crates/tranquil-types/src/lib.rs index 67991ca..ee877e9 100644 --- a/crates/tranquil-types/src/lib.rs +++ b/crates/tranquil-types/src/lib.rs @@ -639,6 +639,24 @@ impl AtUri { pub fn into_inner(self) -> String { self.0 } + + pub fn did(&self) -> Option<&str> { + self.0 + .strip_prefix("at://") + .and_then(|s| s.split('/').next()) + } + + pub fn collection(&self) -> Option<&str> { + self.0 + .strip_prefix("at://") + .and_then(|s| s.split('/').nth(1)) + } + + pub fn rkey(&self) -> Option<&str> { + self.0 + .strip_prefix("at://") + .and_then(|s| s.split('/').nth(2)) + } } impl AsRef for AtUri { @@ -1439,6 +1457,366 @@ impl From for DPoPProofId { } } +#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize, sqlx::Type)] +#[serde(transparent)] +#[sqlx(transparent)] +pub struct TokenId(String); + +impl TokenId { + pub fn new(s: impl Into) -> Self { + Self(s.into()) + } + + pub fn as_str(&self) -> &str { + &self.0 + } + + pub fn into_inner(self) -> String { + self.0 + } +} + +impl AsRef for TokenId { + fn as_ref(&self) -> &str { + &self.0 + } +} + +impl Deref for TokenId { + type Target = str; + + fn deref(&self) -> &Self::Target { + &self.0 + } +} + +impl From for TokenId { + fn from(s: String) -> Self { + Self(s) + } +} + +impl fmt::Display for TokenId { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + write!(f, "{}", self.0) + } +} + +#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize, sqlx::Type)] +#[serde(transparent)] +#[sqlx(transparent)] +pub struct ClientId(String); + +impl ClientId { + pub fn new(s: impl Into) -> Self { + Self(s.into()) + } + + pub fn as_str(&self) -> &str { + &self.0 + } + + pub fn into_inner(self) -> String { + self.0 + } +} + +impl AsRef for ClientId { + fn as_ref(&self) -> &str { + &self.0 + } +} + +impl Deref for ClientId { + type Target = str; + + fn deref(&self) -> &Self::Target { + &self.0 + } +} + +impl From for ClientId { + fn from(s: String) -> Self { + Self(s) + } +} + +impl fmt::Display for ClientId { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + write!(f, "{}", self.0) + } +} + +#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize, sqlx::Type)] +#[serde(transparent)] +#[sqlx(transparent)] +pub struct DeviceId(String); + +impl DeviceId { + pub fn new(s: impl Into) -> Self { + Self(s.into()) + } + + pub fn as_str(&self) -> &str { + &self.0 + } + + pub fn into_inner(self) -> String { + self.0 + } +} + +impl AsRef for DeviceId { + fn as_ref(&self) -> &str { + &self.0 + } +} + +impl Deref for DeviceId { + type Target = str; + + fn deref(&self) -> &Self::Target { + &self.0 + } +} + +impl From for DeviceId { + fn from(s: String) -> Self { + Self(s) + } +} + +impl fmt::Display for DeviceId { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + write!(f, "{}", self.0) + } +} + +#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize, sqlx::Type)] +#[serde(transparent)] +#[sqlx(transparent)] +pub struct RequestId(String); + +impl RequestId { + pub fn new(s: impl Into) -> Self { + Self(s.into()) + } + + pub fn as_str(&self) -> &str { + &self.0 + } + + pub fn into_inner(self) -> String { + self.0 + } +} + +impl AsRef for RequestId { + fn as_ref(&self) -> &str { + &self.0 + } +} + +impl Deref for RequestId { + type Target = str; + + fn deref(&self) -> &Self::Target { + &self.0 + } +} + +impl From for RequestId { + fn from(s: String) -> Self { + Self(s) + } +} + +impl fmt::Display for RequestId { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + write!(f, "{}", self.0) + } +} + +#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize, sqlx::Type)] +#[serde(transparent)] +#[sqlx(transparent)] +pub struct Jti(String); + +impl Jti { + pub fn new(s: impl Into) -> Self { + Self(s.into()) + } + + pub fn as_str(&self) -> &str { + &self.0 + } + + pub fn into_inner(self) -> String { + self.0 + } +} + +impl AsRef for Jti { + fn as_ref(&self) -> &str { + &self.0 + } +} + +impl Deref for Jti { + type Target = str; + + fn deref(&self) -> &Self::Target { + &self.0 + } +} + +impl From for Jti { + fn from(s: String) -> Self { + Self(s) + } +} + +impl fmt::Display for Jti { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + write!(f, "{}", self.0) + } +} + +#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize, sqlx::Type)] +#[serde(transparent)] +#[sqlx(transparent)] +pub struct AuthorizationCode(String); + +impl AuthorizationCode { + pub fn new(s: impl Into) -> Self { + Self(s.into()) + } + + pub fn as_str(&self) -> &str { + &self.0 + } + + pub fn into_inner(self) -> String { + self.0 + } +} + +impl AsRef for AuthorizationCode { + fn as_ref(&self) -> &str { + &self.0 + } +} + +impl Deref for AuthorizationCode { + type Target = str; + + fn deref(&self) -> &Self::Target { + &self.0 + } +} + +impl From for AuthorizationCode { + fn from(s: String) -> Self { + Self(s) + } +} + +impl fmt::Display for AuthorizationCode { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + write!(f, "{}", self.0) + } +} + +#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize, sqlx::Type)] +#[serde(transparent)] +#[sqlx(transparent)] +pub struct RefreshToken(String); + +impl RefreshToken { + pub fn new(s: impl Into) -> Self { + Self(s.into()) + } + + pub fn as_str(&self) -> &str { + &self.0 + } + + pub fn into_inner(self) -> String { + self.0 + } +} + +impl AsRef for RefreshToken { + fn as_ref(&self) -> &str { + &self.0 + } +} + +impl Deref for RefreshToken { + type Target = str; + + fn deref(&self) -> &Self::Target { + &self.0 + } +} + +impl From for RefreshToken { + fn from(s: String) -> Self { + Self(s) + } +} + +impl fmt::Display for RefreshToken { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + write!(f, "{}", self.0) + } +} + +#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize, sqlx::Type)] +#[serde(transparent)] +#[sqlx(transparent)] +pub struct InviteCode(String); + +impl InviteCode { + pub fn new(s: impl Into) -> Self { + Self(s.into()) + } + + pub fn as_str(&self) -> &str { + &self.0 + } + + pub fn into_inner(self) -> String { + self.0 + } +} + +impl AsRef for InviteCode { + fn as_ref(&self) -> &str { + &self.0 + } +} + +impl Deref for InviteCode { + type Target = str; + + fn deref(&self) -> &Self::Target { + &self.0 + } +} + +impl From for InviteCode { + fn from(s: String) -> Self { + Self(s) + } +} + +impl fmt::Display for InviteCode { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + write!(f, "{}", self.0) + } +} + #[cfg(test)] mod tests { use super::*; @@ -1486,6 +1864,7 @@ mod tests { assert!(Handle::new("user.bsky.social").is_ok()); assert!(Handle::new("test.example.com").is_ok()); assert!(Handle::new("invalid handle with spaces").is_err()); + assert!(Handle::new("alice.pds.test").is_ok()); } #[test]