diff --git a/.sqlx/query-13bea39e403ee15f13f877654c6677f7f2ad541edf72324231801ffead506031.json b/.sqlx/query-13bea39e403ee15f13f877654c6677f7f2ad541edf72324231801ffead506031.json new file mode 100644 index 0000000..bacee26 --- /dev/null +++ b/.sqlx/query-13bea39e403ee15f13f877654c6677f7f2ad541edf72324231801ffead506031.json @@ -0,0 +1,64 @@ +{ + "db_name": "PostgreSQL", + "query": "\n SELECT id, did, handle, email, created_at, invites_disabled, email_verified, deactivated_at\n FROM users\n WHERE did = $1\n ", + "describe": { + "columns": [ + { + "ordinal": 0, + "name": "id", + "type_info": "Uuid" + }, + { + "ordinal": 1, + "name": "did", + "type_info": "Text" + }, + { + "ordinal": 2, + "name": "handle", + "type_info": "Text" + }, + { + "ordinal": 3, + "name": "email", + "type_info": "Text" + }, + { + "ordinal": 4, + "name": "created_at", + "type_info": "Timestamptz" + }, + { + "ordinal": 5, + "name": "invites_disabled", + "type_info": "Bool" + }, + { + "ordinal": 6, + "name": "email_verified", + "type_info": "Bool" + }, + { + "ordinal": 7, + "name": "deactivated_at", + "type_info": "Timestamptz" + } + ], + "parameters": { + "Left": [ + "Text" + ] + }, + "nullable": [ + false, + false, + false, + true, + false, + true, + false, + true + ] + }, + "hash": "13bea39e403ee15f13f877654c6677f7f2ad541edf72324231801ffead506031" +} diff --git a/.sqlx/query-176d30f31356a4d128764c9c2eece81f8079a29e40b07ba58adc4380d58068c8.json b/.sqlx/query-176d30f31356a4d128764c9c2eece81f8079a29e40b07ba58adc4380d58068c8.json deleted file mode 100644 index e9b7308..0000000 --- a/.sqlx/query-176d30f31356a4d128764c9c2eece81f8079a29e40b07ba58adc4380d58068c8.json +++ /dev/null @@ -1,40 +0,0 @@ -{ - "db_name": "PostgreSQL", - "query": "\n SELECT did, handle, email, created_at\n FROM users\n WHERE did = $1\n ", - "describe": { - "columns": [ - { - "ordinal": 0, - "name": "did", - "type_info": "Text" - }, - { - "ordinal": 1, - "name": "handle", - "type_info": "Text" - }, - { - "ordinal": 2, - "name": "email", - "type_info": "Text" - }, - { - "ordinal": 3, - "name": "created_at", - "type_info": "Timestamptz" - } - ], - "parameters": { - "Left": [ - "Text" - ] - }, - "nullable": [ - false, - false, - true, - false - ] - }, - "hash": "176d30f31356a4d128764c9c2eece81f8079a29e40b07ba58adc4380d58068c8" -} diff --git a/.sqlx/query-1e034c36940110579d5ba3e6f64b4455a4945b4116dbd561e12269cf1df495b3.json b/.sqlx/query-1e034c36940110579d5ba3e6f64b4455a4945b4116dbd561e12269cf1df495b3.json new file mode 100644 index 0000000..efe5fb4 --- /dev/null +++ b/.sqlx/query-1e034c36940110579d5ba3e6f64b4455a4945b4116dbd561e12269cf1df495b3.json @@ -0,0 +1,22 @@ +{ + "db_name": "PostgreSQL", + "query": "\n SELECT icu.code\n FROM invite_code_uses icu\n WHERE icu.used_by_user = $1\n LIMIT 1\n ", + "describe": { + "columns": [ + { + "ordinal": 0, + "name": "code", + "type_info": "Text" + } + ], + "parameters": { + "Left": [ + "Uuid" + ] + }, + "nullable": [ + false + ] + }, + "hash": "1e034c36940110579d5ba3e6f64b4455a4945b4116dbd561e12269cf1df495b3" +} diff --git a/.sqlx/query-5a98e015997942835800fcd326e69b4f54b9830d0490c4f8841f8435478c57d3.json b/.sqlx/query-5a98e015997942835800fcd326e69b4f54b9830d0490c4f8841f8435478c57d3.json new file mode 100644 index 0000000..5b502fd --- /dev/null +++ b/.sqlx/query-5a98e015997942835800fcd326e69b4f54b9830d0490c4f8841f8435478c57d3.json @@ -0,0 +1,22 @@ +{ + "db_name": "PostgreSQL", + "query": "\n SELECT code FROM invite_codes WHERE created_by_user = $1\n ", + "describe": { + "columns": [ + { + "ordinal": 0, + "name": "code", + "type_info": "Text" + } + ], + "parameters": { + "Left": [ + "Uuid" + ] + }, + "nullable": [ + false + ] + }, + "hash": "5a98e015997942835800fcd326e69b4f54b9830d0490c4f8841f8435478c57d3" +} diff --git a/.sqlx/query-6df413951ea7648c77d8db2fe6e704370869816a3f47c86671dfe000b5961eee.json b/.sqlx/query-6df413951ea7648c77d8db2fe6e704370869816a3f47c86671dfe000b5961eee.json new file mode 100644 index 0000000..2212863 --- /dev/null +++ b/.sqlx/query-6df413951ea7648c77d8db2fe6e704370869816a3f47c86671dfe000b5961eee.json @@ -0,0 +1,64 @@ +{ + "db_name": "PostgreSQL", + "query": "\n SELECT id, did, handle, email, created_at, invites_disabled, email_verified, deactivated_at\n FROM users\n WHERE did = $1\n ", + "describe": { + "columns": [ + { + "ordinal": 0, + "name": "id", + "type_info": "Uuid" + }, + { + "ordinal": 1, + "name": "did", + "type_info": "Text" + }, + { + "ordinal": 2, + "name": "handle", + "type_info": "Text" + }, + { + "ordinal": 3, + "name": "email", + "type_info": "Text" + }, + { + "ordinal": 4, + "name": "created_at", + "type_info": "Timestamptz" + }, + { + "ordinal": 5, + "name": "invites_disabled", + "type_info": "Bool" + }, + { + "ordinal": 6, + "name": "email_verified", + "type_info": "Bool" + }, + { + "ordinal": 7, + "name": "deactivated_at", + "type_info": "Timestamptz" + } + ], + "parameters": { + "Left": [ + "Text" + ] + }, + "nullable": [ + false, + false, + false, + true, + false, + true, + false, + true + ] + }, + "hash": "6df413951ea7648c77d8db2fe6e704370869816a3f47c86671dfe000b5961eee" +} diff --git a/.sqlx/query-c2a90157c47bf1c36f08f4608932d214cc26b4794e0b922b1dae3dad18a7ddc0.json b/.sqlx/query-c2a90157c47bf1c36f08f4608932d214cc26b4794e0b922b1dae3dad18a7ddc0.json deleted file mode 100644 index 813ef1a..0000000 --- a/.sqlx/query-c2a90157c47bf1c36f08f4608932d214cc26b4794e0b922b1dae3dad18a7ddc0.json +++ /dev/null @@ -1,40 +0,0 @@ -{ - "db_name": "PostgreSQL", - "query": "\n SELECT did, handle, email, created_at\n FROM users\n WHERE did = $1\n ", - "describe": { - "columns": [ - { - "ordinal": 0, - "name": "did", - "type_info": "Text" - }, - { - "ordinal": 1, - "name": "handle", - "type_info": "Text" - }, - { - "ordinal": 2, - "name": "email", - "type_info": "Text" - }, - { - "ordinal": 3, - "name": "created_at", - "type_info": "Timestamptz" - } - ], - "parameters": { - "Left": [ - "Text" - ] - }, - "nullable": [ - false, - false, - true, - false - ] - }, - "hash": "c2a90157c47bf1c36f08f4608932d214cc26b4794e0b922b1dae3dad18a7ddc0" -} diff --git a/.sqlx/query-c3139484bba403cd256801e278fe95ae77634e79d14764dd8c3764886cf08eac.json b/.sqlx/query-c3139484bba403cd256801e278fe95ae77634e79d14764dd8c3764886cf08eac.json new file mode 100644 index 0000000..479eee6 --- /dev/null +++ b/.sqlx/query-c3139484bba403cd256801e278fe95ae77634e79d14764dd8c3764886cf08eac.json @@ -0,0 +1,52 @@ +{ + "db_name": "PostgreSQL", + "query": "\n SELECT ic.code, ic.available_uses, ic.disabled, ic.for_account, ic.created_at, u.did as created_by\n FROM invite_codes ic\n JOIN users u ON ic.created_by_user = u.id\n WHERE ic.code = $1\n ", + "describe": { + "columns": [ + { + "ordinal": 0, + "name": "code", + "type_info": "Text" + }, + { + "ordinal": 1, + "name": "available_uses", + "type_info": "Int4" + }, + { + "ordinal": 2, + "name": "disabled", + "type_info": "Bool" + }, + { + "ordinal": 3, + "name": "for_account", + "type_info": "Text" + }, + { + "ordinal": 4, + "name": "created_at", + "type_info": "Timestamptz" + }, + { + "ordinal": 5, + "name": "created_by", + "type_info": "Text" + } + ], + "parameters": { + "Left": [ + "Text" + ] + }, + "nullable": [ + false, + false, + true, + false, + false, + false + ] + }, + "hash": "c3139484bba403cd256801e278fe95ae77634e79d14764dd8c3764886cf08eac" +} diff --git a/.sqlx/query-c9f3d584c161b6492abc082bdbb563d40173a9a4983d6454dba4e02f7e0f8458.json b/.sqlx/query-c9f3d584c161b6492abc082bdbb563d40173a9a4983d6454dba4e02f7e0f8458.json new file mode 100644 index 0000000..3771c48 --- /dev/null +++ b/.sqlx/query-c9f3d584c161b6492abc082bdbb563d40173a9a4983d6454dba4e02f7e0f8458.json @@ -0,0 +1,28 @@ +{ + "db_name": "PostgreSQL", + "query": "\n SELECT u.did as used_by, icu.used_at\n FROM invite_code_uses icu\n JOIN users u ON icu.used_by_user = u.id\n WHERE icu.code = $1\n ", + "describe": { + "columns": [ + { + "ordinal": 0, + "name": "used_by", + "type_info": "Text" + }, + { + "ordinal": 1, + "name": "used_at", + "type_info": "Timestamptz" + } + ], + "parameters": { + "Left": [ + "Text" + ] + }, + "nullable": [ + false, + false + ] + }, + "hash": "c9f3d584c161b6492abc082bdbb563d40173a9a4983d6454dba4e02f7e0f8458" +} diff --git a/scripts/test-infra.sh b/scripts/test-infra.sh index 9fcd79a..d317f13 100755 --- a/scripts/test-infra.sh +++ b/scripts/test-infra.sh @@ -41,7 +41,8 @@ start_infra() { -e POSTGRES_DB=postgres \ -P \ --label tranquil_pds_test=true \ - postgres:18-alpine >/dev/null + postgres:18-alpine \ + -c max_connections=500 >/dev/null echo "Starting MinIO..." $CONTAINER_CMD run -d \ --name "${CONTAINER_PREFIX}-minio" \ diff --git a/src/api/admin/account/info.rs b/src/api/admin/account/info.rs index 2c45baf..801ac05 100644 --- a/src/api/admin/account/info.rs +++ b/src/api/admin/account/info.rs @@ -20,12 +20,39 @@ pub struct GetAccountInfoParams { pub struct AccountInfo { pub did: String, pub handle: String, + #[serde(skip_serializing_if = "Option::is_none")] pub email: Option, pub indexed_at: String, + #[serde(skip_serializing_if = "Option::is_none")] pub invite_note: Option, pub invites_disabled: bool, - pub email_verified_at: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub email_confirmed_at: Option, + #[serde(skip_serializing_if = "Option::is_none")] pub deactivated_at: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub invited_by: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub invites: Option>, +} + +#[derive(Serialize, Clone)] +#[serde(rename_all = "camelCase")] +pub struct InviteCodeInfo { + pub code: String, + pub available: i32, + pub disabled: bool, + pub for_account: String, + pub created_by: String, + pub created_at: String, + pub uses: Vec, +} + +#[derive(Serialize, Clone)] +#[serde(rename_all = "camelCase")] +pub struct InviteCodeUseInfo { + pub used_by: String, + pub used_at: String, } #[derive(Serialize)] @@ -49,7 +76,7 @@ pub async fn get_account_info( } let result = sqlx::query!( r#" - SELECT did, handle, email, created_at + SELECT id, did, handle, email, created_at, invites_disabled, email_verified, deactivated_at FROM users WHERE did = $1 "#, @@ -58,20 +85,30 @@ pub async fn get_account_info( .fetch_optional(&state.db) .await; match result { - Ok(Some(row)) => ( - StatusCode::OK, - Json(AccountInfo { - did: row.did, - handle: row.handle, - email: row.email, - indexed_at: row.created_at.to_rfc3339(), - invite_note: None, - invites_disabled: false, - email_verified_at: None, - deactivated_at: None, - }), - ) - .into_response(), + 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, + handle: row.handle, + 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) => ( StatusCode::NOT_FOUND, Json(json!({"error": "AccountNotFound", "message": "Account not found"})), @@ -88,6 +125,96 @@ pub async fn get_account_info( } } +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 + ) + .fetch_optional(db) + .await + .ok()??; + get_invite_code_info(db, &use_row.code).await +} + +async fn get_invites_for_user( + db: &sqlx::PgPool, + user_id: uuid::Uuid, +) -> Option> { + let codes = sqlx::query_scalar!( + r#" + SELECT code FROM invite_codes WHERE created_by_user = $1 + "#, + user_id + ) + .fetch_all(db) + .await + .ok()?; + if codes.is_empty() { + return None; + } + let mut invites = Vec::new(); + for code in codes { + if let Some(info) = get_invite_code_info(db, &code).await { + invites.push(info); + } + } + if invites.is_empty() { + None + } else { + Some(invites) + } +} + +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()?; + Some(InviteCodeInfo { + code: row.code, + available: row.available_uses, + disabled: row.disabled.unwrap_or(false), + for_account: row.for_account, + created_by: row.created_by, + created_at: row.created_at.to_rfc3339(), + uses: uses + .into_iter() + .map(|u| InviteCodeUseInfo { + used_by: u.used_by, + used_at: u.used_at.to_rfc3339(), + }) + .collect(), + }) +} + pub async fn get_account_infos( State(state): State, _auth: BearerAuthAdmin, @@ -108,7 +235,7 @@ pub async fn get_account_infos( } let result = sqlx::query!( r#" - SELECT did, handle, email, created_at + SELECT id, did, handle, email, created_at, invites_disabled, email_verified, deactivated_at FROM users WHERE did = $1 "#, @@ -117,15 +244,23 @@ pub async fn get_account_infos( .fetch_optional(&state.db) .await; if let Ok(Some(row)) = result { + let invited_by = get_invited_by(&state.db, row.id).await; + let invites = get_invites_for_user(&state.db, row.id).await; infos.push(AccountInfo { did: row.did, handle: row.handle, email: row.email, indexed_at: row.created_at.to_rfc3339(), invite_note: None, - invites_disabled: false, - email_verified_at: None, - deactivated_at: 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, }); } } diff --git a/src/api/admin/account/search.rs b/src/api/admin/account/search.rs index 295a2c4..5bd862c 100644 --- a/src/api/admin/account/search.rs +++ b/src/api/admin/account/search.rs @@ -12,6 +12,7 @@ use tracing::error; #[derive(Deserialize)] pub struct SearchAccountsParams { + pub email: Option, pub handle: Option, pub cursor: Option, #[serde(default = "default_limit")] @@ -31,7 +32,7 @@ pub struct AccountView { pub email: Option, pub indexed_at: String, #[serde(skip_serializing_if = "Option::is_none")] - pub email_verified_at: Option, + pub email_confirmed_at: Option, #[serde(skip_serializing_if = "Option::is_none")] pub deactivated_at: Option, #[serde(skip_serializing_if = "Option::is_none")] @@ -53,6 +54,7 @@ pub async fn search_accounts( ) -> 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::< _, @@ -63,17 +65,21 @@ pub async fn search_accounts( chrono::DateTime, bool, Option>, + Option, ), >( r#" - SELECT did, handle, email, created_at, email_verified, deactivated_at + SELECT did, handle, email, created_at, email_verified, deactivated_at, invites_disabled FROM users - WHERE did > $1 AND ($2::text IS NULL OR handle ILIKE $2) + 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 $3 + LIMIT $4 "#, ) .bind(cursor_did) + .bind(&email_filter) .bind(&handle_filter) .bind(limit + 1) .fetch_all(&state.db) @@ -85,19 +91,19 @@ pub async fn search_accounts( .into_iter() .take(limit as usize) .map( - |(did, handle, email, created_at, email_verified, deactivated_at)| { + |(did, handle, email, created_at, email_verified, deactivated_at, invites_disabled)| { AccountView { did: did.clone(), handle, email, indexed_at: created_at.to_rfc3339(), - email_verified_at: if email_verified { + email_confirmed_at: if email_verified { Some(created_at.to_rfc3339()) } else { None }, deactivated_at: deactivated_at.map(|dt| dt.to_rfc3339()), - invites_disabled: None, + invites_disabled, } }, ) diff --git a/src/api/admin/account/update.rs b/src/api/admin/account/update.rs index 223d0b1..6759af3 100644 --- a/src/api/admin/account/update.rs +++ b/src/api/admin/account/update.rs @@ -8,7 +8,7 @@ use axum::{ }; use serde::Deserialize; use serde_json::json; -use tracing::error; +use tracing::{error, warn}; #[derive(Deserialize)] pub struct UpdateAccountEmailInput { @@ -128,6 +128,15 @@ pub async fn update_account_handle( let _ = state.cache.delete(&format!("handle:{}", old)).await; } let _ = state.cache.delete(&format!("handle:{}", handle)).await; + if let Err(e) = + crate::api::repo::record::sequence_identity_event(&state, did, Some(&handle)).await + { + warn!("Failed to sequence identity event for admin handle update: {}", e); + } + if let Err(e) = crate::api::identity::did::update_plc_handle(&state, did, &handle).await + { + warn!("Failed to update PLC handle for admin handle update: {}", e); + } (StatusCode::OK, Json(json!({}))).into_response() } Err(e) => { diff --git a/src/api/admin/status.rs b/src/api/admin/status.rs index 6b1ded0..b414cc7 100644 --- a/src/api/admin/status.rs +++ b/src/api/admin/status.rs @@ -135,6 +135,16 @@ pub async fn get_subject_status( } } if let Some(blob_cid) = ¶ms.blob { + let did = match ¶ms.did { + Some(d) => d, + None => { + return ( + StatusCode::BAD_REQUEST, + Json(json!({"error": "InvalidRequest", "message": "Must provide a did to request blob state"})), + ) + .into_response(); + } + }; let blob = sqlx::query!( "SELECT cid, takedown_ref FROM blobs WHERE cid = $1", blob_cid @@ -152,7 +162,7 @@ pub async fn get_subject_status( Json(SubjectStatus { subject: json!({ "$type": "com.atproto.admin.defs#repoBlobRef", - "did": "", + "did": did, "cid": row.cid }), takedown, @@ -195,7 +205,7 @@ pub struct UpdateSubjectStatusInput { #[derive(Deserialize)] pub struct StatusAttrInput { - pub apply: bool, + pub applied: bool, pub r#ref: Option, } @@ -221,7 +231,7 @@ pub async fn update_subject_status( } }; if let Some(takedown) = &input.takedown { - let takedown_ref = if takedown.apply { + let takedown_ref = if takedown.applied { takedown.r#ref.clone() } else { None @@ -243,7 +253,7 @@ pub async fn update_subject_status( } } if let Some(deactivated) = &input.deactivated { - let result = if deactivated.apply { + let result = if deactivated.applied { sqlx::query!( "UPDATE users SET deactivated_at = NOW() WHERE did = $1", did @@ -276,7 +286,7 @@ pub async fn update_subject_status( .into_response(); } if let Some(takedown) = &input.takedown { - let status = if takedown.apply { + let status = if takedown.applied { Some("takendown") } else { None @@ -284,7 +294,7 @@ pub async fn update_subject_status( if let Err(e) = crate::api::repo::record::sequence_account_event( &state, did, - !takedown.apply, + !takedown.applied, status, ) .await @@ -293,7 +303,7 @@ pub async fn update_subject_status( } } if let Some(deactivated) = &input.deactivated { - let status = if deactivated.apply { + let status = if deactivated.applied { Some("deactivated") } else { None @@ -301,7 +311,7 @@ pub async fn update_subject_status( if let Err(e) = crate::api::repo::record::sequence_account_event( &state, did, - !deactivated.apply, + !deactivated.applied, status, ) .await @@ -321,11 +331,11 @@ pub async fn update_subject_status( Json(json!({ "subject": input.subject, "takedown": input.takedown.as_ref().map(|t| json!({ - "applied": t.apply, + "applied": t.applied, "ref": t.r#ref })), "deactivated": input.deactivated.as_ref().map(|d| json!({ - "applied": d.apply + "applied": d.applied })) })), ) @@ -336,7 +346,7 @@ pub async fn update_subject_status( let uri = input.subject.get("uri").and_then(|u| u.as_str()); if let Some(uri) = uri { if let Some(takedown) = &input.takedown { - let takedown_ref = if takedown.apply { + let takedown_ref = if takedown.applied { takedown.r#ref.clone() } else { None @@ -365,7 +375,7 @@ pub async fn update_subject_status( Json(json!({ "subject": input.subject, "takedown": input.takedown.as_ref().map(|t| json!({ - "applied": t.apply, + "applied": t.applied, "ref": t.r#ref })) })), @@ -377,7 +387,7 @@ pub async fn update_subject_status( let cid = input.subject.get("cid").and_then(|c| c.as_str()); if let Some(cid) = cid { if let Some(takedown) = &input.takedown { - let takedown_ref = if takedown.apply { + let takedown_ref = if takedown.applied { takedown.r#ref.clone() } else { None @@ -403,7 +413,7 @@ pub async fn update_subject_status( Json(json!({ "subject": input.subject, "takedown": input.takedown.as_ref().map(|t| json!({ - "applied": t.apply, + "applied": t.applied, "ref": t.r#ref })) })), diff --git a/src/api/identity/did.rs b/src/api/identity/did.rs index 4e5c7ca..7d12505 100644 --- a/src/api/identity/did.rs +++ b/src/api/identity/did.rs @@ -780,7 +780,7 @@ pub async fn update_handle( } } -async fn update_plc_handle( +pub async fn update_plc_handle( state: &AppState, did: &str, new_handle: &str, diff --git a/tests/account_notifications.rs b/tests/account_notifications.rs index a8cb906..1c7a599 100644 --- a/tests/account_notifications.rs +++ b/tests/account_notifications.rs @@ -1,27 +1,17 @@ mod common; -use common::{base_url, client, create_account_and_login, get_db_connection_string}; +use common::{base_url, client, create_account_and_login, get_test_db_pool}; use serde_json::{Value, json}; -use sqlx::PgPool; use tranquil_pds::comms::{CommsType, NewComms, enqueue_comms}; -async fn get_pool() -> PgPool { - let conn_str = get_db_connection_string().await; - sqlx::postgres::PgPoolOptions::new() - .max_connections(5) - .connect(&conn_str) - .await - .expect("Failed to connect to test database") -} - #[tokio::test] async fn test_get_notification_history() { let client = client(); let base = base_url().await; - let pool = get_pool().await; + 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) - .fetch_one(&pool) + .fetch_one(pool) .await .expect("User not found"); @@ -33,7 +23,7 @@ async fn test_get_notification_history() { format!("Subject {}", i), format!("Body {}", i), ); - enqueue_comms(&pool, comms) + enqueue_comms(pool, comms) .await .expect("Failed to enqueue"); } @@ -86,9 +76,9 @@ async fn test_verify_channel_discord() { .contains(&json!("discord")) ); - let pool = get_pool().await; + let pool = get_test_db_pool().await; let user_id: uuid::Uuid = sqlx::query_scalar!("SELECT id FROM users WHERE did = $1", did) - .fetch_one(&pool) + .fetch_one(pool) .await .expect("User not found"); @@ -96,7 +86,7 @@ async fn test_verify_channel_discord() { "SELECT body, metadata FROM comms_queue WHERE user_id = $1 AND comms_type = 'channel_verification' ORDER BY created_at DESC LIMIT 1", user_id ) - .fetch_one(&pool) + .fetch_one(pool) .await .expect("Verification code not found"); @@ -213,7 +203,7 @@ async fn test_verify_channel_not_set() { async fn test_update_email_via_notification_prefs() { let client = client(); let base = base_url().await; - let pool = get_pool().await; + let pool = get_test_db_pool().await; let (token, did) = create_account_and_login(&client).await; let unique_email = format!("newemail_{}@example.com", uuid::Uuid::new_v4()); @@ -240,7 +230,7 @@ async fn test_update_email_via_notification_prefs() { ); let user_id: uuid::Uuid = sqlx::query_scalar!("SELECT id FROM users WHERE did = $1", did) - .fetch_one(&pool) + .fetch_one(pool) .await .expect("User not found"); @@ -248,7 +238,7 @@ async fn test_update_email_via_notification_prefs() { "SELECT body FROM comms_queue WHERE user_id = $1 AND comms_type = 'email_update' ORDER BY created_at DESC LIMIT 1", user_id ) - .fetch_one(&pool) + .fetch_one(pool) .await .expect("Verification code not found"); diff --git a/tests/admin_email.rs b/tests/admin_email.rs index 5dd2809..001ba6b 100644 --- a/tests/admin_email.rs +++ b/tests/admin_email.rs @@ -2,22 +2,12 @@ mod common; use reqwest::StatusCode; use serde_json::{Value, json}; -use sqlx::PgPool; - -async fn get_pool() -> PgPool { - let conn_str = common::get_db_connection_string().await; - sqlx::postgres::PgPoolOptions::new() - .max_connections(5) - .connect(&conn_str) - .await - .expect("Failed to connect to test database") -} #[tokio::test] async fn test_send_email_success() { let client = common::client(); let base_url = common::base_url().await; - let pool = get_pool().await; + let pool = common::get_test_db_pool().await; let (access_jwt, did) = common::create_admin_account_and_login(&client).await; let res = client .post(format!("{}/xrpc/com.atproto.admin.sendEmail", base_url)) @@ -35,14 +25,14 @@ async fn test_send_email_success() { let body: Value = res.json().await.expect("Invalid JSON"); assert_eq!(body["sent"], true); let user = sqlx::query!("SELECT id FROM users WHERE did = $1", did) - .fetch_one(&pool) + .fetch_one(pool) .await .expect("User not found"); let notification = sqlx::query!( "SELECT subject, body, comms_type as \"comms_type: String\" FROM comms_queue WHERE user_id = $1 AND comms_type = 'admin_email' ORDER BY created_at DESC LIMIT 1", user.id ) - .fetch_one(&pool) + .fetch_one(pool) .await .expect("Notification not found"); assert_eq!(notification.subject.as_deref(), Some("Test Admin Email")); @@ -57,7 +47,7 @@ async fn test_send_email_success() { async fn test_send_email_default_subject() { let client = common::client(); let base_url = common::base_url().await; - let pool = get_pool().await; + let pool = common::get_test_db_pool().await; let (access_jwt, did) = common::create_admin_account_and_login(&client).await; let res = client .post(format!("{}/xrpc/com.atproto.admin.sendEmail", base_url)) @@ -74,14 +64,14 @@ async fn test_send_email_default_subject() { let body: Value = res.json().await.expect("Invalid JSON"); assert_eq!(body["sent"], true); let user = sqlx::query!("SELECT id FROM users WHERE did = $1", did) - .fetch_one(&pool) + .fetch_one(pool) .await .expect("User not found"); let notification = sqlx::query!( "SELECT subject FROM comms_queue WHERE user_id = $1 AND comms_type = 'admin_email' AND body = 'Email without subject' LIMIT 1", user.id ) - .fetch_one(&pool) + .fetch_one(pool) .await .expect("Notification not found"); assert!(notification.subject.is_some()); diff --git a/tests/admin_moderation.rs b/tests/admin_moderation.rs index 9d73379..ec29cb1 100644 --- a/tests/admin_moderation.rs +++ b/tests/admin_moderation.rs @@ -88,7 +88,7 @@ async fn test_update_subject_status_takedown_user() { "did": target_did }, "takedown": { - "apply": true, + "applied": true, "ref": "mod-action-123" } }); @@ -134,7 +134,7 @@ async fn test_update_subject_status_remove_takedown() { "did": target_did }, "takedown": { - "apply": true, + "applied": true, "ref": "mod-action-456" } }); @@ -153,7 +153,7 @@ async fn test_update_subject_status_remove_takedown() { "did": target_did }, "takedown": { - "apply": false + "applied": false } }); let res = client @@ -197,7 +197,7 @@ async fn test_update_subject_status_deactivate_user() { "did": target_did }, "deactivated": { - "apply": true + "applied": true } }); let res = client @@ -236,7 +236,7 @@ async fn test_update_subject_status_invalid_type() { "did": "did:plc:test" }, "takedown": { - "apply": true + "applied": true } }); let res = client @@ -263,7 +263,7 @@ async fn test_update_subject_status_no_auth() { "did": "did:plc:test" }, "takedown": { - "apply": true + "applied": true } }); let res = client diff --git a/tests/common/mod.rs b/tests/common/mod.rs index fb01e0c..a2a3c76 100644 --- a/tests/common/mod.rs +++ b/tests/common/mod.rs @@ -18,6 +18,7 @@ use wiremock::{Mock, MockServer, ResponseTemplate}; static SERVER_URL: OnceLock = OnceLock::new(); static APP_PORT: OnceLock = OnceLock::new(); static MOCK_APPVIEW: OnceLock = OnceLock::new(); +static TEST_DB_POOL: OnceLock = OnceLock::new(); #[cfg(not(feature = "external-infra"))] use testcontainers::core::ContainerPort; @@ -237,7 +238,8 @@ async fn setup_mock_appview(_mock_server: &MockServer) {} async fn spawn_app(database_url: String) -> String { use tranquil_pds::rate_limit::RateLimiters; let pool = PgPoolOptions::new() - .max_connections(50) + .max_connections(3) + .acquire_timeout(std::time::Duration::from_secs(30)) .connect(&database_url) .await .expect("Failed to connect to Postgres. Make sure the database is running."); @@ -245,6 +247,13 @@ async fn spawn_app(database_url: String) -> String { .run(&pool) .await .expect("Failed to run migrations"); + let test_pool = PgPoolOptions::new() + .max_connections(5) + .acquire_timeout(std::time::Duration::from_secs(30)) + .connect(&database_url) + .await + .expect("Failed to create test pool"); + TEST_DB_POOL.set(test_pool).ok(); let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); let addr = listener.local_addr().unwrap(); APP_PORT.set(addr.port()).ok(); @@ -291,19 +300,20 @@ pub async fn get_db_connection_string() -> String { } } +#[allow(dead_code)] +pub async fn get_test_db_pool() -> &'static sqlx::PgPool { + base_url().await; + TEST_DB_POOL.get().expect("TEST_DB_POOL not initialized") +} + #[allow(dead_code)] pub async fn verify_new_account(client: &Client, did: &str) -> String { - let conn_str = get_db_connection_string().await; - let pool = sqlx::postgres::PgPoolOptions::new() - .max_connections(2) - .connect(&conn_str) - .await - .expect("Failed to connect to test database"); + let pool = get_test_db_pool().await; let body_text: String = sqlx::query_scalar!( "SELECT body FROM comms_queue WHERE user_id = (SELECT id FROM users WHERE did = $1) AND comms_type = 'email_verification' ORDER BY created_at DESC LIMIT 1", did ) - .fetch_one(&pool) + .fetch_one(pool) .await .expect("Failed to get verification code"); @@ -454,15 +464,10 @@ async fn create_account_and_login_internal(client: &Client, make_admin: bool) -> if res.status() == StatusCode::OK { let body: Value = res.json().await.expect("Invalid JSON"); let did = body["did"].as_str().expect("No did").to_string(); - let conn_str = get_db_connection_string().await; - let pool = sqlx::postgres::PgPoolOptions::new() - .max_connections(2) - .connect(&conn_str) - .await - .expect("Failed to connect to test database"); + let pool = get_test_db_pool().await; if make_admin { sqlx::query!("UPDATE users SET is_admin = TRUE WHERE did = $1", &did) - .execute(&pool) + .execute(pool) .await .expect("Failed to mark user as admin"); } @@ -476,7 +481,7 @@ async fn create_account_and_login_internal(client: &Client, make_admin: bool) -> "SELECT body FROM comms_queue WHERE user_id = (SELECT id FROM users WHERE did = $1) AND comms_type = 'email_verification' ORDER BY created_at DESC LIMIT 1", &did ) - .fetch_one(&pool) + .fetch_one(pool) .await .expect("Failed to get verification from comms_queue"); let lines: Vec<&str> = body_text.lines().collect(); diff --git a/tests/delete_account.rs b/tests/delete_account.rs index 5b0b678..7132539 100644 --- a/tests/delete_account.rs +++ b/tests/delete_account.rs @@ -4,16 +4,6 @@ use chrono::Utc; use common::*; use reqwest::StatusCode; use serde_json::{Value, json}; -use sqlx::PgPool; - -async fn get_pool() -> PgPool { - let conn_str = get_db_connection_string().await; - sqlx::postgres::PgPoolOptions::new() - .max_connections(5) - .connect(&conn_str) - .await - .expect("Failed to connect to test database") -} async fn create_verified_account( client: &reqwest::Client, @@ -61,12 +51,12 @@ async fn test_delete_account_full_flow() { .await .expect("Failed to request account deletion"); assert_eq!(request_delete_res.status(), StatusCode::OK); - let pool = get_pool().await; + let pool = get_test_db_pool().await; let row = sqlx::query!( "SELECT token FROM account_deletion_requests WHERE did = $1", did ) - .fetch_one(&pool) + .fetch_one(pool) .await .expect("Failed to query deletion token"); let token = row.token; @@ -86,7 +76,7 @@ async fn test_delete_account_full_flow() { .expect("Failed to delete account"); assert_eq!(delete_res.status(), StatusCode::OK); let user_row = sqlx::query!("SELECT id FROM users WHERE did = $1", did) - .fetch_optional(&pool) + .fetch_optional(pool) .await .expect("Failed to query user"); assert!(user_row.is_none(), "User should be deleted from database"); @@ -118,12 +108,12 @@ async fn test_delete_account_wrong_password() { .await .expect("Failed to request account deletion"); assert_eq!(request_delete_res.status(), StatusCode::OK); - let pool = get_pool().await; + let pool = get_test_db_pool().await; let row = sqlx::query!( "SELECT token FROM account_deletion_requests WHERE did = $1", did ) - .fetch_one(&pool) + .fetch_one(pool) .await .expect("Failed to query deletion token"); let token = row.token; @@ -208,12 +198,12 @@ async fn test_delete_account_expired_token() { .await .expect("Failed to request account deletion"); assert_eq!(request_delete_res.status(), StatusCode::OK); - let pool = get_pool().await; + let pool = get_test_db_pool().await; let row = sqlx::query!( "SELECT token FROM account_deletion_requests WHERE did = $1", did ) - .fetch_one(&pool) + .fetch_one(pool) .await .expect("Failed to query deletion token"); let token = row.token; @@ -221,7 +211,7 @@ async fn test_delete_account_expired_token() { "UPDATE account_deletion_requests SET expires_at = NOW() - INTERVAL '1 hour' WHERE token = $1", token ) - .execute(&pool) + .execute(pool) .await .expect("Failed to expire token"); let delete_payload = json!({ @@ -267,12 +257,12 @@ async fn test_delete_account_token_mismatch() { .await .expect("Failed to request account deletion"); assert_eq!(request_delete_res.status(), StatusCode::OK); - let pool = get_pool().await; + let pool = get_test_db_pool().await; let row = sqlx::query!( "SELECT token FROM account_deletion_requests WHERE did = $1", did1 ) - .fetch_one(&pool) + .fetch_one(pool) .await .expect("Failed to query deletion token"); let token = row.token; @@ -328,12 +318,12 @@ async fn test_delete_account_with_app_password() { .await .expect("Failed to request account deletion"); assert_eq!(request_delete_res.status(), StatusCode::OK); - let pool = get_pool().await; + let pool = get_test_db_pool().await; let row = sqlx::query!( "SELECT token FROM account_deletion_requests WHERE did = $1", did ) - .fetch_one(&pool) + .fetch_one(pool) .await .expect("Failed to query deletion token"); let token = row.token; @@ -353,7 +343,7 @@ async fn test_delete_account_with_app_password() { .expect("Failed to delete account"); assert_eq!(delete_res.status(), StatusCode::OK); let user_row = sqlx::query!("SELECT id FROM users WHERE did = $1", did) - .fetch_optional(&pool) + .fetch_optional(pool) .await .expect("Failed to query user"); assert!(user_row.is_none(), "User should be deleted from database"); diff --git a/tests/email_update.rs b/tests/email_update.rs index d2ab608..1bfd5b6 100644 --- a/tests/email_update.rs +++ b/tests/email_update.rs @@ -3,15 +3,6 @@ use reqwest::StatusCode; use serde_json::{Value, json}; use sqlx::PgPool; -async fn get_pool() -> PgPool { - let conn_str = common::get_db_connection_string().await; - sqlx::postgres::PgPoolOptions::new() - .max_connections(5) - .connect(&conn_str) - .await - .expect("Failed to connect to test database") -} - async fn get_email_update_token(pool: &PgPool, did: &str) -> String { let body_text: String = sqlx::query_scalar!( "SELECT body FROM comms_queue WHERE user_id = (SELECT id FROM users WHERE did = $1) AND comms_type = 'email_update' ORDER BY created_at DESC LIMIT 1", @@ -88,7 +79,7 @@ async fn test_request_email_update_returns_token_required() { async fn test_update_email_flow_success() { let client = common::client(); let base_url = common::base_url().await; - let pool = get_pool().await; + let pool = common::get_test_db_pool().await; let handle = format!("emailup-{}", uuid::Uuid::new_v4()); let email = format!("{}@example.com", handle); let (access_jwt, did) = create_verified_account(&client, &base_url, &handle, &email).await; @@ -107,7 +98,7 @@ async fn test_update_email_flow_success() { let body: Value = res.json().await.expect("Invalid JSON"); assert_eq!(body["tokenRequired"], true); - let code = get_email_update_token(&pool, &did).await; + let code = get_email_update_token(pool, &did).await; let res = client .post(format!("{}/xrpc/com.atproto.server.updateEmail", base_url)) @@ -122,7 +113,7 @@ async fn test_update_email_flow_success() { assert_eq!(res.status(), StatusCode::OK); let user_email: Option = sqlx::query_scalar!("SELECT email FROM users WHERE did = $1", did) - .fetch_one(&pool) + .fetch_one(pool) .await .expect("User not found"); assert_eq!(user_email, Some(new_email)); @@ -244,7 +235,7 @@ async fn test_update_email_invalid_format() { async fn test_confirm_email_confirms_existing_email() { let client = common::client(); let base_url = common::base_url().await; - let pool = get_pool().await; + let pool = common::get_test_db_pool().await; let handle = format!("emailconfirm-{}", uuid::Uuid::new_v4()); let email = format!("{}@example.com", handle); @@ -270,7 +261,7 @@ async fn test_confirm_email_confirms_existing_email() { "SELECT body FROM comms_queue WHERE user_id = (SELECT id FROM users WHERE did = $1) AND comms_type = 'email_verification' ORDER BY created_at DESC LIMIT 1", did ) - .fetch_one(&pool) + .fetch_one(pool) .await .expect("Verification email not found"); @@ -296,7 +287,7 @@ async fn test_confirm_email_confirms_existing_email() { "SELECT email_verified FROM users WHERE did = $1", did ) - .fetch_one(&pool) + .fetch_one(pool) .await .expect("User not found"); assert!(verified); @@ -306,7 +297,7 @@ async fn test_confirm_email_confirms_existing_email() { async fn test_confirm_email_rejects_wrong_email() { let client = common::client(); let base_url = common::base_url().await; - let pool = get_pool().await; + let pool = common::get_test_db_pool().await; let handle = format!("emailconf-wrong-{}", uuid::Uuid::new_v4()); let email = format!("{}@example.com", handle); @@ -332,7 +323,7 @@ async fn test_confirm_email_rejects_wrong_email() { "SELECT body FROM comms_queue WHERE user_id = (SELECT id FROM users WHERE did = $1) AND comms_type = 'email_verification' ORDER BY created_at DESC LIMIT 1", did ) - .fetch_one(&pool) + .fetch_one(pool) .await .expect("Verification email not found"); @@ -400,7 +391,7 @@ async fn test_confirm_email_invalid_token() { async fn test_unverified_account_can_update_email_without_token() { let client = common::client(); let base_url = common::base_url().await; - let pool = get_pool().await; + let pool = common::get_test_db_pool().await; let handle = format!("emailup-unverified-{}", uuid::Uuid::new_v4()); let email = format!("{}@example.com", handle); @@ -454,7 +445,7 @@ async fn test_unverified_account_can_update_email_without_token() { let user_email: Option = sqlx::query_scalar!("SELECT email FROM users WHERE did = $1", did) - .fetch_one(&pool) + .fetch_one(pool) .await .expect("User not found"); assert_eq!(user_email, Some(new_email)); @@ -464,7 +455,7 @@ async fn test_unverified_account_can_update_email_without_token() { async fn test_update_email_taken_by_another_user() { let client = common::client(); let base_url = common::base_url().await; - let pool = get_pool().await; + let pool = common::get_test_db_pool().await; let handle1 = format!("emailup-dup1-{}", uuid::Uuid::new_v4()); let email1 = format!("{}@example.com", handle1); @@ -485,7 +476,7 @@ async fn test_update_email_taken_by_another_user() { .expect("Failed to request email update"); assert_eq!(res.status(), StatusCode::OK); - let code = get_email_update_token(&pool, &did2).await; + let code = get_email_update_token(pool, &did2).await; let res = client .post(format!("{}/xrpc/com.atproto.server.updateEmail", base_url)) diff --git a/tests/helpers/mod.rs b/tests/helpers/mod.rs index 489145b..a487ea6 100644 --- a/tests/helpers/mod.rs +++ b/tests/helpers/mod.rs @@ -217,30 +217,20 @@ pub async fn create_repost( #[allow(dead_code)] pub async fn set_account_takedown(did: &str, takedown_ref: Option<&str>) { - let conn_str = get_db_connection_string().await; - let pool = sqlx::postgres::PgPoolOptions::new() - .max_connections(2) - .connect(&conn_str) - .await - .expect("Failed to connect to test database"); + let pool = get_test_db_pool().await; sqlx::query!( "UPDATE users SET takedown_ref = $1 WHERE did = $2", takedown_ref, did ) - .execute(&pool) + .execute(pool) .await .expect("Failed to update takedown_ref"); } #[allow(dead_code)] pub async fn set_account_deactivated(did: &str, deactivated: bool) { - let conn_str = get_db_connection_string().await; - let pool = sqlx::postgres::PgPoolOptions::new() - .max_connections(2) - .connect(&conn_str) - .await - .expect("Failed to connect to test database"); + let pool = get_test_db_pool().await; let deactivated_at: Option> = if deactivated { Some(Utc::now()) } else { None }; sqlx::query!( @@ -248,7 +238,7 @@ pub async fn set_account_deactivated(did: &str, deactivated: bool) { deactivated_at, did ) - .execute(&pool) + .execute(pool) .await .expect("Failed to update deactivated_at"); } diff --git a/tests/jwt_security.rs b/tests/jwt_security.rs index 324a38d..f3c9540 100644 --- a/tests/jwt_security.rs +++ b/tests/jwt_security.rs @@ -2,7 +2,7 @@ mod common; use base64::{Engine as _, engine::general_purpose::URL_SAFE_NO_PAD}; use chrono::{Duration, Utc}; -use common::{base_url, client, create_account_and_login, get_db_connection_string}; +use common::{base_url, client, create_account_and_login, get_test_db_pool}; use k256::SecretKey; use k256::ecdsa::{Signature, SigningKey, signature::Signer}; use rand::rngs::OsRng; @@ -683,15 +683,11 @@ async fn test_refresh_token_replay_protection() { let account: Value = create_res.json().await.unwrap(); let did = account["did"].as_str().unwrap(); - let pool = sqlx::postgres::PgPoolOptions::new() - .max_connections(2) - .connect(&get_db_connection_string().await) - .await - .unwrap(); + let pool = get_test_db_pool().await; let body_text: String = sqlx::query_scalar!( "SELECT body FROM comms_queue WHERE user_id = (SELECT id FROM users WHERE did = $1) AND comms_type = 'email_verification' ORDER BY created_at DESC LIMIT 1", did - ).fetch_one(&pool).await.unwrap(); + ).fetch_one(pool).await.unwrap(); let lines: Vec<&str> = body_text.lines().collect(); let code = lines .iter() diff --git a/tests/notifications.rs b/tests/notifications.rs index 187118b..4837a31 100644 --- a/tests/notifications.rs +++ b/tests/notifications.rs @@ -1,24 +1,14 @@ mod common; -use sqlx::PgPool; use tranquil_pds::comms::{ CommsChannel, CommsStatus, CommsType, NewComms, enqueue_comms, enqueue_welcome, }; -async fn get_pool() -> PgPool { - let conn_str = common::get_db_connection_string().await; - sqlx::postgres::PgPoolOptions::new() - .max_connections(5) - .connect(&conn_str) - .await - .expect("Failed to connect to test database") -} - #[tokio::test] async fn test_enqueue_comms() { - let pool = get_pool().await; + 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) - .fetch_one(&pool) + .fetch_one(pool) .await .expect("User not found"); let item = NewComms::email( @@ -28,7 +18,7 @@ async fn test_enqueue_comms() { "Test Subject".to_string(), "Test body".to_string(), ); - let comms_id = enqueue_comms(&pool, item) + let comms_id = enqueue_comms(pool, item) .await .expect("Failed to enqueue comms"); let row = sqlx::query!( @@ -43,7 +33,7 @@ async fn test_enqueue_comms() { "#, comms_id ) - .fetch_one(&pool) + .fetch_one(pool) .await .expect("Comms not found"); assert_eq!(row.user_id, user_id); @@ -57,13 +47,13 @@ async fn test_enqueue_comms() { #[tokio::test] async fn test_enqueue_welcome() { - let pool = get_pool().await; + 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) + .fetch_one(pool) .await .expect("User not found"); - let comms_id = enqueue_welcome(&pool, user_row.id, "example.com") + let comms_id = enqueue_welcome(pool, user_row.id, "example.com") .await .expect("Failed to enqueue welcome comms"); let row = sqlx::query!( @@ -76,7 +66,7 @@ async fn test_enqueue_welcome() { "#, comms_id ) - .fetch_one(&pool) + .fetch_one(pool) .await .expect("Comms not found"); assert_eq!(Some(row.recipient), user_row.email); @@ -87,17 +77,17 @@ async fn test_enqueue_welcome() { #[tokio::test] async fn test_comms_queue_status_index() { - let pool = get_pool().await; + 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) - .fetch_one(&pool) + .fetch_one(pool) .await .expect("User not found"); let initial_count: i64 = sqlx::query_scalar!( "SELECT COUNT(*) FROM comms_queue WHERE status = 'pending' AND user_id = $1", user_id ) - .fetch_one(&pool) + .fetch_one(pool) .await .expect("Failed to count") .unwrap_or(0); @@ -109,13 +99,13 @@ async fn test_comms_queue_status_index() { "Test".to_string(), "Body".to_string(), ); - enqueue_comms(&pool, item).await.expect("Failed to enqueue"); + enqueue_comms(pool, item).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 ) - .fetch_one(&pool) + .fetch_one(pool) .await .expect("Failed to count") .unwrap_or(0); diff --git a/tests/oauth.rs b/tests/oauth.rs index 6b2c17c..752d2ca 100644 --- a/tests/oauth.rs +++ b/tests/oauth.rs @@ -2,7 +2,7 @@ mod common; mod helpers; use base64::{Engine as _, engine::general_purpose::URL_SAFE_NO_PAD}; use chrono::Utc; -use common::{base_url, client, get_db_connection_string}; +use common::{base_url, client, get_test_db_pool}; use helpers::verify_new_account; use reqwest::{StatusCode, redirect}; use serde_json::{Value, json}; @@ -449,15 +449,10 @@ async fn test_oauth_2fa_flow() { let account: Value = create_res.json().await.unwrap(); let user_did = account["did"].as_str().unwrap(); verify_new_account(&http_client, user_did).await; - let db_url = get_db_connection_string().await; - let pool = sqlx::postgres::PgPoolOptions::new() - .max_connections(1) - .connect(&db_url) - .await - .unwrap(); + let pool = get_test_db_pool().await; sqlx::query("UPDATE users SET two_factor_enabled = true WHERE did = $1") .bind(user_did) - .execute(&pool) + .execute(pool) .await .unwrap(); let redirect_uri = "https://example.com/2fa-callback"; @@ -516,7 +511,7 @@ async fn test_oauth_2fa_flow() { let twofa_code: String = sqlx::query_scalar("SELECT code FROM oauth_2fa_challenge WHERE request_uri = $1") .bind(request_uri) - .fetch_one(&pool) + .fetch_one(pool) .await .unwrap(); let twofa_res = http_client @@ -575,15 +570,10 @@ async fn test_oauth_2fa_lockout() { let account: Value = create_res.json().await.unwrap(); let user_did = account["did"].as_str().unwrap(); verify_new_account(&http_client, user_did).await; - let db_url = get_db_connection_string().await; - let pool = sqlx::postgres::PgPoolOptions::new() - .max_connections(1) - .connect(&db_url) - .await - .unwrap(); + let pool = get_test_db_pool().await; sqlx::query("UPDATE users SET two_factor_enabled = true WHERE did = $1") .bind(user_did) - .execute(&pool) + .execute(pool) .await .unwrap(); let redirect_uri = "https://example.com/2fa-lockout-callback"; @@ -754,15 +744,10 @@ async fn test_account_selector_with_2fa() { .json::() .await .unwrap(); - let db_url = get_db_connection_string().await; - let pool = sqlx::postgres::PgPoolOptions::new() - .max_connections(1) - .connect(&db_url) - .await - .unwrap(); + let pool = get_test_db_pool().await; sqlx::query("UPDATE users SET two_factor_enabled = true WHERE did = $1") .bind(&user_did) - .execute(&pool) + .execute(pool) .await .unwrap(); let (code_verifier2, code_challenge2) = generate_pkce(); @@ -803,7 +788,7 @@ async fn test_account_selector_with_2fa() { let twofa_code: String = sqlx::query_scalar("SELECT code FROM oauth_2fa_challenge WHERE request_uri = $1") .bind(request_uri2) - .fetch_one(&pool) + .fetch_one(pool) .await .unwrap(); let twofa_res = http_client diff --git a/tests/password_reset.rs b/tests/password_reset.rs index d72fb09..3db9810 100644 --- a/tests/password_reset.rs +++ b/tests/password_reset.rs @@ -3,22 +3,12 @@ mod helpers; use helpers::verify_new_account; use reqwest::StatusCode; use serde_json::{Value, json}; -use sqlx::PgPool; - -async fn get_pool() -> PgPool { - let conn_str = common::get_db_connection_string().await; - sqlx::postgres::PgPoolOptions::new() - .max_connections(5) - .connect(&conn_str) - .await - .expect("Failed to connect to test database") -} #[tokio::test] async fn test_request_password_reset_creates_code() { let client = common::client(); let base_url = common::base_url().await; - let pool = get_pool().await; + let pool = common::get_test_db_pool().await; let handle = format!("pwreset-{}", uuid::Uuid::new_v4()); let email = format!("{}@example.com", handle); let payload = json!({ @@ -50,7 +40,7 @@ async fn test_request_password_reset_creates_code() { "SELECT password_reset_code, password_reset_code_expires_at FROM users WHERE email = $1", email ) - .fetch_one(&pool) + .fetch_one(pool) .await .expect("User not found"); assert!(user.password_reset_code.is_some()); @@ -80,7 +70,7 @@ async fn test_request_password_reset_unknown_email_returns_ok() { async fn test_reset_password_with_valid_token() { let client = common::client(); let base_url = common::base_url().await; - let pool = get_pool().await; + let pool = common::get_test_db_pool().await; let handle = format!("pwreset2-{}", uuid::Uuid::new_v4()); let email = format!("{}@example.com", handle); let old_password = "Oldpass123!"; @@ -117,7 +107,7 @@ async fn test_reset_password_with_valid_token() { "SELECT password_reset_code FROM users WHERE email = $1", email ) - .fetch_one(&pool) + .fetch_one(pool) .await .expect("User not found"); let token = user.password_reset_code.expect("No reset code"); @@ -138,7 +128,7 @@ async fn test_reset_password_with_valid_token() { "SELECT password_reset_code, password_reset_code_expires_at FROM users WHERE email = $1", email ) - .fetch_one(&pool) + .fetch_one(pool) .await .expect("User not found"); assert!(user.password_reset_code.is_none()); @@ -196,7 +186,7 @@ async fn test_reset_password_with_invalid_token() { async fn test_reset_password_with_expired_token() { let client = common::client(); let base_url = common::base_url().await; - let pool = get_pool().await; + let pool = common::get_test_db_pool().await; let handle = format!("pwreset3-{}", uuid::Uuid::new_v4()); let email = format!("{}@example.com", handle); let payload = json!({ @@ -228,7 +218,7 @@ async fn test_reset_password_with_expired_token() { "SELECT password_reset_code FROM users WHERE email = $1", email ) - .fetch_one(&pool) + .fetch_one(pool) .await .expect("User not found"); let token = user.password_reset_code.expect("No reset code"); @@ -236,7 +226,7 @@ async fn test_reset_password_with_expired_token() { "UPDATE users SET password_reset_code_expires_at = NOW() - INTERVAL '1 hour' WHERE email = $1", email ) - .execute(&pool) + .execute(pool) .await .expect("Failed to expire token"); let res = client @@ -260,7 +250,7 @@ async fn test_reset_password_with_expired_token() { async fn test_reset_password_invalidates_sessions() { let client = common::client(); let base_url = common::base_url().await; - let pool = get_pool().await; + let pool = common::get_test_db_pool().await; let handle = format!("pwreset4-{}", uuid::Uuid::new_v4()); let email = format!("{}@example.com", handle); let payload = json!({ @@ -302,7 +292,7 @@ async fn test_reset_password_invalidates_sessions() { "SELECT password_reset_code FROM users WHERE email = $1", email ) - .fetch_one(&pool) + .fetch_one(pool) .await .expect("User not found"); let token = user.password_reset_code.expect("No reset code"); @@ -348,7 +338,7 @@ async fn test_request_password_reset_empty_email() { #[tokio::test] async fn test_reset_password_creates_notification() { - let pool = get_pool().await; + let pool = common::get_test_db_pool().await; let client = common::client(); let base_url = common::base_url().await; let handle = format!("pwreset5-{}", uuid::Uuid::new_v4()); @@ -369,14 +359,14 @@ async fn test_reset_password_creates_notification() { .expect("Failed to create account"); assert_eq!(res.status(), StatusCode::OK); let user = sqlx::query!("SELECT id FROM users WHERE email = $1", email) - .fetch_one(&pool) + .fetch_one(pool) .await .expect("User not found"); let initial_count: i64 = sqlx::query_scalar!( "SELECT COUNT(*) FROM comms_queue WHERE user_id = $1 AND comms_type = 'password_reset'", user.id ) - .fetch_one(&pool) + .fetch_one(pool) .await .expect("Failed to count") .unwrap_or(0); @@ -394,7 +384,7 @@ async fn test_reset_password_creates_notification() { "SELECT COUNT(*) FROM comms_queue WHERE user_id = $1 AND comms_type = 'password_reset'", user.id ) - .fetch_one(&pool) + .fetch_one(pool) .await .expect("Failed to count") .unwrap_or(0); diff --git a/tests/signing_key.rs b/tests/signing_key.rs index 34ec8a2..79e5470 100644 --- a/tests/signing_key.rs +++ b/tests/signing_key.rs @@ -3,16 +3,6 @@ mod helpers; use helpers::verify_new_account; use reqwest::StatusCode; use serde_json::{Value, json}; -use sqlx::PgPool; - -async fn get_pool() -> PgPool { - let conn_str = common::get_db_connection_string().await; - sqlx::postgres::PgPoolOptions::new() - .max_connections(5) - .connect(&conn_str) - .await - .expect("Failed to connect to test database") -} #[tokio::test] async fn test_reserve_signing_key_without_did() { @@ -41,7 +31,7 @@ async fn test_reserve_signing_key_without_did() { async fn test_reserve_signing_key_with_did() { let client = common::client(); let base_url = common::base_url().await; - let pool = get_pool().await; + let pool = common::get_test_db_pool().await; let target_did = "did:plc:test123456"; let res = client .post(format!( @@ -60,7 +50,7 @@ async fn test_reserve_signing_key_with_did() { "SELECT did, public_key_did_key FROM reserved_signing_keys WHERE public_key_did_key = $1", signing_key ) - .fetch_one(&pool) + .fetch_one(pool) .await .expect("Reserved key not found in database"); assert_eq!(row.did.as_deref(), Some(target_did)); @@ -71,7 +61,7 @@ async fn test_reserve_signing_key_with_did() { async fn test_reserve_signing_key_stores_private_key() { let client = common::client(); let base_url = common::base_url().await; - let pool = get_pool().await; + let pool = common::get_test_db_pool().await; let res = client .post(format!( "{}/xrpc/com.atproto.server.reserveSigningKey", @@ -88,7 +78,7 @@ async fn test_reserve_signing_key_stores_private_key() { "SELECT private_key_bytes, expires_at, used_at FROM reserved_signing_keys WHERE public_key_did_key = $1", signing_key ) - .fetch_one(&pool) + .fetch_one(pool) .await .expect("Reserved key not found in database"); assert_eq!( @@ -161,7 +151,7 @@ async fn test_reserve_signing_key_is_public() { async fn test_create_account_with_reserved_signing_key() { let client = common::client(); let base_url = common::base_url().await; - let pool = get_pool().await; + let pool = common::get_test_db_pool().await; let res = client .post(format!( "{}/xrpc/com.atproto.server.reserveSigningKey", @@ -199,7 +189,7 @@ async fn test_create_account_with_reserved_signing_key() { "SELECT used_at FROM reserved_signing_keys WHERE public_key_did_key = $1", signing_key ) - .fetch_one(&pool) + .fetch_one(pool) .await .expect("Reserved key not found"); assert!(