diff --git a/.sqlx/query-297fcbb356d65aae3faae5430000b6c6fbec8566a4adbb595c91606fdfa3bedc.json b/.sqlx/query-297fcbb356d65aae3faae5430000b6c6fbec8566a4adbb595c91606fdfa3bedc.json new file mode 100644 index 0000000..d14cb47 --- /dev/null +++ b/.sqlx/query-297fcbb356d65aae3faae5430000b6c6fbec8566a4adbb595c91606fdfa3bedc.json @@ -0,0 +1,112 @@ +{ + "db_name": "PostgreSQL", + "query": "\n SELECT handle, email, email_verified, is_admin, deactivated_at, takedown_ref,\n preferred_locale,\n preferred_comms_channel as \"preferred_comms_channel!: CommsChannel\",\n discord_verified, telegram_verified, signal_verified,\n migrated_to_pds, migrated_at,\n (SELECT verified FROM user_totp WHERE did = users.did) as totp_enabled\n FROM users\n WHERE did = $1\n ", + "describe": { + "columns": [ + { + "ordinal": 0, + "name": "handle", + "type_info": "Text" + }, + { + "ordinal": 1, + "name": "email", + "type_info": "Text" + }, + { + "ordinal": 2, + "name": "email_verified", + "type_info": "Bool" + }, + { + "ordinal": 3, + "name": "is_admin", + "type_info": "Bool" + }, + { + "ordinal": 4, + "name": "deactivated_at", + "type_info": "Timestamptz" + }, + { + "ordinal": 5, + "name": "takedown_ref", + "type_info": "Text" + }, + { + "ordinal": 6, + "name": "preferred_locale", + "type_info": "Varchar" + }, + { + "ordinal": 7, + "name": "preferred_comms_channel!: CommsChannel", + "type_info": { + "Custom": { + "name": "comms_channel", + "kind": { + "Enum": [ + "email", + "discord", + "telegram", + "signal" + ] + } + } + } + }, + { + "ordinal": 8, + "name": "discord_verified", + "type_info": "Bool" + }, + { + "ordinal": 9, + "name": "telegram_verified", + "type_info": "Bool" + }, + { + "ordinal": 10, + "name": "signal_verified", + "type_info": "Bool" + }, + { + "ordinal": 11, + "name": "migrated_to_pds", + "type_info": "Text" + }, + { + "ordinal": 12, + "name": "migrated_at", + "type_info": "Timestamptz" + }, + { + "ordinal": 13, + "name": "totp_enabled", + "type_info": "Bool" + } + ], + "parameters": { + "Left": [ + "Text" + ] + }, + "nullable": [ + false, + true, + false, + false, + true, + true, + true, + false, + false, + false, + false, + true, + true, + null + ] + }, + "hash": "297fcbb356d65aae3faae5430000b6c6fbec8566a4adbb595c91606fdfa3bedc" +} diff --git a/.sqlx/query-3b056b9e79847c8bbb8507f283213e7209b417e7933f5b2277a83cae7e1c7888.json b/.sqlx/query-3b056b9e79847c8bbb8507f283213e7209b417e7933f5b2277a83cae7e1c7888.json new file mode 100644 index 0000000..77b3c66 --- /dev/null +++ b/.sqlx/query-3b056b9e79847c8bbb8507f283213e7209b417e7933f5b2277a83cae7e1c7888.json @@ -0,0 +1,15 @@ +{ + "db_name": "PostgreSQL", + "query": "DELETE FROM account_preferences WHERE user_id = $1 AND name = $2", + "describe": { + "columns": [], + "parameters": { + "Left": [ + "Uuid", + "Text" + ] + }, + "nullable": [] + }, + "hash": "3b056b9e79847c8bbb8507f283213e7209b417e7933f5b2277a83cae7e1c7888" +} diff --git a/.sqlx/query-a960b981a146a0e422ef53601dfc31e29cf777aa194227c48c6ebc6905ea3249.json b/.sqlx/query-a960b981a146a0e422ef53601dfc31e29cf777aa194227c48c6ebc6905ea3249.json new file mode 100644 index 0000000..61db6fb --- /dev/null +++ b/.sqlx/query-a960b981a146a0e422ef53601dfc31e29cf777aa194227c48c6ebc6905ea3249.json @@ -0,0 +1,136 @@ +{ + "db_name": "PostgreSQL", + "query": "SELECT\n u.id, u.did, u.handle, u.password_hash, u.email, u.deactivated_at, u.takedown_ref,\n u.email_verified, u.discord_verified, u.telegram_verified, u.signal_verified,\n u.allow_legacy_login, u.migrated_to_pds,\n u.preferred_comms_channel as \"preferred_comms_channel: CommsChannel\",\n k.key_bytes, k.encryption_version,\n (SELECT verified FROM user_totp WHERE did = u.did) as totp_enabled,\n COALESCE((SELECT (value_json)::boolean FROM account_preferences WHERE user_id = u.id AND name = 'email_auth_factor' ORDER BY created_at DESC LIMIT 1), false) as \"email_2fa_enabled!\"\n FROM users u\n JOIN user_keys k ON u.id = k.user_id\n WHERE u.handle = $1 OR u.email = $1 OR u.did = $1", + "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": "password_hash", + "type_info": "Text" + }, + { + "ordinal": 4, + "name": "email", + "type_info": "Text" + }, + { + "ordinal": 5, + "name": "deactivated_at", + "type_info": "Timestamptz" + }, + { + "ordinal": 6, + "name": "takedown_ref", + "type_info": "Text" + }, + { + "ordinal": 7, + "name": "email_verified", + "type_info": "Bool" + }, + { + "ordinal": 8, + "name": "discord_verified", + "type_info": "Bool" + }, + { + "ordinal": 9, + "name": "telegram_verified", + "type_info": "Bool" + }, + { + "ordinal": 10, + "name": "signal_verified", + "type_info": "Bool" + }, + { + "ordinal": 11, + "name": "allow_legacy_login", + "type_info": "Bool" + }, + { + "ordinal": 12, + "name": "migrated_to_pds", + "type_info": "Text" + }, + { + "ordinal": 13, + "name": "preferred_comms_channel: CommsChannel", + "type_info": { + "Custom": { + "name": "comms_channel", + "kind": { + "Enum": [ + "email", + "discord", + "telegram", + "signal" + ] + } + } + } + }, + { + "ordinal": 14, + "name": "key_bytes", + "type_info": "Bytea" + }, + { + "ordinal": 15, + "name": "encryption_version", + "type_info": "Int4" + }, + { + "ordinal": 16, + "name": "totp_enabled", + "type_info": "Bool" + }, + { + "ordinal": 17, + "name": "email_2fa_enabled!", + "type_info": "Bool" + } + ], + "parameters": { + "Left": [ + "Text" + ] + }, + "nullable": [ + false, + false, + false, + true, + true, + true, + true, + false, + false, + false, + false, + false, + true, + false, + false, + true, + null, + null + ] + }, + "hash": "a960b981a146a0e422ef53601dfc31e29cf777aa194227c48c6ebc6905ea3249" +} diff --git a/.sqlx/query-c8728a1247c535e941e2b3bcb4100d7b3610f31c7acfdc1f8c072e1c5ca0ea18.json b/.sqlx/query-c8728a1247c535e941e2b3bcb4100d7b3610f31c7acfdc1f8c072e1c5ca0ea18.json new file mode 100644 index 0000000..cc1a28d --- /dev/null +++ b/.sqlx/query-c8728a1247c535e941e2b3bcb4100d7b3610f31c7acfdc1f8c072e1c5ca0ea18.json @@ -0,0 +1,118 @@ +{ + "db_name": "PostgreSQL", + "query": "\n SELECT u.handle, u.email, u.email_verified, u.is_admin, u.deactivated_at, u.takedown_ref,\n u.preferred_locale,\n u.preferred_comms_channel as \"preferred_comms_channel!: CommsChannel\",\n u.discord_verified, u.telegram_verified, u.signal_verified,\n u.migrated_to_pds, u.migrated_at,\n (SELECT verified FROM user_totp WHERE did = u.did) as totp_enabled,\n COALESCE((SELECT (value_json)::boolean FROM account_preferences WHERE user_id = u.id AND name = 'email_auth_factor' ORDER BY created_at DESC LIMIT 1), false) as \"email_2fa_enabled!\"\n FROM users u\n WHERE u.did = $1\n ", + "describe": { + "columns": [ + { + "ordinal": 0, + "name": "handle", + "type_info": "Text" + }, + { + "ordinal": 1, + "name": "email", + "type_info": "Text" + }, + { + "ordinal": 2, + "name": "email_verified", + "type_info": "Bool" + }, + { + "ordinal": 3, + "name": "is_admin", + "type_info": "Bool" + }, + { + "ordinal": 4, + "name": "deactivated_at", + "type_info": "Timestamptz" + }, + { + "ordinal": 5, + "name": "takedown_ref", + "type_info": "Text" + }, + { + "ordinal": 6, + "name": "preferred_locale", + "type_info": "Varchar" + }, + { + "ordinal": 7, + "name": "preferred_comms_channel!: CommsChannel", + "type_info": { + "Custom": { + "name": "comms_channel", + "kind": { + "Enum": [ + "email", + "discord", + "telegram", + "signal" + ] + } + } + } + }, + { + "ordinal": 8, + "name": "discord_verified", + "type_info": "Bool" + }, + { + "ordinal": 9, + "name": "telegram_verified", + "type_info": "Bool" + }, + { + "ordinal": 10, + "name": "signal_verified", + "type_info": "Bool" + }, + { + "ordinal": 11, + "name": "migrated_to_pds", + "type_info": "Text" + }, + { + "ordinal": 12, + "name": "migrated_at", + "type_info": "Timestamptz" + }, + { + "ordinal": 13, + "name": "totp_enabled", + "type_info": "Bool" + }, + { + "ordinal": 14, + "name": "email_2fa_enabled!", + "type_info": "Bool" + } + ], + "parameters": { + "Left": [ + "Text" + ] + }, + "nullable": [ + false, + true, + false, + false, + true, + true, + true, + false, + false, + false, + false, + true, + true, + null, + null + ] + }, + "hash": "c8728a1247c535e941e2b3bcb4100d7b3610f31c7acfdc1f8c072e1c5ca0ea18" +} diff --git a/crates/tranquil-cache/src/lib.rs b/crates/tranquil-cache/src/lib.rs index dbcad63..b8552cf 100644 --- a/crates/tranquil-cache/src/lib.rs +++ b/crates/tranquil-cache/src/lib.rs @@ -91,6 +91,10 @@ impl Cache for NoOpCache { async fn set_bytes(&self, _key: &str, _value: &[u8], _ttl: Duration) -> Result<(), CacheError> { Ok(()) } + + fn is_available(&self) -> bool { + false + } } #[derive(Clone)] diff --git a/crates/tranquil-comms/src/locale.rs b/crates/tranquil-comms/src/locale.rs index 5900f13..c0bf56a 100644 --- a/crates/tranquil-comms/src/locale.rs +++ b/crates/tranquil-comms/src/locale.rs @@ -16,6 +16,7 @@ pub struct NotificationStrings { pub password_reset_body: &'static str, pub email_update_subject: &'static str, pub email_update_body: &'static str, + pub short_token_body: &'static str, pub account_deletion_subject: &'static str, pub account_deletion_body: &'static str, pub plc_operation_subject: &'static str, @@ -50,6 +51,7 @@ static STRINGS_EN: NotificationStrings = NotificationStrings { password_reset_body: "Hello @{handle},\n\nYour password reset code is: {code}\n\nThis code will expire in 10 minutes.\n\nIf you did not request this, please ignore this message.", email_update_subject: "Confirm your new email - {hostname}", email_update_body: "Hello @{handle},\n\nYour verification code is:\n{code}\n\nCopy the code above and enter it at:\n{verify_page}\n\nThis code will expire in 10 minutes.\n\nOr if you like to live dangerously:\n{verify_link}\n\nIf you did not request this, please ignore this email.", + short_token_body: "Hello @{handle},\n\nYour verification code is:\n{code}\n\nThis code will expire in 15 minutes.\n\nIf you did not request this, please ignore this email.", account_deletion_subject: "Account Deletion Request - {hostname}", account_deletion_body: "Hello @{handle},\n\nYour account deletion confirmation code is: {code}\n\nThis code will expire in 10 minutes.\n\nIf you did not request this, please secure your account immediately.", plc_operation_subject: "{hostname} - PLC Operation Token", @@ -73,6 +75,7 @@ static STRINGS_ZH: NotificationStrings = NotificationStrings { password_reset_body: "您好 @{handle},\n\n您的密码重置验证码是:{code}\n\n此验证码将在10分钟后过期。\n\n如果这不是您的操作,请忽略此消息。", email_update_subject: "确认您的新邮箱 - {hostname}", email_update_body: "您好 @{handle},\n\n您的验证码是:\n{code}\n\n复制上述验证码并在此输入:\n{verify_page}\n\n此验证码将在10分钟后过期。\n\n或者直接点击链接:\n{verify_link}\n\n如果这不是您的操作,请忽略此邮件。", + short_token_body: "您好 @{handle},\n\n您的验证码是:\n{code}\n\n此验证码将在15分钟后过期。\n\n如果这不是您的操作,请忽略此邮件。", account_deletion_subject: "账户删除请求 - {hostname}", account_deletion_body: "您好 @{handle},\n\n您的账户删除确认码是:{code}\n\n此验证码将在10分钟后过期。\n\n如果这不是您的操作,请立即保护您的账户。", plc_operation_subject: "{hostname} - PLC 操作令牌", @@ -96,6 +99,7 @@ static STRINGS_JA: NotificationStrings = NotificationStrings { password_reset_body: "@{handle} 様\n\nパスワードリセットコードは:{code}\n\nこのコードは10分後に期限切れとなります。\n\nこの操作に心当たりがない場合は、このメッセージを無視してください。", email_update_subject: "新しいメールアドレスの確認 - {hostname}", email_update_body: "@{handle} 様\n\n確認コードは:\n{code}\n\n上記のコードをコピーして、こちらで入力してください:\n{verify_page}\n\nこのコードは10分後に期限切れとなります。\n\n自己責任でワンクリック認証:\n{verify_link}\n\nこの操作に心当たりがない場合は、このメールを無視してください。", + short_token_body: "@{handle} 様\n\n確認コードは:\n{code}\n\nこのコードは15分後に期限切れとなります。\n\nこの操作に心当たりがない場合は、このメールを無視してください。", account_deletion_subject: "アカウント削除リクエスト - {hostname}", account_deletion_body: "@{handle} 様\n\nアカウント削除の確認コードは:{code}\n\nこのコードは10分後に期限切れとなります。\n\nこの操作に心当たりがない場合は、直ちにアカウントを保護してください。", plc_operation_subject: "{hostname} - PLC 操作トークン", @@ -119,6 +123,7 @@ static STRINGS_KO: NotificationStrings = NotificationStrings { password_reset_body: "안녕하세요 @{handle}님,\n\n비밀번호 재설정 코드는: {code}\n\n이 코드는 10분 후에 만료됩니다.\n\n요청하지 않으셨다면 이 메시지를 무시하세요.", email_update_subject: "새 이메일 주소 확인 - {hostname}", email_update_body: "안녕하세요 @{handle}님,\n\n인증 코드는:\n{code}\n\n위 코드를 복사하여 여기에 입력하세요:\n{verify_page}\n\n이 코드는 10분 후에 만료됩니다.\n\n위험을 감수하고 원클릭 인증:\n{verify_link}\n\n요청하지 않으셨다면 이 이메일을 무시하세요.", + short_token_body: "안녕하세요 @{handle}님,\n\n인증 코드는:\n{code}\n\n이 코드는 15분 후에 만료됩니다.\n\n요청하지 않으셨다면 이 이메일을 무시하세요.", account_deletion_subject: "계정 삭제 요청 - {hostname}", account_deletion_body: "안녕하세요 @{handle}님,\n\n계정 삭제 확인 코드는: {code}\n\n이 코드는 10분 후에 만료됩니다.\n\n요청하지 않으셨다면 즉시 계정을 보호하세요.", plc_operation_subject: "{hostname} - PLC 작업 토큰", @@ -142,6 +147,7 @@ static STRINGS_SV: NotificationStrings = NotificationStrings { password_reset_body: "Hej @{handle},\n\nDin kod för lösenordsåterställning är: {code}\n\nDenna kod upphör om 10 minuter.\n\nOm du inte begärde detta kan du ignorera detta meddelande.", email_update_subject: "Bekräfta din nya e-post - {hostname}", email_update_body: "Hej @{handle},\n\nDin verifieringskod är:\n{code}\n\nKopiera koden ovan och ange den på:\n{verify_page}\n\nDenna kod upphör om 10 minuter.\n\nEller om du gillar att leva farligt:\n{verify_link}\n\nOm du inte begärde detta kan du ignorera detta meddelande.", + short_token_body: "Hej @{handle},\n\nDin verifieringskod är:\n{code}\n\nDenna kod upphör om 15 minuter.\n\nOm du inte begärde detta kan du ignorera detta meddelande.", account_deletion_subject: "Begäran om kontoradering - {hostname}", account_deletion_body: "Hej @{handle},\n\nDin bekräftelsekod för kontoradering är: {code}\n\nDenna kod upphör om 10 minuter.\n\nOm du inte begärde detta, skydda ditt konto omedelbart.", plc_operation_subject: "{hostname} - PLC-operationstoken", @@ -165,6 +171,7 @@ static STRINGS_FI: NotificationStrings = NotificationStrings { password_reset_body: "Hei @{handle},\n\nSalasanan palautuskoodisi on: {code}\n\nTämä koodi vanhenee 10 minuutissa.\n\nJos et pyytänyt tätä, voit jättää tämän viestin huomiotta.", email_update_subject: "Vahvista uusi sähköpostisi - {hostname}", email_update_body: "Hei @{handle},\n\nVahvistuskoodisi on:\n{code}\n\nKopioi koodi yllä ja syötä se osoitteessa:\n{verify_page}\n\nTämä koodi vanhenee 10 minuutissa.\n\nTai jos pidät vaarallisesta elämästä:\n{verify_link}\n\nJos et pyytänyt tätä, voit jättää tämän viestin huomiotta.", + short_token_body: "Hei @{handle},\n\nVahvistuskoodisi on:\n{code}\n\nTämä koodi vanhenee 15 minuutissa.\n\nJos et pyytänyt tätä, voit jättää tämän viestin huomiotta.", account_deletion_subject: "Tilin poistopyyntö - {hostname}", account_deletion_body: "Hei @{handle},\n\nTilin poiston vahvistuskoodisi on: {code}\n\nTämä koodi vanhenee 10 minuutissa.\n\nJos et pyytänyt tätä, suojaa tilisi välittömästi.", plc_operation_subject: "{hostname} - PLC-toimintotunniste", diff --git a/crates/tranquil-db-traits/src/user.rs b/crates/tranquil-db-traits/src/user.rs index 58a0ca4..c7e7f28 100644 --- a/crates/tranquil-db-traits/src/user.rs +++ b/crates/tranquil-db-traits/src/user.rs @@ -768,6 +768,8 @@ pub struct UserSessionInfo { pub channel_verification: ChannelVerificationStatus, pub migrated_to_pds: Option, pub migrated_at: Option>, + pub totp_enabled: bool, + pub email_2fa_enabled: bool, } #[derive(Debug, Clone)] @@ -792,6 +794,7 @@ pub struct UserLoginFull { pub key_bytes: Vec, pub encryption_version: Option, pub totp_enabled: bool, + pub email_2fa_enabled: bool, } #[derive(Debug, Clone)] diff --git a/crates/tranquil-db/src/postgres/infra.rs b/crates/tranquil-db/src/postgres/infra.rs index 240050b..c79df0e 100644 --- a/crates/tranquil-db/src/postgres/infra.rs +++ b/crates/tranquil-db/src/postgres/infra.rs @@ -661,17 +661,29 @@ impl InfraRepository for PostgresInfraRepository { name: &str, value_json: serde_json::Value, ) -> Result<(), DbError> { + let mut tx = self.pool.begin().await.map_err(map_sqlx_error)?; + sqlx::query!( - r#"INSERT INTO account_preferences (user_id, name, value_json) VALUES ($1, $2, $3) - ON CONFLICT (user_id, name) DO UPDATE SET value_json = $3"#, + r#"DELETE FROM account_preferences WHERE user_id = $1 AND name = $2"#, + user_id, + name + ) + .execute(&mut *tx) + .await + .map_err(map_sqlx_error)?; + + sqlx::query!( + r#"INSERT INTO account_preferences (user_id, name, value_json) VALUES ($1, $2, $3)"#, user_id, name, value_json ) - .execute(&self.pool) + .execute(&mut *tx) .await .map_err(map_sqlx_error)?; + tx.commit().await.map_err(map_sqlx_error)?; + Ok(()) } diff --git a/crates/tranquil-db/src/postgres/user.rs b/crates/tranquil-db/src/postgres/user.rs index 2115489..df63f75 100644 --- a/crates/tranquil-db/src/postgres/user.rs +++ b/crates/tranquil-db/src/postgres/user.rs @@ -1374,13 +1374,15 @@ impl UserRepository for PostgresUserRepository { async fn get_session_info_by_did(&self, did: &Did) -> Result, DbError> { sqlx::query!( r#" - SELECT handle, email, email_verified, is_admin, deactivated_at, takedown_ref, - preferred_locale, - preferred_comms_channel as "preferred_comms_channel!: CommsChannel", - discord_verified, telegram_verified, signal_verified, - migrated_to_pds, migrated_at - FROM users - WHERE did = $1 + SELECT u.handle, u.email, u.email_verified, u.is_admin, u.deactivated_at, u.takedown_ref, + u.preferred_locale, + u.preferred_comms_channel as "preferred_comms_channel!: CommsChannel", + u.discord_verified, u.telegram_verified, u.signal_verified, + u.migrated_to_pds, u.migrated_at, + (SELECT verified FROM user_totp WHERE did = u.did) as totp_enabled, + COALESCE((SELECT (value_json)::boolean FROM account_preferences WHERE user_id = u.id AND name = 'email_auth_factor' ORDER BY created_at DESC LIMIT 1), false) as "email_2fa_enabled!" + FROM users u + WHERE u.did = $1 "#, did.as_str() ) @@ -1404,6 +1406,8 @@ impl UserRepository for PostgresUserRepository { ), migrated_to_pds: row.migrated_to_pds, migrated_at: row.migrated_at, + totp_enabled: row.totp_enabled.unwrap_or(false), + email_2fa_enabled: row.email_2fa_enabled, }) }) } @@ -1468,7 +1472,8 @@ impl UserRepository for PostgresUserRepository { u.allow_legacy_login, u.migrated_to_pds, u.preferred_comms_channel as "preferred_comms_channel: CommsChannel", k.key_bytes, k.encryption_version, - (SELECT verified FROM user_totp WHERE did = u.did) as totp_enabled + (SELECT verified FROM user_totp WHERE did = u.did) as totp_enabled, + COALESCE((SELECT (value_json)::boolean FROM account_preferences WHERE user_id = u.id AND name = 'email_auth_factor' ORDER BY created_at DESC LIMIT 1), false) as "email_2fa_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"#, @@ -1498,6 +1503,7 @@ impl UserRepository for PostgresUserRepository { key_bytes: row.key_bytes, encryption_version: row.encryption_version, totp_enabled: row.totp_enabled.unwrap_or(false), + email_2fa_enabled: row.email_2fa_enabled, }) }) } diff --git a/crates/tranquil-infra/src/lib.rs b/crates/tranquil-infra/src/lib.rs index 9463dd7..fc03575 100644 --- a/crates/tranquil-infra/src/lib.rs +++ b/crates/tranquil-infra/src/lib.rs @@ -73,6 +73,9 @@ pub trait Cache: Send + Sync { async fn delete(&self, key: &str) -> Result<(), CacheError>; async fn get_bytes(&self, key: &str) -> Option>; async fn set_bytes(&self, key: &str, value: &[u8], ttl: Duration) -> Result<(), CacheError>; + fn is_available(&self) -> bool { + true + } } #[async_trait] diff --git a/crates/tranquil-pds/src/api/error.rs b/crates/tranquil-pds/src/api/error.rs index b5a8278..a1d0196 100644 --- a/crates/tranquil-pds/src/api/error.rs +++ b/crates/tranquil-pds/src/api/error.rs @@ -114,6 +114,8 @@ pub enum ApiError { SsoSessionExpired, SsoAlreadyLinked, SsoLinkNotFound, + AuthFactorTokenRequired, + LegacyLoginBlocked, } impl ApiError { @@ -132,11 +134,11 @@ impl ApiError { | Self::AuthenticationFailed(_) | Self::AccountDeactivated | Self::AccountTakedown - | Self::InvalidCode(_) | Self::InvalidPassword(_) | Self::InvalidToken(_) | Self::PasskeyCounterAnomaly | Self::OAuthExpiredToken(_) => StatusCode::UNAUTHORIZED, + Self::InvalidCode(_) => StatusCode::BAD_REQUEST, Self::ExpiredToken(_) => StatusCode::BAD_REQUEST, Self::Forbidden | Self::AdminRequired @@ -210,7 +212,9 @@ impl ApiError { | Self::SsoInvalidAction | Self::SsoNotAuthenticated | Self::SsoSessionExpired - | Self::SsoAlreadyLinked => StatusCode::BAD_REQUEST, + | Self::SsoAlreadyLinked + | Self::AuthFactorTokenRequired + | Self::LegacyLoginBlocked => StatusCode::BAD_REQUEST, Self::PasskeyNotFound | Self::SsoLinkNotFound => StatusCode::NOT_FOUND, } } @@ -313,6 +317,8 @@ impl ApiError { Self::SsoSessionExpired => Cow::Borrowed("SsoSessionExpired"), Self::SsoAlreadyLinked => Cow::Borrowed("SsoAlreadyLinked"), Self::SsoLinkNotFound => Cow::Borrowed("SsoLinkNotFound"), + Self::AuthFactorTokenRequired => Cow::Borrowed("AuthFactorTokenRequired"), + Self::LegacyLoginBlocked => Cow::Borrowed("MfaRequired"), } } fn message(&self) -> Option { @@ -436,6 +442,12 @@ impl ApiError { Self::InvalidEmail => Some("Please provide a valid email address".to_string()), Self::InvalidInviteCode => Some("The invite code provided is invalid".to_string()), Self::DuplicateCreate => Some("Account creation failed: duplicate request".to_string()), + Self::LegacyLoginBlocked => Some( + "This account requires MFA. Please use an OAuth client that supports TOTP verification.".to_string(), + ), + Self::AuthFactorTokenRequired => { + Some("A sign in code has been sent to your email address".to_string()) + } _ => None, } } diff --git a/crates/tranquil-pds/src/api/server/email.rs b/crates/tranquil-pds/src/api/server/email.rs index 11688b0..3f51602 100644 --- a/crates/tranquil-pds/src/api/server/email.rs +++ b/crates/tranquil-pds/src/api/server/email.rs @@ -66,7 +66,7 @@ pub async fn request_email_update( .log_db_err("getting email info")? .ok_or(ApiError::AccountNotFound)?; - let Some(current_email) = user.email else { + let Some(_current_email) = user.email else { return Err(ApiError::InvalidRequest( "account does not have an email address".into(), )); @@ -75,12 +75,16 @@ pub async fn request_email_update( let token_required = user.email_verified; if token_required { - let code = crate::auth::verification_token::generate_channel_update_token( - &auth.did, - "email_update", - ¤t_email.to_lowercase(), - ); - let formatted_code = crate::auth::verification_token::format_token_for_display(&code); + let token = crate::auth::email_token::create_email_token( + state.cache.as_ref(), + auth.did.as_str(), + crate::auth::email_token::EmailTokenPurpose::UpdateEmail, + ) + .await + .map_err(|e| { + error!("Failed to create email update token: {:?}", e); + ApiError::InternalError(Some("Failed to generate verification code".into())) + })?; if let Some(Json(ref inp)) = input && let Some(ref new_email) = inp.new_email @@ -89,7 +93,7 @@ pub async fn request_email_update( if !new_email.is_empty() && crate::api::validation::is_valid_email(&new_email) { let pending = PendingEmailUpdate { new_email, - token_hash: hash_token(&code), + token_hash: hash_token(&token), authorized: false, }; if let Ok(json) = serde_json::to_string(&pending) { @@ -102,12 +106,12 @@ pub async fn request_email_update( } let hostname = pds_hostname(); - if let Err(e) = crate::comms::comms_repo::enqueue_email_update_token( + if let Err(e) = crate::comms::comms_repo::enqueue_short_token_email( state.user_repo.as_ref(), state.infra_repo.as_ref(), user.id, - &code, - &formatted_code, + &token, + "email_update", hostname, ) .await @@ -239,9 +243,44 @@ pub async fn update_email( )); } - if let Some(ref current) = current_email - && new_email == current.to_lowercase() - { + let email_unchanged = current_email + .as_ref() + .map(|c| new_email == c.to_lowercase()) + .unwrap_or(false); + + if email_unchanged { + if let Some(email_auth_factor) = input.email_auth_factor { + if email_verified { + let token = input + .token + .as_ref() + .filter(|t| !t.is_empty()) + .ok_or(ApiError::TokenRequired)?; + + crate::auth::email_token::validate_email_token( + state.cache.as_ref(), + did.as_str(), + crate::auth::email_token::EmailTokenPurpose::UpdateEmail, + token, + ) + .await + .map_err(|e| match e { + crate::auth::email_token::TokenError::ExpiredToken => { + ApiError::ExpiredToken(None) + } + _ => ApiError::InvalidToken(None), + })?; + } + + state + .infra_repo + .upsert_account_preference(user_id, "email_auth_factor", json!(email_auth_factor)) + .await + .map_err(|e| { + error!("Failed to update email_auth_factor preference: {}", e); + ApiError::InternalError(Some("Failed to update 2FA setting".into())) + })?; + } return Ok(EmptyResponse::ok().into_response()); } @@ -260,34 +299,57 @@ pub async fn update_email( } if !authorized_via_link { - let Some(ref t) = input.token else { - return Err(ApiError::TokenRequired); - }; - let confirmation_token = - crate::auth::verification_token::normalize_token_input(t.trim()); - - let current_email_lower = current_email + let token = input + .token .as_ref() - .map(|e| e.to_lowercase()) - .unwrap_or_default(); + .filter(|t| !t.is_empty()) + .ok_or(ApiError::TokenRequired)?; - let verified = crate::auth::verification_token::verify_channel_update_token( - &confirmation_token, - "email_update", - ¤t_email_lower, - ); + let short_token_result = crate::auth::email_token::validate_email_token( + state.cache.as_ref(), + did.as_str(), + crate::auth::email_token::EmailTokenPurpose::UpdateEmail, + token, + ) + .await; - match verified { - Ok(token_data) => { - if token_data.did != did.as_str() { - return Err(ApiError::InvalidToken(None)); + if let Err(e) = short_token_result { + let confirmation_token = + crate::auth::verification_token::normalize_token_input(token.trim()); + + let current_email_lower = current_email + .as_ref() + .map(|e| e.to_lowercase()) + .unwrap_or_default(); + + let verified = crate::auth::verification_token::verify_channel_update_token( + &confirmation_token, + "email_update", + ¤t_email_lower, + ); + + match verified { + Ok(token_data) => { + if token_data.did != did.as_str() { + return Err(ApiError::InvalidToken(None)); + } + } + Err(crate::auth::verification_token::VerifyError::Expired) => { + return Err(match e { + crate::auth::email_token::TokenError::ExpiredToken => { + ApiError::ExpiredToken(None) + } + _ => ApiError::InvalidToken(None), + }); + } + Err(_) => { + return Err(match e { + crate::auth::email_token::TokenError::ExpiredToken => { + ApiError::ExpiredToken(None) + } + _ => ApiError::InvalidToken(None), + }); } - } - Err(crate::auth::verification_token::VerifyError::Expired) => { - return Err(ApiError::ExpiredToken(None)); - } - Err(_) => { - return Err(ApiError::InvalidToken(None)); } } } diff --git a/crates/tranquil-pds/src/api/server/session.rs b/crates/tranquil-pds/src/api/server/session.rs index 495e861..98f540c 100644 --- a/crates/tranquil-pds/src/api/server/session.rs +++ b/crates/tranquil-pds/src/api/server/session.rs @@ -32,6 +32,7 @@ pub struct CreateSessionInput { pub password: PlainPassword, #[serde(default)] pub allow_takendown: bool, + pub auth_factor_token: Option, } #[derive(Serialize)] @@ -48,6 +49,8 @@ pub struct CreateSessionOutput { #[serde(skip_serializing_if = "Option::is_none")] pub email_confirmed: Option, #[serde(skip_serializing_if = "Option::is_none")] + pub email_auth_factor: Option, + #[serde(skip_serializing_if = "Option::is_none")] pub active: Option, #[serde(skip_serializing_if = "Option::is_none")] pub status: Option, @@ -158,18 +161,82 @@ pub async fn create_session( .into_response(); } 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); - return ( - StatusCode::FORBIDDEN, - Json(json!({ - "error": "MfaRequired", - "message": "This account requires MFA. Please use an OAuth client that supports TOTP verification.", - "did": row.did - })), - ) - .into_response(); + let email_2fa_enabled = row.email_2fa_enabled; + let is_legacy_login = has_totp || email_2fa_enabled; + let twofa_ctx = crate::auth::legacy_2fa::Legacy2faContext { + email_2fa_enabled, + has_totp, + allow_legacy_login: row.allow_legacy_login, + }; + match crate::auth::legacy_2fa::process_legacy_2fa( + state.cache.as_ref(), + &row.did, + &twofa_ctx, + input.auth_factor_token.as_deref(), + ) + .await + { + Ok(crate::auth::legacy_2fa::Legacy2faOutcome::NotRequired) => {} + Ok(crate::auth::legacy_2fa::Legacy2faOutcome::Blocked) => { + warn!("Legacy login blocked for TOTP-enabled account: {}", row.did); + return ApiError::LegacyLoginBlocked.into_response(); + } + Ok(crate::auth::legacy_2fa::Legacy2faOutcome::ChallengeSent(code)) => { + let hostname = pds_hostname(); + if let Err(e) = crate::comms::comms_repo::enqueue_2fa_code( + state.user_repo.as_ref(), + state.infra_repo.as_ref(), + row.id, + code.as_str(), + hostname, + ) + .await + { + error!("Failed to send 2FA code: {:?}", e); + crate::auth::legacy_2fa::clear_challenge(state.cache.as_ref(), &row.did).await; + return ApiError::InternalError(Some( + "Failed to send verification code. Please try again.".into(), + )) + .into_response(); + } + return ApiError::AuthFactorTokenRequired.into_response(); + } + Ok(crate::auth::legacy_2fa::Legacy2faOutcome::Verified) => {} + Err(crate::auth::legacy_2fa::Legacy2faFlowError::Challenge(e)) => { + use crate::auth::legacy_2fa::ChallengeError; + return match e { + ChallengeError::CacheUnavailable => { + error!("Cache unavailable for 2FA, blocking legacy login"); + ApiError::ServiceUnavailable(Some( + "2FA service temporarily unavailable. Please try again later or use an OAuth client.".into(), + )) + .into_response() + } + ChallengeError::RateLimited => ApiError::RateLimitExceeded(Some( + "Please wait before requesting a new verification code.".into(), + )) + .into_response(), + ChallengeError::CacheError => { + error!("Cache error during 2FA challenge creation"); + ApiError::InternalError(None).into_response() + } + }; + } + Err(crate::auth::legacy_2fa::Legacy2faFlowError::Validation(e)) => { + use crate::auth::legacy_2fa::ValidationError; + warn!("Invalid 2FA code for {}: {:?}", row.did, e); + let msg = match e { + ValidationError::TooManyAttempts => "Too many attempts. Please request a new code.", + ValidationError::ChallengeExpired => "Code has expired. Please request a new code.", + ValidationError::CacheUnavailable => { + "2FA service temporarily unavailable. Please try again later." + } + ValidationError::ChallengeNotFound + | ValidationError::InvalidCode + | ValidationError::CacheError => "Invalid verification code", + }; + return ApiError::InvalidCode(Some(msg.into())).into_response(); + } } let access_meta = match crate::auth::create_access_token_with_delegation( &row.did, @@ -236,6 +303,11 @@ pub async fn create_session( let handle = full_handle(&row.handle, pds_host); let is_active = account_state.is_active(); let status = account_state.status_for_session().map(String::from); + let email_auth_factor_out = if email_2fa_enabled || has_totp { + Some(true) + } else { + None + }; Json(CreateSessionOutput { access_jwt: access_meta.token, refresh_jwt: refresh_meta.token, @@ -244,6 +316,7 @@ pub async fn create_session( did_doc, email: row.email, email_confirmed: Some(row.channel_verification.email), + email_auth_factor: email_auth_factor_out, active: Some(is_active), status, }) @@ -301,6 +374,9 @@ pub async fn get_session( response["email"] = json!(email_value); response["emailConfirmed"] = json!(email_confirmed_value); } + if row.email_2fa_enabled || row.totp_enabled { + response["emailAuthFactor"] = json!(true); + } if let Some(status) = account_state.status_for_session() { response["status"] = json!(status); } diff --git a/crates/tranquil-pds/src/api/server/totp.rs b/crates/tranquil-pds/src/api/server/totp.rs index d05d4cf..a7f8336 100644 --- a/crates/tranquil-pds/src/api/server/totp.rs +++ b/crates/tranquil-pds/src/api/server/totp.rs @@ -187,6 +187,8 @@ pub async fn disable_totp( .await .log_db_err("deleting TOTP")?; + crate::auth::legacy_2fa::clear_challenge(state.cache.as_ref(), &auth.did).await; + info!(did = %session_mfa.did(), "TOTP disabled (verified via {} and {})", password_mfa.method(), totp_mfa.method()); Ok(EmptyResponse::ok().into_response()) diff --git a/crates/tranquil-pds/src/auth/email_token.rs b/crates/tranquil-pds/src/auth/email_token.rs new file mode 100644 index 0000000..dfed756 --- /dev/null +++ b/crates/tranquil-pds/src/auth/email_token.rs @@ -0,0 +1,303 @@ +use rand::Rng; +use serde::{Deserialize, Serialize}; +use std::time::Duration; + +use crate::cache::Cache; + +const TOKEN_TTL_SECS: u64 = 900; +const BASE32_CHARS: &[u8] = b"ABCDEFGHIJKLMNOPQRSTUVWXYZ234567"; + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum EmailTokenPurpose { + UpdateEmail, + ConfirmEmail, + DeleteAccount, + ResetPassword, + PlcOperation, +} + +impl EmailTokenPurpose { + fn as_str(&self) -> &'static str { + match self { + Self::UpdateEmail => "update_email", + Self::ConfirmEmail => "confirm_email", + Self::DeleteAccount => "delete_account", + Self::ResetPassword => "reset_password", + Self::PlcOperation => "plc_operation", + } + } +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +struct TokenData { + token: String, + created_at: u64, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum TokenError { + CacheUnavailable, + CacheError, + InvalidToken, + ExpiredToken, +} + +fn cache_key(did: &str, purpose: EmailTokenPurpose) -> String { + format!("email_token:{}:{}", purpose.as_str(), did) +} + +fn generate_short_token() -> String { + let mut rng = rand::thread_rng(); + let token: String = (0..10) + .map(|_| BASE32_CHARS[rng.gen_range(0..BASE32_CHARS.len())] as char) + .collect(); + format!("{}-{}", &token[0..5], &token[5..10]) +} + +fn current_timestamp() -> u64 { + chrono::Utc::now().timestamp().max(0) as u64 +} + +pub async fn create_email_token( + cache: &dyn Cache, + did: &str, + purpose: EmailTokenPurpose, +) -> Result { + if !cache.is_available() { + return Err(TokenError::CacheUnavailable); + } + + let token = generate_short_token(); + let data = TokenData { + token: token.clone(), + created_at: current_timestamp(), + }; + + let json = serde_json::to_string(&data).map_err(|_| TokenError::CacheError)?; + + cache + .set( + &cache_key(did, purpose), + &json, + Duration::from_secs(TOKEN_TTL_SECS), + ) + .await + .map_err(|_| TokenError::CacheError)?; + + Ok(token) +} + +pub async fn validate_email_token( + cache: &dyn Cache, + did: &str, + purpose: EmailTokenPurpose, + token: &str, +) -> Result<(), TokenError> { + if !cache.is_available() { + return Err(TokenError::CacheUnavailable); + } + + let key = cache_key(did, purpose); + let json = cache.get(&key).await.ok_or(TokenError::InvalidToken)?; + + let data: TokenData = serde_json::from_str(&json).map_err(|_| TokenError::InvalidToken)?; + + let elapsed = current_timestamp().saturating_sub(data.created_at); + if elapsed > TOKEN_TTL_SECS { + let _ = cache.delete(&key).await; + return Err(TokenError::ExpiredToken); + } + + let normalized_input = token.to_uppercase().replace('-', ""); + let normalized_stored = data.token.to_uppercase().replace('-', ""); + + if !constant_time_eq(normalized_input.as_bytes(), normalized_stored.as_bytes()) { + return Err(TokenError::InvalidToken); + } + + let _ = cache.delete(&key).await; + + Ok(()) +} + +pub async fn delete_email_token(cache: &dyn Cache, did: &str, purpose: EmailTokenPurpose) { + let _ = cache.delete(&cache_key(did, purpose)).await; +} + +fn constant_time_eq(a: &[u8], b: &[u8]) -> bool { + if a.len() != b.len() { + return false; + } + a.iter() + .zip(b.iter()) + .fold(0u8, |acc, (x, y)| acc | (x ^ y)) + == 0 +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::cache::CacheError; + use async_trait::async_trait; + use std::collections::HashMap; + use std::sync::Mutex; + + struct MockCache { + data: Mutex>, + } + + impl MockCache { + fn new() -> Self { + Self { + data: Mutex::new(HashMap::new()), + } + } + } + + #[async_trait] + impl Cache for MockCache { + async fn get(&self, key: &str) -> Option { + let data = self.data.lock().unwrap(); + let now = current_timestamp(); + data.get(key) + .filter(|(_, exp)| *exp > now) + .map(|(v, _)| v.clone()) + } + + async fn set(&self, key: &str, value: &str, ttl: Duration) -> Result<(), CacheError> { + let mut data = self.data.lock().unwrap(); + let expires = current_timestamp() + ttl.as_secs(); + data.insert(key.to_string(), (value.to_string(), expires)); + Ok(()) + } + + async fn delete(&self, key: &str) -> Result<(), CacheError> { + let mut data = self.data.lock().unwrap(); + data.remove(key); + Ok(()) + } + + async fn get_bytes(&self, _key: &str) -> Option> { + None + } + + async fn set_bytes( + &self, + _key: &str, + _value: &[u8], + _ttl: Duration, + ) -> Result<(), CacheError> { + Ok(()) + } + + fn is_available(&self) -> bool { + true + } + } + + #[tokio::test] + async fn test_create_and_validate_token() { + let cache = MockCache::new(); + let did = "did:plc:test123"; + + let token = create_email_token(&cache, did, EmailTokenPurpose::UpdateEmail) + .await + .unwrap(); + + assert_eq!(token.len(), 11); + assert!(token.contains('-')); + + let result = + validate_email_token(&cache, did, EmailTokenPurpose::UpdateEmail, &token).await; + assert!(result.is_ok()); + } + + #[tokio::test] + async fn test_token_consumed_after_use() { + let cache = MockCache::new(); + let did = "did:plc:test123"; + + let token = create_email_token(&cache, did, EmailTokenPurpose::UpdateEmail) + .await + .unwrap(); + + validate_email_token(&cache, did, EmailTokenPurpose::UpdateEmail, &token) + .await + .unwrap(); + + let result = + validate_email_token(&cache, did, EmailTokenPurpose::UpdateEmail, &token).await; + assert_eq!(result.unwrap_err(), TokenError::InvalidToken); + } + + #[tokio::test] + async fn test_invalid_token_rejected() { + let cache = MockCache::new(); + let did = "did:plc:test123"; + + let _token = create_email_token(&cache, did, EmailTokenPurpose::UpdateEmail) + .await + .unwrap(); + + let result = + validate_email_token(&cache, did, EmailTokenPurpose::UpdateEmail, "XXXXX-XXXXX").await; + assert_eq!(result.unwrap_err(), TokenError::InvalidToken); + } + + #[tokio::test] + async fn test_wrong_purpose_rejected() { + let cache = MockCache::new(); + let did = "did:plc:test123"; + + let token = create_email_token(&cache, did, EmailTokenPurpose::UpdateEmail) + .await + .unwrap(); + + let result = + validate_email_token(&cache, did, EmailTokenPurpose::ConfirmEmail, &token).await; + assert_eq!(result.unwrap_err(), TokenError::InvalidToken); + } + + #[tokio::test] + async fn test_token_format() { + (0..100).for_each(|_| { + let token = generate_short_token(); + assert_eq!(token.len(), 11); + assert_eq!(&token[5..6], "-"); + assert!( + token[0..5] + .chars() + .all(|c| BASE32_CHARS.contains(&(c as u8))) + ); + assert!( + token[6..11] + .chars() + .all(|c| BASE32_CHARS.contains(&(c as u8))) + ); + }); + } + + #[tokio::test] + async fn test_case_insensitive_validation() { + let cache = MockCache::new(); + let did = "did:plc:test123"; + + let token = create_email_token(&cache, did, EmailTokenPurpose::UpdateEmail) + .await + .unwrap(); + + let lowercase = token.to_lowercase(); + let result = + validate_email_token(&cache, did, EmailTokenPurpose::UpdateEmail, &lowercase).await; + assert!(result.is_ok()); + } + + #[tokio::test] + async fn test_noop_cache_returns_unavailable() { + let cache = crate::cache::NoOpCache; + let did = "did:plc:test"; + + let result = create_email_token(&cache, did, EmailTokenPurpose::UpdateEmail).await; + assert_eq!(result.unwrap_err(), TokenError::CacheUnavailable); + } +} diff --git a/crates/tranquil-pds/src/auth/legacy_2fa.rs b/crates/tranquil-pds/src/auth/legacy_2fa.rs new file mode 100644 index 0000000..d6906d3 --- /dev/null +++ b/crates/tranquil-pds/src/auth/legacy_2fa.rs @@ -0,0 +1,514 @@ +use chrono::Utc; +use rand::Rng; +use serde::{Deserialize, Serialize}; +use std::time::Duration; + +use crate::cache::Cache; +use crate::types::Did; + +const CHALLENGE_TTL_SECS: u64 = 300; +const MIN_REMAINING_TTL_SECS: u64 = 10; +const MAX_ATTEMPTS: u8 = 5; +const CODE_LENGTH: usize = 8; +const COOLDOWN_SECS: u64 = 60; + +#[derive(Debug, Clone, Serialize, Deserialize)] +struct ChallengeData { + code: String, + attempts: u8, + created_at: u64, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum ChallengeError { + CacheUnavailable, + RateLimited, + CacheError, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum ValidationError { + InvalidCode, + TooManyAttempts, + ChallengeNotFound, + ChallengeExpired, + CacheUnavailable, + CacheError, +} + +#[derive(Debug)] +pub struct ChallengeCode(String); + +impl ChallengeCode { + pub fn as_str(&self) -> &str { + &self.0 + } +} + +impl std::fmt::Display for ChallengeCode { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + write!(f, "{}", self.0) + } +} + +pub async fn create_challenge( + cache: &dyn Cache, + did: &Did, +) -> Result { + create_challenge_code(cache, did).await +} + +pub async fn clear_challenge(cache: &dyn Cache, did: &Did) { + let _ = cache.delete(&challenge_key(did.as_str())).await; + let _ = cache.delete(&cooldown_key(did.as_str())).await; +} + +async fn validate_challenge_internal( + cache: &dyn Cache, + did: &str, + code: &str, +) -> Result<(), ValidationError> { + if !cache.is_available() { + return Err(ValidationError::CacheUnavailable); + } + + let challenge_k = challenge_key(did); + + let json = cache + .get(&challenge_k) + .await + .ok_or(ValidationError::ChallengeNotFound)?; + + let data: ChallengeData = + serde_json::from_str(&json).map_err(|_| ValidationError::ChallengeNotFound)?; + + if data.attempts >= MAX_ATTEMPTS { + let _ = cache.delete(&challenge_k).await; + return Err(ValidationError::TooManyAttempts); + } + + let elapsed = current_timestamp().saturating_sub(data.created_at); + let remaining_ttl = CHALLENGE_TTL_SECS.saturating_sub(elapsed); + if remaining_ttl < MIN_REMAINING_TTL_SECS { + let _ = cache.delete(&challenge_k).await; + return Err(ValidationError::ChallengeExpired); + } + + if !constant_time_eq(code.as_bytes(), data.code.as_bytes()) { + let updated = ChallengeData { + code: data.code, + attempts: data.attempts + 1, + created_at: data.created_at, + }; + let updated_json = + serde_json::to_string(&updated).map_err(|_| ValidationError::CacheError)?; + cache + .set( + &challenge_k, + &updated_json, + Duration::from_secs(remaining_ttl), + ) + .await + .map_err(|_| ValidationError::CacheError)?; + return Err(ValidationError::InvalidCode); + } + + let _ = cache.delete(&challenge_k).await; + let _ = cache.delete(&cooldown_key(did)).await; + + Ok(()) +} + +fn challenge_key(did: &str) -> String { + format!("legacy_2fa:{}", did) +} + +fn cooldown_key(did: &str) -> String { + format!("legacy_2fa_cooldown:{}", did) +} + +fn generate_code() -> String { + let mut rng = rand::thread_rng(); + (0..CODE_LENGTH) + .map(|_| rng.gen_range(0..10).to_string()) + .collect() +} + +fn current_timestamp() -> u64 { + Utc::now().timestamp().max(0) as u64 +} + +fn constant_time_eq(a: &[u8], b: &[u8]) -> bool { + if a.len() != b.len() { + return false; + } + a.iter() + .zip(b.iter()) + .fold(0u8, |acc, (x, y)| acc | (x ^ y)) + == 0 +} + +pub enum Legacy2faOutcome { + NotRequired, + Blocked, + ChallengeSent(ChallengeCode), + Verified, +} + +pub struct Legacy2faContext { + pub email_2fa_enabled: bool, + pub has_totp: bool, + pub allow_legacy_login: bool, +} + +impl Legacy2faContext { + pub fn requires_2fa(&self) -> bool { + self.email_2fa_enabled || self.has_totp + } + + pub fn is_blocked(&self) -> bool { + self.has_totp && !self.allow_legacy_login && !self.email_2fa_enabled + } +} + +pub async fn process_legacy_2fa( + cache: &dyn Cache, + did: &Did, + ctx: &Legacy2faContext, + auth_factor_token: Option<&str>, +) -> Result { + if !ctx.requires_2fa() { + return Ok(Legacy2faOutcome::NotRequired); + } + + if ctx.is_blocked() { + return Ok(Legacy2faOutcome::Blocked); + } + + match auth_factor_token.filter(|t| !t.is_empty()) { + None => { + let code = create_challenge_code(cache, did).await?; + Ok(Legacy2faOutcome::ChallengeSent(code)) + } + Some(token) => { + validate_challenge(cache, did, token).await?; + Ok(Legacy2faOutcome::Verified) + } + } +} + +pub async fn validate_challenge( + cache: &dyn Cache, + did: &Did, + code: &str, +) -> Result<(), ValidationError> { + validate_challenge_internal(cache, did.as_str(), code).await +} + +async fn create_challenge_code( + cache: &dyn Cache, + did: &Did, +) -> Result { + if !cache.is_available() { + return Err(ChallengeError::CacheUnavailable); + } + + let cooldown = cooldown_key(did.as_str()); + if cache.get(&cooldown).await.is_some() { + return Err(ChallengeError::RateLimited); + } + + let code = generate_code(); + let now = current_timestamp(); + + let data = ChallengeData { + code: code.clone(), + attempts: 0, + created_at: now, + }; + + let json = serde_json::to_string(&data).map_err(|_| ChallengeError::CacheError)?; + + cache + .set( + &challenge_key(did.as_str()), + &json, + Duration::from_secs(CHALLENGE_TTL_SECS), + ) + .await + .map_err(|_| ChallengeError::CacheError)?; + + cache + .set(&cooldown, "1", Duration::from_secs(COOLDOWN_SECS)) + .await + .map_err(|_| ChallengeError::CacheError)?; + + Ok(ChallengeCode(code)) +} + +#[derive(Debug)] +pub enum Legacy2faFlowError { + Challenge(ChallengeError), + Validation(ValidationError), +} + +impl From for Legacy2faFlowError { + fn from(e: ChallengeError) -> Self { + Self::Challenge(e) + } +} + +impl From for Legacy2faFlowError { + fn from(e: ValidationError) -> Self { + Self::Validation(e) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::cache::CacheError; + use async_trait::async_trait; + use std::collections::HashMap; + use std::sync::Mutex; + + struct MockCache { + data: Mutex>, + } + + impl MockCache { + fn new() -> Self { + Self { + data: Mutex::new(HashMap::new()), + } + } + } + + #[async_trait] + impl Cache for MockCache { + async fn get(&self, key: &str) -> Option { + let data = self.data.lock().unwrap(); + let now = current_timestamp(); + data.get(key) + .filter(|(_, exp)| *exp > now) + .map(|(v, _)| v.clone()) + } + + async fn set(&self, key: &str, value: &str, ttl: Duration) -> Result<(), CacheError> { + let mut data = self.data.lock().unwrap(); + let expires = current_timestamp() + ttl.as_secs(); + data.insert(key.to_string(), (value.to_string(), expires)); + Ok(()) + } + + async fn delete(&self, key: &str) -> Result<(), CacheError> { + let mut data = self.data.lock().unwrap(); + data.remove(key); + Ok(()) + } + + async fn get_bytes(&self, _key: &str) -> Option> { + None + } + + async fn set_bytes( + &self, + _key: &str, + _value: &[u8], + _ttl: Duration, + ) -> Result<(), CacheError> { + Ok(()) + } + + fn is_available(&self) -> bool { + true + } + } + + #[tokio::test] + async fn test_create_and_validate_challenge() { + let cache = MockCache::new(); + let did = Did::new("did:plc:test123".to_string()).unwrap(); + + let code = create_challenge(&cache, &did).await.unwrap(); + assert_eq!(code.as_str().len(), CODE_LENGTH); + + let result = validate_challenge(&cache, &did, code.as_str()).await; + assert!(result.is_ok()); + } + + #[tokio::test] + async fn test_invalid_code_rejected() { + let cache = MockCache::new(); + let did = Did::new("did:plc:test123".to_string()).unwrap(); + + let _code = create_challenge(&cache, &did).await.unwrap(); + let result = validate_challenge(&cache, &did, "00000000").await; + assert_eq!(result.unwrap_err(), ValidationError::InvalidCode); + } + + #[tokio::test] + async fn test_challenge_consumed_on_success() { + let cache = MockCache::new(); + let did = Did::new("did:plc:test123".to_string()).unwrap(); + + let code = create_challenge(&cache, &did).await.unwrap(); + validate_challenge(&cache, &did, code.as_str()) + .await + .unwrap(); + + let result = validate_challenge(&cache, &did, code.as_str()).await; + assert_eq!(result.unwrap_err(), ValidationError::ChallengeNotFound); + } + + #[tokio::test] + async fn test_max_attempts_exceeded() { + let cache = MockCache::new(); + let did = Did::new("did:plc:test123".to_string()).unwrap(); + + let _code = create_challenge(&cache, &did).await.unwrap(); + + (0..MAX_ATTEMPTS).for_each(|_| { + let _ = futures::executor::block_on(validate_challenge(&cache, &did, "wrong123")); + }); + + let result = validate_challenge(&cache, &did, "anything").await; + assert_eq!(result.unwrap_err(), ValidationError::TooManyAttempts); + } + + #[tokio::test] + async fn test_rate_limiting() { + let cache = MockCache::new(); + let did = Did::new("did:plc:test123".to_string()).unwrap(); + + let _first = create_challenge(&cache, &did).await.unwrap(); + let result = create_challenge(&cache, &did).await; + assert_eq!(result.unwrap_err(), ChallengeError::RateLimited); + } + + #[tokio::test] + async fn test_noop_cache_returns_unavailable() { + let cache = crate::cache::NoOpCache; + let did = Did::new("did:plc:test".to_string()).unwrap(); + + let result = create_challenge(&cache, &did).await; + assert_eq!(result.unwrap_err(), ChallengeError::CacheUnavailable); + } + + #[tokio::test] + async fn test_code_generation_is_numeric() { + (0..100).for_each(|_| { + let code = generate_code(); + assert!(code.chars().all(|c| c.is_ascii_digit())); + assert_eq!(code.len(), CODE_LENGTH); + }); + } + + #[tokio::test] + async fn test_constant_time_eq() { + assert!(constant_time_eq(b"12345678", b"12345678")); + assert!(!constant_time_eq(b"12345678", b"12345679")); + assert!(!constant_time_eq(b"12345678", b"1234567")); + assert!(!constant_time_eq(b"", b"1")); + assert!(constant_time_eq(b"", b"")); + } + + #[tokio::test] + async fn test_process_flow_not_required() { + let cache = MockCache::new(); + let did = Did::new("did:plc:test".to_string()).unwrap(); + let ctx = Legacy2faContext { + email_2fa_enabled: false, + has_totp: false, + allow_legacy_login: true, + }; + + let outcome = process_legacy_2fa(&cache, &did, &ctx, None).await.unwrap(); + assert!(matches!(outcome, Legacy2faOutcome::NotRequired)); + } + + #[tokio::test] + async fn test_process_flow_blocked() { + let cache = MockCache::new(); + let did = Did::new("did:plc:test".to_string()).unwrap(); + let ctx = Legacy2faContext { + email_2fa_enabled: false, + has_totp: true, + allow_legacy_login: false, + }; + + let outcome = process_legacy_2fa(&cache, &did, &ctx, None).await.unwrap(); + assert!(matches!(outcome, Legacy2faOutcome::Blocked)); + } + + #[tokio::test] + async fn test_process_flow_challenge_sent_totp() { + let cache = MockCache::new(); + let did = Did::new("did:plc:test".to_string()).unwrap(); + let ctx = Legacy2faContext { + email_2fa_enabled: false, + has_totp: true, + allow_legacy_login: true, + }; + + let outcome = process_legacy_2fa(&cache, &did, &ctx, None).await.unwrap(); + assert!(matches!(outcome, Legacy2faOutcome::ChallengeSent(_))); + } + + #[tokio::test] + async fn test_process_flow_challenge_sent_email_2fa_enabled() { + let cache = MockCache::new(); + let did = Did::new("did:plc:test2".to_string()).unwrap(); + let ctx = Legacy2faContext { + email_2fa_enabled: true, + has_totp: false, + allow_legacy_login: false, + }; + + let outcome = process_legacy_2fa(&cache, &did, &ctx, None).await.unwrap(); + assert!(matches!(outcome, Legacy2faOutcome::ChallengeSent(_))); + } + + #[tokio::test] + async fn test_process_flow_verified() { + let cache = MockCache::new(); + let did = Did::new("did:plc:test".to_string()).unwrap(); + let ctx = Legacy2faContext { + email_2fa_enabled: true, + has_totp: false, + allow_legacy_login: false, + }; + + let code = create_challenge(&cache, &did).await.unwrap(); + + let outcome = process_legacy_2fa(&cache, &did, &ctx, Some(code.as_str())) + .await + .unwrap(); + assert!(matches!(outcome, Legacy2faOutcome::Verified)); + } + + #[tokio::test] + async fn test_attempts_persist_across_failures() { + let cache = MockCache::new(); + let did = Did::new("did:plc:test123".to_string()).unwrap(); + + let code = create_challenge(&cache, &did).await.unwrap(); + + (0..3).for_each(|_| { + let result = futures::executor::block_on(validate_challenge(&cache, &did, "wrong123")); + assert_eq!(result.unwrap_err(), ValidationError::InvalidCode); + }); + + let result = validate_challenge(&cache, &did, code.as_str()).await; + assert!(result.is_ok()); + } + + #[tokio::test] + async fn test_validation_on_noop_cache_returns_unavailable() { + let cache = crate::cache::NoOpCache; + let did = Did::new("did:plc:test".to_string()).unwrap(); + + let result = validate_challenge(&cache, &did, "12345678").await; + assert_eq!(result.unwrap_err(), ValidationError::CacheUnavailable); + } +} diff --git a/crates/tranquil-pds/src/auth/mod.rs b/crates/tranquil-pds/src/auth/mod.rs index 05e91e6..3a0af7e 100644 --- a/crates/tranquil-pds/src/auth/mod.rs +++ b/crates/tranquil-pds/src/auth/mod.rs @@ -11,7 +11,9 @@ use tranquil_db::UserRepository; use tranquil_db_traits::OAuthRepository; pub mod account_verified; +pub mod email_token; pub mod extractor; +pub mod legacy_2fa; pub mod login_identifier; pub mod mfa_verified; pub mod scope_check; diff --git a/crates/tranquil-pds/src/comms/service.rs b/crates/tranquil-pds/src/comms/service.rs index 0dd0792..e5d23d9 100644 --- a/crates/tranquil-pds/src/comms/service.rs +++ b/crates/tranquil-pds/src/comms/service.rs @@ -403,6 +403,57 @@ pub mod repo { .await } + pub async fn enqueue_short_token_email( + user_repo: &dyn UserRepository, + infra_repo: &dyn InfraRepository, + user_id: Uuid, + token: &str, + purpose: &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.clone().unwrap_or_default(); + + let (subject_template, body_template, comms_type) = match purpose { + "email_update" => ( + strings.email_update_subject, + strings.short_token_body, + CommsType::EmailUpdate, + ), + _ => ( + strings.email_update_subject, + strings.short_token_body, + CommsType::EmailUpdate, + ), + }; + + let verify_page = format!("https://{}/app/settings", hostname); + let body = format_message( + body_template, + &[ + ("handle", &prefs.handle), + ("code", token), + ("verify_page", &verify_page), + ], + ); + let subject = format_message(subject_template, &[("hostname", hostname)]); + infra_repo + .enqueue_comms( + Some(user_id), + tranquil_db_traits::CommsChannel::Email, + comms_type, + ¤t_email, + Some(&subject), + &body, + None, + ) + .await + } + pub async fn enqueue_account_deletion( user_repo: &dyn UserRepository, infra_repo: &dyn InfraRepository, diff --git a/crates/tranquil-pds/tests/actor.rs b/crates/tranquil-pds/tests/actor.rs index 331011d..a54e654 100644 --- a/crates/tranquil-pds/tests/actor.rs +++ b/crates/tranquil-pds/tests/actor.rs @@ -115,14 +115,16 @@ async fn test_put_preferences_multiple_same_type() { let body: Value = resp.json().await.unwrap(); let prefs_arr = body["preferences"].as_array().unwrap(); assert_eq!(prefs_arr.len(), 3); - let adult_pref = prefs_arr - .iter() - .find(|p| p.get("$type").and_then(|t| t.as_str()) == Some("app.bsky.actor.defs#adultContentPref")); + let adult_pref = prefs_arr.iter().find(|p| { + p.get("$type").and_then(|t| t.as_str()) == Some("app.bsky.actor.defs#adultContentPref") + }); assert!(adult_pref.is_some()); assert_eq!(adult_pref.unwrap()["enabled"], false); let content_label_prefs: Vec<&Value> = prefs_arr .iter() - .filter(|p| p.get("$type").and_then(|t| t.as_str()) == Some("app.bsky.actor.defs#contentLabelPref")) + .filter(|p| { + p.get("$type").and_then(|t| t.as_str()) == Some("app.bsky.actor.defs#contentLabelPref") + }) .collect(); assert_eq!(content_label_prefs.len(), 2); let dogs_pref = content_label_prefs diff --git a/crates/tranquil-pds/tests/legacy_2fa.rs b/crates/tranquil-pds/tests/legacy_2fa.rs new file mode 100644 index 0000000..b917cf6 --- /dev/null +++ b/crates/tranquil-pds/tests/legacy_2fa.rs @@ -0,0 +1,481 @@ +mod common; + +use common::{base_url, client, create_account_and_login, get_test_db_pool}; +use reqwest::StatusCode; +use serde_json::{Value, json}; + +async fn enable_totp_for_user(did: &str) { + let pool = get_test_db_pool().await; + let secret = vec![0u8; 20]; + sqlx::query( + r#"INSERT INTO user_totp (did, secret_encrypted, encryption_version, verified, created_at) + VALUES ($1, $2, 1, TRUE, NOW()) + ON CONFLICT (did) DO UPDATE SET verified = TRUE"#, + ) + .bind(did) + .bind(&secret) + .execute(pool) + .await + .expect("Failed to enable TOTP"); +} + +async fn set_allow_legacy_login(did: &str, allow: bool) { + let pool = get_test_db_pool().await; + sqlx::query("UPDATE users SET allow_legacy_login = $1 WHERE did = $2") + .bind(allow) + .bind(did) + .execute(pool) + .await + .expect("Failed to set allow_legacy_login"); +} + +async fn get_2fa_code_from_queue(did: &str) -> Option { + let pool = get_test_db_pool().await; + let row: Option<(String,)> = sqlx::query_as( + r#"SELECT body FROM comms_queue + WHERE user_id = (SELECT id FROM users WHERE did = $1) + AND comms_type = 'two_factor_code' + ORDER BY created_at DESC LIMIT 1"#, + ) + .bind(did) + .fetch_optional(pool) + .await + .ok() + .flatten(); + + row.and_then(|(body,)| { + body.lines() + .find(|line: &&str| line.chars().all(|c: char| c.is_ascii_digit()) && line.len() == 8) + .map(|s: &str| s.to_string()) + .or_else(|| { + body.split_whitespace() + .find(|word: &&str| { + word.chars().all(|c: char| c.is_ascii_digit()) && word.len() == 8 + }) + .map(|s: &str| s.to_string()) + }) + }) +} + +async fn clear_2fa_challenges_for_user(did: &str) { + let pool = get_test_db_pool().await; + let _ = sqlx::query( + "DELETE FROM comms_queue WHERE user_id = (SELECT id FROM users WHERE did = $1) AND comms_type = 'two_factor_code'", + ) + .bind(did) + .execute(pool) + .await; +} + +async fn set_email_auth_factor(did: &str, enabled: bool) { + let pool = get_test_db_pool().await; + let user_id: uuid::Uuid = + sqlx::query_scalar::<_, uuid::Uuid>("SELECT id FROM users WHERE did = $1") + .bind(did) + .fetch_one(pool) + .await + .expect("Failed to get user id"); + let pool = get_test_db_pool().await; + let _ = sqlx::query( + "DELETE FROM account_preferences WHERE user_id = $1 AND name = 'email_auth_factor'", + ) + .bind(user_id) + .execute(pool) + .await; + let pool = get_test_db_pool().await; + sqlx::query( + "INSERT INTO account_preferences (user_id, name, value_json) VALUES ($1, 'email_auth_factor', $2::jsonb)", + ) + .bind(user_id) + .bind(serde_json::json!(enabled)) + .execute(pool) + .await + .expect("Failed to set email_auth_factor"); +} + +#[tokio::test] +async fn test_legacy_2fa_auth_factor_required() { + let client = client(); + let base = base_url().await; + let (_token, did) = create_account_and_login(&client).await; + + enable_totp_for_user(&did).await; + set_allow_legacy_login(&did, true).await; + + let pool = get_test_db_pool().await; + let handle: String = sqlx::query_scalar::<_, String>("SELECT handle FROM users WHERE did = $1") + .bind(&did) + .fetch_one(pool) + .await + .expect("Failed to get handle"); + + let login_payload = json!({ + "identifier": handle, + "password": "Testpass123!" + }); + let resp = client + .post(format!("{}/xrpc/com.atproto.server.createSession", base)) + .json(&login_payload) + .send() + .await + .unwrap(); + + assert_eq!(resp.status(), StatusCode::BAD_REQUEST); + let body: Value = resp.json().await.unwrap(); + assert_eq!(body["error"], "AuthFactorTokenRequired"); + assert!( + body["message"] + .as_str() + .unwrap_or("") + .contains("sign in code") + ); +} + +#[tokio::test] +async fn test_legacy_2fa_valid_code_succeeds() { + let client = client(); + let base = base_url().await; + let (_token, did) = create_account_and_login(&client).await; + + enable_totp_for_user(&did).await; + set_allow_legacy_login(&did, true).await; + clear_2fa_challenges_for_user(&did).await; + + let pool = get_test_db_pool().await; + let handle: String = sqlx::query_scalar::<_, String>("SELECT handle FROM users WHERE did = $1") + .bind(&did) + .fetch_one(pool) + .await + .expect("Failed to get handle"); + + let login_payload = json!({ + "identifier": handle, + "password": "Testpass123!" + }); + let resp = client + .post(format!("{}/xrpc/com.atproto.server.createSession", base)) + .json(&login_payload) + .send() + .await + .unwrap(); + assert_eq!(resp.status(), StatusCode::BAD_REQUEST); + + tokio::time::sleep(tokio::time::Duration::from_millis(100)).await; + let code = get_2fa_code_from_queue(&did) + .await + .expect("2FA code should be in queue"); + + let login_with_code = json!({ + "identifier": handle, + "password": "Testpass123!", + "authFactorToken": code + }); + let resp = client + .post(format!("{}/xrpc/com.atproto.server.createSession", base)) + .json(&login_with_code) + .send() + .await + .unwrap(); + + assert_eq!(resp.status(), StatusCode::OK); + let body: Value = resp.json().await.unwrap(); + assert!(body.get("accessJwt").is_some()); + assert!(body.get("refreshJwt").is_some()); + assert_eq!(body["did"], did); +} + +#[tokio::test] +async fn test_legacy_2fa_invalid_code_rejected() { + let client = client(); + let base = base_url().await; + let (_token, did) = create_account_and_login(&client).await; + + enable_totp_for_user(&did).await; + set_allow_legacy_login(&did, true).await; + clear_2fa_challenges_for_user(&did).await; + + let pool = get_test_db_pool().await; + let handle: String = sqlx::query_scalar::<_, String>("SELECT handle FROM users WHERE did = $1") + .bind(&did) + .fetch_one(pool) + .await + .expect("Failed to get handle"); + + let resp = client + .post(format!("{}/xrpc/com.atproto.server.createSession", base)) + .json(&json!({ + "identifier": handle, + "password": "Testpass123!" + })) + .send() + .await + .unwrap(); + assert_eq!(resp.status(), StatusCode::BAD_REQUEST); + + let login_with_bad_code = json!({ + "identifier": handle, + "password": "Testpass123!", + "authFactorToken": "00000000" + }); + let resp = client + .post(format!("{}/xrpc/com.atproto.server.createSession", base)) + .json(&login_with_bad_code) + .send() + .await + .unwrap(); + + let status = resp.status(); + let body: Value = resp.json().await.unwrap(); + assert_eq!( + status, + StatusCode::BAD_REQUEST, + "Expected 400, got {}. Response: {:?}", + status, + body + ); + assert_eq!(body["error"], "InvalidCode"); +} + +#[tokio::test] +async fn test_legacy_2fa_blocked_when_disabled() { + let client = client(); + let base = base_url().await; + let (_token, did) = create_account_and_login(&client).await; + + enable_totp_for_user(&did).await; + set_allow_legacy_login(&did, false).await; + + let pool = get_test_db_pool().await; + let handle: String = sqlx::query_scalar::<_, String>("SELECT handle FROM users WHERE did = $1") + .bind(&did) + .fetch_one(pool) + .await + .expect("Failed to get handle"); + + let login_payload = json!({ + "identifier": handle, + "password": "Testpass123!" + }); + let resp = client + .post(format!("{}/xrpc/com.atproto.server.createSession", base)) + .json(&login_payload) + .send() + .await + .unwrap(); + + assert_eq!(resp.status(), StatusCode::BAD_REQUEST); + let body: Value = resp.json().await.unwrap(); + assert_eq!(body["error"], "MfaRequired"); +} + +#[tokio::test] +async fn test_legacy_2fa_no_totp_no_challenge() { + let client = client(); + let base = base_url().await; + let (_token, did) = create_account_and_login(&client).await; + + let pool = get_test_db_pool().await; + let handle: String = sqlx::query_scalar::<_, String>("SELECT handle FROM users WHERE did = $1") + .bind(&did) + .fetch_one(pool) + .await + .expect("Failed to get handle"); + + let login_payload = json!({ + "identifier": handle, + "password": "Testpass123!" + }); + let resp = client + .post(format!("{}/xrpc/com.atproto.server.createSession", base)) + .json(&login_payload) + .send() + .await + .unwrap(); + + assert_eq!(resp.status(), StatusCode::OK); + let body: Value = resp.json().await.unwrap(); + assert!(body.get("accessJwt").is_some()); +} + +#[tokio::test] +async fn test_legacy_2fa_code_consumed_after_use() { + let client = client(); + let base = base_url().await; + let (_token, did) = create_account_and_login(&client).await; + + enable_totp_for_user(&did).await; + set_allow_legacy_login(&did, true).await; + clear_2fa_challenges_for_user(&did).await; + + let pool = get_test_db_pool().await; + let handle: String = sqlx::query_scalar::<_, String>("SELECT handle FROM users WHERE did = $1") + .bind(&did) + .fetch_one(pool) + .await + .expect("Failed to get handle"); + + let resp = client + .post(format!("{}/xrpc/com.atproto.server.createSession", base)) + .json(&json!({ + "identifier": handle, + "password": "Testpass123!" + })) + .send() + .await + .unwrap(); + assert_eq!(resp.status(), StatusCode::BAD_REQUEST); + + tokio::time::sleep(tokio::time::Duration::from_millis(100)).await; + let code = get_2fa_code_from_queue(&did) + .await + .expect("2FA code should be in queue"); + + let resp = client + .post(format!("{}/xrpc/com.atproto.server.createSession", base)) + .json(&json!({ + "identifier": handle, + "password": "Testpass123!", + "authFactorToken": code + })) + .send() + .await + .unwrap(); + assert_eq!(resp.status(), StatusCode::OK); + + clear_2fa_challenges_for_user(&did).await; + let resp = client + .post(format!("{}/xrpc/com.atproto.server.createSession", base)) + .json(&json!({ + "identifier": handle, + "password": "Testpass123!" + })) + .send() + .await + .unwrap(); + assert_eq!(resp.status(), StatusCode::BAD_REQUEST); + let body: Value = resp.json().await.unwrap(); + assert_eq!(body["error"], "AuthFactorTokenRequired"); + + tokio::time::sleep(tokio::time::Duration::from_millis(100)).await; + let new_code = get_2fa_code_from_queue(&did) + .await + .expect("New 2FA code should be in queue"); + + let resp = client + .post(format!("{}/xrpc/com.atproto.server.createSession", base)) + .json(&json!({ + "identifier": handle, + "password": "Testpass123!", + "authFactorToken": code + })) + .send() + .await + .unwrap(); + let status = resp.status(); + let body: Value = resp.json().await.unwrap(); + assert_eq!( + status, + StatusCode::BAD_REQUEST, + "Expected 400 for old code, got {}. Response: {:?}", + status, + body + ); + assert_eq!(body["error"], "InvalidCode"); + + let resp = client + .post(format!("{}/xrpc/com.atproto.server.createSession", base)) + .json(&json!({ + "identifier": handle, + "password": "Testpass123!", + "authFactorToken": new_code + })) + .send() + .await + .unwrap(); + assert_eq!(resp.status(), StatusCode::OK); +} + +#[tokio::test] +async fn test_email_auth_factor_requires_code() { + let client = client(); + let base = base_url().await; + let (_token, did) = create_account_and_login(&client).await; + + set_email_auth_factor(&did, true).await; + clear_2fa_challenges_for_user(&did).await; + + let pool = get_test_db_pool().await; + let handle: String = sqlx::query_scalar::<_, String>("SELECT handle FROM users WHERE did = $1") + .bind(&did) + .fetch_one(pool) + .await + .expect("Failed to get handle"); + + let login_payload = json!({ + "identifier": handle, + "password": "Testpass123!" + }); + let resp = client + .post(format!("{}/xrpc/com.atproto.server.createSession", base)) + .json(&login_payload) + .send() + .await + .unwrap(); + + assert_eq!(resp.status(), StatusCode::BAD_REQUEST); + let body: Value = resp.json().await.unwrap(); + assert_eq!(body["error"], "AuthFactorTokenRequired"); + + tokio::time::sleep(tokio::time::Duration::from_millis(100)).await; + let code = get_2fa_code_from_queue(&did) + .await + .expect("2FA code should be in queue"); + + let login_with_code = json!({ + "identifier": handle, + "password": "Testpass123!", + "authFactorToken": code + }); + let resp = client + .post(format!("{}/xrpc/com.atproto.server.createSession", base)) + .json(&login_with_code) + .send() + .await + .unwrap(); + + assert_eq!(resp.status(), StatusCode::OK); + let body: Value = resp.json().await.unwrap(); + assert!(body.get("accessJwt").is_some()); + assert_eq!(body["emailAuthFactor"], true); +} + +#[tokio::test] +async fn test_email_auth_factor_disabled_no_challenge() { + let client = client(); + let base = base_url().await; + let (_token, did) = create_account_and_login(&client).await; + + set_email_auth_factor(&did, false).await; + + let pool = get_test_db_pool().await; + let handle: String = sqlx::query_scalar::<_, String>("SELECT handle FROM users WHERE did = $1") + .bind(&did) + .fetch_one(pool) + .await + .expect("Failed to get handle"); + + let login_payload = json!({ + "identifier": handle, + "password": "Testpass123!" + }); + let resp = client + .post(format!("{}/xrpc/com.atproto.server.createSession", base)) + .json(&login_payload) + .send() + .await + .unwrap(); + + assert_eq!(resp.status(), StatusCode::OK); + let body: Value = resp.json().await.unwrap(); + assert!(body.get("accessJwt").is_some()); +} diff --git a/crates/tranquil-pds/tests/shutdown_unit.rs b/crates/tranquil-pds/tests/shutdown_unit.rs index 4067943..8eca3de 100644 --- a/crates/tranquil-pds/tests/shutdown_unit.rs +++ b/crates/tranquil-pds/tests/shutdown_unit.rs @@ -60,11 +60,7 @@ async fn test_cancelled_future_completes_on_cancel() { shutdown.cancel(); - let result = tokio::time::timeout( - std::time::Duration::from_millis(100), - handle, - ) - .await; + let result = tokio::time::timeout(std::time::Duration::from_millis(100), handle).await; assert!(result.is_ok()); assert!(result.unwrap().unwrap());