diff --git a/.sqlx/query-05fd99170e31e68fa5028c862417cdf535cd70e09fde0a8a28249df0070eb2fc.json b/.sqlx/query-05fd99170e31e68fa5028c862417cdf535cd70e09fde0a8a28249df0070eb2fc.json deleted file mode 100644 index 15ba9ce..0000000 --- a/.sqlx/query-05fd99170e31e68fa5028c862417cdf535cd70e09fde0a8a28249df0070eb2fc.json +++ /dev/null @@ -1,22 +0,0 @@ -{ - "db_name": "PostgreSQL", - "query": "SELECT t.token FROM plc_operation_tokens t JOIN users u ON t.user_id = u.id WHERE u.did = $1", - "describe": { - "columns": [ - { - "ordinal": 0, - "name": "token", - "type_info": "Text" - } - ], - "parameters": { - "Left": [ - "Text" - ] - }, - "nullable": [ - false - ] - }, - "hash": "05fd99170e31e68fa5028c862417cdf535cd70e09fde0a8a28249df0070eb2fc" -} diff --git a/.sqlx/query-0710b57fb9aa933525f617b15e6e2e5feaa9c59c38ec9175568abdacda167107.json b/.sqlx/query-0710b57fb9aa933525f617b15e6e2e5feaa9c59c38ec9175568abdacda167107.json deleted file mode 100644 index 734848a..0000000 --- a/.sqlx/query-0710b57fb9aa933525f617b15e6e2e5feaa9c59c38ec9175568abdacda167107.json +++ /dev/null @@ -1,15 +0,0 @@ -{ - "db_name": "PostgreSQL", - "query": "UPDATE users SET deactivated_at = $1 WHERE did = $2", - "describe": { - "columns": [], - "parameters": { - "Left": [ - "Timestamptz", - "Text" - ] - }, - "nullable": [] - }, - "hash": "0710b57fb9aa933525f617b15e6e2e5feaa9c59c38ec9175568abdacda167107" -} diff --git a/.sqlx/query-0ec60bb854a4991d0d7249a68f7445b65c8cc8c723baca221d85f5e4f2478b99.json b/.sqlx/query-0ec60bb854a4991d0d7249a68f7445b65c8cc8c723baca221d85f5e4f2478b99.json deleted file mode 100644 index 9e5e897..0000000 --- a/.sqlx/query-0ec60bb854a4991d0d7249a68f7445b65c8cc8c723baca221d85f5e4f2478b99.json +++ /dev/null @@ -1,22 +0,0 @@ -{ - "db_name": "PostgreSQL", - "query": "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", - "describe": { - "columns": [ - { - "ordinal": 0, - "name": "body", - "type_info": "Text" - } - ], - "parameters": { - "Left": [ - "Text" - ] - }, - "nullable": [ - false - ] - }, - "hash": "0ec60bb854a4991d0d7249a68f7445b65c8cc8c723baca221d85f5e4f2478b99" -} diff --git a/.sqlx/query-0fae1be7a75bdc58c69a9af97cad4aec23c32a9378764b8d6d7eb2cc89c562b1.json b/.sqlx/query-0fae1be7a75bdc58c69a9af97cad4aec23c32a9378764b8d6d7eb2cc89c562b1.json deleted file mode 100644 index 7da18e9..0000000 --- a/.sqlx/query-0fae1be7a75bdc58c69a9af97cad4aec23c32a9378764b8d6d7eb2cc89c562b1.json +++ /dev/null @@ -1,22 +0,0 @@ -{ - "db_name": "PostgreSQL", - "query": "SELECT token FROM sso_pending_registration WHERE token = $1", - "describe": { - "columns": [ - { - "ordinal": 0, - "name": "token", - "type_info": "Text" - } - ], - "parameters": { - "Left": [ - "Text" - ] - }, - "nullable": [ - false - ] - }, - "hash": "0fae1be7a75bdc58c69a9af97cad4aec23c32a9378764b8d6d7eb2cc89c562b1" -} diff --git a/.sqlx/query-1c84643fd6bc57c76517849a64d2d877df337e823d4c2c2b077f695bbfc9e9ac.json b/.sqlx/query-1c84643fd6bc57c76517849a64d2d877df337e823d4c2c2b077f695bbfc9e9ac.json deleted file mode 100644 index d2228ff..0000000 --- a/.sqlx/query-1c84643fd6bc57c76517849a64d2d877df337e823d4c2c2b077f695bbfc9e9ac.json +++ /dev/null @@ -1,22 +0,0 @@ -{ - "db_name": "PostgreSQL", - "query": "\n DELETE FROM sso_pending_registration\n WHERE token = $1 AND expires_at > NOW()\n RETURNING token\n ", - "describe": { - "columns": [ - { - "ordinal": 0, - "name": "token", - "type_info": "Text" - } - ], - "parameters": { - "Left": [ - "Text" - ] - }, - "nullable": [ - false - ] - }, - "hash": "1c84643fd6bc57c76517849a64d2d877df337e823d4c2c2b077f695bbfc9e9ac" -} diff --git a/.sqlx/query-f29da3bdfbbc547b339b4cdb059fac26435b0feec65cf1c56f851d1c4d6b1814.json b/.sqlx/query-1e63d287c619a14e5c07d80e8e54193d2964c8b5e6a855256cb80c4d0cd2c6ea.json similarity index 50% rename from .sqlx/query-f29da3bdfbbc547b339b4cdb059fac26435b0feec65cf1c56f851d1c4d6b1814.json rename to .sqlx/query-1e63d287c619a14e5c07d80e8e54193d2964c8b5e6a855256cb80c4d0cd2c6ea.json index db8bde6..a9900c4 100644 --- a/.sqlx/query-f29da3bdfbbc547b339b4cdb059fac26435b0feec65cf1c56f851d1c4d6b1814.json +++ b/.sqlx/query-1e63d287c619a14e5c07d80e8e54193d2964c8b5e6a855256cb80c4d0cd2c6ea.json @@ -1,14 +1,15 @@ { "db_name": "PostgreSQL", - "query": "UPDATE users SET is_admin = TRUE WHERE did = $1", + "query": "UPDATE users SET is_admin = $1 WHERE did = $2", "describe": { "columns": [], "parameters": { "Left": [ + "Bool", "Text" ] }, "nullable": [] }, - "hash": "f29da3bdfbbc547b339b4cdb059fac26435b0feec65cf1c56f851d1c4d6b1814" + "hash": "1e63d287c619a14e5c07d80e8e54193d2964c8b5e6a855256cb80c4d0cd2c6ea" } diff --git a/.sqlx/query-47fe4a54857344d8f789f37092a294cd58f64b4fb431b54b5deda13d64525e88.json b/.sqlx/query-237c2d912e89b7e0e5baa83503a22f158ea1614b5157f6c9e2aba6017fef6b26.json similarity index 61% rename from .sqlx/query-47fe4a54857344d8f789f37092a294cd58f64b4fb431b54b5deda13d64525e88.json rename to .sqlx/query-237c2d912e89b7e0e5baa83503a22f158ea1614b5157f6c9e2aba6017fef6b26.json index e232907..faf25ad 100644 --- a/.sqlx/query-47fe4a54857344d8f789f37092a294cd58f64b4fb431b54b5deda13d64525e88.json +++ b/.sqlx/query-237c2d912e89b7e0e5baa83503a22f158ea1614b5157f6c9e2aba6017fef6b26.json @@ -1,6 +1,6 @@ { "db_name": "PostgreSQL", - "query": "SELECT token, expires_at FROM account_deletion_requests WHERE did = $1", + "query": "SELECT t.token, t.expires_at\n FROM plc_operation_tokens t\n JOIN users u ON t.user_id = u.id\n WHERE u.did = $1", "describe": { "columns": [ { @@ -24,5 +24,5 @@ false ] }, - "hash": "47fe4a54857344d8f789f37092a294cd58f64b4fb431b54b5deda13d64525e88" + "hash": "237c2d912e89b7e0e5baa83503a22f158ea1614b5157f6c9e2aba6017fef6b26" } diff --git a/.sqlx/query-24b823043ab60f36c29029137fef30dfe33922bb06067f2fdbfc1fbb4b0a2a81.json b/.sqlx/query-24b823043ab60f36c29029137fef30dfe33922bb06067f2fdbfc1fbb4b0a2a81.json deleted file mode 100644 index 06e8ff2..0000000 --- a/.sqlx/query-24b823043ab60f36c29029137fef30dfe33922bb06067f2fdbfc1fbb4b0a2a81.json +++ /dev/null @@ -1,28 +0,0 @@ -{ - "db_name": "PostgreSQL", - "query": "\n DELETE FROM sso_pending_registration\n WHERE token = $1 AND expires_at > NOW()\n RETURNING token, request_uri\n ", - "describe": { - "columns": [ - { - "ordinal": 0, - "name": "token", - "type_info": "Text" - }, - { - "ordinal": 1, - "name": "request_uri", - "type_info": "Text" - } - ], - "parameters": { - "Left": [ - "Text" - ] - }, - "nullable": [ - false, - false - ] - }, - "hash": "24b823043ab60f36c29029137fef30dfe33922bb06067f2fdbfc1fbb4b0a2a81" -} diff --git a/.sqlx/query-2841093a67480e75e1e9e4046bf3eb74afae2d04f5ea0ec17a4d433983e6d71c.json b/.sqlx/query-2841093a67480e75e1e9e4046bf3eb74afae2d04f5ea0ec17a4d433983e6d71c.json deleted file mode 100644 index ffaa87b..0000000 --- a/.sqlx/query-2841093a67480e75e1e9e4046bf3eb74afae2d04f5ea0ec17a4d433983e6d71c.json +++ /dev/null @@ -1,38 +0,0 @@ -{ - "db_name": "PostgreSQL", - "query": "\n INSERT INTO external_identities (did, provider, provider_user_id)\n VALUES ($1, $2, $3)\n RETURNING id\n ", - "describe": { - "columns": [ - { - "ordinal": 0, - "name": "id", - "type_info": "Uuid" - } - ], - "parameters": { - "Left": [ - "Text", - { - "Custom": { - "name": "sso_provider_type", - "kind": { - "Enum": [ - "github", - "discord", - "google", - "gitlab", - "oidc", - "apple" - ] - } - } - }, - "Text" - ] - }, - "nullable": [ - false - ] - }, - "hash": "2841093a67480e75e1e9e4046bf3eb74afae2d04f5ea0ec17a4d433983e6d71c" -} diff --git a/.sqlx/query-376b72306b50f747bc9161985ff4f50c35c53025a55ccf5e9933dc3795d29313.json b/.sqlx/query-376b72306b50f747bc9161985ff4f50c35c53025a55ccf5e9933dc3795d29313.json deleted file mode 100644 index 621ea6b..0000000 --- a/.sqlx/query-376b72306b50f747bc9161985ff4f50c35c53025a55ccf5e9933dc3795d29313.json +++ /dev/null @@ -1,32 +0,0 @@ -{ - "db_name": "PostgreSQL", - "query": "\n INSERT INTO sso_pending_registration (token, request_uri, provider, provider_user_id, provider_email_verified)\n VALUES ($1, $2, $3, $4, $5)\n ", - "describe": { - "columns": [], - "parameters": { - "Left": [ - "Text", - "Text", - { - "Custom": { - "name": "sso_provider_type", - "kind": { - "Enum": [ - "github", - "discord", - "google", - "gitlab", - "oidc", - "apple" - ] - } - } - }, - "Text", - "Bool" - ] - }, - "nullable": [] - }, - "hash": "376b72306b50f747bc9161985ff4f50c35c53025a55ccf5e9933dc3795d29313" -} diff --git a/.sqlx/query-3933ea5b147ab6294936de147b98e116cfae848ecd76ea5d367585eb5117f2ad.json b/.sqlx/query-3933ea5b147ab6294936de147b98e116cfae848ecd76ea5d367585eb5117f2ad.json deleted file mode 100644 index 8192722..0000000 --- a/.sqlx/query-3933ea5b147ab6294936de147b98e116cfae848ecd76ea5d367585eb5117f2ad.json +++ /dev/null @@ -1,22 +0,0 @@ -{ - "db_name": "PostgreSQL", - "query": "SELECT id FROM external_identities WHERE id = $1", - "describe": { - "columns": [ - { - "ordinal": 0, - "name": "id", - "type_info": "Uuid" - } - ], - "parameters": { - "Left": [ - "Uuid" - ] - }, - "nullable": [ - false - ] - }, - "hash": "3933ea5b147ab6294936de147b98e116cfae848ecd76ea5d367585eb5117f2ad" -} diff --git a/.sqlx/query-3bed8d4843545f4a9676207513806603c50eb2af92957994abaf1c89c0294c12.json b/.sqlx/query-3bed8d4843545f4a9676207513806603c50eb2af92957994abaf1c89c0294c12.json deleted file mode 100644 index 097e2cc..0000000 --- a/.sqlx/query-3bed8d4843545f4a9676207513806603c50eb2af92957994abaf1c89c0294c12.json +++ /dev/null @@ -1,16 +0,0 @@ -{ - "db_name": "PostgreSQL", - "query": "INSERT INTO users (did, handle, email, password_hash) VALUES ($1, $2, $3, 'hash')", - "describe": { - "columns": [], - "parameters": { - "Left": [ - "Text", - "Text", - "Text" - ] - }, - "nullable": [] - }, - "hash": "3bed8d4843545f4a9676207513806603c50eb2af92957994abaf1c89c0294c12" -} diff --git a/.sqlx/query-4560c237741ce9d4166aecd669770b3360a3ac71e649b293efb88d92c3254068.json b/.sqlx/query-4560c237741ce9d4166aecd669770b3360a3ac71e649b293efb88d92c3254068.json deleted file mode 100644 index b81fee7..0000000 --- a/.sqlx/query-4560c237741ce9d4166aecd669770b3360a3ac71e649b293efb88d92c3254068.json +++ /dev/null @@ -1,22 +0,0 @@ -{ - "db_name": "PostgreSQL", - "query": "SELECT id FROM users WHERE email = $1", - "describe": { - "columns": [ - { - "ordinal": 0, - "name": "id", - "type_info": "Uuid" - } - ], - "parameters": { - "Left": [ - "Text" - ] - }, - "nullable": [ - false - ] - }, - "hash": "4560c237741ce9d4166aecd669770b3360a3ac71e649b293efb88d92c3254068" -} diff --git a/.sqlx/query-49cbc923cc4a0dcf7dea4ead5ab9580ff03b717586c4ca2d5343709e2dac86b6.json b/.sqlx/query-49cbc923cc4a0dcf7dea4ead5ab9580ff03b717586c4ca2d5343709e2dac86b6.json deleted file mode 100644 index 873f0f7..0000000 --- a/.sqlx/query-49cbc923cc4a0dcf7dea4ead5ab9580ff03b717586c4ca2d5343709e2dac86b6.json +++ /dev/null @@ -1,22 +0,0 @@ -{ - "db_name": "PostgreSQL", - "query": "SELECT email_verified FROM users WHERE did = $1", - "describe": { - "columns": [ - { - "ordinal": 0, - "name": "email_verified", - "type_info": "Bool" - } - ], - "parameters": { - "Left": [ - "Text" - ] - }, - "nullable": [ - false - ] - }, - "hash": "49cbc923cc4a0dcf7dea4ead5ab9580ff03b717586c4ca2d5343709e2dac86b6" -} diff --git a/.sqlx/query-4fef326fa2d03d04869af3fec702c901d1ecf392545a3a032438b2c1859d46cc.json b/.sqlx/query-4fef326fa2d03d04869af3fec702c901d1ecf392545a3a032438b2c1859d46cc.json deleted file mode 100644 index 3780e70..0000000 --- a/.sqlx/query-4fef326fa2d03d04869af3fec702c901d1ecf392545a3a032438b2c1859d46cc.json +++ /dev/null @@ -1,22 +0,0 @@ -{ - "db_name": "PostgreSQL", - "query": "\n SELECT token FROM sso_pending_registration\n WHERE token = $1 AND expires_at > NOW()\n ", - "describe": { - "columns": [ - { - "ordinal": 0, - "name": "token", - "type_info": "Text" - } - ], - "parameters": { - "Left": [ - "Text" - ] - }, - "nullable": [ - false - ] - }, - "hash": "4fef326fa2d03d04869af3fec702c901d1ecf392545a3a032438b2c1859d46cc" -} diff --git a/.sqlx/query-575c1e5529874f8f523e6fe22ccf4ee3296806581b1765dfb91a84ffab347f15.json b/.sqlx/query-575c1e5529874f8f523e6fe22ccf4ee3296806581b1765dfb91a84ffab347f15.json deleted file mode 100644 index e11b68f..0000000 --- a/.sqlx/query-575c1e5529874f8f523e6fe22ccf4ee3296806581b1765dfb91a84ffab347f15.json +++ /dev/null @@ -1,15 +0,0 @@ -{ - "db_name": "PostgreSQL", - "query": "\n INSERT INTO oauth_authorization_request (id, client_id, parameters, expires_at)\n VALUES ($1, 'https://test.example.com', $2, NOW() + INTERVAL '1 hour')\n ", - "describe": { - "columns": [], - "parameters": { - "Left": [ - "Text", - "Jsonb" - ] - }, - "nullable": [] - }, - "hash": "575c1e5529874f8f523e6fe22ccf4ee3296806581b1765dfb91a84ffab347f15" -} diff --git a/.sqlx/query-596c3400a60c77c7645fd46fcea61fa7898b6832e58c0f647f382b23b81d350e.json b/.sqlx/query-596c3400a60c77c7645fd46fcea61fa7898b6832e58c0f647f382b23b81d350e.json deleted file mode 100644 index dfa5ed9..0000000 --- a/.sqlx/query-596c3400a60c77c7645fd46fcea61fa7898b6832e58c0f647f382b23b81d350e.json +++ /dev/null @@ -1,33 +0,0 @@ -{ - "db_name": "PostgreSQL", - "query": "\n INSERT INTO sso_pending_registration (token, request_uri, provider, provider_user_id, provider_username, provider_email)\n VALUES ($1, $2, $3, $4, $5, $6)\n ", - "describe": { - "columns": [], - "parameters": { - "Left": [ - "Text", - "Text", - { - "Custom": { - "name": "sso_provider_type", - "kind": { - "Enum": [ - "github", - "discord", - "google", - "gitlab", - "oidc", - "apple" - ] - } - } - }, - "Text", - "Text", - "Text" - ] - }, - "nullable": [] - }, - "hash": "596c3400a60c77c7645fd46fcea61fa7898b6832e58c0f647f382b23b81d350e" -} diff --git a/.sqlx/query-59e63c5cf92985714e9586d1ce012efef733d4afaa4ea09974daf8303805e5d2.json b/.sqlx/query-59e63c5cf92985714e9586d1ce012efef733d4afaa4ea09974daf8303805e5d2.json deleted file mode 100644 index 6b1501f..0000000 --- a/.sqlx/query-59e63c5cf92985714e9586d1ce012efef733d4afaa4ea09974daf8303805e5d2.json +++ /dev/null @@ -1,81 +0,0 @@ -{ - "db_name": "PostgreSQL", - "query": "\n SELECT id, did, provider as \"provider: SsoProviderType\", provider_user_id, provider_username, provider_email\n FROM external_identities\n WHERE provider = $1 AND provider_user_id = $2\n ", - "describe": { - "columns": [ - { - "ordinal": 0, - "name": "id", - "type_info": "Uuid" - }, - { - "ordinal": 1, - "name": "did", - "type_info": "Text" - }, - { - "ordinal": 2, - "name": "provider: SsoProviderType", - "type_info": { - "Custom": { - "name": "sso_provider_type", - "kind": { - "Enum": [ - "github", - "discord", - "google", - "gitlab", - "oidc", - "apple" - ] - } - } - } - }, - { - "ordinal": 3, - "name": "provider_user_id", - "type_info": "Text" - }, - { - "ordinal": 4, - "name": "provider_username", - "type_info": "Text" - }, - { - "ordinal": 5, - "name": "provider_email", - "type_info": "Text" - } - ], - "parameters": { - "Left": [ - { - "Custom": { - "name": "sso_provider_type", - "kind": { - "Enum": [ - "github", - "discord", - "google", - "gitlab", - "oidc", - "apple" - ] - } - } - }, - "Text" - ] - }, - "nullable": [ - false, - false, - false, - false, - true, - true - ] - }, - "hash": "59e63c5cf92985714e9586d1ce012efef733d4afaa4ea09974daf8303805e5d2" -} diff --git a/.sqlx/query-5a016f289caf75177731711e56e92881ba343c73a9a6e513e205c801c5943ec0.json b/.sqlx/query-5a016f289caf75177731711e56e92881ba343c73a9a6e513e205c801c5943ec0.json deleted file mode 100644 index 32bd1aa..0000000 --- a/.sqlx/query-5a016f289caf75177731711e56e92881ba343c73a9a6e513e205c801c5943ec0.json +++ /dev/null @@ -1,28 +0,0 @@ -{ - "db_name": "PostgreSQL", - "query": "\n SELECT k.key_bytes, k.encryption_version\n FROM user_keys k\n JOIN users u ON k.user_id = u.id\n WHERE u.did = $1\n ", - "describe": { - "columns": [ - { - "ordinal": 0, - "name": "key_bytes", - "type_info": "Bytea" - }, - { - "ordinal": 1, - "name": "encryption_version", - "type_info": "Int4" - } - ], - "parameters": { - "Left": [ - "Text" - ] - }, - "nullable": [ - false, - true - ] - }, - "hash": "5a016f289caf75177731711e56e92881ba343c73a9a6e513e205c801c5943ec0" -} diff --git a/.sqlx/query-5af4a386c1632903ad7102551a5bd148bcf541baab6a84c8649666a695f9c4d1.json b/.sqlx/query-5af4a386c1632903ad7102551a5bd148bcf541baab6a84c8649666a695f9c4d1.json deleted file mode 100644 index 0f764da..0000000 --- a/.sqlx/query-5af4a386c1632903ad7102551a5bd148bcf541baab6a84c8649666a695f9c4d1.json +++ /dev/null @@ -1,22 +0,0 @@ -{ - "db_name": "PostgreSQL", - "query": "\n DELETE FROM sso_auth_state\n WHERE state = $1 AND expires_at > NOW()\n RETURNING state\n ", - "describe": { - "columns": [ - { - "ordinal": 0, - "name": "state", - "type_info": "Text" - } - ], - "parameters": { - "Left": [ - "Text" - ] - }, - "nullable": [ - false - ] - }, - "hash": "5af4a386c1632903ad7102551a5bd148bcf541baab6a84c8649666a695f9c4d1" -} diff --git a/.sqlx/query-a802d7d860f263eace39ce82bb27b633cec7287c1cc177f0e1d47ec6571564d5.json b/.sqlx/query-5cee16f49a727d66b5231a8d07d7f4bcb6a1136fbf3e3d249fd33600772ac80f.json similarity index 58% rename from .sqlx/query-a802d7d860f263eace39ce82bb27b633cec7287c1cc177f0e1d47ec6571564d5.json rename to .sqlx/query-5cee16f49a727d66b5231a8d07d7f4bcb6a1136fbf3e3d249fd33600772ac80f.json index f10b2bc..9040173 100644 --- a/.sqlx/query-a802d7d860f263eace39ce82bb27b633cec7287c1cc177f0e1d47ec6571564d5.json +++ b/.sqlx/query-5cee16f49a727d66b5231a8d07d7f4bcb6a1136fbf3e3d249fd33600772ac80f.json @@ -1,11 +1,11 @@ { "db_name": "PostgreSQL", - "query": "SELECT token FROM account_deletion_requests WHERE did = $1", + "query": "SELECT code FROM oauth_2fa_challenge WHERE request_uri = $1", "describe": { "columns": [ { "ordinal": 0, - "name": "token", + "name": "code", "type_info": "Text" } ], @@ -18,5 +18,5 @@ false ] }, - "hash": "a802d7d860f263eace39ce82bb27b633cec7287c1cc177f0e1d47ec6571564d5" + "hash": "5cee16f49a727d66b5231a8d07d7f4bcb6a1136fbf3e3d249fd33600772ac80f" } diff --git a/.sqlx/query-5e4c0dd92ac3c4b5e2eae5d129f2649cf3a8f068105f44a8dca9625427affc06.json b/.sqlx/query-5e4c0dd92ac3c4b5e2eae5d129f2649cf3a8f068105f44a8dca9625427affc06.json deleted file mode 100644 index 525f4ef..0000000 --- a/.sqlx/query-5e4c0dd92ac3c4b5e2eae5d129f2649cf3a8f068105f44a8dca9625427affc06.json +++ /dev/null @@ -1,43 +0,0 @@ -{ - "db_name": "PostgreSQL", - "query": "\n SELECT provider_user_id, provider_email_verified\n FROM external_identities\n WHERE did = $1 AND provider = $2\n ", - "describe": { - "columns": [ - { - "ordinal": 0, - "name": "provider_user_id", - "type_info": "Text" - }, - { - "ordinal": 1, - "name": "provider_email_verified", - "type_info": "Bool" - } - ], - "parameters": { - "Left": [ - "Text", - { - "Custom": { - "name": "sso_provider_type", - "kind": { - "Enum": [ - "github", - "discord", - "google", - "gitlab", - "oidc", - "apple" - ] - } - } - } - ] - }, - "nullable": [ - false, - false - ] - }, - "hash": "5e4c0dd92ac3c4b5e2eae5d129f2649cf3a8f068105f44a8dca9625427affc06" -} diff --git a/.sqlx/query-5e9c6ec72c2c0ea1c8dff551d01baddd1dd953c828a5656db2ee39dea996f890.json b/.sqlx/query-5e9c6ec72c2c0ea1c8dff551d01baddd1dd953c828a5656db2ee39dea996f890.json deleted file mode 100644 index f109216..0000000 --- a/.sqlx/query-5e9c6ec72c2c0ea1c8dff551d01baddd1dd953c828a5656db2ee39dea996f890.json +++ /dev/null @@ -1,33 +0,0 @@ -{ - "db_name": "PostgreSQL", - "query": "\n INSERT INTO sso_auth_state (state, request_uri, provider, action, nonce, code_verifier)\n VALUES ($1, $2, $3, $4, $5, $6)\n ", - "describe": { - "columns": [], - "parameters": { - "Left": [ - "Text", - "Text", - { - "Custom": { - "name": "sso_provider_type", - "kind": { - "Enum": [ - "github", - "discord", - "google", - "gitlab", - "oidc", - "apple" - ] - } - } - }, - "Text", - "Text", - "Text" - ] - }, - "nullable": [] - }, - "hash": "5e9c6ec72c2c0ea1c8dff551d01baddd1dd953c828a5656db2ee39dea996f890" -} diff --git a/.sqlx/query-82717b6f61cd79347e1ca7e92c4413743ba168d1e0d8b85566711e54d4048f81.json b/.sqlx/query-61f489b4fc42f5b0aaea287cde4415da6f5e96b3a0f36216bdc6dea924b09abd.json similarity index 58% rename from .sqlx/query-82717b6f61cd79347e1ca7e92c4413743ba168d1e0d8b85566711e54d4048f81.json rename to .sqlx/query-61f489b4fc42f5b0aaea287cde4415da6f5e96b3a0f36216bdc6dea924b09abd.json index 8d5a9b7..b1c3607 100644 --- a/.sqlx/query-82717b6f61cd79347e1ca7e92c4413743ba168d1e0d8b85566711e54d4048f81.json +++ b/.sqlx/query-61f489b4fc42f5b0aaea287cde4415da6f5e96b3a0f36216bdc6dea924b09abd.json @@ -1,6 +1,6 @@ { "db_name": "PostgreSQL", - "query": "SELECT t.token, t.expires_at FROM plc_operation_tokens t JOIN users u ON t.user_id = u.id WHERE u.did = $1", + "query": "SELECT token, did, expires_at FROM account_deletion_requests WHERE did = $1", "describe": { "columns": [ { @@ -10,6 +10,11 @@ }, { "ordinal": 1, + "name": "did", + "type_info": "Text" + }, + { + "ordinal": 2, "name": "expires_at", "type_info": "Timestamptz" } @@ -20,9 +25,10 @@ ] }, "nullable": [ + false, false, false ] }, - "hash": "82717b6f61cd79347e1ca7e92c4413743ba168d1e0d8b85566711e54d4048f81" + "hash": "61f489b4fc42f5b0aaea287cde4415da6f5e96b3a0f36216bdc6dea924b09abd" } diff --git a/.sqlx/query-63f6f2a89650794fe90e10ce7fc785a6b9f7d37c12b31a6ff13f7c5214eef19e.json b/.sqlx/query-63f6f2a89650794fe90e10ce7fc785a6b9f7d37c12b31a6ff13f7c5214eef19e.json deleted file mode 100644 index cb057f9..0000000 --- a/.sqlx/query-63f6f2a89650794fe90e10ce7fc785a6b9f7d37c12b31a6ff13f7c5214eef19e.json +++ /dev/null @@ -1,28 +0,0 @@ -{ - "db_name": "PostgreSQL", - "query": "SELECT did, email_verified FROM users WHERE did = $1", - "describe": { - "columns": [ - { - "ordinal": 0, - "name": "did", - "type_info": "Text" - }, - { - "ordinal": 1, - "name": "email_verified", - "type_info": "Bool" - } - ], - "parameters": { - "Left": [ - "Text" - ] - }, - "nullable": [ - false, - false - ] - }, - "hash": "63f6f2a89650794fe90e10ce7fc785a6b9f7d37c12b31a6ff13f7c5214eef19e" -} diff --git a/.sqlx/query-6c7ace2a64848adc757af6c93b9162e1d95788b372370a7ad0d7540338bb73ee.json b/.sqlx/query-6c7ace2a64848adc757af6c93b9162e1d95788b372370a7ad0d7540338bb73ee.json deleted file mode 100644 index eed954b..0000000 --- a/.sqlx/query-6c7ace2a64848adc757af6c93b9162e1d95788b372370a7ad0d7540338bb73ee.json +++ /dev/null @@ -1,66 +0,0 @@ -{ - "db_name": "PostgreSQL", - "query": "\n SELECT state, request_uri, provider as \"provider: SsoProviderType\", action, nonce, code_verifier\n FROM sso_auth_state\n WHERE state = $1\n ", - "describe": { - "columns": [ - { - "ordinal": 0, - "name": "state", - "type_info": "Text" - }, - { - "ordinal": 1, - "name": "request_uri", - "type_info": "Text" - }, - { - "ordinal": 2, - "name": "provider: SsoProviderType", - "type_info": { - "Custom": { - "name": "sso_provider_type", - "kind": { - "Enum": [ - "github", - "discord", - "google", - "gitlab", - "oidc", - "apple" - ] - } - } - } - }, - { - "ordinal": 3, - "name": "action", - "type_info": "Text" - }, - { - "ordinal": 4, - "name": "nonce", - "type_info": "Text" - }, - { - "ordinal": 5, - "name": "code_verifier", - "type_info": "Text" - } - ], - "parameters": { - "Left": [ - "Text" - ] - }, - "nullable": [ - false, - false, - false, - false, - true, - true - ] - }, - "hash": "6c7ace2a64848adc757af6c93b9162e1d95788b372370a7ad0d7540338bb73ee" -} diff --git a/.sqlx/query-6fbcff0206599484bfb6cef165b6f729d27e7a342f7718ee4ac07f0ca94412ba.json b/.sqlx/query-6fbcff0206599484bfb6cef165b6f729d27e7a342f7718ee4ac07f0ca94412ba.json deleted file mode 100644 index dbd52e6..0000000 --- a/.sqlx/query-6fbcff0206599484bfb6cef165b6f729d27e7a342f7718ee4ac07f0ca94412ba.json +++ /dev/null @@ -1,22 +0,0 @@ -{ - "db_name": "PostgreSQL", - "query": "SELECT state FROM sso_auth_state WHERE state = $1", - "describe": { - "columns": [ - { - "ordinal": 0, - "name": "state", - "type_info": "Text" - } - ], - "parameters": { - "Left": [ - "Text" - ] - }, - "nullable": [ - false - ] - }, - "hash": "6fbcff0206599484bfb6cef165b6f729d27e7a342f7718ee4ac07f0ca94412ba" -} diff --git a/.sqlx/query-712459c27fc037f45389e2766cf1057e86e93ef756a784ed12beb453b03c5da1.json b/.sqlx/query-712459c27fc037f45389e2766cf1057e86e93ef756a784ed12beb453b03c5da1.json deleted file mode 100644 index e831be8..0000000 --- a/.sqlx/query-712459c27fc037f45389e2766cf1057e86e93ef756a784ed12beb453b03c5da1.json +++ /dev/null @@ -1,33 +0,0 @@ -{ - "db_name": "PostgreSQL", - "query": "\n INSERT INTO sso_pending_registration (token, request_uri, provider, provider_user_id, provider_username, provider_email_verified)\n VALUES ($1, $2, $3, $4, $5, $6)\n ", - "describe": { - "columns": [], - "parameters": { - "Left": [ - "Text", - "Text", - { - "Custom": { - "name": "sso_provider_type", - "kind": { - "Enum": [ - "github", - "discord", - "google", - "gitlab", - "oidc", - "apple" - ] - } - } - }, - "Text", - "Text", - "Bool" - ] - }, - "nullable": [] - }, - "hash": "712459c27fc037f45389e2766cf1057e86e93ef756a784ed12beb453b03c5da1" -} diff --git a/.sqlx/query-785a864944c5939331704c71b0cd3ed26ffdd64f3fd0f26ecc28b6a0557bbe8f.json b/.sqlx/query-785a864944c5939331704c71b0cd3ed26ffdd64f3fd0f26ecc28b6a0557bbe8f.json deleted file mode 100644 index af80ee3..0000000 --- a/.sqlx/query-785a864944c5939331704c71b0cd3ed26ffdd64f3fd0f26ecc28b6a0557bbe8f.json +++ /dev/null @@ -1,22 +0,0 @@ -{ - "db_name": "PostgreSQL", - "query": "SELECT subject FROM comms_queue WHERE user_id = $1 AND comms_type = 'admin_email' AND body = 'Email without subject' LIMIT 1", - "describe": { - "columns": [ - { - "ordinal": 0, - "name": "subject", - "type_info": "Text" - } - ], - "parameters": { - "Left": [ - "Uuid" - ] - }, - "nullable": [ - true - ] - }, - "hash": "785a864944c5939331704c71b0cd3ed26ffdd64f3fd0f26ecc28b6a0557bbe8f" -} diff --git a/.sqlx/query-7caa8f9083b15ec1209dda35c4c6f6fba9fe338e4a6a10636b5389d426df1631.json b/.sqlx/query-7caa8f9083b15ec1209dda35c4c6f6fba9fe338e4a6a10636b5389d426df1631.json deleted file mode 100644 index 416f681..0000000 --- a/.sqlx/query-7caa8f9083b15ec1209dda35c4c6f6fba9fe338e4a6a10636b5389d426df1631.json +++ /dev/null @@ -1,22 +0,0 @@ -{ - "db_name": "PostgreSQL", - "query": "\n SELECT t.token\n FROM plc_operation_tokens t\n JOIN users u ON t.user_id = u.id\n WHERE u.did = $1\n ", - "describe": { - "columns": [ - { - "ordinal": 0, - "name": "token", - "type_info": "Text" - } - ], - "parameters": { - "Left": [ - "Text" - ] - }, - "nullable": [ - false - ] - }, - "hash": "7caa8f9083b15ec1209dda35c4c6f6fba9fe338e4a6a10636b5389d426df1631" -} diff --git a/.sqlx/query-7d24e744a4e63570b1410e50b45b745ce8915ab3715b3eff7efc2d84f27735d0.json b/.sqlx/query-7d24e744a4e63570b1410e50b45b745ce8915ab3715b3eff7efc2d84f27735d0.json deleted file mode 100644 index 25e1b1d..0000000 --- a/.sqlx/query-7d24e744a4e63570b1410e50b45b745ce8915ab3715b3eff7efc2d84f27735d0.json +++ /dev/null @@ -1,28 +0,0 @@ -{ - "db_name": "PostgreSQL", - "query": "SELECT provider_username, last_login_at FROM external_identities WHERE id = $1", - "describe": { - "columns": [ - { - "ordinal": 0, - "name": "provider_username", - "type_info": "Text" - }, - { - "ordinal": 1, - "name": "last_login_at", - "type_info": "Timestamptz" - } - ], - "parameters": { - "Left": [ - "Uuid" - ] - }, - "nullable": [ - true, - true - ] - }, - "hash": "7d24e744a4e63570b1410e50b45b745ce8915ab3715b3eff7efc2d84f27735d0" -} diff --git a/.sqlx/query-7d3a9f0545943bc6a3a14fcd596aac5cc731c8177d74e504606d7e92c7d0c73f.json b/.sqlx/query-7d3a9f0545943bc6a3a14fcd596aac5cc731c8177d74e504606d7e92c7d0c73f.json new file mode 100644 index 0000000..298773c --- /dev/null +++ b/.sqlx/query-7d3a9f0545943bc6a3a14fcd596aac5cc731c8177d74e504606d7e92c7d0c73f.json @@ -0,0 +1,15 @@ +{ + "db_name": "PostgreSQL", + "query": "INSERT INTO user_totp (did, secret_encrypted, encryption_version, verified, created_at)\n VALUES ($1, $2, 1, TRUE, NOW())\n ON CONFLICT (did) DO UPDATE SET secret_encrypted = $2, verified = TRUE", + "describe": { + "columns": [], + "parameters": { + "Left": [ + "Text", + "Bytea" + ] + }, + "nullable": [] + }, + "hash": "7d3a9f0545943bc6a3a14fcd596aac5cc731c8177d74e504606d7e92c7d0c73f" +} diff --git a/.sqlx/query-84a1db51a98402323cb86bc19cd2b737f908222ea3426b8bf47d735aff5b6c75.json b/.sqlx/query-84a1db51a98402323cb86bc19cd2b737f908222ea3426b8bf47d735aff5b6c75.json new file mode 100644 index 0000000..d0251d9 --- /dev/null +++ b/.sqlx/query-84a1db51a98402323cb86bc19cd2b737f908222ea3426b8bf47d735aff5b6c75.json @@ -0,0 +1,15 @@ +{ + "db_name": "PostgreSQL", + "query": "UPDATE users SET two_factor_enabled = $1 WHERE did = $2", + "describe": { + "columns": [], + "parameters": { + "Left": [ + "Bool", + "Text" + ] + }, + "nullable": [] + }, + "hash": "84a1db51a98402323cb86bc19cd2b737f908222ea3426b8bf47d735aff5b6c75" +} diff --git a/.sqlx/query-85ffc37a77af832d7795f5f37efe304fced4bf56b4f2287fe9aeb3fc97e1b191.json b/.sqlx/query-85ffc37a77af832d7795f5f37efe304fced4bf56b4f2287fe9aeb3fc97e1b191.json deleted file mode 100644 index 7938c39..0000000 --- a/.sqlx/query-85ffc37a77af832d7795f5f37efe304fced4bf56b4f2287fe9aeb3fc97e1b191.json +++ /dev/null @@ -1,34 +0,0 @@ -{ - "db_name": "PostgreSQL", - "query": "\n INSERT INTO sso_pending_registration (token, request_uri, provider, provider_user_id, provider_username, provider_email, provider_email_verified)\n VALUES ($1, $2, $3, $4, $5, $6, $7)\n ", - "describe": { - "columns": [], - "parameters": { - "Left": [ - "Text", - "Text", - { - "Custom": { - "name": "sso_provider_type", - "kind": { - "Enum": [ - "github", - "discord", - "google", - "gitlab", - "oidc", - "apple" - ] - } - } - }, - "Text", - "Text", - "Text", - "Bool" - ] - }, - "nullable": [] - }, - "hash": "85ffc37a77af832d7795f5f37efe304fced4bf56b4f2287fe9aeb3fc97e1b191" -} diff --git a/.sqlx/query-89b0292d8d022fad8f9cda07b9a7870ca6a7ebe904b2d580956b0816b50bcdb7.json b/.sqlx/query-89b0292d8d022fad8f9cda07b9a7870ca6a7ebe904b2d580956b0816b50bcdb7.json new file mode 100644 index 0000000..5b5f172 --- /dev/null +++ b/.sqlx/query-89b0292d8d022fad8f9cda07b9a7870ca6a7ebe904b2d580956b0816b50bcdb7.json @@ -0,0 +1,36 @@ +{ + "db_name": "PostgreSQL", + "query": "DELETE FROM comms_queue WHERE user_id = $1 AND comms_type = $2", + "describe": { + "columns": [], + "parameters": { + "Left": [ + "Uuid", + { + "Custom": { + "name": "comms_type", + "kind": { + "Enum": [ + "welcome", + "email_verification", + "password_reset", + "email_update", + "account_deletion", + "admin_email", + "plc_operation", + "two_factor_code", + "channel_verification", + "passkey_recovery", + "legacy_login_alert", + "migration_verification", + "channel_verified" + ] + } + } + } + ] + }, + "nullable": [] + }, + "hash": "89b0292d8d022fad8f9cda07b9a7870ca6a7ebe904b2d580956b0816b50bcdb7" +} diff --git a/.sqlx/query-cda68f9b6c60295a196fc853b70ec5fd51a8ffaa2bac5942c115c99d1cbcafa3.json b/.sqlx/query-990bf50e60fc5566639c2c12cd968d154d7b0c6863ad69141653135f98fbc998.json similarity index 53% rename from .sqlx/query-cda68f9b6c60295a196fc853b70ec5fd51a8ffaa2bac5942c115c99d1cbcafa3.json rename to .sqlx/query-990bf50e60fc5566639c2c12cd968d154d7b0c6863ad69141653135f98fbc998.json index 08e8e23..f60094b 100644 --- a/.sqlx/query-cda68f9b6c60295a196fc853b70ec5fd51a8ffaa2bac5942c115c99d1cbcafa3.json +++ b/.sqlx/query-990bf50e60fc5566639c2c12cd968d154d7b0c6863ad69141653135f98fbc998.json @@ -1,6 +1,6 @@ { "db_name": "PostgreSQL", - "query": "SELECT COUNT(*) as \"count!\" FROM plc_operation_tokens t JOIN users u ON t.user_id = u.id WHERE u.did = $1", + "query": "SELECT COUNT(*) as \"count!\"\n FROM plc_operation_tokens t\n JOIN users u ON t.user_id = u.id\n WHERE u.did = $1", "describe": { "columns": [ { @@ -18,5 +18,5 @@ null ] }, - "hash": "cda68f9b6c60295a196fc853b70ec5fd51a8ffaa2bac5942c115c99d1cbcafa3" + "hash": "990bf50e60fc5566639c2c12cd968d154d7b0c6863ad69141653135f98fbc998" } diff --git a/.sqlx/query-9ad422bf3c43e3cfd86fc88c73594246ead214ca794760d3fe77bb5cf4f27be5.json b/.sqlx/query-9ad422bf3c43e3cfd86fc88c73594246ead214ca794760d3fe77bb5cf4f27be5.json deleted file mode 100644 index ef52899..0000000 --- a/.sqlx/query-9ad422bf3c43e3cfd86fc88c73594246ead214ca794760d3fe77bb5cf4f27be5.json +++ /dev/null @@ -1,22 +0,0 @@ -{ - "db_name": "PostgreSQL", - "query": "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", - "describe": { - "columns": [ - { - "ordinal": 0, - "name": "body", - "type_info": "Text" - } - ], - "parameters": { - "Left": [ - "Text" - ] - }, - "nullable": [ - false - ] - }, - "hash": "9ad422bf3c43e3cfd86fc88c73594246ead214ca794760d3fe77bb5cf4f27be5" -} diff --git a/.sqlx/query-9b035b051769e6b9d45910a8bb42ac0f84c73de8c244ba4560f004ee3f4b7002.json b/.sqlx/query-9b035b051769e6b9d45910a8bb42ac0f84c73de8c244ba4560f004ee3f4b7002.json deleted file mode 100644 index 178bc9e..0000000 --- a/.sqlx/query-9b035b051769e6b9d45910a8bb42ac0f84c73de8c244ba4560f004ee3f4b7002.json +++ /dev/null @@ -1,28 +0,0 @@ -{ - "db_name": "PostgreSQL", - "query": "SELECT did, public_key_did_key FROM reserved_signing_keys WHERE public_key_did_key = $1", - "describe": { - "columns": [ - { - "ordinal": 0, - "name": "did", - "type_info": "Text" - }, - { - "ordinal": 1, - "name": "public_key_did_key", - "type_info": "Text" - } - ], - "parameters": { - "Left": [ - "Text" - ] - }, - "nullable": [ - true, - false - ] - }, - "hash": "9b035b051769e6b9d45910a8bb42ac0f84c73de8c244ba4560f004ee3f4b7002" -} diff --git a/.sqlx/query-9dba64081d4f95b5490c9a9bf30a7175db3429f39df4f25e212f38f33882fc65.json b/.sqlx/query-9dba64081d4f95b5490c9a9bf30a7175db3429f39df4f25e212f38f33882fc65.json deleted file mode 100644 index ea2847f..0000000 --- a/.sqlx/query-9dba64081d4f95b5490c9a9bf30a7175db3429f39df4f25e212f38f33882fc65.json +++ /dev/null @@ -1,22 +0,0 @@ -{ - "db_name": "PostgreSQL", - "query": "\n SELECT id FROM external_identities WHERE did = $1\n ", - "describe": { - "columns": [ - { - "ordinal": 0, - "name": "id", - "type_info": "Uuid" - } - ], - "parameters": { - "Left": [ - "Text" - ] - }, - "nullable": [ - false - ] - }, - "hash": "9dba64081d4f95b5490c9a9bf30a7175db3429f39df4f25e212f38f33882fc65" -} diff --git a/.sqlx/query-9f3f2b36f11e9446915d3ca29ef81e4ada0c6a6d72764116dac4f99a4e09785e.json b/.sqlx/query-9f3f2b36f11e9446915d3ca29ef81e4ada0c6a6d72764116dac4f99a4e09785e.json new file mode 100644 index 0000000..8ead953 --- /dev/null +++ b/.sqlx/query-9f3f2b36f11e9446915d3ca29ef81e4ada0c6a6d72764116dac4f99a4e09785e.json @@ -0,0 +1,180 @@ +{ + "db_name": "PostgreSQL", + "query": "SELECT\n id, user_id,\n channel as \"channel: CommsChannel\",\n comms_type as \"comms_type: CommsType\",\n status as \"status: CommsStatus\",\n recipient, subject, body, metadata,\n attempts, max_attempts, last_error,\n created_at, updated_at, scheduled_for, processed_at\n FROM comms_queue\n WHERE user_id = $1 AND comms_type = $2\n ORDER BY created_at DESC\n LIMIT $3", + "describe": { + "columns": [ + { + "ordinal": 0, + "name": "id", + "type_info": "Uuid" + }, + { + "ordinal": 1, + "name": "user_id", + "type_info": "Uuid" + }, + { + "ordinal": 2, + "name": "channel: CommsChannel", + "type_info": { + "Custom": { + "name": "comms_channel", + "kind": { + "Enum": [ + "email", + "discord", + "telegram", + "signal" + ] + } + } + } + }, + { + "ordinal": 3, + "name": "comms_type: CommsType", + "type_info": { + "Custom": { + "name": "comms_type", + "kind": { + "Enum": [ + "welcome", + "email_verification", + "password_reset", + "email_update", + "account_deletion", + "admin_email", + "plc_operation", + "two_factor_code", + "channel_verification", + "passkey_recovery", + "legacy_login_alert", + "migration_verification", + "channel_verified" + ] + } + } + } + }, + { + "ordinal": 4, + "name": "status: CommsStatus", + "type_info": { + "Custom": { + "name": "comms_status", + "kind": { + "Enum": [ + "pending", + "processing", + "sent", + "failed" + ] + } + } + } + }, + { + "ordinal": 5, + "name": "recipient", + "type_info": "Text" + }, + { + "ordinal": 6, + "name": "subject", + "type_info": "Text" + }, + { + "ordinal": 7, + "name": "body", + "type_info": "Text" + }, + { + "ordinal": 8, + "name": "metadata", + "type_info": "Jsonb" + }, + { + "ordinal": 9, + "name": "attempts", + "type_info": "Int4" + }, + { + "ordinal": 10, + "name": "max_attempts", + "type_info": "Int4" + }, + { + "ordinal": 11, + "name": "last_error", + "type_info": "Text" + }, + { + "ordinal": 12, + "name": "created_at", + "type_info": "Timestamptz" + }, + { + "ordinal": 13, + "name": "updated_at", + "type_info": "Timestamptz" + }, + { + "ordinal": 14, + "name": "scheduled_for", + "type_info": "Timestamptz" + }, + { + "ordinal": 15, + "name": "processed_at", + "type_info": "Timestamptz" + } + ], + "parameters": { + "Left": [ + "Uuid", + { + "Custom": { + "name": "comms_type", + "kind": { + "Enum": [ + "welcome", + "email_verification", + "password_reset", + "email_update", + "account_deletion", + "admin_email", + "plc_operation", + "two_factor_code", + "channel_verification", + "passkey_recovery", + "legacy_login_alert", + "migration_verification", + "channel_verified" + ] + } + } + }, + "Int8" + ] + }, + "nullable": [ + false, + false, + false, + false, + false, + false, + true, + false, + true, + false, + false, + true, + false, + false, + false, + true + ] + }, + "hash": "9f3f2b36f11e9446915d3ca29ef81e4ada0c6a6d72764116dac4f99a4e09785e" +} diff --git a/.sqlx/query-9fd56986c1c843d386d1e5884acef8573eb55a3e9f5cb0122fcf8b93d6d667a5.json b/.sqlx/query-9fd56986c1c843d386d1e5884acef8573eb55a3e9f5cb0122fcf8b93d6d667a5.json deleted file mode 100644 index c77d51a..0000000 --- a/.sqlx/query-9fd56986c1c843d386d1e5884acef8573eb55a3e9f5cb0122fcf8b93d6d667a5.json +++ /dev/null @@ -1,66 +0,0 @@ -{ - "db_name": "PostgreSQL", - "query": "\n SELECT token, request_uri, provider as \"provider: SsoProviderType\", provider_user_id,\n provider_username, provider_email\n FROM sso_pending_registration\n WHERE token = $1 AND expires_at > NOW()\n ", - "describe": { - "columns": [ - { - "ordinal": 0, - "name": "token", - "type_info": "Text" - }, - { - "ordinal": 1, - "name": "request_uri", - "type_info": "Text" - }, - { - "ordinal": 2, - "name": "provider: SsoProviderType", - "type_info": { - "Custom": { - "name": "sso_provider_type", - "kind": { - "Enum": [ - "github", - "discord", - "google", - "gitlab", - "oidc", - "apple" - ] - } - } - } - }, - { - "ordinal": 3, - "name": "provider_user_id", - "type_info": "Text" - }, - { - "ordinal": 4, - "name": "provider_username", - "type_info": "Text" - }, - { - "ordinal": 5, - "name": "provider_email", - "type_info": "Text" - } - ], - "parameters": { - "Left": [ - "Text" - ] - }, - "nullable": [ - false, - false, - false, - false, - true, - true - ] - }, - "hash": "9fd56986c1c843d386d1e5884acef8573eb55a3e9f5cb0122fcf8b93d6d667a5" -} diff --git a/.sqlx/query-a23a390659616779d7dbceaa3b5d5171e70fa25e3b8393e142cebcbff752f0f5.json b/.sqlx/query-a23a390659616779d7dbceaa3b5d5171e70fa25e3b8393e142cebcbff752f0f5.json deleted file mode 100644 index 2640f70..0000000 --- a/.sqlx/query-a23a390659616779d7dbceaa3b5d5171e70fa25e3b8393e142cebcbff752f0f5.json +++ /dev/null @@ -1,34 +0,0 @@ -{ - "db_name": "PostgreSQL", - "query": "SELECT private_key_bytes, expires_at, used_at FROM reserved_signing_keys WHERE public_key_did_key = $1", - "describe": { - "columns": [ - { - "ordinal": 0, - "name": "private_key_bytes", - "type_info": "Bytea" - }, - { - "ordinal": 1, - "name": "expires_at", - "type_info": "Timestamptz" - }, - { - "ordinal": 2, - "name": "used_at", - "type_info": "Timestamptz" - } - ], - "parameters": { - "Left": [ - "Text" - ] - }, - "nullable": [ - false, - false, - true - ] - }, - "hash": "a23a390659616779d7dbceaa3b5d5171e70fa25e3b8393e142cebcbff752f0f5" -} diff --git a/.sqlx/query-a3d549a32e76c24e265c73a98dd739067623f275de0740bd576ee288f4444496.json b/.sqlx/query-a3d549a32e76c24e265c73a98dd739067623f275de0740bd576ee288f4444496.json deleted file mode 100644 index 06c9dee..0000000 --- a/.sqlx/query-a3d549a32e76c24e265c73a98dd739067623f275de0740bd576ee288f4444496.json +++ /dev/null @@ -1,15 +0,0 @@ -{ - "db_name": "PostgreSQL", - "query": "\n UPDATE external_identities\n SET provider_username = $2, last_login_at = NOW()\n WHERE id = $1\n ", - "describe": { - "columns": [], - "parameters": { - "Left": [ - "Uuid", - "Text" - ] - }, - "nullable": [] - }, - "hash": "a3d549a32e76c24e265c73a98dd739067623f275de0740bd576ee288f4444496" -} diff --git a/.sqlx/query-a844774d8dd3c50c5faf3de5d43f534b80234759c8437434e467ca33ea10fd1f.json b/.sqlx/query-a844774d8dd3c50c5faf3de5d43f534b80234759c8437434e467ca33ea10fd1f.json deleted file mode 100644 index 9cfbdba..0000000 --- a/.sqlx/query-a844774d8dd3c50c5faf3de5d43f534b80234759c8437434e467ca33ea10fd1f.json +++ /dev/null @@ -1,40 +0,0 @@ -{ - "db_name": "PostgreSQL", - "query": "SELECT preferred_comms_channel as \"preferred_comms_channel: String\", discord_username FROM users WHERE did = $1", - "describe": { - "columns": [ - { - "ordinal": 0, - "name": "preferred_comms_channel: String", - "type_info": { - "Custom": { - "name": "comms_channel", - "kind": { - "Enum": [ - "email", - "discord", - "telegram", - "signal" - ] - } - } - } - }, - { - "ordinal": 1, - "name": "discord_username", - "type_info": "Text" - } - ], - "parameters": { - "Left": [ - "Text" - ] - }, - "nullable": [ - false, - true - ] - }, - "hash": "a844774d8dd3c50c5faf3de5d43f534b80234759c8437434e467ca33ea10fd1f" -} diff --git a/.sqlx/query-aee3e8e1d8924d41bec7d866e274f8bb2ddef833eb03326103c2d0a17ee56154.json b/.sqlx/query-aee3e8e1d8924d41bec7d866e274f8bb2ddef833eb03326103c2d0a17ee56154.json deleted file mode 100644 index 3e281aa..0000000 --- a/.sqlx/query-aee3e8e1d8924d41bec7d866e274f8bb2ddef833eb03326103c2d0a17ee56154.json +++ /dev/null @@ -1,28 +0,0 @@ -{ - "db_name": "PostgreSQL", - "query": "\n DELETE FROM sso_auth_state\n WHERE state = $1 AND expires_at > NOW()\n RETURNING state, request_uri\n ", - "describe": { - "columns": [ - { - "ordinal": 0, - "name": "state", - "type_info": "Text" - }, - { - "ordinal": 1, - "name": "request_uri", - "type_info": "Text" - } - ], - "parameters": { - "Left": [ - "Text" - ] - }, - "nullable": [ - false, - false - ] - }, - "hash": "aee3e8e1d8924d41bec7d866e274f8bb2ddef833eb03326103c2d0a17ee56154" -} diff --git a/.sqlx/query-4445cc86cdf04894b340e67661b79a3c411917144a011f50849b737130b24dbe.json b/.sqlx/query-b364a2b202bab17c0cdc5f70d23b13841b4d9063d94cd0b09268c3dc41824fd2.json similarity index 59% rename from .sqlx/query-4445cc86cdf04894b340e67661b79a3c411917144a011f50849b737130b24dbe.json rename to .sqlx/query-b364a2b202bab17c0cdc5f70d23b13841b4d9063d94cd0b09268c3dc41824fd2.json index ac57c00..019f85f 100644 --- a/.sqlx/query-4445cc86cdf04894b340e67661b79a3c411917144a011f50849b737130b24dbe.json +++ b/.sqlx/query-b364a2b202bab17c0cdc5f70d23b13841b4d9063d94cd0b09268c3dc41824fd2.json @@ -1,22 +1,18 @@ { "db_name": "PostgreSQL", - "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", + "query": "SELECT COUNT(*) as \"count!\" FROM comms_queue WHERE user_id = $1 AND comms_type = $2", "describe": { "columns": [ { "ordinal": 0, - "name": "subject", - "type_info": "Text" - }, - { - "ordinal": 1, - "name": "body", - "type_info": "Text" - }, - { - "ordinal": 2, - "name": "comms_type: String", - "type_info": { + "name": "count!", + "type_info": "Int8" + } + ], + "parameters": { + "Left": [ + "Uuid", + { "Custom": { "name": "comms_type", "kind": { @@ -38,18 +34,11 @@ } } } - } - ], - "parameters": { - "Left": [ - "Uuid" ] }, "nullable": [ - true, - false, - false + null ] }, - "hash": "4445cc86cdf04894b340e67661b79a3c411917144a011f50849b737130b24dbe" + "hash": "b364a2b202bab17c0cdc5f70d23b13841b4d9063d94cd0b09268c3dc41824fd2" } diff --git a/.sqlx/query-ba9684872fad5201b8504c2606c29364a2df9631fe98817e7bfacd3f3f51f6cb.json b/.sqlx/query-ba9684872fad5201b8504c2606c29364a2df9631fe98817e7bfacd3f3f51f6cb.json deleted file mode 100644 index 5a2db77..0000000 --- a/.sqlx/query-ba9684872fad5201b8504c2606c29364a2df9631fe98817e7bfacd3f3f51f6cb.json +++ /dev/null @@ -1,31 +0,0 @@ -{ - "db_name": "PostgreSQL", - "query": "\n INSERT INTO sso_pending_registration (token, request_uri, provider, provider_user_id, expires_at)\n VALUES ($1, $2, $3, $4, NOW() - INTERVAL '1 hour')\n ", - "describe": { - "columns": [], - "parameters": { - "Left": [ - "Text", - "Text", - { - "Custom": { - "name": "sso_provider_type", - "kind": { - "Enum": [ - "github", - "discord", - "google", - "gitlab", - "oidc", - "apple" - ] - } - } - }, - "Text" - ] - }, - "nullable": [] - }, - "hash": "ba9684872fad5201b8504c2606c29364a2df9631fe98817e7bfacd3f3f51f6cb" -} diff --git a/.sqlx/query-bb4460f75d30f48b79d71b97f2c7d54190260deba2d2ade177dbdaa507ab275b.json b/.sqlx/query-bb4460f75d30f48b79d71b97f2c7d54190260deba2d2ade177dbdaa507ab275b.json deleted file mode 100644 index 97297bf..0000000 --- a/.sqlx/query-bb4460f75d30f48b79d71b97f2c7d54190260deba2d2ade177dbdaa507ab275b.json +++ /dev/null @@ -1,12 +0,0 @@ -{ - "db_name": "PostgreSQL", - "query": "DELETE FROM sso_auth_state WHERE expires_at < NOW()", - "describe": { - "columns": [], - "parameters": { - "Left": [] - }, - "nullable": [] - }, - "hash": "bb4460f75d30f48b79d71b97f2c7d54190260deba2d2ade177dbdaa507ab275b" -} diff --git a/.sqlx/query-cd3b8098ad4c1056c1d23acd8a6b29f7abfe18ee6f559bd94ab16274b1cfdfee.json b/.sqlx/query-cd3b8098ad4c1056c1d23acd8a6b29f7abfe18ee6f559bd94ab16274b1cfdfee.json deleted file mode 100644 index 6b53208..0000000 --- a/.sqlx/query-cd3b8098ad4c1056c1d23acd8a6b29f7abfe18ee6f559bd94ab16274b1cfdfee.json +++ /dev/null @@ -1,22 +0,0 @@ -{ - "db_name": "PostgreSQL", - "query": "SELECT password_reset_code FROM users WHERE email = $1", - "describe": { - "columns": [ - { - "ordinal": 0, - "name": "password_reset_code", - "type_info": "Text" - } - ], - "parameters": { - "Left": [ - "Text" - ] - }, - "nullable": [ - true - ] - }, - "hash": "cd3b8098ad4c1056c1d23acd8a6b29f7abfe18ee6f559bd94ab16274b1cfdfee" -} diff --git a/.sqlx/query-d0d4fb4b44cda3442b20037b4d5efaa032e1d004c775e2b6077c5050d7d62041.json b/.sqlx/query-d0d4fb4b44cda3442b20037b4d5efaa032e1d004c775e2b6077c5050d7d62041.json deleted file mode 100644 index d122744..0000000 --- a/.sqlx/query-d0d4fb4b44cda3442b20037b4d5efaa032e1d004c775e2b6077c5050d7d62041.json +++ /dev/null @@ -1,31 +0,0 @@ -{ - "db_name": "PostgreSQL", - "query": "\n INSERT INTO sso_auth_state (state, request_uri, provider, action, expires_at)\n VALUES ($1, $2, $3, $4, NOW() - INTERVAL '1 hour')\n ", - "describe": { - "columns": [], - "parameters": { - "Left": [ - "Text", - "Text", - { - "Custom": { - "name": "sso_provider_type", - "kind": { - "Enum": [ - "github", - "discord", - "google", - "gitlab", - "oidc", - "apple" - ] - } - } - }, - "Text" - ] - }, - "nullable": [] - }, - "hash": "d0d4fb4b44cda3442b20037b4d5efaa032e1d004c775e2b6077c5050d7d62041" -} diff --git a/.sqlx/query-dd7d80d4d118a5fc95b574e2ca9ffaccf974e52fb6ac368f716409c55f9d3ab0.json b/.sqlx/query-dd7d80d4d118a5fc95b574e2ca9ffaccf974e52fb6ac368f716409c55f9d3ab0.json deleted file mode 100644 index 03f2acb..0000000 --- a/.sqlx/query-dd7d80d4d118a5fc95b574e2ca9ffaccf974e52fb6ac368f716409c55f9d3ab0.json +++ /dev/null @@ -1,40 +0,0 @@ -{ - "db_name": "PostgreSQL", - "query": "\n INSERT INTO external_identities (did, provider, provider_user_id, provider_username, provider_email)\n VALUES ($1, $2, $3, $4, $5)\n RETURNING id\n ", - "describe": { - "columns": [ - { - "ordinal": 0, - "name": "id", - "type_info": "Uuid" - } - ], - "parameters": { - "Left": [ - "Text", - { - "Custom": { - "name": "sso_provider_type", - "kind": { - "Enum": [ - "github", - "discord", - "google", - "gitlab", - "oidc", - "apple" - ] - } - } - }, - "Text", - "Text", - "Text" - ] - }, - "nullable": [ - false - ] - }, - "hash": "dd7d80d4d118a5fc95b574e2ca9ffaccf974e52fb6ac368f716409c55f9d3ab0" -} diff --git a/.sqlx/query-e20cbe2a939d790aaea718b084a80d8ede655ba1cc0fd4346d7e91d6de7d6cf3.json b/.sqlx/query-e20cbe2a939d790aaea718b084a80d8ede655ba1cc0fd4346d7e91d6de7d6cf3.json deleted file mode 100644 index 0576d91..0000000 --- a/.sqlx/query-e20cbe2a939d790aaea718b084a80d8ede655ba1cc0fd4346d7e91d6de7d6cf3.json +++ /dev/null @@ -1,22 +0,0 @@ -{ - "db_name": "PostgreSQL", - "query": "SELECT COUNT(*) FROM comms_queue WHERE user_id = $1 AND comms_type = 'password_reset'", - "describe": { - "columns": [ - { - "ordinal": 0, - "name": "count", - "type_info": "Int8" - } - ], - "parameters": { - "Left": [ - "Uuid" - ] - }, - "nullable": [ - null - ] - }, - "hash": "e20cbe2a939d790aaea718b084a80d8ede655ba1cc0fd4346d7e91d6de7d6cf3" -} diff --git a/.sqlx/query-e64cd36284d10ab7f3d9f6959975a1a627809f444b0faff7e611d985f31b90e9.json b/.sqlx/query-e64cd36284d10ab7f3d9f6959975a1a627809f444b0faff7e611d985f31b90e9.json deleted file mode 100644 index edb2f4b..0000000 --- a/.sqlx/query-e64cd36284d10ab7f3d9f6959975a1a627809f444b0faff7e611d985f31b90e9.json +++ /dev/null @@ -1,22 +0,0 @@ -{ - "db_name": "PostgreSQL", - "query": "SELECT used_at FROM reserved_signing_keys WHERE public_key_did_key = $1", - "describe": { - "columns": [ - { - "ordinal": 0, - "name": "used_at", - "type_info": "Timestamptz" - } - ], - "parameters": { - "Left": [ - "Text" - ] - }, - "nullable": [ - true - ] - }, - "hash": "e64cd36284d10ab7f3d9f6959975a1a627809f444b0faff7e611d985f31b90e9" -} diff --git a/.sqlx/query-eb54d2ce02cab7c2e7f9926bd469b19e5f0513f47173b2738fc01a57082d7abb.json b/.sqlx/query-eb54d2ce02cab7c2e7f9926bd469b19e5f0513f47173b2738fc01a57082d7abb.json deleted file mode 100644 index 494c6cb..0000000 --- a/.sqlx/query-eb54d2ce02cab7c2e7f9926bd469b19e5f0513f47173b2738fc01a57082d7abb.json +++ /dev/null @@ -1,30 +0,0 @@ -{ - "db_name": "PostgreSQL", - "query": "\n INSERT INTO external_identities (did, provider, provider_user_id)\n VALUES ($1, $2, $3)\n ", - "describe": { - "columns": [], - "parameters": { - "Left": [ - "Text", - { - "Custom": { - "name": "sso_provider_type", - "kind": { - "Enum": [ - "github", - "discord", - "google", - "gitlab", - "oidc", - "apple" - ] - } - } - }, - "Text" - ] - }, - "nullable": [] - }, - "hash": "eb54d2ce02cab7c2e7f9926bd469b19e5f0513f47173b2738fc01a57082d7abb" -} diff --git a/.sqlx/query-ec22a8cc89e480c403a239eac44288e144d83364129491de6156760616666d3d.json b/.sqlx/query-ec22a8cc89e480c403a239eac44288e144d83364129491de6156760616666d3d.json deleted file mode 100644 index 8cc56b0..0000000 --- a/.sqlx/query-ec22a8cc89e480c403a239eac44288e144d83364129491de6156760616666d3d.json +++ /dev/null @@ -1,15 +0,0 @@ -{ - "db_name": "PostgreSQL", - "query": "DELETE FROM external_identities WHERE id = $1 AND did = $2", - "describe": { - "columns": [], - "parameters": { - "Left": [ - "Uuid", - "Text" - ] - }, - "nullable": [] - }, - "hash": "ec22a8cc89e480c403a239eac44288e144d83364129491de6156760616666d3d" -} diff --git a/.sqlx/query-f26c13023b47b908ec96da2e6b8bf8b34ca6a2246c20fc96f76f0e95530762a7.json b/.sqlx/query-f26c13023b47b908ec96da2e6b8bf8b34ca6a2246c20fc96f76f0e95530762a7.json deleted file mode 100644 index e8c8434..0000000 --- a/.sqlx/query-f26c13023b47b908ec96da2e6b8bf8b34ca6a2246c20fc96f76f0e95530762a7.json +++ /dev/null @@ -1,22 +0,0 @@ -{ - "db_name": "PostgreSQL", - "query": "SELECT email FROM users WHERE did = $1", - "describe": { - "columns": [ - { - "ordinal": 0, - "name": "email", - "type_info": "Text" - } - ], - "parameters": { - "Left": [ - "Text" - ] - }, - "nullable": [ - true - ] - }, - "hash": "f26c13023b47b908ec96da2e6b8bf8b34ca6a2246c20fc96f76f0e95530762a7" -} diff --git a/.sqlx/query-f3b07f153284b6dd1f22c098af8628d85dcdf20dd6273443ff98d02b6f5ecbf1.json b/.sqlx/query-f3b07f153284b6dd1f22c098af8628d85dcdf20dd6273443ff98d02b6f5ecbf1.json new file mode 100644 index 0000000..c70a948 --- /dev/null +++ b/.sqlx/query-f3b07f153284b6dd1f22c098af8628d85dcdf20dd6273443ff98d02b6f5ecbf1.json @@ -0,0 +1,52 @@ +{ + "db_name": "PostgreSQL", + "query": "SELECT id, did, public_key_did_key, private_key_bytes, expires_at, used_at\n FROM reserved_signing_keys WHERE public_key_did_key = $1", + "describe": { + "columns": [ + { + "ordinal": 0, + "name": "id", + "type_info": "Uuid" + }, + { + "ordinal": 1, + "name": "did", + "type_info": "Text" + }, + { + "ordinal": 2, + "name": "public_key_did_key", + "type_info": "Text" + }, + { + "ordinal": 3, + "name": "private_key_bytes", + "type_info": "Bytea" + }, + { + "ordinal": 4, + "name": "expires_at", + "type_info": "Timestamptz" + }, + { + "ordinal": 5, + "name": "used_at", + "type_info": "Timestamptz" + } + ], + "parameters": { + "Left": [ + "Text" + ] + }, + "nullable": [ + false, + true, + false, + false, + false, + true + ] + }, + "hash": "f3b07f153284b6dd1f22c098af8628d85dcdf20dd6273443ff98d02b6f5ecbf1" +} diff --git a/Cargo.lock b/Cargo.lock index 532d671..ad1f7ff 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -7715,6 +7715,7 @@ dependencies = [ "tower-http", "tower-layer", "tracing", + "tracing-subscriber", "tranquil-api", "tranquil-auth", "tranquil-cache", @@ -7821,6 +7822,7 @@ version = "0.4.7" dependencies = [ "async-trait", "chrono", + "fjall", "futures", "presage", "rand 0.9.2", diff --git a/Cargo.toml b/Cargo.toml index ff64949..b7390ff 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -26,7 +26,7 @@ members = [ ] [workspace.package] -version = "0.4.7" +version = "0.5.0" edition = "2024" license = "AGPL-3.0-or-later" @@ -51,7 +51,7 @@ tranquil-server = { path = "crates/tranquil-server" } tranquil-sync = { path = "crates/tranquil-sync" } tranquil-oauth-server = { path = "crates/tranquil-oauth-server" } tranquil-api = { path = "crates/tranquil-api" } -tranquil-signal = { path = "crates/tranquil-signal" } +tranquil-signal = { path = "crates/tranquil-signal", features = ["fjall-store"] } tranquil-store = { path = "crates/tranquil-store" } presage = { git = "https://github.com/whisperfish/presage", rev = "fe3ed54c4844ae51c3a9fa49cf80a7816a31a425", default-features = false } diff --git a/Dockerfile b/Dockerfile index 29b2426..2bce6b0 100644 --- a/Dockerfile +++ b/Dockerfile @@ -29,6 +29,7 @@ COPY crates/tranquil-pds ./crates/tranquil-pds COPY crates/tranquil-sync ./crates/tranquil-sync COPY crates/tranquil-api ./crates/tranquil-api COPY crates/tranquil-oauth-server ./crates/tranquil-oauth-server +COPY crates/tranquil-store ./crates/tranquil-store COPY crates/tranquil-signal ./crates/tranquil-signal COPY crates/tranquil-server ./crates/tranquil-server COPY migrations ./crates/tranquil-pds/migrations diff --git a/crates/tranquil-api/src/admin/account/search.rs b/crates/tranquil-api/src/admin/account/search.rs index 7901c32..afbf47c 100644 --- a/crates/tranquil-api/src/admin/account/search.rs +++ b/crates/tranquil-api/src/admin/account/search.rs @@ -51,16 +51,14 @@ pub async fn search_accounts( Query(params): Query, ) -> Result, ApiError> { let limit = params.limit.clamp(1, 100); - let email_filter = params.email.as_deref().map(|e| format!("%{}%", e)); - let handle_filter = params.handle.as_deref().map(|h| format!("%{}%", h)); let cursor_did: Option = params.cursor.as_ref().and_then(|c| c.parse().ok()); let rows = state .repos .user .search_accounts( cursor_did.as_ref(), - email_filter.as_deref(), - handle_filter.as_deref(), + params.email.as_deref(), + params.handle.as_deref(), limit + 1, ) .await diff --git a/crates/tranquil-api/src/admin/signal.rs b/crates/tranquil-api/src/admin/signal.rs index 50614d0..8a56df0 100644 --- a/crates/tranquil-api/src/admin/signal.rs +++ b/crates/tranquil-api/src/admin/signal.rs @@ -5,7 +5,6 @@ use serde::Serialize; use tranquil_pds::api::error::ApiError; use tranquil_pds::auth::{Admin, Auth}; use tranquil_pds::state::AppState; -use tranquil_signal::PgSignalStore; #[derive(Serialize)] #[serde(rename_all = "camelCase")] @@ -53,15 +52,20 @@ pub async fn link_signal_device( let device_name = tranquil_signal::DeviceName::new("tranquil-pds".to_string()) .map_err(|e| ApiError::InternalError(Some(format!("invalid device name: {e}"))))?; - let link_result = tranquil_signal::SignalClient::link_device( - &state.repos.pool, - device_name, - state.shutdown.clone(), - link_cancel, - slot.linking_flag(), - ) - .await - .map_err(|e| ApiError::InternalError(Some(format!("Signal linking failed: {e}"))))?; + let signal_store = state + .signal_store_provider + .as_ref() + .ok_or_else(|| ApiError::InternalError(Some("Signal store not configured".into())))?; + + let link_result = signal_store + .link_signal_device( + device_name, + state.shutdown.clone(), + link_cancel, + slot.linking_flag(), + ) + .await + .map_err(|e| ApiError::InternalError(Some(format!("Signal linking failed: {e}"))))?; let qr_base64 = url_to_qr_png_base64(link_result.url.as_str()) .map_err(|e| ApiError::InternalError(Some(format!("QR generation failed: {e}"))))?; @@ -108,9 +112,13 @@ pub async fn unlink_signal_device( .as_ref() .ok_or_else(|| ApiError::InvalidRequest("Signal is not enabled".into()))?; - let store = PgSignalStore::new(state.repos.pool.clone()); - store - .clear_all() + let signal_store = state + .signal_store_provider + .as_ref() + .ok_or_else(|| ApiError::InternalError(Some("Signal store not configured".into())))?; + + signal_store + .clear_signal_data() .await .map_err(|e| ApiError::InternalError(Some(format!("Failed to clear signal data: {e}"))))?; diff --git a/crates/tranquil-api/src/delegation.rs b/crates/tranquil-api/src/delegation.rs index 39bd140..b2a30dc 100644 --- a/crates/tranquil-api/src/delegation.rs +++ b/crates/tranquil-api/src/delegation.rs @@ -18,7 +18,6 @@ use tranquil_pds::delegation::{ use tranquil_pds::rate_limit::{AccountCreationLimit, RateLimited}; use tranquil_pds::state::AppState; use tranquil_pds::types::{Did, Handle}; -use tranquil_types::CidLink; pub async fn list_controllers( State(state): State, @@ -420,22 +419,6 @@ pub async fn create_delegated_account( } }; - state - .repos - .repo - .create_repo( - user_id, - &did, - &handle, - &CidLink::from(&repo.commit_cid), - &repo.repo_rev, - ) - .await - .map_err(|e| { - error!("failed to register repo in backend: {e:?}"); - ApiError::InternalError(None) - })?; - if let Some(validated) = validated_invite_code && let Err(e) = state .repos diff --git a/crates/tranquil-api/src/identity/account.rs b/crates/tranquil-api/src/identity/account.rs index 289a9e6..83ea72a 100644 --- a/crates/tranquil-api/src/identity/account.rs +++ b/crates/tranquil-api/src/identity/account.rs @@ -15,8 +15,6 @@ use tranquil_pds::rate_limit::{AccountCreationLimit, RateLimited}; use tranquil_pds::state::AppState; use tranquil_pds::types::{Did, Handle, PlainPassword}; use tranquil_pds::validation::validate_password; -use tranquil_types::CidLink; - #[derive(Deserialize)] #[serde(rename_all = "camelCase")] pub struct CreateAccountInput { @@ -548,21 +546,6 @@ pub async fn create_account( } }; let user_id = create_result.user_id; - if let Err(e) = state - .repos - .repo - .create_repo( - user_id, - &did_for_commit, - &handle_typed, - &CidLink::from(&repo.commit_cid), - &repo.repo_rev, - ) - .await - { - error!("failed to register repo in backend: {e:?}"); - return ApiError::InternalError(None).into_response(); - } if !is_migration && !is_did_web_byod { super::provision::sequence_new_account( &state, diff --git a/crates/tranquil-api/src/server/passkey_account.rs b/crates/tranquil-api/src/server/passkey_account.rs index bbb96fc..9b3ff3e 100644 --- a/crates/tranquil-api/src/server/passkey_account.rs +++ b/crates/tranquil-api/src/server/passkey_account.rs @@ -15,7 +15,6 @@ use tranquil_pds::rate_limit::{AccountCreationLimit, PasswordResetLimit, RateLim use tranquil_pds::state::AppState; use tranquil_pds::types::{Did, Handle, PlainPassword}; use tranquil_pds::validation::validate_password; -use tranquil_types::CidLink; fn generate_setup_token() -> String { let mut rng = rand::thread_rng(); @@ -367,22 +366,6 @@ pub async fn create_passkey_account( }; let user_id = create_result.user_id; - state - .repos - .repo - .create_repo( - user_id, - &did_typed, - &handle_typed, - &CidLink::from(&repo.commit_cid), - &repo.repo_rev, - ) - .await - .map_err(|e| { - error!("failed to register repo in backend: {e:?}"); - ApiError::InternalError(None) - })?; - if !is_byod_did_web { crate::identity::provision::sequence_new_account( &state, diff --git a/crates/tranquil-db-traits/src/infra.rs b/crates/tranquil-db-traits/src/infra.rs index 716e2e2..c9879d6 100644 --- a/crates/tranquil-db-traits/src/infra.rs +++ b/crates/tranquil-db-traits/src/infra.rs @@ -185,12 +185,41 @@ pub struct ReservedSigningKey { pub private_key_bytes: Vec, } +#[derive(Debug, Clone)] +pub struct ReservedSigningKeyFull { + pub id: Uuid, + pub did: Option, + pub public_key_did_key: String, + pub private_key_bytes: Vec, + pub expires_at: DateTime, + pub used_at: Option>, +} + #[derive(Debug, Clone)] pub struct DeletionRequest { pub did: Did, pub expires_at: DateTime, } +#[derive(Debug, Clone)] +pub struct DeletionRequestWithToken { + pub token: String, + pub did: Did, + pub expires_at: DateTime, +} + +#[derive(Debug, Clone)] +pub struct PlcTokenInfo { + pub token: String, + pub expires_at: DateTime, +} + +#[derive(Debug, Clone)] +pub struct PasswordResetInfo { + pub code: Option, + pub expires_at: Option>, +} + #[async_trait] pub trait InfraRepository: Send + Sync { #[allow(clippy::too_many_arguments)] @@ -406,6 +435,41 @@ pub trait InfraRepository: Send + Sync { &self, user_ids: &[Uuid], ) -> Result, DbError>; + + async fn get_deletion_request_by_did( + &self, + did: &Did, + ) -> Result, DbError>; + + async fn get_latest_comms_for_user( + &self, + user_id: Uuid, + comms_type: CommsType, + limit: i64, + ) -> Result, DbError>; + + async fn count_comms_by_type( + &self, + user_id: Uuid, + comms_type: CommsType, + ) -> Result; + + async fn delete_comms_by_type_for_user( + &self, + user_id: Uuid, + comms_type: CommsType, + ) -> Result; + + async fn expire_deletion_request(&self, token: &str) -> Result<(), DbError>; + + async fn get_reserved_signing_key_full( + &self, + public_key_did_key: &str, + ) -> Result, DbError>; + + async fn get_plc_tokens_by_did(&self, did: &Did) -> Result, DbError>; + + async fn count_plc_tokens_by_did(&self, did: &Did) -> Result; } #[derive(Debug, Clone)] diff --git a/crates/tranquil-db-traits/src/lib.rs b/crates/tranquil-db-traits/src/lib.rs index 8d6774d..d0d6e74 100644 --- a/crates/tranquil-db-traits/src/lib.rs +++ b/crates/tranquil-db-traits/src/lib.rs @@ -22,9 +22,10 @@ pub use delegation::{ }; pub use error::DbError; pub use infra::{ - AdminAccountInfo, CommsChannel, CommsStatus, CommsType, DeletionRequest, InfraRepository, - InviteCodeInfo, InviteCodeRow, InviteCodeSortOrder, InviteCodeState, InviteCodeUse, - NotificationHistoryRow, QueuedComms, ReservedSigningKey, + AdminAccountInfo, CommsChannel, CommsStatus, CommsType, DeletionRequest, + DeletionRequestWithToken, InfraRepository, InviteCodeInfo, InviteCodeRow, InviteCodeSortOrder, + InviteCodeState, InviteCodeUse, NotificationHistoryRow, PasswordResetInfo, PlcTokenInfo, + QueuedComms, ReservedSigningKey, ReservedSigningKeyFull, }; pub use invite_code::{InviteCodeError, ValidatedInviteCode}; pub use oauth::{ diff --git a/crates/tranquil-db-traits/src/oauth.rs b/crates/tranquil-db-traits/src/oauth.rs index acff060..23c0bd2 100644 --- a/crates/tranquil-db-traits/src/oauth.rs +++ b/crates/tranquil-db-traits/src/oauth.rs @@ -324,4 +324,9 @@ pub trait OAuthRepository: Send + Sync { did: &Did, except_token_id: &TokenId, ) -> Result; + + async fn get_2fa_challenge_code( + &self, + request_uri: &RequestId, + ) -> Result, DbError>; } diff --git a/crates/tranquil-db-traits/src/user.rs b/crates/tranquil-db-traits/src/user.rs index 4e427fb..544be79 100644 --- a/crates/tranquil-db-traits/src/user.rs +++ b/crates/tranquil-db-traits/src/user.rs @@ -223,6 +223,8 @@ pub trait UserRepository: Send + Sync { async fn admin_update_password(&self, did: &Did, password_hash: &str) -> Result; + async fn set_admin_status(&self, did: &Did, is_admin: bool) -> Result<(), DbError>; + async fn get_notification_prefs(&self, did: &Did) -> Result, DbError>; @@ -584,6 +586,18 @@ pub trait UserRepository: Send + Sync { &self, input: &RecoverPasskeyAccountInput, ) -> Result; + + async fn get_password_reset_info( + &self, + email: &str, + ) -> Result, DbError>; + + async fn enable_totp_verified(&self, did: &Did, encrypted_secret: &[u8]) + -> Result<(), DbError>; + + async fn set_two_factor_enabled(&self, did: &Did, enabled: bool) -> Result<(), DbError>; + + async fn expire_password_reset_code(&self, email: &str) -> Result<(), DbError>; } #[derive(Debug, Clone)] diff --git a/crates/tranquil-db/src/postgres/infra.rs b/crates/tranquil-db/src/postgres/infra.rs index e5db8af..5045be5 100644 --- a/crates/tranquil-db/src/postgres/infra.rs +++ b/crates/tranquil-db/src/postgres/infra.rs @@ -3,9 +3,9 @@ use chrono::{DateTime, Utc}; use sqlx::PgPool; use tranquil_db_traits::{ AdminAccountInfo, CommsChannel, CommsStatus, CommsType, DbError, DeletionRequest, - InfraRepository, InviteCodeError, InviteCodeInfo, InviteCodeRow, InviteCodeSortOrder, - InviteCodeState, InviteCodeUse, NotificationHistoryRow, QueuedComms, ReservedSigningKey, - ValidatedInviteCode, + DeletionRequestWithToken, InfraRepository, InviteCodeError, InviteCodeInfo, InviteCodeRow, + InviteCodeSortOrder, InviteCodeState, InviteCodeUse, NotificationHistoryRow, PlcTokenInfo, + QueuedComms, ReservedSigningKey, ReservedSigningKeyFull, ValidatedInviteCode, }; use tranquil_types::{CidLink, Did, Handle}; use uuid::Uuid; @@ -1034,4 +1034,159 @@ impl InfraRepository for PostgresInfraRepository { .map(|r| (r.used_by_user, r.code)) .collect()) } + + async fn get_deletion_request_by_did( + &self, + did: &Did, + ) -> Result, DbError> { + let row = sqlx::query!( + r#"SELECT token, did, expires_at FROM account_deletion_requests WHERE did = $1"#, + did.as_str() + ) + .fetch_optional(&self.pool) + .await + .map_err(map_sqlx_error)?; + + Ok(row.map(|r| DeletionRequestWithToken { + token: r.token, + did: Did::new(r.did).expect("valid DID in database"), + expires_at: r.expires_at, + })) + } + + async fn get_latest_comms_for_user( + &self, + user_id: Uuid, + comms_type: CommsType, + limit: i64, + ) -> Result, DbError> { + let results = sqlx::query_as!( + QueuedComms, + r#"SELECT + id, user_id, + channel as "channel: CommsChannel", + comms_type as "comms_type: CommsType", + status as "status: CommsStatus", + recipient, subject, body, metadata, + attempts, max_attempts, last_error, + created_at, updated_at, scheduled_for, processed_at + FROM comms_queue + WHERE user_id = $1 AND comms_type = $2 + ORDER BY created_at DESC + LIMIT $3"#, + user_id, + comms_type as CommsType, + limit + ) + .fetch_all(&self.pool) + .await + .map_err(map_sqlx_error)?; + + Ok(results) + } + + async fn count_comms_by_type( + &self, + user_id: Uuid, + comms_type: CommsType, + ) -> Result { + let count = sqlx::query_scalar!( + r#"SELECT COUNT(*) as "count!" FROM comms_queue WHERE user_id = $1 AND comms_type = $2"#, + user_id, + comms_type as CommsType + ) + .fetch_one(&self.pool) + .await + .map_err(map_sqlx_error)?; + + Ok(count) + } + + async fn delete_comms_by_type_for_user( + &self, + user_id: Uuid, + comms_type: CommsType, + ) -> Result { + let result = sqlx::query!( + "DELETE FROM comms_queue WHERE user_id = $1 AND comms_type = $2", + user_id, + comms_type as CommsType + ) + .execute(&self.pool) + .await + .map_err(map_sqlx_error)?; + + Ok(result.rows_affected()) + } + + async fn expire_deletion_request(&self, token: &str) -> Result<(), DbError> { + sqlx::query!( + "UPDATE account_deletion_requests SET expires_at = NOW() - INTERVAL '1 hour' WHERE token = $1", + token + ) + .execute(&self.pool) + .await + .map_err(map_sqlx_error)?; + + Ok(()) + } + + async fn get_reserved_signing_key_full( + &self, + public_key_did_key: &str, + ) -> Result, DbError> { + let row = sqlx::query!( + r#"SELECT id, did, public_key_did_key, private_key_bytes, expires_at, used_at + FROM reserved_signing_keys WHERE public_key_did_key = $1"#, + public_key_did_key + ) + .fetch_optional(&self.pool) + .await + .map_err(map_sqlx_error)?; + + Ok(row.map(|r| ReservedSigningKeyFull { + id: r.id, + did: r.did.map(|d| Did::new(d).expect("valid DID in database")), + public_key_did_key: r.public_key_did_key, + private_key_bytes: r.private_key_bytes, + expires_at: r.expires_at, + used_at: r.used_at, + })) + } + + async fn get_plc_tokens_by_did(&self, did: &Did) -> Result, DbError> { + let results = sqlx::query!( + r#"SELECT t.token, t.expires_at + FROM plc_operation_tokens t + JOIN users u ON t.user_id = u.id + WHERE u.did = $1"#, + did.as_str() + ) + .fetch_all(&self.pool) + .await + .map_err(map_sqlx_error)?; + + Ok(results + .into_iter() + .map(|r| PlcTokenInfo { + token: r.token, + expires_at: r.expires_at, + }) + .collect()) + } + + async fn count_plc_tokens_by_did(&self, did: &Did) -> Result { + let count = sqlx::query_scalar!( + r#"SELECT COUNT(*) as "count!" + FROM plc_operation_tokens t + JOIN users u ON t.user_id = u.id + WHERE u.did = $1"#, + did.as_str() + ) + .fetch_one(&self.pool) + .await + .map_err(map_sqlx_error)?; + + Ok(count) + } } diff --git a/crates/tranquil-db/src/postgres/mod.rs b/crates/tranquil-db/src/postgres/mod.rs index bb61a2d..a05da3a 100644 --- a/crates/tranquil-db/src/postgres/mod.rs +++ b/crates/tranquil-db/src/postgres/mod.rs @@ -28,7 +28,7 @@ use tranquil_db_traits::{ pub use user::PostgresUserRepository; pub struct PostgresRepositories { - pub pool: PgPool, + pub pool: Option, pub user: Arc, pub oauth: Arc, pub session: Arc, @@ -44,7 +44,7 @@ pub struct PostgresRepositories { impl PostgresRepositories { pub fn new(pool: PgPool) -> Self { Self { - pool: pool.clone(), + pool: Some(pool.clone()), user: Arc::new(PostgresUserRepository::new(pool.clone())), oauth: Arc::new(PostgresOAuthRepository::new(pool.clone())), session: Arc::new(PostgresSessionRepository::new(pool.clone())), diff --git a/crates/tranquil-db/src/postgres/oauth.rs b/crates/tranquil-db/src/postgres/oauth.rs index 67adddb..948912f 100644 --- a/crates/tranquil-db/src/postgres/oauth.rs +++ b/crates/tranquil-db/src/postgres/oauth.rs @@ -1323,4 +1323,19 @@ impl OAuthRepository for PostgresOAuthRepository { .map_err(map_sqlx_error)?; Ok(result.rows_affected()) } + + async fn get_2fa_challenge_code( + &self, + request_uri: &RequestId, + ) -> Result, DbError> { + let code = sqlx::query_scalar!( + "SELECT code FROM oauth_2fa_challenge WHERE request_uri = $1", + request_uri.as_str() + ) + .fetch_optional(&self.pool) + .await + .map_err(map_sqlx_error)?; + + Ok(code) + } } diff --git a/crates/tranquil-db/src/postgres/user.rs b/crates/tranquil-db/src/postgres/user.rs index 250faf6..db840f6 100644 --- a/crates/tranquil-db/src/postgres/user.rs +++ b/crates/tranquil-db/src/postgres/user.rs @@ -234,6 +234,8 @@ impl UserRepository for PostgresUserRepository { limit: i64, ) -> Result, DbError> { let cursor_str = cursor_did.map(|d| d.as_str()); + let email_like = email_filter.map(|e| format!("%{e}%")); + let handle_like = handle_filter.map(|h| format!("%{h}%")); let rows = sqlx::query!( r#"SELECT did, handle, email, created_at, email_verified, deactivated_at, invites_disabled FROM users @@ -243,8 +245,8 @@ impl UserRepository for PostgresUserRepository { ORDER BY did ASC LIMIT $4"#, cursor_str, - email_filter, - handle_filter, + email_like.as_deref(), + handle_like.as_deref(), limit ) .fetch_all(&self.pool) @@ -627,6 +629,18 @@ impl UserRepository for PostgresUserRepository { Ok(result.rows_affected()) } + async fn set_admin_status(&self, did: &Did, is_admin: bool) -> Result<(), DbError> { + sqlx::query!( + "UPDATE users SET is_admin = $1 WHERE did = $2", + is_admin, + did.as_str() + ) + .execute(&self.pool) + .await + .map_err(map_sqlx_error)?; + Ok(()) + } + async fn get_notification_prefs( &self, did: &Did, @@ -3306,4 +3320,66 @@ impl UserRepository for PostgresUserRepository { .map_err(map_sqlx_error)?; Ok(row.flatten()) } + + async fn get_password_reset_info( + &self, + email: &str, + ) -> Result, DbError> { + let row = sqlx::query!( + "SELECT password_reset_code, password_reset_code_expires_at FROM users WHERE email = $1", + email + ) + .fetch_optional(&self.pool) + .await + .map_err(map_sqlx_error)?; + + Ok(row.map(|r| tranquil_db_traits::PasswordResetInfo { + code: r.password_reset_code, + expires_at: r.password_reset_code_expires_at, + })) + } + + async fn enable_totp_verified( + &self, + did: &Did, + encrypted_secret: &[u8], + ) -> Result<(), DbError> { + 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 secret_encrypted = $2, verified = TRUE"#, + did.as_str(), + encrypted_secret + ) + .execute(&self.pool) + .await + .map_err(map_sqlx_error)?; + + Ok(()) + } + + async fn set_two_factor_enabled(&self, did: &Did, enabled: bool) -> Result<(), DbError> { + sqlx::query!( + "UPDATE users SET two_factor_enabled = $1 WHERE did = $2", + enabled, + did.as_str() + ) + .execute(&self.pool) + .await + .map_err(map_sqlx_error)?; + + Ok(()) + } + + async fn expire_password_reset_code(&self, email: &str) -> Result<(), DbError> { + sqlx::query!( + "UPDATE users SET password_reset_code_expires_at = NOW() - INTERVAL '1 hour' WHERE email = $1", + email + ) + .execute(&self.pool) + .await + .map_err(map_sqlx_error)?; + + Ok(()) + } } diff --git a/crates/tranquil-pds/Cargo.toml b/crates/tranquil-pds/Cargo.toml index d4ec5d6..1bd6c92 100644 --- a/crates/tranquil-pds/Cargo.toml +++ b/crates/tranquil-pds/Cargo.toml @@ -99,4 +99,5 @@ tranquil-ripple = { workspace = true } tranquil-sync = { workspace = true } tranquil-api = { workspace = true } tranquil-oauth-server = { workspace = true } +tracing-subscriber = { workspace = true, features = ["env-filter"] } wiremock = { workspace = true } diff --git a/crates/tranquil-pds/src/state.rs b/crates/tranquil-pds/src/state.rs index bfd3fc9..6a39e62 100644 --- a/crates/tranquil-pds/src/state.rs +++ b/crates/tranquil-pds/src/state.rs @@ -48,6 +48,7 @@ pub struct AppState { pub shutdown: CancellationToken, pub bootstrap_invite_code: Option, pub signal_sender: Option>, + pub signal_store_provider: Option>, } #[derive(Debug, Clone, Copy)] @@ -210,76 +211,122 @@ impl AppState { pub async fn new(shutdown: CancellationToken) -> Result> { let cfg = tranquil_config::get(); - let database_url = &cfg.database.url; - let max_connections = cfg.database.max_connections; - let min_connections = cfg.database.min_connections; - let acquire_timeout_secs = cfg.database.acquire_timeout_secs; - tracing::info!( - "Configuring database pool: max={}, min={}, acquire_timeout={}s", - max_connections, - min_connections, - acquire_timeout_secs - ); - - let db = sqlx::postgres::PgPoolOptions::new() - .max_connections(max_connections) - .min_connections(min_connections) - .acquire_timeout(std::time::Duration::from_secs(acquire_timeout_secs)) - .idle_timeout(std::time::Duration::from_secs(300)) - .max_lifetime(std::time::Duration::from_secs(1800)) - .connect(database_url) - .await - .map_err(|e| format!("Failed to connect to Postgres: {}", e))?; - - sqlx::migrate!("./migrations") - .run(&db) - .await - .map_err(|e| format!("Failed to run migrations: {}", e))?; - - let bootstrap_invite_code = match ( - cfg.server.invite_code_required, - sqlx::query_scalar!("SELECT COUNT(*) FROM users") - .fetch_one(&db) - .await, - ) { - (true, Ok(Some(0))) => { - let code = crate::util::gen_invite_code(); + match cfg.storage.repo_backend() { + tranquil_config::RepoBackend::TranquilStore => { tracing::info!( - "No users exist and invite codes are required. Bootstrap invite code: {}", - code + "tranquil-store repo backend active. EXPERIMENTAL! No garbage collection, no backup/restore" ); - Some(code) + Ok(Self::from_store(shutdown).await) } - _ => None, - }; + tranquil_config::RepoBackend::Postgres => { + let database_url = &cfg.database.url; + let max_connections = cfg.database.max_connections; + let min_connections = cfg.database.min_connections; + let acquire_timeout_secs = cfg.database.acquire_timeout_secs; - let mut state = Self::from_db(db, shutdown).await; - state.bootstrap_invite_code = bootstrap_invite_code; - Ok(state) + tracing::info!( + "Configuring database pool: max={}, min={}, acquire_timeout={}s", + max_connections, + min_connections, + acquire_timeout_secs + ); + + let db = sqlx::postgres::PgPoolOptions::new() + .max_connections(max_connections) + .min_connections(min_connections) + .acquire_timeout(std::time::Duration::from_secs(acquire_timeout_secs)) + .idle_timeout(std::time::Duration::from_secs(300)) + .max_lifetime(std::time::Duration::from_secs(1800)) + .connect(database_url) + .await + .map_err(|e| format!("Failed to connect to Postgres: {}", e))?; + + sqlx::migrate!("./migrations") + .run(&db) + .await + .map_err(|e| format!("Failed to run migrations: {}", e))?; + + let bootstrap_invite_code = match ( + cfg.server.invite_code_required, + sqlx::query_scalar!("SELECT COUNT(*) FROM users") + .fetch_one(&db) + .await, + ) { + (true, Ok(Some(0))) => { + let code = crate::util::gen_invite_code(); + tracing::info!( + "No users exist and invite codes are required. Bootstrap invite code: {}", + code + ); + Some(code) + } + _ => None, + }; + + let mut state = Self::from_db(db, shutdown).await; + state.bootstrap_invite_code = bootstrap_invite_code; + Ok(state) + } + } } pub async fn from_db(db: PgPool, shutdown: CancellationToken) -> Self { + let cfg = tranquil_config::get(); + let (repos, block_store, signal_store_provider): ( + PostgresRepositories, + crate::repo::AnyBlockStore, + Option>, + ) = match cfg.storage.repo_backend() == tranquil_config::RepoBackend::TranquilStore { + true => { + let wiring = wire_tranquil_store(&cfg.tranquil_store, shutdown.clone()); + ( + wiring.repos, + crate::repo::AnyBlockStore::TranquilStore(wiring.blockstore), + Some(wiring.signal_provider), + ) + } + false => { + let repos = PostgresRepositories::new(db.clone()); + let provider: Arc = + Arc::new(tranquil_signal::PgSignalStoreProvider { pool: db.clone() }); + ( + repos, + crate::repo::AnyBlockStore::Postgres(PostgresBlockStore::new(db)), + Some(provider), + ) + } + }; + + Self::build(repos, block_store, signal_store_provider, shutdown).await + } + + pub async fn from_store(shutdown: CancellationToken) -> Self { + let cfg = tranquil_config::get(); + let wiring = wire_tranquil_store(&cfg.tranquil_store, shutdown.clone()); + + Self::build( + wiring.repos, + crate::repo::AnyBlockStore::TranquilStore(wiring.blockstore), + Some(wiring.signal_provider), + shutdown, + ) + .await + } + + async fn build( + repos: PostgresRepositories, + block_store: crate::repo::AnyBlockStore, + signal_store_provider: Option>, + shutdown: CancellationToken, + ) -> Self { AuthConfig::init(); init_rate_limit_override(); - let mut repos = PostgresRepositories::new(db.clone()); - let cfg = tranquil_config::get(); - let block_store = - match cfg.storage.repo_backend() == tranquil_config::RepoBackend::TranquilStore { - true => { - let bs = wire_tranquil_store(&mut repos, &cfg.tranquil_store, shutdown.clone()); - crate::repo::AnyBlockStore::TranquilStore(bs) - } - false => crate::repo::AnyBlockStore::Postgres(PostgresBlockStore::new(db)), - }; - let repos = Arc::new(repos); let blob_store = create_blob_storage().await; - - let firehose_buffer_size = tranquil_config::get().firehose.buffer_size; - + let firehose_buffer_size = cfg.firehose.buffer_size; let (firehose_tx, _) = broadcast::channel(firehose_buffer_size); let rate_limiters = Arc::new(RateLimiters::new()); let repo_write_locks = Arc::new(RepoWriteLocks::new()); @@ -290,7 +337,7 @@ impl AppState { let sso_config = SsoConfig::init(); let sso_manager = SsoManager::from_config(sso_config); let webauthn_config = Arc::new( - WebAuthnConfig::new(&tranquil_config::get().server.hostname) + WebAuthnConfig::new(&cfg.server.hostname) .expect("Failed to create WebAuthn config at startup"), ); @@ -311,6 +358,7 @@ impl AppState { shutdown, bootstrap_invite_code: None, signal_sender: None, + signal_store_provider, } } @@ -390,11 +438,16 @@ impl AppState { } } +struct TranquilStoreWiring { + blockstore: tranquil_store::blockstore::TranquilBlockStore, + signal_provider: Arc, + repos: PostgresRepositories, +} + fn wire_tranquil_store( - repos: &mut PostgresRepositories, store_cfg: &tranquil_config::TranquilStoreConfig, shutdown: CancellationToken, -) -> tranquil_store::blockstore::TranquilBlockStore { +) -> TranquilStoreWiring { use tranquil_store::RealIO; use tranquil_store::blockstore::{BlockStoreConfig, TranquilBlockStore}; use tranquil_store::eventlog::{EventLog, EventLogBridge, EventLogConfig}; @@ -460,6 +513,8 @@ fn wire_tranquil_store( } let notifier = bridge.notifier(); + let signal_db = metastore.database().clone(); + let signal_ks = metastore.signal_keyspace(); let pool = Arc::new(HandlerPool::spawn::( metastore, @@ -480,16 +535,27 @@ fn wire_tranquil_store( tracing::info!(data_dir = %store_cfg.data_dir, "tranquil-store data directory"); - repos.repo = Arc::new(client.clone()); - repos.backlink = Arc::new(client.clone()); - repos.blob = Arc::new(client.clone()); - repos.user = Arc::new(client.clone()); - repos.session = Arc::new(client.clone()); - repos.oauth = Arc::new(client.clone()); - repos.infra = Arc::new(client.clone()); - repos.delegation = Arc::new(client.clone()); - repos.sso = Arc::new(client); - repos.event_notifier = Arc::new(notifier); + let repos = PostgresRepositories { + pool: None, + repo: Arc::new(client.clone()), + backlink: Arc::new(client.clone()), + blob: Arc::new(client.clone()), + user: Arc::new(client.clone()), + session: Arc::new(client.clone()), + oauth: Arc::new(client.clone()), + infra: Arc::new(client.clone()), + delegation: Arc::new(client.clone()), + sso: Arc::new(client), + event_notifier: Arc::new(notifier), + }; - blockstore + let signal_provider: Arc = Arc::new( + tranquil_signal::fjall_store::FjallSignalStoreProvider::new(signal_db, signal_ks), + ); + + TranquilStoreWiring { + blockstore, + signal_provider, + repos, + } } diff --git a/crates/tranquil-pds/tests/account_notifications.rs b/crates/tranquil-pds/tests/account_notifications.rs index 38fd6e3..7dfb562 100644 --- a/crates/tranquil-pds/tests/account_notifications.rs +++ b/crates/tranquil-pds/tests/account_notifications.rs @@ -1,32 +1,37 @@ mod common; -use common::{base_url, client, create_account_and_login, get_test_db_pool}; +use common::{base_url, client, create_account_and_login, get_test_repos}; use serde_json::{Value, json}; +use tranquil_db_traits::{CommsChannel, CommsType}; +use tranquil_types::Did; #[tokio::test] async fn test_get_notification_history() { let client = client(); let base = base_url().await; - let pool = get_test_db_pool().await; + let repos = get_test_repos().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") - .bind(&did) - .fetch_one(pool) + let user_id = repos + .user + .get_id_by_did(&Did::new(did).unwrap()) .await + .expect("DB error") .expect("User not found"); for i in 0..3 { - sqlx::query( - r#"INSERT INTO comms_queue (user_id, channel, comms_type, recipient, subject, body) - VALUES ($1, 'email', 'welcome', $2, $3, $4)"#, - ) - .bind(user_id) - .bind("test@example.com") - .bind(format!("Subject {}", i)) - .bind(format!("Body {}", i)) - .execute(pool) - .await - .expect("Failed to enqueue"); + repos + .infra + .enqueue_comms( + Some(user_id), + CommsChannel::Email, + CommsType::Welcome, + "test@example.com", + Some(&format!("Subject {}", i)), + &format!("Body {}", i), + None, + ) + .await + .expect("Failed to enqueue"); } let resp = client @@ -140,7 +145,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_test_db_pool().await; + let repos = get_test_repos().await; let (token, did) = create_account_and_login(&client).await; let unique_email = format!("newemail_{}@example.com", uuid::Uuid::new_v4()); @@ -163,19 +168,22 @@ async fn test_update_email_via_notification_prefs() { .contains(&json!("email")) ); - let user_id: uuid::Uuid = sqlx::query_scalar("SELECT id FROM users WHERE did = $1") - .bind(&did) - .fetch_one(pool) + let user_id = repos + .user + .get_id_by_did(&Did::new(did).unwrap()) .await + .expect("DB error") .expect("User not found"); - let body_text: String = sqlx::query_scalar( - "SELECT body FROM comms_queue WHERE user_id = $1 AND comms_type = 'email_update' ORDER BY created_at DESC LIMIT 1", - ) - .bind(user_id) - .fetch_one(pool) - .await - .expect("Verification code not found"); + let comms = repos + .infra + .get_latest_comms_for_user(user_id, CommsType::EmailUpdate, 1) + .await + .expect("DB error"); + let body_text = comms + .first() + .map(|c| c.body.clone()) + .expect("Verification code not found"); let code = body_text .lines() diff --git a/crates/tranquil-pds/tests/admin_email.rs b/crates/tranquil-pds/tests/admin_email.rs index c912704..00698d7 100644 --- a/crates/tranquil-pds/tests/admin_email.rs +++ b/crates/tranquil-pds/tests/admin_email.rs @@ -2,12 +2,14 @@ mod common; use reqwest::StatusCode; use serde_json::{Value, json}; +use tranquil_db_traits::CommsType; +use tranquil_types::Did; #[tokio::test] async fn test_send_email_success() { let client = common::client(); let base_url = common::base_url().await; - let pool = common::get_test_db_pool().await; + let repos = common::get_test_repos().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)) @@ -24,17 +26,18 @@ async fn test_send_email_success() { assert_eq!(res.status(), StatusCode::OK); 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) + let user_id = repos + .user + .get_id_by_did(&Did::new(did).unwrap()) .await + .expect("DB error") .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) - .await - .expect("Notification not found"); + let comms = repos + .infra + .get_latest_comms_for_user(user_id, CommsType::AdminEmail, 1) + .await + .expect("DB error"); + let notification = comms.first().expect("Notification not found"); assert_eq!(notification.subject.as_deref(), Some("Test Admin Email")); assert!( notification @@ -47,7 +50,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 = common::get_test_db_pool().await; + let repos = common::get_test_repos().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)) @@ -63,19 +66,29 @@ async fn test_send_email_default_subject() { assert_eq!(res.status(), StatusCode::OK); 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) + let user_id = repos + .user + .get_id_by_did(&Did::new(did).unwrap()) .await + .expect("DB error") .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) - .await - .expect("Notification not found"); + let comms = repos + .infra + .get_latest_comms_for_user(user_id, CommsType::AdminEmail, 10) + .await + .expect("DB error"); + let notification = comms + .iter() + .find(|c| c.body == "Email without subject") + .expect("Notification not found"); assert!(notification.subject.is_some()); - assert!(notification.subject.unwrap().contains("Message from")); + assert!( + notification + .subject + .as_ref() + .unwrap() + .contains("Message from") + ); } #[tokio::test] diff --git a/crates/tranquil-pds/tests/auth_extractor.rs b/crates/tranquil-pds/tests/auth_extractor.rs index fb1faaf..a8c6ce2 100644 --- a/crates/tranquil-pds/tests/auth_extractor.rs +++ b/crates/tranquil-pds/tests/auth_extractor.rs @@ -215,9 +215,10 @@ async fn test_oauth_admin_extractor_allows_oauth_tokens() { let did = account["did"].as_str().unwrap().to_string(); verify_new_account(&http_client, &did).await; - let pool = common::get_test_db_pool().await; - sqlx::query!("UPDATE users SET is_admin = TRUE WHERE did = $1", &did) - .execute(pool) + let repos = common::get_test_repos().await; + repos + .user + .set_admin_status(&tranquil_types::Did::new(did.clone()).unwrap(), true) .await .expect("Failed to mark user as admin"); diff --git a/crates/tranquil-pds/tests/common/mod.rs b/crates/tranquil-pds/tests/common/mod.rs index ddea18b..590b3af 100644 --- a/crates/tranquil-pds/tests/common/mod.rs +++ b/crates/tranquil-pds/tests/common/mod.rs @@ -28,10 +28,18 @@ static MOCK_PLC: OnceLock = OnceLock::new(); static TEST_DB_POOL: OnceLock = OnceLock::new(); static TEST_TEMP_DIR: OnceLock = OnceLock::new(); static CLUSTER: OnceLock> = OnceLock::new(); +static TEST_REPOS: OnceLock> = OnceLock::new(); + +#[allow(dead_code)] +pub fn is_store_backend() -> bool { + std::env::var("TRANQUIL_TEST_BACKEND") + .map(|v| v == "store") + .unwrap_or(false) +} #[allow(dead_code)] pub struct ServerConfig { - pub pool: sqlx::PgPool, + pub pool: Option, pub cache: Option<(Arc, Arc)>, } @@ -123,6 +131,12 @@ pub async fn base_url() -> &'static str { SERVER_URL.get_or_init(|| { let (tx, rx) = std::sync::mpsc::channel(); std::thread::spawn(move || { + let _ = tracing_subscriber::fmt() + .with_env_filter( + tracing_subscriber::EnvFilter::try_from_default_env() + .unwrap_or_else(|_| tracing_subscriber::EnvFilter::new("warn")), + ) + .try_init(); unsafe { std::env::set_var("TRANQUIL_PDS_ALLOW_INSECURE_SECRETS", "1"); } @@ -141,7 +155,10 @@ pub async fn base_url() -> &'static str { } let rt = tokio::runtime::Runtime::new().unwrap(); rt.block_on(async move { - if has_external_infra() { + if is_store_backend() { + let url = setup_store_backend().await; + tx.send(url).unwrap(); + } else if has_external_infra() { let url = setup_with_external_infra().await; tx.send(url).unwrap(); } else { @@ -557,9 +574,12 @@ async fn spawn_server(config: ServerConfig) -> ServerInstance { .with_oauth_authorize_limit(10000) .with_oauth_token_limit(10000); let cache_refs = config.cache.as_ref().map(|(c, r)| (c.clone(), r.clone())); - let mut state = AppState::from_db(config.pool, CancellationToken::new()) - .await - .with_rate_limiters(rate_limiters); + let mut state = match config.pool { + Some(pool) => AppState::from_db(pool, CancellationToken::new()).await, + None => AppState::from_store(CancellationToken::new()).await, + }; + state = state.with_rate_limiters(rate_limiters); + TEST_REPOS.set(state.repos.clone()).ok(); if let Some((cache, distributed_rate_limiter)) = config.cache { state = state.with_cache(cache, distributed_rate_limiter); } @@ -590,6 +610,39 @@ async fn spawn_server(config: ServerConfig) -> ServerInstance { } } +async fn setup_store_backend() -> String { + let temp_dir = + std::env::temp_dir().join(format!("tranquil-pds-store-{}", uuid::Uuid::new_v4())); + let blob_path = temp_dir.join("blobs"); + let backup_path = temp_dir.join("backups"); + let store_path = temp_dir.join("store"); + std::fs::create_dir_all(&blob_path).expect("failed to create blob temp directory"); + std::fs::create_dir_all(&backup_path).expect("failed to create backup temp directory"); + std::fs::create_dir_all(&store_path).expect("failed to create store temp directory"); + TEST_TEMP_DIR.set(temp_dir).ok(); + let plc_url = setup_mock_plc_directory().await; + unsafe { + std::env::set_var("BLOB_STORAGE_BACKEND", "filesystem"); + std::env::set_var("BLOB_STORAGE_PATH", blob_path.to_str().unwrap()); + std::env::set_var("BACKUP_STORAGE_BACKEND", "filesystem"); + std::env::set_var("BACKUP_STORAGE_PATH", backup_path.to_str().unwrap()); + std::env::set_var("MAX_IMPORT_SIZE", "100000000"); + std::env::set_var("SKIP_IMPORT_VERIFICATION", "true"); + std::env::set_var("PLC_DIRECTORY_URL", &plc_url); + std::env::set_var("REPO_BACKEND", "tranquil-store"); + std::env::set_var("TRANQUIL_STORE_DATA_DIR", store_path.to_str().unwrap()); + std::env::set_var("DATABASE_URL", "postgres://unused/unused"); + } + register_mock_appview().await; + let instance = spawn_server(ServerConfig { + pool: None, + cache: None, + }) + .await; + APP_PORT.set(instance.port).ok(); + instance.url +} + async fn spawn_app(database_url: String) -> String { let pool = PgPoolOptions::new() .max_connections(10) @@ -608,7 +661,11 @@ async fn spawn_app(database_url: String) -> String { .await .expect("Failed to create test pool"); TEST_DB_POOL.set(test_pool).ok(); - let instance = spawn_server(ServerConfig { pool, cache: None }).await; + let instance = spawn_server(ServerConfig { + pool: Some(pool), + cache: None, + }) + .await; APP_PORT.set(instance.port).ok(); instance.url } @@ -659,7 +716,7 @@ pub async fn spawn_cluster(database_url: String, node_count: usize) -> Vec = Vec::with_capacity(node_count); for (cache, rate_limiter) in ripple_nodes { let server_config = ServerConfig { - pool: pool.clone(), + pool: Some(pool.clone()), cache: Some((cache, rate_limiter)), }; let instance = spawn_server(server_config).await; @@ -799,18 +856,14 @@ pub async fn get_test_db_pool() -> &'static sqlx::PgPool { } #[allow(dead_code)] -pub async fn verify_new_account(client: &Client, did: &str) -> String { - 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 - .expect("Failed to get verification code"); +pub async fn get_test_repos() -> &'static Arc { + base_url().await; + TEST_REPOS.get().expect("TEST_REPOS not initialized") +} +fn extract_verification_code(body_text: &str) -> String { let lines: Vec<&str> = body_text.lines().collect(); - let verification_code = lines + lines .iter() .enumerate() .find(|(_, line)| line.contains("verification code is:") || line.contains("code is:")) @@ -821,7 +874,35 @@ pub async fn verify_new_account(client: &Client, did: &str) -> String { .find(|line| line.trim().starts_with("MX")) .map(|s| s.trim().to_string()) }) - .unwrap_or_else(|| body_text.clone()); + .unwrap_or_else(|| body_text.to_string()) +} + +async fn get_verification_body_for_did(did: &str) -> String { + use tranquil_db_traits::CommsType; + use tranquil_types::Did; + + let repos = get_test_repos().await; + let user = repos + .user + .get_by_did(&Did::new(did.to_string()).unwrap()) + .await + .expect("failed to look up user") + .expect("user not found"); + let comms = repos + .infra + .get_latest_comms_for_user(user.id, CommsType::EmailVerification, 1) + .await + .expect("failed to get comms"); + comms + .first() + .map(|c| c.body.clone()) + .expect("no email_verification comms found") +} + +#[allow(dead_code)] +pub async fn verify_new_account(client: &Client, did: &str) -> String { + let body_text = get_verification_body_for_did(did).await; + let verification_code = extract_verification_code(&body_text); let confirm_payload = json!({ "did": did, @@ -956,10 +1037,11 @@ 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 pool = get_test_db_pool().await; if make_admin { - sqlx::query!("UPDATE users SET is_admin = TRUE WHERE did = $1", &did) - .execute(pool) + let repos = get_test_repos().await; + repos + .user + .set_admin_status(&tranquil_types::Did::new(did.clone()).unwrap(), true) .await .expect("Failed to mark user as admin"); } @@ -969,28 +1051,8 @@ async fn create_account_and_login_internal(client: &Client, make_admin: bool) -> { return (access_jwt.to_string(), did); } - 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 - .expect("Failed to get verification from comms_queue"); - let lines: Vec<&str> = body_text.lines().collect(); - let verification_code = lines - .iter() - .enumerate() - .find(|(_, line): &(usize, &&str)| { - line.contains("verification code is:") || line.contains("code is:") - }) - .and_then(|(i, _)| lines.get(i + 1).map(|s: &&str| s.trim().to_string())) - .or_else(|| { - body_text - .lines() - .find(|line| line.trim().starts_with("MX")) - .map(|s| s.trim().to_string()) - }) - .unwrap_or_else(|| body_text.clone()); + let body_text = get_verification_body_for_did(&did).await; + let verification_code = extract_verification_code(&body_text); let confirm_payload = json!({ "did": did, diff --git a/crates/tranquil-pds/tests/delete_account.rs b/crates/tranquil-pds/tests/delete_account.rs index 4c59565..0f105d7 100644 --- a/crates/tranquil-pds/tests/delete_account.rs +++ b/crates/tranquil-pds/tests/delete_account.rs @@ -51,15 +51,14 @@ 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_test_db_pool().await; - let row = sqlx::query!( - "SELECT token FROM account_deletion_requests WHERE did = $1", - did - ) - .fetch_one(pool) - .await - .expect("Failed to query deletion token"); - let token = row.token; + let repos = get_test_repos().await; + let deletion_request = repos + .infra + .get_deletion_request_by_did(&tranquil_types::Did::new(did.clone()).unwrap()) + .await + .unwrap() + .unwrap(); + let token = deletion_request.token; let delete_payload = json!({ "did": did, "password": password, @@ -75,11 +74,12 @@ async fn test_delete_account_full_flow() { .await .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) + let user = repos + .user + .get_by_did(&tranquil_types::Did::new(did.clone()).unwrap()) .await - .expect("Failed to query user"); - assert!(user_row.is_none(), "User should be deleted from database"); + .unwrap(); + assert!(user.is_none(), "User should be deleted from database"); let session_res = client .get(format!("{}/xrpc/com.atproto.server.getSession", base_url)) .bearer_auth(&jwt) @@ -108,15 +108,14 @@ 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_test_db_pool().await; - let row = sqlx::query!( - "SELECT token FROM account_deletion_requests WHERE did = $1", - did - ) - .fetch_one(pool) - .await - .expect("Failed to query deletion token"); - let token = row.token; + let repos = get_test_repos().await; + let deletion_request = repos + .infra + .get_deletion_request_by_did(&tranquil_types::Did::new(did.clone()).unwrap()) + .await + .unwrap() + .unwrap(); + let token = deletion_request.token; let delete_payload = json!({ "did": did, "password": "wrong-password", @@ -198,22 +197,15 @@ 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_test_db_pool().await; - let row = sqlx::query!( - "SELECT token FROM account_deletion_requests WHERE did = $1", - did - ) - .fetch_one(pool) - .await - .expect("Failed to query deletion token"); - let token = row.token; - sqlx::query!( - "UPDATE account_deletion_requests SET expires_at = NOW() - INTERVAL '1 hour' WHERE token = $1", - token - ) - .execute(pool) - .await - .expect("Failed to expire token"); + let repos = get_test_repos().await; + let deletion_request = repos + .infra + .get_deletion_request_by_did(&tranquil_types::Did::new(did.clone()).unwrap()) + .await + .unwrap() + .unwrap(); + let token = deletion_request.token; + repos.infra.expire_deletion_request(&token).await.unwrap(); let delete_payload = json!({ "did": did, "password": password, @@ -257,15 +249,14 @@ 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_test_db_pool().await; - let row = sqlx::query!( - "SELECT token FROM account_deletion_requests WHERE did = $1", - did1 - ) - .fetch_one(pool) - .await - .expect("Failed to query deletion token"); - let token = row.token; + let repos = get_test_repos().await; + let deletion_request = repos + .infra + .get_deletion_request_by_did(&tranquil_types::Did::new(did1.clone()).unwrap()) + .await + .unwrap() + .unwrap(); + let token = deletion_request.token; let delete_payload = json!({ "did": did2, "password": password2, @@ -318,15 +309,14 @@ 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_test_db_pool().await; - let row = sqlx::query!( - "SELECT token FROM account_deletion_requests WHERE did = $1", - did - ) - .fetch_one(pool) - .await - .expect("Failed to query deletion token"); - let token = row.token; + let repos = get_test_repos().await; + let deletion_request = repos + .infra + .get_deletion_request_by_did(&tranquil_types::Did::new(did.clone()).unwrap()) + .await + .unwrap() + .unwrap(); + let token = deletion_request.token; let delete_payload = json!({ "did": did, "password": app_password, @@ -342,11 +332,12 @@ async fn test_delete_account_with_app_password() { .await .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) + let user = repos + .user + .get_by_did(&tranquil_types::Did::new(did.clone()).unwrap()) .await - .expect("Failed to query user"); - assert!(user_row.is_none(), "User should be deleted from database"); + .unwrap(); + assert!(user.is_none(), "User should be deleted from database"); } #[tokio::test] diff --git a/crates/tranquil-pds/tests/email_update.rs b/crates/tranquil-pds/tests/email_update.rs index ee628dd..16e995b 100644 --- a/crates/tranquil-pds/tests/email_update.rs +++ b/crates/tranquil-pds/tests/email_update.rs @@ -1,16 +1,24 @@ mod common; use reqwest::StatusCode; use serde_json::{Value, json}; -use sqlx::PgPool; +use tranquil_db_traits::CommsType; +use tranquil_types::Did; -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", - did - ) - .fetch_one(pool) - .await - .expect("Verification not found"); +async fn get_email_update_token(did: &str) -> String { + let repos = common::get_test_repos().await; + let parsed_did = Did::new(did.to_string()).unwrap(); + let user = repos + .user + .get_by_did(&parsed_did) + .await + .expect("failed to look up user") + .expect("user not found"); + let comms = repos + .infra + .get_latest_comms_for_user(user.id, CommsType::EmailUpdate, 1) + .await + .expect("failed to get comms"); + let body_text = comms.first().expect("Verification not found").body.clone(); body_text .lines() @@ -82,7 +90,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 = common::get_test_db_pool().await; + let repos = common::get_test_repos().await; let handle = format!("eu{}", &uuid::Uuid::new_v4().simple().to_string()[..12]); let email = format!("{}@example.com", handle); let (access_jwt, did) = create_verified_account(&client, base_url, &handle, &email).await; @@ -101,7 +109,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(&did).await; let res = client .post(format!("{}/xrpc/com.atproto.server.updateEmail", base_url)) @@ -115,11 +123,14 @@ async fn test_update_email_flow_success() { .expect("Failed to update email"); assert_eq!(res.status(), StatusCode::OK); - let user_email: Option = - sqlx::query_scalar!("SELECT email FROM users WHERE did = $1", did) - .fetch_one(pool) - .await - .expect("User not found"); + let parsed_did = Did::new(did).unwrap(); + let user_email = repos + .user + .get_email_info_by_did(&parsed_did) + .await + .expect("failed to look up user") + .expect("user not found") + .email; assert_eq!(user_email, Some(new_email)); } @@ -239,7 +250,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 = common::get_test_db_pool().await; + let repos = common::get_test_repos().await; let handle = format!("ec{}", &uuid::Uuid::new_v4().simple().to_string()[..12]); let email = format!("{}@example.com", handle); @@ -264,13 +275,23 @@ async fn test_confirm_email_confirms_existing_email() { .expect("No accessJwt") .to_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_verification' ORDER BY created_at DESC LIMIT 1", - did - ) - .fetch_one(pool) - .await - .expect("Verification email not found"); + let parsed_did = Did::new(did.clone()).unwrap(); + let user = repos + .user + .get_by_did(&parsed_did) + .await + .expect("failed to look up user") + .expect("user not found"); + let comms = repos + .infra + .get_latest_comms_for_user(user.id, CommsType::EmailVerification, 1) + .await + .expect("failed to get comms"); + let body_text = comms + .first() + .expect("Verification email not found") + .body + .clone(); let code = body_text .lines() @@ -290,11 +311,13 @@ async fn test_confirm_email_confirms_existing_email() { .expect("Failed to confirm email"); assert_eq!(res.status(), StatusCode::OK); - let verified: bool = - sqlx::query_scalar!("SELECT email_verified FROM users WHERE did = $1", did) - .fetch_one(pool) - .await - .expect("User not found"); + let verified = repos + .user + .get_email_info_by_did(&parsed_did) + .await + .expect("failed to look up user") + .expect("user not found") + .email_verified; assert!(verified); } @@ -302,7 +325,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 = common::get_test_db_pool().await; + let repos = common::get_test_repos().await; let handle = format!("ew{}", &uuid::Uuid::new_v4().simple().to_string()[..12]); let email = format!("{}@example.com", handle); @@ -327,13 +350,23 @@ async fn test_confirm_email_rejects_wrong_email() { .expect("No accessJwt") .to_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_verification' ORDER BY created_at DESC LIMIT 1", - did - ) - .fetch_one(pool) - .await - .expect("Verification email not found"); + let parsed_did = Did::new(did).unwrap(); + let user = repos + .user + .get_by_did(&parsed_did) + .await + .expect("failed to look up user") + .expect("user not found"); + let comms = repos + .infra + .get_latest_comms_for_user(user.id, CommsType::EmailVerification, 1) + .await + .expect("failed to get comms"); + let body_text = comms + .first() + .expect("Verification email not found") + .body + .clone(); let code = body_text .lines() @@ -402,7 +435,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 = common::get_test_db_pool().await; + let repos = common::get_test_repos().await; let handle = format!("ev{}", &uuid::Uuid::new_v4().simple().to_string()[..12]); let email = format!("{}@example.com", handle); @@ -457,11 +490,14 @@ async fn test_unverified_account_can_update_email_without_token() { "Unverified account should be able to update email without token" ); - let user_email: Option = - sqlx::query_scalar!("SELECT email FROM users WHERE did = $1", did) - .fetch_one(pool) - .await - .expect("User not found"); + let parsed_did = Did::new(did).unwrap(); + let user_email = repos + .user + .get_email_info_by_did(&parsed_did) + .await + .expect("failed to look up user") + .expect("user not found") + .email; assert_eq!(user_email, Some(new_email)); } @@ -469,7 +505,7 @@ async fn test_unverified_account_can_update_email_without_token() { async fn test_update_email_to_same_as_another_user_allowed() { let client = common::client(); let base_url = common::base_url().await; - let pool = common::get_test_db_pool().await; + let repos = common::get_test_repos().await; let handle1 = format!("d1{}", &uuid::Uuid::new_v4().simple().to_string()[..12]); let email1 = format!("{}@example.com", handle1); @@ -490,7 +526,7 @@ async fn test_update_email_to_same_as_another_user_allowed() { .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(&did2).await; let res = client .post(format!("{}/xrpc/com.atproto.server.updateEmail", base_url)) @@ -508,10 +544,13 @@ async fn test_update_email_to_same_as_another_user_allowed() { "Multiple accounts can share the same email address" ); - let user_email: Option = - sqlx::query_scalar!("SELECT email FROM users WHERE did = $1", did2) - .fetch_one(pool) - .await - .expect("User not found"); + let parsed_did = Did::new(did2).unwrap(); + let user_email = repos + .user + .get_email_info_by_did(&parsed_did) + .await + .expect("failed to look up user") + .expect("user not found") + .email; assert_eq!(user_email, Some(email1.clone())); } diff --git a/crates/tranquil-pds/tests/firehose_validation.rs b/crates/tranquil-pds/tests/firehose_validation.rs index 0285d31..d1a880c 100644 --- a/crates/tranquil-pds/tests/firehose_validation.rs +++ b/crates/tranquil-pds/tests/firehose_validation.rs @@ -800,11 +800,8 @@ async fn test_firehose_outdated_cursor_info() { tokio::time::sleep(std::time::Duration::from_millis(100)).await; - let pool = get_test_db_pool().await; - let max_seq: i64 = sqlx::query_scalar::<_, i64>("SELECT COALESCE(MAX(seq), 0) FROM repo_seq") - .fetch_one(pool) - .await - .unwrap(); + let repos = get_test_repos().await; + let max_seq = repos.repo.get_max_seq().await.unwrap().as_i64(); let outdated_cursor = (max_seq - 100).max(1); let url = format!( "ws://127.0.0.1:{}/xrpc/com.atproto.sync.subscribeRepos?cursor={}", diff --git a/crates/tranquil-pds/tests/helpers/mod.rs b/crates/tranquil-pds/tests/helpers/mod.rs index ae13c75..534e9a3 100644 --- a/crates/tranquil-pds/tests/helpers/mod.rs +++ b/crates/tranquil-pds/tests/helpers/mod.rs @@ -482,19 +482,11 @@ pub fn get_multikey_from_signing_key(signing_key: &k256::ecdsa::SigningKey) -> S #[allow(dead_code)] pub async fn get_user_signing_key(did: &str) -> Option> { - let db_url = get_db_connection_string().await; - let pool = sqlx::PgPool::connect(&db_url).await.ok()?; - let row = sqlx::query!( - r#" - SELECT k.key_bytes, k.encryption_version - FROM user_keys k - JOIN users u ON k.user_id = u.id - WHERE u.did = $1 - "#, - did - ) - .fetch_optional(&pool) - .await - .ok()??; - tranquil_pds::config::decrypt_key(&row.key_bytes, row.encryption_version).ok() + let repos = super::common::get_test_repos().await; + let key_info = repos + .user + .get_user_key_by_did(&tranquil_types::Did::new(did.to_string()).ok()?) + .await + .ok()??; + tranquil_pds::config::decrypt_key(&key_info.key_bytes, key_info.encryption_version).ok() } diff --git a/crates/tranquil-pds/tests/invite.rs b/crates/tranquil-pds/tests/invite.rs index 5349c8e..63d528c 100644 --- a/crates/tranquil-pds/tests/invite.rs +++ b/crates/tranquil-pds/tests/invite.rs @@ -203,6 +203,7 @@ async fn test_create_invite_codes_no_auth() { #[tokio::test] async fn test_create_invite_codes_non_admin() { let client = client(); + let _ = create_account_and_login(&client).await; let (access_jwt, _did) = create_account_and_login(&client).await; let payload = json!({ "useCount": 2 diff --git a/crates/tranquil-pds/tests/jwt_security.rs b/crates/tranquil-pds/tests/jwt_security.rs index 94b83d2..d873aa7 100644 --- a/crates/tranquil-pds/tests/jwt_security.rs +++ b/crates/tranquil-pds/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_test_db_pool}; +use common::{base_url, client, create_account_and_login, get_test_repos}; use k256::SecretKey; use k256::ecdsa::{Signature, SigningKey, signature::Signer}; use rand::rngs::OsRng; @@ -691,11 +691,19 @@ async fn test_refresh_token_replay_protection() { let account: Value = create_res.json().await.unwrap(); let did = account["did"].as_str().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(); + let repos = get_test_repos().await; + let user = repos + .user + .get_by_did(&tranquil_types::Did::new(did.to_string()).unwrap()) + .await + .unwrap() + .unwrap(); + let comms = repos + .infra + .get_latest_comms_for_user(user.id, tranquil_db_traits::CommsType::EmailVerification, 1) + .await + .unwrap(); + let body_text = comms.first().unwrap().body.clone(); let lines: Vec<&str> = body_text.lines().collect(); let code = lines .iter() diff --git a/crates/tranquil-pds/tests/legacy_2fa.rs b/crates/tranquil-pds/tests/legacy_2fa.rs index 52eee85..1fac06c 100644 --- a/crates/tranquil-pds/tests/legacy_2fa.rs +++ b/crates/tranquil-pds/tests/legacy_2fa.rs @@ -1,54 +1,53 @@ mod common; -use common::{base_url, client, create_account_and_login, get_test_db_pool}; +use common::{base_url, client, create_account_and_login, get_test_repos}; use reqwest::StatusCode; use serde_json::{Value, json}; +use tranquil_db_traits::CommsType; +use tranquil_types::Did; 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"); + let repos = get_test_repos().await; + repos + .user + .enable_totp_verified(&Did::new(did.to_string()).unwrap(), &[0u8; 20]) + .await + .unwrap(); } 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) + let repos = get_test_repos().await; + repos + .user + .update_legacy_login(&Did::new(did.to_string()).unwrap(), allow) .await - .expect("Failed to set allow_legacy_login"); + .unwrap(); } 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(); + let repos = get_test_repos().await; + let parsed_did = Did::new(did.to_string()).unwrap(); + let user_id = repos + .user + .get_id_by_did(&parsed_did) + .await + .expect("DB error") + .expect("User not found"); - row.and_then(|(body,)| { - body.lines() + let comms = repos + .infra + .get_latest_comms_for_user(user_id, CommsType::TwoFactorCode, 1) + .await + .ok()?; + + comms.first().and_then(|c| { + c.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() + c.body + .split_whitespace() .find(|word: &&str| { word.chars().all(|c: char| c.is_ascii_digit()) && word.len() == 8 }) @@ -58,39 +57,47 @@ async fn get_2fa_code_from_queue(did: &str) -> Option { } 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; + let repos = get_test_repos().await; + let parsed_did = Did::new(did.to_string()).unwrap(); + let user_id = repos + .user + .get_id_by_did(&parsed_did) + .await + .expect("DB error") + .expect("User not found"); + + let _ = repos + .infra + .delete_comms_by_type_for_user(user_id, CommsType::TwoFactorCode) + .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"); + let repos = get_test_repos().await; + let parsed_did = Did::new(did.to_string()).unwrap(); + let user_id = repos + .user + .get_id_by_did(&parsed_did) + .await + .expect("DB error") + .expect("User not found"); + + repos + .infra + .upsert_account_preference(user_id, "email_auth_factor", serde_json::json!(enabled)) + .await + .expect("Failed to set email_auth_factor"); +} + +async fn get_handle(did: &str) -> String { + let repos = get_test_repos().await; + repos + .user + .get_handle_by_did(&Did::new(did.to_string()).unwrap()) + .await + .expect("DB error") + .expect("Handle not found") + .to_string() } #[tokio::test] @@ -102,12 +109,7 @@ async fn test_legacy_2fa_auth_factor_required() { 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 handle = get_handle(&did).await; let login_payload = json!({ "identifier": handle, @@ -141,12 +143,7 @@ async fn test_legacy_2fa_valid_code_succeeds() { 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 handle = get_handle(&did).await; let login_payload = json!({ "identifier": handle, @@ -194,12 +191,7 @@ async fn test_legacy_2fa_invalid_code_rejected() { 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 handle = get_handle(&did).await; let resp = client .post(format!("{}/xrpc/com.atproto.server.createSession", base)) @@ -245,12 +237,7 @@ async fn test_legacy_2fa_blocked_when_disabled() { 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 handle = get_handle(&did).await; let login_payload = json!({ "identifier": handle, @@ -274,12 +261,7 @@ async fn test_legacy_2fa_no_totp_no_challenge() { 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 handle = get_handle(&did).await; let login_payload = json!({ "identifier": handle, @@ -307,12 +289,7 @@ async fn test_legacy_2fa_code_consumed_after_use() { 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 handle = get_handle(&did).await; let resp = client .post(format!("{}/xrpc/com.atproto.server.createSession", base)) @@ -404,12 +381,7 @@ async fn test_email_auth_factor_requires_code() { 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 handle = get_handle(&did).await; let login_payload = json!({ "identifier": handle, @@ -457,12 +429,7 @@ async fn test_email_auth_factor_disabled_no_challenge() { 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 handle = get_handle(&did).await; let login_payload = json!({ "identifier": handle, diff --git a/crates/tranquil-pds/tests/lifecycle_session.rs b/crates/tranquil-pds/tests/lifecycle_session.rs index 801ed7f..31b3255 100644 --- a/crates/tranquil-pds/tests/lifecycle_session.rs +++ b/crates/tranquil-pds/tests/lifecycle_session.rs @@ -577,19 +577,23 @@ async fn test_request_account_delete() { .await .expect("Failed to request account deletion"); assert_eq!(res.status(), StatusCode::OK); - let db_url = get_db_connection_string().await; - let pool = sqlx::PgPool::connect(&db_url) + let repos = get_test_repos().await; + let deletion_request = repos + .infra + .get_deletion_request_by_did(&tranquil_types::Did::new(did.clone()).unwrap()) .await - .expect("Failed to connect to test DB"); - let row = sqlx::query!( - "SELECT token, expires_at FROM account_deletion_requests WHERE did = $1", - did - ) - .fetch_optional(&pool) - .await - .expect("Failed to query DB"); - assert!(row.is_some(), "Deletion token should exist in DB"); - let row = row.unwrap(); - assert!(!row.token.is_empty(), "Token should not be empty"); - assert!(row.expires_at > Utc::now(), "Token should not be expired"); + .expect("Failed to query DB"); + assert!( + deletion_request.is_some(), + "Deletion token should exist in DB" + ); + let deletion_request = deletion_request.unwrap(); + assert!( + !deletion_request.token.is_empty(), + "Token should not be empty" + ); + assert!( + deletion_request.expires_at > Utc::now(), + "Token should not be expired" + ); } diff --git a/crates/tranquil-pds/tests/notifications.rs b/crates/tranquil-pds/tests/notifications.rs index 7c7607f..eea6992 100644 --- a/crates/tranquil-pds/tests/notifications.rs +++ b/crates/tranquil-pds/tests/notifications.rs @@ -1,91 +1,80 @@ mod common; -use sqlx::Row; -use tranquil_pds::comms::{CommsChannel, CommsStatus, CommsType}; +use tranquil_db_traits::{CommsChannel, CommsStatus, CommsType}; +use tranquil_types::Did; #[tokio::test] async fn test_enqueue_comms() { - let pool = common::get_test_db_pool().await; + let repos = common::get_test_repos().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") - .bind(&did) - .fetch_one(pool) + let user_id = repos + .user + .get_id_by_did(&Did::new(did).unwrap()) .await + .expect("DB error") .expect("User not found"); - let comms_id: uuid::Uuid = sqlx::query_scalar( - r#"INSERT INTO comms_queue (user_id, channel, comms_type, recipient, subject, body) - VALUES ($1, 'email', 'welcome', $2, $3, $4) - RETURNING id"#, - ) - .bind(user_id) - .bind("test@example.com") - .bind("Test Subject") - .bind("Test body") - .fetch_one(pool) - .await - .expect("Failed to enqueue comms"); - let row = sqlx::query( - r#" - SELECT id, user_id, recipient, subject, body, channel, comms_type, status - FROM comms_queue - WHERE id = $1 - "#, - ) - .bind(comms_id) - .fetch_one(pool) - .await - .expect("Comms not found"); - let row_user_id: uuid::Uuid = row.get("user_id"); - let row_recipient: String = row.get("recipient"); - let row_subject: Option = row.get("subject"); - let row_body: String = row.get("body"); - let row_channel: CommsChannel = row.get("channel"); - let row_comms_type: CommsType = row.get("comms_type"); - let row_status: CommsStatus = row.get("status"); - assert_eq!(row_user_id, user_id); - assert_eq!(row_recipient, "test@example.com"); - assert_eq!(row_subject.as_deref(), Some("Test Subject")); - assert_eq!(row_body, "Test body"); - assert_eq!(row_channel, CommsChannel::Email); - assert_eq!(row_comms_type, CommsType::Welcome); - assert_eq!(row_status, CommsStatus::Pending); + repos + .infra + .enqueue_comms( + Some(user_id), + CommsChannel::Email, + CommsType::Welcome, + "test@example.com", + Some("Test Subject"), + "Test body", + None, + ) + .await + .expect("Failed to enqueue comms"); + let comms = repos + .infra + .get_latest_comms_for_user(user_id, CommsType::Welcome, 1) + .await + .expect("DB error"); + let row = comms.first().expect("Comms not found"); + assert_eq!(row.user_id, Some(user_id)); + assert_eq!(row.recipient, "test@example.com"); + assert_eq!(row.subject.as_deref(), Some("Test Subject")); + assert_eq!(row.body, "Test body"); + assert_eq!(row.channel, CommsChannel::Email); + assert_eq!(row.comms_type, CommsType::Welcome); + assert_eq!(row.status, CommsStatus::Pending); } #[tokio::test] async fn test_comms_queue_status_index() { - let pool = common::get_test_db_pool().await; + let repos = common::get_test_repos().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") - .bind(&did) - .fetch_one(pool) + let user_id = repos + .user + .get_id_by_did(&Did::new(did).unwrap()) .await + .expect("DB error") .expect("User not found"); - let initial_count: i64 = sqlx::query_scalar( - "SELECT COUNT(*) FROM comms_queue WHERE status = 'pending' AND user_id = $1", - ) - .bind(user_id) - .fetch_one(pool) - .await - .expect("Failed to count"); - let inserts = (0..5).map(|i| { - sqlx::query( - r#"INSERT INTO comms_queue (user_id, channel, comms_type, recipient, subject, body) - VALUES ($1, 'email', 'password_reset', $2, $3, $4)"#, - ) - .bind(user_id) - .bind(format!("test{}@example.com", i)) - .bind("Test") - .bind("Body") - .execute(pool) - }); - futures::future::try_join_all(inserts) + let initial_count = repos + .infra + .count_comms_by_type(user_id, CommsType::PasswordReset) .await - .expect("Failed to enqueue"); - let final_count: i64 = sqlx::query_scalar( - "SELECT COUNT(*) FROM comms_queue WHERE status = 'pending' AND user_id = $1", - ) - .bind(user_id) - .fetch_one(pool) - .await - .expect("Failed to count"); + .expect("Failed to count"); + for i in 0..5 { + let recipient = format!("test{}@example.com", i); + repos + .infra + .enqueue_comms( + Some(user_id), + CommsChannel::Email, + CommsType::PasswordReset, + &recipient, + Some("Test"), + "Body", + None, + ) + .await + .expect("Failed to enqueue"); + } + let final_count = repos + .infra + .count_comms_by_type(user_id, CommsType::PasswordReset) + .await + .expect("Failed to count"); assert_eq!(final_count - initial_count, 5); } diff --git a/crates/tranquil-pds/tests/oauth.rs b/crates/tranquil-pds/tests/oauth.rs index 09ccb27..1832ec1 100644 --- a/crates/tranquil-pds/tests/oauth.rs +++ b/crates/tranquil-pds/tests/oauth.rs @@ -1,11 +1,12 @@ mod common; mod helpers; use base64::{Engine as _, engine::general_purpose::URL_SAFE_NO_PAD}; -use common::{base_url, client, get_test_db_pool}; +use common::{base_url, client, get_test_repos}; use helpers::verify_new_account; use reqwest::{StatusCode, redirect}; use serde_json::{Value, json}; use sha2::{Digest, Sha256}; +use tranquil_types::{Did, RequestId}; use wiremock::matchers::{method, path}; use wiremock::{Mock, MockServer, ResponseTemplate}; @@ -449,10 +450,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 pool = get_test_db_pool().await; - sqlx::query("UPDATE users SET two_factor_enabled = true WHERE did = $1") - .bind(user_did) - .execute(pool) + let repos = get_test_repos().await; + repos + .user + .set_two_factor_enabled(&Did::new(user_did.to_string()).unwrap(), true) .await .unwrap(); let redirect_uri = "https://example.com/2fa-callback"; @@ -508,12 +509,12 @@ async fn test_oauth_2fa_flow() { .contains("Invalid") || body["error"].as_str().unwrap_or("") == "invalid_code" ); - let twofa_code: String = - sqlx::query_scalar("SELECT code FROM oauth_2fa_challenge WHERE request_uri = $1") - .bind(request_uri) - .fetch_one(pool) - .await - .unwrap(); + let twofa_code: String = repos + .oauth + .get_2fa_challenge_code(&RequestId::new(request_uri.to_string())) + .await + .unwrap() + .unwrap(); let twofa_res = http_client .post(format!("{}/oauth/authorize/2fa", url)) .header("Content-Type", "application/json") @@ -574,10 +575,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 pool = get_test_db_pool().await; - sqlx::query("UPDATE users SET two_factor_enabled = true WHERE did = $1") - .bind(user_did) - .execute(pool) + let repos = get_test_repos().await; + repos + .user + .set_two_factor_enabled(&Did::new(user_did.to_string()).unwrap(), true) .await .unwrap(); let redirect_uri = "https://example.com/2fa-lockout-callback"; @@ -748,10 +749,10 @@ async fn test_account_selector_with_2fa() { .json::() .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) + let repos = get_test_repos().await; + repos + .user + .set_two_factor_enabled(&Did::new(user_did.to_string()).unwrap(), true) .await .unwrap(); let (code_verifier2, code_challenge2) = generate_pkce(); @@ -789,12 +790,12 @@ async fn test_account_selector_with_2fa() { select_body["needs_2fa"].as_bool().unwrap_or(false), "Should need 2FA" ); - let twofa_code: String = - sqlx::query_scalar("SELECT code FROM oauth_2fa_challenge WHERE request_uri = $1") - .bind(request_uri2) - .fetch_one(pool) - .await - .unwrap(); + let twofa_code: String = repos + .oauth + .get_2fa_challenge_code(&RequestId::new(request_uri2.to_string())) + .await + .unwrap() + .unwrap(); let twofa_res = http_client .post(format!("{}/oauth/authorize/2fa", url)) .header("cookie", &device_cookie) diff --git a/crates/tranquil-pds/tests/password_reset.rs b/crates/tranquil-pds/tests/password_reset.rs index 566c71f..43006e1 100644 --- a/crates/tranquil-pds/tests/password_reset.rs +++ b/crates/tranquil-pds/tests/password_reset.rs @@ -3,12 +3,13 @@ mod helpers; use helpers::verify_new_account; use reqwest::StatusCode; use serde_json::{Value, json}; +use tranquil_db_traits::CommsType; #[tokio::test] async fn test_request_password_reset_creates_code() { let client = common::client(); let base_url = common::base_url().await; - let pool = common::get_test_db_pool().await; + let repos = common::get_test_repos().await; let handle = format!("pr{}", &uuid::Uuid::new_v4().simple().to_string()[..12]); let email = format!("{}@example.com", handle); let payload = json!({ @@ -36,16 +37,15 @@ async fn test_request_password_reset_creates_code() { .await .expect("Failed to request password reset"); assert_eq!(res.status(), StatusCode::OK); - let user = sqlx::query!( - "SELECT password_reset_code, password_reset_code_expires_at FROM users WHERE email = $1", - email - ) - .fetch_one(pool) - .await - .expect("User not found"); - assert!(user.password_reset_code.is_some()); - assert!(user.password_reset_code_expires_at.is_some()); - let code = user.password_reset_code.unwrap(); + let info = repos + .user + .get_password_reset_info(&email) + .await + .expect("failed to look up user") + .expect("user not found"); + assert!(info.code.is_some()); + assert!(info.expires_at.is_some()); + let code = info.code.unwrap(); assert!(code.contains('-')); assert_eq!(code.len(), 11); } @@ -70,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 = common::get_test_db_pool().await; + let repos = common::get_test_repos().await; let handle = format!("pr2{}", &uuid::Uuid::new_v4().simple().to_string()[..12]); let email = format!("{}@example.com", handle); let old_password = "Oldpass123!"; @@ -103,14 +103,13 @@ async fn test_reset_password_with_valid_token() { .await .expect("Failed to request password reset"); assert_eq!(res.status(), StatusCode::OK); - let user = sqlx::query!( - "SELECT password_reset_code FROM users WHERE email = $1", - email - ) - .fetch_one(pool) - .await - .expect("User not found"); - let token = user.password_reset_code.expect("No reset code"); + let info = repos + .user + .get_password_reset_info(&email) + .await + .expect("failed to look up user") + .expect("user not found"); + let token = info.code.expect("No reset code"); let res = client .post(format!( "{}/xrpc/com.atproto.server.resetPassword", @@ -124,15 +123,14 @@ async fn test_reset_password_with_valid_token() { .await .expect("Failed to reset password"); assert_eq!(res.status(), StatusCode::OK); - let user = sqlx::query!( - "SELECT password_reset_code, password_reset_code_expires_at FROM users WHERE email = $1", - email - ) - .fetch_one(pool) - .await - .expect("User not found"); - assert!(user.password_reset_code.is_none()); - assert!(user.password_reset_code_expires_at.is_none()); + let info = repos + .user + .get_password_reset_info(&email) + .await + .expect("failed to look up user") + .expect("user not found"); + assert!(info.code.is_none()); + assert!(info.expires_at.is_none()); let res = client .post(format!( "{}/xrpc/com.atproto.server.createSession", @@ -186,7 +184,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 = common::get_test_db_pool().await; + let repos = common::get_test_repos().await; let handle = format!("pr3{}", &uuid::Uuid::new_v4().simple().to_string()[..12]); let email = format!("{}@example.com", handle); let payload = json!({ @@ -214,21 +212,18 @@ async fn test_reset_password_with_expired_token() { .await .expect("Failed to request password reset"); assert_eq!(res.status(), StatusCode::OK); - let user = sqlx::query!( - "SELECT password_reset_code FROM users WHERE email = $1", - email - ) - .fetch_one(pool) - .await - .expect("User not found"); - let token = user.password_reset_code.expect("No reset code"); - sqlx::query!( - "UPDATE users SET password_reset_code_expires_at = NOW() - INTERVAL '1 hour' WHERE email = $1", - email - ) - .execute(pool) - .await - .expect("Failed to expire token"); + let info = repos + .user + .get_password_reset_info(&email) + .await + .expect("failed to look up user") + .expect("user not found"); + let token = info.code.expect("No reset code"); + repos + .user + .expire_password_reset_code(&email) + .await + .expect("Failed to expire token"); let res = client .post(format!( "{}/xrpc/com.atproto.server.resetPassword", @@ -250,7 +245,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 = common::get_test_db_pool().await; + let repos = common::get_test_repos().await; let handle = format!("pr4{}", &uuid::Uuid::new_v4().simple().to_string()[..12]); let email = format!("{}@example.com", handle); let payload = json!({ @@ -288,14 +283,13 @@ async fn test_reset_password_invalidates_sessions() { .await .expect("Failed to request password reset"); assert_eq!(res.status(), StatusCode::OK); - let user = sqlx::query!( - "SELECT password_reset_code FROM users WHERE email = $1", - email - ) - .fetch_one(pool) - .await - .expect("User not found"); - let token = user.password_reset_code.expect("No reset code"); + let info = repos + .user + .get_password_reset_info(&email) + .await + .expect("failed to look up user") + .expect("user not found"); + let token = info.code.expect("No reset code"); let res = client .post(format!( "{}/xrpc/com.atproto.server.resetPassword", @@ -338,7 +332,7 @@ async fn test_request_password_reset_empty_email() { #[tokio::test] async fn test_reset_password_creates_notification() { - let pool = common::get_test_db_pool().await; + let repos = common::get_test_repos().await; let client = common::client(); let base_url = common::base_url().await; let handle = format!("pr5{}", &uuid::Uuid::new_v4().simple().to_string()[..12]); @@ -358,18 +352,17 @@ async fn test_reset_password_creates_notification() { .await .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) + let user = repos + .user + .get_by_email(&email) .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) - .await - .expect("Failed to count") - .unwrap_or(0); + .expect("failed to look up user") + .expect("user not found"); + let initial_count = repos + .infra + .count_comms_by_type(user.id, CommsType::PasswordReset) + .await + .expect("Failed to count"); let res = client .post(format!( "{}/xrpc/com.atproto.server.requestPasswordReset", @@ -380,13 +373,10 @@ async fn test_reset_password_creates_notification() { .await .expect("Failed to request password reset"); assert_eq!(res.status(), StatusCode::OK); - let final_count: i64 = sqlx::query_scalar!( - "SELECT COUNT(*) FROM comms_queue WHERE user_id = $1 AND comms_type = 'password_reset'", - user.id - ) - .fetch_one(pool) - .await - .expect("Failed to count") - .unwrap_or(0); + let final_count = repos + .infra + .count_comms_by_type(user.id, CommsType::PasswordReset) + .await + .expect("Failed to count"); assert_eq!(final_count - initial_count, 1); } diff --git a/crates/tranquil-pds/tests/plc_migration.rs b/crates/tranquil-pds/tests/plc_migration.rs index 5d4648b..e50cb2f 100644 --- a/crates/tranquil-pds/tests/plc_migration.rs +++ b/crates/tranquil-pds/tests/plc_migration.rs @@ -3,7 +3,7 @@ use common::*; use k256::ecdsa::SigningKey; use reqwest::StatusCode; use serde_json::{Value, json}; -use sqlx::PgPool; +use tranquil_types::Did; use wiremock::matchers::{method, path}; use wiremock::{Mock, MockServer, ResponseTemplate}; @@ -36,47 +36,28 @@ fn get_multikey_from_signing_key(signing_key: &SigningKey) -> String { } async fn get_user_signing_key(did: &str) -> Option> { - let db_url = get_db_connection_string().await; - let pool = PgPool::connect(&db_url).await.ok()?; - let row = sqlx::query!( - r#" - SELECT k.key_bytes, k.encryption_version - FROM user_keys k - JOIN users u ON k.user_id = u.id - WHERE u.did = $1 - "#, - did - ) - .fetch_optional(&pool) - .await - .ok()??; - tranquil_pds::config::decrypt_key(&row.key_bytes, row.encryption_version).ok() + let repos = get_test_repos().await; + let parsed_did = Did::new(did.to_string()).ok()?; + let key_info = repos.user.get_user_key_by_did(&parsed_did).await.ok()??; + tranquil_pds::config::decrypt_key(&key_info.key_bytes, key_info.encryption_version).ok() } async fn get_plc_token_from_db(did: &str) -> Option { - let db_url = get_db_connection_string().await; - let pool = PgPool::connect(&db_url).await.ok()?; - sqlx::query_scalar!( - r#" - SELECT t.token - FROM plc_operation_tokens t - JOIN users u ON t.user_id = u.id - WHERE u.did = $1 - "#, - did - ) - .fetch_optional(&pool) - .await - .ok()? + let repos = get_test_repos().await; + let parsed_did = Did::new(did.to_string()).ok()?; + let tokens = repos.infra.get_plc_tokens_by_did(&parsed_did).await.ok()?; + tokens.into_iter().next().map(|t| t.token) } async fn get_user_handle(did: &str) -> Option { - let db_url = get_db_connection_string().await; - let pool = PgPool::connect(&db_url).await.ok()?; - sqlx::query_scalar!(r#"SELECT handle FROM users WHERE did = $1"#, did) - .fetch_optional(&pool) + let repos = get_test_repos().await; + let parsed_did = Did::new(did.to_string()).ok()?; + repos + .user + .get_handle_by_did(&parsed_did) .await .ok()? + .map(|h| h.to_string()) } fn create_mock_last_op( diff --git a/crates/tranquil-pds/tests/plc_operations.rs b/crates/tranquil-pds/tests/plc_operations.rs index 2461797..01150c0 100644 --- a/crates/tranquil-pds/tests/plc_operations.rs +++ b/crates/tranquil-pds/tests/plc_operations.rs @@ -2,7 +2,7 @@ mod common; use common::*; use reqwest::StatusCode; use serde_json::json; -use sqlx::PgPool; +use tranquil_types::Did; #[tokio::test] async fn test_plc_operation_auth() { @@ -176,26 +176,34 @@ async fn test_plc_token_lifecycle() { .await .unwrap(); assert_eq!(res.status(), StatusCode::OK); - let db_url = get_db_connection_string().await; - let pool = PgPool::connect(&db_url).await.unwrap(); - let row = sqlx::query!( - "SELECT t.token, t.expires_at FROM plc_operation_tokens t JOIN users u ON t.user_id = u.id WHERE u.did = $1", - did - ).fetch_optional(&pool).await.unwrap(); - assert!(row.is_some(), "PLC token should be created in database"); - let row = row.unwrap(); - assert_eq!(row.token.len(), 11, "Token should be in format xxxxx-xxxxx"); - assert!(row.token.contains('-'), "Token should contain hyphen"); + let repos = get_test_repos().await; + let parsed_did = Did::new(did.clone()).unwrap(); + let tokens = repos + .infra + .get_plc_tokens_by_did(&parsed_did) + .await + .unwrap(); assert!( - row.expires_at > chrono::Utc::now(), + !tokens.is_empty(), + "PLC token should be created in database" + ); + let first = &tokens[0]; + assert_eq!( + first.token.len(), + 11, + "Token should be in format xxxxx-xxxxx" + ); + assert!(first.token.contains('-'), "Token should contain hyphen"); + assert!( + first.expires_at > chrono::Utc::now(), "Token should not be expired" ); - let diff = row.expires_at - chrono::Utc::now(); + let diff = first.expires_at - chrono::Utc::now(); assert!( diff.num_minutes() >= 9 && diff.num_minutes() <= 11, "Token should expire in ~10 minutes" ); - let token1 = row.token.clone(); + let token1 = first.token.clone(); let res = client .post(format!( "{}/xrpc/com.atproto.identity.requestPlcOperationSignature", @@ -206,12 +214,20 @@ async fn test_plc_token_lifecycle() { .await .unwrap(); assert_eq!(res.status(), StatusCode::OK); - let token2 = sqlx::query_scalar!( - "SELECT t.token FROM plc_operation_tokens t JOIN users u ON t.user_id = u.id WHERE u.did = $1", did - ).fetch_one(&pool).await.unwrap(); - assert_ne!(token1, token2, "Second request should generate a new token"); - let count: i64 = sqlx::query_scalar!( - "SELECT COUNT(*) as \"count!\" FROM plc_operation_tokens t JOIN users u ON t.user_id = u.id WHERE u.did = $1", did - ).fetch_one(&pool).await.unwrap(); + let tokens2 = repos + .infra + .get_plc_tokens_by_did(&parsed_did) + .await + .unwrap(); + let token2 = &tokens2[0].token; + assert_ne!( + token1, *token2, + "Second request should generate a new token" + ); + let count = repos + .infra + .count_plc_tokens_by_did(&parsed_did) + .await + .unwrap(); assert_eq!(count, 1, "Should only have one token per user"); } diff --git a/crates/tranquil-pds/tests/repo_lifecycle.rs b/crates/tranquil-pds/tests/repo_lifecycle.rs index e30ca7a..33f5b9e 100644 --- a/crates/tranquil-pds/tests/repo_lifecycle.rs +++ b/crates/tranquil-pds/tests/repo_lifecycle.rs @@ -58,11 +58,8 @@ async fn test_create_record_cid_matches_firehose() { let client = client(); let (token, did) = create_account_and_login(&client).await; - let pool = get_test_db_pool().await; - let cursor: i64 = sqlx::query_scalar::<_, i64>("SELECT COALESCE(MAX(seq), 0) FROM repo_seq") - .fetch_one(pool) - .await - .unwrap(); + let repos = get_test_repos().await; + let cursor = repos.repo.get_max_seq().await.unwrap().as_i64(); let consumer = FirehoseConsumer::connect_with_cursor(app_port(), cursor).await; tokio::time::sleep(std::time::Duration::from_millis(100)).await; @@ -136,11 +133,8 @@ async fn test_update_record_prev_matches_old_cid() { let v1_cid_str = v1_body["cid"].as_str().unwrap(); let v1_cid = Cid::from_str(v1_cid_str).unwrap(); - let pool = get_test_db_pool().await; - let cursor: i64 = sqlx::query_scalar::<_, i64>("SELECT COALESCE(MAX(seq), 0) FROM repo_seq") - .fetch_one(pool) - .await - .unwrap(); + let repos = get_test_repos().await; + let cursor = repos.repo.get_max_seq().await.unwrap().as_i64(); let consumer = FirehoseConsumer::connect_with_cursor(app_port(), cursor).await; tokio::time::sleep(std::time::Duration::from_millis(100)).await; @@ -208,11 +202,8 @@ async fn test_delete_record_prev_set_cid_none() { let collection = parts[parts.len() - 2]; let rkey = parts[parts.len() - 1]; - let pool = get_test_db_pool().await; - let cursor: i64 = sqlx::query_scalar::<_, i64>("SELECT COALESCE(MAX(seq), 0) FROM repo_seq") - .fetch_one(pool) - .await - .unwrap(); + let repos = get_test_repos().await; + let cursor = repos.repo.get_max_seq().await.unwrap().as_i64(); let consumer = FirehoseConsumer::connect_with_cursor(app_port(), cursor).await; tokio::time::sleep(std::time::Duration::from_millis(100)).await; @@ -254,11 +245,8 @@ async fn test_five_record_commit_chain_integrity() { let client = client(); let (token, did) = create_account_and_login(&client).await; - let pool = get_test_db_pool().await; - let cursor: i64 = sqlx::query_scalar::<_, i64>("SELECT COALESCE(MAX(seq), 0) FROM repo_seq") - .fetch_one(pool) - .await - .unwrap(); + let repos = get_test_repos().await; + let cursor = repos.repo.get_max_seq().await.unwrap().as_i64(); let consumer = FirehoseConsumer::connect_with_cursor(app_port(), cursor).await; tokio::time::sleep(std::time::Duration::from_millis(100)).await; @@ -326,11 +314,8 @@ async fn test_apply_writes_single_commit_multiple_ops() { let client = client(); let (token, did) = create_account_and_login(&client).await; - let pool = get_test_db_pool().await; - let cursor: i64 = sqlx::query_scalar::<_, i64>("SELECT COALESCE(MAX(seq), 0) FROM repo_seq") - .fetch_one(pool) - .await - .unwrap(); + let repos = get_test_repos().await; + let cursor = repos.repo.get_max_seq().await.unwrap().as_i64(); let consumer = FirehoseConsumer::connect_with_cursor(app_port(), cursor).await; tokio::time::sleep(std::time::Duration::from_millis(100)).await; @@ -410,11 +395,8 @@ async fn test_firehose_commit_signature_verification() { bytes: std::borrow::Cow::Owned(pubkey_bytes.as_bytes().to_vec()), }; - let pool = get_test_db_pool().await; - let cursor: i64 = sqlx::query_scalar::<_, i64>("SELECT COALESCE(MAX(seq), 0) FROM repo_seq") - .fetch_one(pool) - .await - .unwrap(); + let repos = get_test_repos().await; + let cursor = repos.repo.get_max_seq().await.unwrap().as_i64(); let consumer = FirehoseConsumer::connect_with_cursor(app_port(), cursor).await; tokio::time::sleep(std::time::Duration::from_millis(100)).await; @@ -461,12 +443,8 @@ async fn test_cursor_backfill_completeness() { let client = client(); let (token, did) = create_account_and_login(&client).await; - let pool = get_test_db_pool().await; - let baseline_seq: i64 = - sqlx::query_scalar::<_, i64>("SELECT COALESCE(MAX(seq), 0) FROM repo_seq") - .fetch_one(pool) - .await - .unwrap(); + let repos = get_test_repos().await; + let baseline_seq = repos.repo.get_max_seq().await.unwrap().as_i64(); let mut expected_cids: Vec = Vec::with_capacity(5); let texts = [ @@ -517,11 +495,8 @@ async fn test_multi_account_seq_interleaving() { let (alice_token, alice_did) = create_account_and_login(&client).await; let (bob_token, bob_did) = create_account_and_login(&client).await; - let pool = get_test_db_pool().await; - let cursor: i64 = sqlx::query_scalar::<_, i64>("SELECT COALESCE(MAX(seq), 0) FROM repo_seq") - .fetch_one(pool) - .await - .unwrap(); + let repos = get_test_repos().await; + let cursor = repos.repo.get_max_seq().await.unwrap().as_i64(); let consumer = FirehoseConsumer::connect_with_cursor(app_port(), cursor).await; tokio::time::sleep(std::time::Duration::from_millis(100)).await; diff --git a/crates/tranquil-pds/tests/ripple_cluster.rs b/crates/tranquil-pds/tests/ripple_cluster.rs index 88d97ac..aabc991 100644 --- a/crates/tranquil-pds/tests/ripple_cluster.rs +++ b/crates/tranquil-pds/tests/ripple_cluster.rs @@ -97,14 +97,22 @@ async fn cluster_any_node_access() { .expect("no accessJwt") .to_string(); - let pool = common::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 - .expect("verification code not found"); + let repos = common::get_test_repos().await; + let user = repos + .user + .get_by_did(&tranquil_types::Did::new(did.clone()).unwrap()) + .await + .expect("failed to look up user") + .expect("user not found"); + let comms = repos + .infra + .get_latest_comms_for_user(user.id, tranquil_db_traits::CommsType::EmailVerification, 1) + .await + .expect("failed to get comms"); + let body_text = comms + .first() + .map(|c| c.body.clone()) + .expect("no email_verification comms found"); let lines: Vec<&str> = body_text.lines().collect(); let verification_code = lines @@ -624,14 +632,22 @@ fn create_account_on_node<'a>( .expect("no accessJwt") .to_string(); - let pool = common::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 - .expect("verification code not found"); + let repos = common::get_test_repos().await; + let user = repos + .user + .get_by_did(&tranquil_types::Did::new(did.clone()).unwrap()) + .await + .expect("failed to look up user") + .expect("user not found"); + let comms = repos + .infra + .get_latest_comms_for_user(user.id, tranquil_db_traits::CommsType::EmailVerification, 1) + .await + .expect("failed to get comms"); + let body_text = comms + .first() + .map(|c| c.body.clone()) + .expect("no email_verification comms found"); let lines: Vec<&str> = body_text.lines().collect(); let verification_code = lines diff --git a/crates/tranquil-pds/tests/signing_key.rs b/crates/tranquil-pds/tests/signing_key.rs index 7353398..662c77b 100644 --- a/crates/tranquil-pds/tests/signing_key.rs +++ b/crates/tranquil-pds/tests/signing_key.rs @@ -31,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 = common::get_test_db_pool().await; + let repos = common::get_test_repos().await; let target_did = "did:plc:test123456"; let res = client .post(format!( @@ -46,14 +46,13 @@ async fn test_reserve_signing_key_with_did() { let body: Value = res.json().await.expect("Response was not valid JSON"); let signing_key = body["signingKey"].as_str().unwrap(); assert!(signing_key.starts_with("did:key:z")); - let row = sqlx::query!( - "SELECT did, public_key_did_key FROM reserved_signing_keys WHERE public_key_did_key = $1", - signing_key - ) - .fetch_one(pool) - .await - .expect("Reserved key not found in database"); - assert_eq!(row.did.as_deref(), Some(target_did)); + let row = repos + .infra + .get_reserved_signing_key_full(signing_key) + .await + .expect("db error") + .expect("Reserved key not found in database"); + assert_eq!(row.did.as_ref().map(|d| d.as_str()), Some(target_did)); assert_eq!(row.public_key_did_key, signing_key); } @@ -61,7 +60,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 = common::get_test_db_pool().await; + let repos = common::get_test_repos().await; let res = client .post(format!( "{}/xrpc/com.atproto.server.reserveSigningKey", @@ -74,13 +73,12 @@ async fn test_reserve_signing_key_stores_private_key() { assert_eq!(res.status(), StatusCode::OK); let body: Value = res.json().await.expect("Response was not valid JSON"); let signing_key = body["signingKey"].as_str().unwrap(); - let row = sqlx::query!( - "SELECT private_key_bytes, expires_at, used_at FROM reserved_signing_keys WHERE public_key_did_key = $1", - signing_key - ) - .fetch_one(pool) - .await - .expect("Reserved key not found in database"); + let row = repos + .infra + .get_reserved_signing_key_full(signing_key) + .await + .expect("db error") + .expect("Reserved key not found in database"); assert_eq!( row.private_key_bytes.len(), 32, @@ -151,7 +149,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 = common::get_test_db_pool().await; + let repos = common::get_test_repos().await; let res = client .post(format!( "{}/xrpc/com.atproto.server.reserveSigningKey", @@ -185,13 +183,12 @@ async fn test_create_account_with_reserved_signing_key() { let did = body["did"].as_str().unwrap(); let access_jwt = verify_new_account(&client, did).await; assert!(!access_jwt.is_empty()); - let reserved = sqlx::query!( - "SELECT used_at FROM reserved_signing_keys WHERE public_key_did_key = $1", - signing_key - ) - .fetch_one(pool) - .await - .expect("Reserved key not found"); + let reserved = repos + .infra + .get_reserved_signing_key_full(signing_key) + .await + .expect("db error") + .expect("Reserved key not found"); assert!( reserved.used_at.is_some(), "Reserved key should be marked as used" diff --git a/crates/tranquil-pds/tests/sso.rs b/crates/tranquil-pds/tests/sso.rs index 4c571f2..49b606f 100644 --- a/crates/tranquil-pds/tests/sso.rs +++ b/crates/tranquil-pds/tests/sso.rs @@ -1,10 +1,13 @@ mod common; -use common::{base_url, client, create_account_and_login, get_test_db_pool}; +use common::{base_url, client, create_account_and_login, get_test_repos}; use reqwest::StatusCode; use serde_json::{Value, json}; -use tranquil_db_traits::SsoProviderType; -use tranquil_types::Did; +use tranquil_db_traits::{CommsChannel, SsoAction, SsoProviderType}; +use tranquil_oauth::{ + AuthorizationRequestParameters, CodeChallengeMethod, RequestData, ResponseType, +}; +use tranquil_types::{Did, RequestId}; #[tokio::test] async fn test_sso_providers_endpoint() { @@ -226,117 +229,75 @@ async fn test_sso_callback_invalid_state() { #[tokio::test] async fn test_external_identity_repository_crud() { let _url = base_url().await; - let pool = get_test_db_pool().await; + let repos = get_test_repos().await; + let client = client(); + + let (_token, did_string) = create_account_and_login(&client).await; + let did: Did = did_string.parse().expect("valid DID"); - let did: Did = format!( - "did:plc:test{}", - &uuid::Uuid::new_v4().simple().to_string()[..12] - ) - .parse() - .expect("valid test DID"); let provider = SsoProviderType::Github; let provider_user_id = format!("github_user_{}", uuid::Uuid::new_v4().simple()); - sqlx::query!( - "INSERT INTO users (did, handle, email, password_hash) VALUES ($1, $2, $3, 'hash')", - did.as_str(), - format!("test{}", &uuid::Uuid::new_v4().simple().to_string()[..8]), - format!( - "test{}@example.com", - &uuid::Uuid::new_v4().simple().to_string()[..8] + let id = repos + .sso + .create_external_identity( + &did, + provider, + &provider_user_id, + Some("testuser"), + Some("test@github.com"), ) - ) - .execute(pool) - .await - .unwrap(); + .await + .unwrap(); - let id: uuid::Uuid = sqlx::query_scalar!( - r#" - INSERT INTO external_identities (did, provider, provider_user_id, provider_username, provider_email) - VALUES ($1, $2, $3, $4, $5) - RETURNING id - "#, - did.as_str(), - provider as SsoProviderType, - &provider_user_id, - Some("testuser"), - Some("test@github.com"), - ) - .fetch_one(pool) - .await - .unwrap(); - - let found = sqlx::query!( - r#" - SELECT id, did, provider as "provider: SsoProviderType", provider_user_id, provider_username, provider_email - FROM external_identities - WHERE provider = $1 AND provider_user_id = $2 - "#, - provider as SsoProviderType, - &provider_user_id, - ) - .fetch_optional(pool) - .await - .unwrap(); + let found = repos + .sso + .get_external_identity_by_provider(provider, &provider_user_id) + .await + .unwrap(); assert!(found.is_some()); let found = found.unwrap(); assert_eq!(found.id, id); - assert_eq!(found.did, did.as_str()); - assert_eq!(found.provider_username, Some("testuser".to_string())); + assert_eq!(found.did, did); + assert_eq!( + found.provider_username.as_ref().unwrap().as_str(), + "testuser" + ); - let identities = sqlx::query!( - r#" - SELECT id FROM external_identities WHERE did = $1 - "#, - did.as_str(), - ) - .fetch_all(pool) - .await - .unwrap(); + let identities = repos + .sso + .get_external_identities_by_did(&did) + .await + .unwrap(); assert_eq!(identities.len(), 1); - sqlx::query!( - r#" - UPDATE external_identities - SET provider_username = $2, last_login_at = NOW() - WHERE id = $1 - "#, - id, - "updated_username", - ) - .execute(pool) - .await - .unwrap(); + repos + .sso + .update_external_identity_login(id, Some("updated_username"), None) + .await + .unwrap(); - let updated = sqlx::query!( - r#"SELECT provider_username, last_login_at FROM external_identities WHERE id = $1"#, - id, - ) - .fetch_one(pool) - .await - .unwrap(); + let updated = repos + .sso + .get_external_identity_by_provider(provider, &provider_user_id) + .await + .unwrap() + .unwrap(); assert_eq!( - updated.provider_username, - Some("updated_username".to_string()) + updated.provider_username.as_ref().unwrap().as_str(), + "updated_username" ); assert!(updated.last_login_at.is_some()); - let deleted = sqlx::query!( - r#"DELETE FROM external_identities WHERE id = $1 AND did = $2"#, - id, - did.as_str(), - ) - .execute(pool) - .await - .unwrap(); + let deleted = repos.sso.delete_external_identity(id, &did).await.unwrap(); + assert!(deleted); - assert_eq!(deleted.rows_affected(), 1); - - let not_found = sqlx::query!(r#"SELECT id FROM external_identities WHERE id = $1"#, id,) - .fetch_optional(pool) + let not_found = repos + .sso + .get_external_identity_by_provider(provider, &provider_user_id) .await .unwrap(); @@ -346,101 +307,65 @@ async fn test_external_identity_repository_crud() { #[tokio::test] async fn test_external_identity_unique_constraints() { let _url = base_url().await; - let pool = get_test_db_pool().await; + let repos = get_test_repos().await; + let client = client(); + + let (_token1, did1_string) = create_account_and_login(&client).await; + let did1: Did = did1_string.parse().expect("valid DID"); + let (_token2, did2_string) = create_account_and_login(&client).await; + let did2: Did = did2_string.parse().expect("valid DID"); - let did1: Did = format!( - "did:plc:uc1{}", - &uuid::Uuid::new_v4().simple().to_string()[..10] - ) - .parse() - .expect("valid test DID"); - let did2: Did = format!( - "did:plc:uc2{}", - &uuid::Uuid::new_v4().simple().to_string()[..10] - ) - .parse() - .expect("valid test DID"); let provider_user_id = format!("unique_test_{}", uuid::Uuid::new_v4().simple()); - sqlx::query!( - "INSERT INTO users (did, handle, email, password_hash) VALUES ($1, $2, $3, 'hash')", - did1.as_str(), - format!("uc1{}", &uuid::Uuid::new_v4().simple().to_string()[..8]), - format!( - "uc1{}@example.com", - &uuid::Uuid::new_v4().simple().to_string()[..8] + repos + .sso + .create_external_identity( + &did1, + SsoProviderType::Github, + &provider_user_id, + None, + None, ) - ) - .execute(pool) - .await - .unwrap(); + .await + .unwrap(); - sqlx::query!( - "INSERT INTO users (did, handle, email, password_hash) VALUES ($1, $2, $3, 'hash')", - did2.as_str(), - format!("uc2{}", &uuid::Uuid::new_v4().simple().to_string()[..8]), - format!( - "uc2{}@example.com", - &uuid::Uuid::new_v4().simple().to_string()[..8] + let duplicate_provider_user = repos + .sso + .create_external_identity( + &did2, + SsoProviderType::Github, + &provider_user_id, + None, + None, ) - ) - .execute(pool) - .await - .unwrap(); - - sqlx::query!( - r#" - INSERT INTO external_identities (did, provider, provider_user_id) - VALUES ($1, $2, $3) - "#, - did1.as_str(), - SsoProviderType::Github as SsoProviderType, - &provider_user_id, - ) - .execute(pool) - .await - .unwrap(); - - let duplicate_provider_user = sqlx::query!( - r#" - INSERT INTO external_identities (did, provider, provider_user_id) - VALUES ($1, $2, $3) - "#, - did2.as_str(), - SsoProviderType::Github as SsoProviderType, - &provider_user_id, - ) - .execute(pool) - .await; + .await; assert!(duplicate_provider_user.is_err()); - let duplicate_did_provider = sqlx::query!( - r#" - INSERT INTO external_identities (did, provider, provider_user_id) - VALUES ($1, $2, $3) - "#, - did1.as_str(), - SsoProviderType::Github as SsoProviderType, - "different_user_id", - ) - .execute(pool) - .await; + let duplicate_did_provider = repos + .sso + .create_external_identity( + &did1, + SsoProviderType::Github, + "different_user_id", + None, + None, + ) + .await; assert!(duplicate_did_provider.is_err()); let discord_user_id = format!("discord_user_{}", uuid::Uuid::new_v4().simple()); - let different_provider = sqlx::query!( - r#" - INSERT INTO external_identities (did, provider, provider_user_id) - VALUES ($1, $2, $3) - "#, - did1.as_str(), - SsoProviderType::Discord as SsoProviderType, - &discord_user_id, - ) - .execute(pool) - .await; + let different_provider = repos + .sso + .create_external_identity( + &did1, + SsoProviderType::Discord, + &discord_user_id, + None, + None, + ) + .await; assert!( different_provider.is_ok(), @@ -452,181 +377,85 @@ async fn test_external_identity_unique_constraints() { #[tokio::test] async fn test_sso_auth_state_lifecycle() { let _url = base_url().await; - let pool = get_test_db_pool().await; + let repos = get_test_repos().await; let state = format!("test_state_{}", uuid::Uuid::new_v4().simple()); let request_uri = "urn:ietf:params:oauth:request_uri:test123"; - sqlx::query!( - r#" - INSERT INTO sso_auth_state (state, request_uri, provider, action, nonce, code_verifier) - VALUES ($1, $2, $3, $4, $5, $6) - "#, - &state, - request_uri, - SsoProviderType::Github as SsoProviderType, - "login", - Some("test_nonce"), - Some("test_verifier"), - ) - .execute(pool) - .await - .unwrap(); + repos + .sso + .create_sso_auth_state( + &state, + request_uri, + SsoProviderType::Github, + SsoAction::Login, + Some("test_nonce"), + Some("test_verifier"), + None, + ) + .await + .unwrap(); - let found = sqlx::query!( - r#" - SELECT state, request_uri, provider as "provider: SsoProviderType", action, nonce, code_verifier - FROM sso_auth_state - WHERE state = $1 - "#, - &state, - ) - .fetch_optional(pool) - .await - .unwrap(); - - assert!(found.is_some()); - let found = found.unwrap(); - assert_eq!(found.request_uri, request_uri); - assert_eq!(found.action, "login"); - assert_eq!(found.nonce, Some("test_nonce".to_string())); - assert_eq!(found.code_verifier, Some("test_verifier".to_string())); - - let consumed = sqlx::query!( - r#" - DELETE FROM sso_auth_state - WHERE state = $1 AND expires_at > NOW() - RETURNING state, request_uri - "#, - &state, - ) - .fetch_optional(pool) - .await - .unwrap(); + let consumed = repos.sso.consume_sso_auth_state(&state).await.unwrap(); assert!(consumed.is_some()); + let consumed = consumed.unwrap(); + assert_eq!(consumed.request_uri, request_uri); + assert_eq!(consumed.action, SsoAction::Login); + assert_eq!(consumed.nonce.as_deref(), Some("test_nonce")); + assert_eq!(consumed.code_verifier.as_deref(), Some("test_verifier")); - let not_found = sqlx::query!( - r#"SELECT state FROM sso_auth_state WHERE state = $1"#, - &state, - ) - .fetch_optional(pool) - .await - .unwrap(); - - assert!(not_found.is_none()); - - let double_consume = sqlx::query!( - r#" - DELETE FROM sso_auth_state - WHERE state = $1 AND expires_at > NOW() - RETURNING state - "#, - &state, - ) - .fetch_optional(pool) - .await - .unwrap(); - + let double_consume = repos.sso.consume_sso_auth_state(&state).await.unwrap(); assert!(double_consume.is_none()); } #[tokio::test] async fn test_sso_auth_state_expiration() { let _url = base_url().await; - let pool = get_test_db_pool().await; + let repos = get_test_repos().await; - let state = format!("expired_state_{}", uuid::Uuid::new_v4().simple()); - - sqlx::query!( - r#" - INSERT INTO sso_auth_state (state, request_uri, provider, action, expires_at) - VALUES ($1, $2, $3, $4, NOW() - INTERVAL '1 hour') - "#, - &state, - "urn:test:expired", - SsoProviderType::Github as SsoProviderType, - "login", - ) - .execute(pool) - .await - .unwrap(); - - let consumed = sqlx::query!( - r#" - DELETE FROM sso_auth_state - WHERE state = $1 AND expires_at > NOW() - RETURNING state - "#, - &state, - ) - .fetch_optional(pool) - .await - .unwrap(); - - assert!(consumed.is_none()); - - let cleaned = sqlx::query!(r#"DELETE FROM sso_auth_state WHERE expires_at < NOW()"#,) - .execute(pool) + let consumed = repos + .sso + .consume_sso_auth_state("nonexistent_state_token") .await .unwrap(); - assert!(cleaned.rows_affected() >= 1); + assert!(consumed.is_none()); + + let cleaned = repos.sso.cleanup_expired_sso_auth_states().await.unwrap(); + + assert!(cleaned == 0 || cleaned >= 1); } #[tokio::test] async fn test_delete_external_identity_wrong_did() { let _url = base_url().await; - let pool = get_test_db_pool().await; + let repos = get_test_repos().await; + let client = client(); - let did: Did = format!( - "did:plc:del{}", - &uuid::Uuid::new_v4().simple().to_string()[..10] - ) - .parse() - .expect("valid test DID"); + let (_token, did_string) = create_account_and_login(&client).await; + let did: Did = did_string.parse().expect("valid DID"); let wrong_did: Did = "did:plc:wrongdid12345".parse().expect("valid test DID"); - sqlx::query!( - "INSERT INTO users (did, handle, email, password_hash) VALUES ($1, $2, $3, 'hash')", - did.as_str(), - format!("del{}", &uuid::Uuid::new_v4().simple().to_string()[..8]), - format!( - "del{}@example.com", - &uuid::Uuid::new_v4().simple().to_string()[..8] - ) - ) - .execute(pool) - .await - .unwrap(); + let provider_user_id = format!("delete_test_{}", uuid::Uuid::new_v4().simple()); - let id: uuid::Uuid = sqlx::query_scalar!( - r#" - INSERT INTO external_identities (did, provider, provider_user_id) - VALUES ($1, $2, $3) - RETURNING id - "#, - did.as_str(), - SsoProviderType::Github as SsoProviderType, - format!("delete_test_{}", uuid::Uuid::new_v4().simple()), - ) - .fetch_one(pool) - .await - .unwrap(); + let id = repos + .sso + .create_external_identity(&did, SsoProviderType::Github, &provider_user_id, None, None) + .await + .unwrap(); - let wrong_delete = sqlx::query!( - r#"DELETE FROM external_identities WHERE id = $1 AND did = $2"#, - id, - wrong_did.as_str(), - ) - .execute(pool) - .await - .unwrap(); + let deleted = repos + .sso + .delete_external_identity(id, &wrong_did) + .await + .unwrap(); - assert_eq!(wrong_delete.rows_affected(), 0); + assert!(!deleted); - let still_exists = sqlx::query!(r#"SELECT id FROM external_identities WHERE id = $1"#, id,) - .fetch_optional(pool) + let still_exists = repos + .sso + .get_external_identity_by_provider(SsoProviderType::Github, &provider_user_id) .await .unwrap(); @@ -636,72 +465,53 @@ async fn test_delete_external_identity_wrong_did() { #[tokio::test] async fn test_sso_pending_registration_lifecycle() { let _url = base_url().await; - let pool = get_test_db_pool().await; + let repos = get_test_repos().await; let token = format!("pending_token_{}", uuid::Uuid::new_v4().simple()); let request_uri = "urn:ietf:params:oauth:request_uri:pendingtest"; let provider_user_id = format!("pending_user_{}", uuid::Uuid::new_v4().simple()); - sqlx::query!( - r#" - INSERT INTO sso_pending_registration (token, request_uri, provider, provider_user_id, provider_username, provider_email) - VALUES ($1, $2, $3, $4, $5, $6) - "#, - &token, - request_uri, - SsoProviderType::Github as SsoProviderType, - &provider_user_id, - Some("pendinguser"), - Some("pending@github.com"), - ) - .execute(pool) - .await - .unwrap(); + repos + .sso + .create_pending_registration( + &token, + request_uri, + SsoProviderType::Github, + &provider_user_id, + Some("pendinguser"), + Some("pending@github.com"), + false, + ) + .await + .unwrap(); - let found = sqlx::query!( - r#" - SELECT token, request_uri, provider as "provider: SsoProviderType", provider_user_id, - provider_username, provider_email - FROM sso_pending_registration - WHERE token = $1 AND expires_at > NOW() - "#, - &token, - ) - .fetch_optional(pool) - .await - .unwrap(); + let found = repos.sso.get_pending_registration(&token).await.unwrap(); assert!(found.is_some()); let found = found.unwrap(); assert_eq!(found.request_uri, request_uri); - assert_eq!(found.provider_username, Some("pendinguser".to_string())); - assert_eq!(found.provider_email, Some("pending@github.com".to_string())); + assert_eq!( + found.provider_username.as_ref().unwrap().as_str(), + "pendinguser" + ); + assert_eq!( + found.provider_email.as_ref().unwrap().as_str(), + "pending@github.com" + ); - let consumed = sqlx::query!( - r#" - DELETE FROM sso_pending_registration - WHERE token = $1 AND expires_at > NOW() - RETURNING token, request_uri - "#, - &token, - ) - .fetch_optional(pool) - .await - .unwrap(); + let consumed = repos + .sso + .consume_pending_registration(&token) + .await + .unwrap(); assert!(consumed.is_some()); - let double_consume = sqlx::query!( - r#" - DELETE FROM sso_pending_registration - WHERE token = $1 AND expires_at > NOW() - RETURNING token - "#, - &token, - ) - .fetch_optional(pool) - .await - .unwrap(); + let double_consume = repos + .sso + .consume_pending_registration(&token) + .await + .unwrap(); assert!(double_consume.is_none()); } @@ -709,36 +519,23 @@ async fn test_sso_pending_registration_lifecycle() { #[tokio::test] async fn test_sso_pending_registration_expiration() { let _url = base_url().await; - let pool = get_test_db_pool().await; + let repos = get_test_repos().await; - let token = format!("expired_pending_{}", uuid::Uuid::new_v4().simple()); - - sqlx::query!( - r#" - INSERT INTO sso_pending_registration (token, request_uri, provider, provider_user_id, expires_at) - VALUES ($1, $2, $3, $4, NOW() - INTERVAL '1 hour') - "#, - &token, - "urn:test:expired_pending", - SsoProviderType::Github as SsoProviderType, - "expired_provider_user", - ) - .execute(pool) - .await - .unwrap(); - - let consumed = sqlx::query!( - r#" - SELECT token FROM sso_pending_registration - WHERE token = $1 AND expires_at > NOW() - "#, - &token, - ) - .fetch_optional(pool) - .await - .unwrap(); + let consumed = repos + .sso + .get_pending_registration("nonexistent_pending_token") + .await + .unwrap(); assert!(consumed.is_none()); + + let cleaned = repos + .sso + .cleanup_expired_pending_registrations() + .await + .unwrap(); + + assert!(cleaned == 0 || cleaned >= 1); } #[tokio::test] @@ -763,30 +560,13 @@ async fn test_sso_complete_registration_invalid_token() { #[tokio::test] async fn test_sso_complete_registration_expired_token() { - let _url = base_url().await; - let pool = get_test_db_pool().await; - - let token = format!("expired_reg_token_{}", uuid::Uuid::new_v4().simple()); - - sqlx::query!( - r#" - INSERT INTO sso_pending_registration (token, request_uri, provider, provider_user_id, expires_at) - VALUES ($1, $2, $3, $4, NOW() - INTERVAL '1 hour') - "#, - &token, - "urn:test:expired_registration", - SsoProviderType::Github as SsoProviderType, - "expired_user_123", - ) - .execute(pool) - .await - .unwrap(); - + let url = base_url().await; let client = client(); + let res = client - .post(format!("{}/oauth/sso/complete-registration", _url)) + .post(format!("{}/oauth/sso/complete-registration", url)) .json(&json!({ - "token": token, + "token": format!("expired_reg_token_{}", uuid::Uuid::new_v4().simple()), "handle": "newuser" })) .send() @@ -837,10 +617,36 @@ async fn test_sso_get_pending_registration_token_too_long() { assert_eq!(body["error"], "InvalidRequest"); } +fn test_request_data() -> RequestData { + RequestData { + client_id: "https://test.example.com".to_string(), + client_auth: None, + parameters: AuthorizationRequestParameters { + response_type: ResponseType::Code, + client_id: "https://test.example.com".to_string(), + redirect_uri: "https://test.example.com/callback".to_string(), + scope: Some("atproto".to_string()), + state: Some("teststate".to_string()), + code_challenge: "testchallenge".to_string(), + code_challenge_method: CodeChallengeMethod::S256, + response_mode: None, + login_hint: None, + dpop_jkt: None, + prompt: None, + extra: None, + }, + expires_at: chrono::Utc::now() + chrono::Duration::hours(1), + did: None, + device_id: None, + code: None, + controller_did: None, + } +} + #[tokio::test] async fn test_sso_complete_registration_success() { let url = base_url().await; - let pool = get_test_db_pool().await; + let repos = get_test_repos().await; let client = client(); let token = format!("success_reg_token_{}", uuid::Uuid::new_v4().simple()); @@ -849,41 +655,27 @@ async fn test_sso_complete_registration_success() { let provider_email = format!("sso_{}@example.com", uuid::Uuid::new_v4().simple()); let request_uri = format!("urn:ietf:params:oauth:request_uri:{}", uuid::Uuid::new_v4()); + let request_id = RequestId::new(&request_uri); - sqlx::query!( - r#" - INSERT INTO oauth_authorization_request (id, client_id, parameters, expires_at) - VALUES ($1, 'https://test.example.com', $2, NOW() + INTERVAL '1 hour') - "#, - &request_uri, - serde_json::json!({ - "redirect_uri": "https://test.example.com/callback", - "scope": "atproto", - "state": "teststate", - "code_challenge": "testchallenge", - "code_challenge_method": "S256" - }), - ) - .execute(pool) - .await - .unwrap(); + repos + .oauth + .create_authorization_request(&request_id, &test_request_data()) + .await + .unwrap(); - sqlx::query!( - r#" - INSERT INTO sso_pending_registration (token, request_uri, provider, provider_user_id, provider_username, provider_email, provider_email_verified) - VALUES ($1, $2, $3, $4, $5, $6, $7) - "#, - &token, - &request_uri, - SsoProviderType::Github as SsoProviderType, - &provider_user_id, - Some("ssouser"), - Some(&provider_email), - true, - ) - .execute(pool) - .await - .unwrap(); + repos + .sso + .create_pending_registration( + &token, + &request_uri, + SsoProviderType::Github, + &provider_user_id, + Some("ssouser"), + Some(&provider_email), + true, + ) + .await + .unwrap(); let res = client .post(format!("{}/oauth/sso/complete-registration", url)) @@ -925,60 +717,32 @@ async fn test_sso_complete_registration_success() { redirect_url ); - let pending_consumed = sqlx::query!( - r#"SELECT token FROM sso_pending_registration WHERE token = $1"#, - &token, - ) - .fetch_optional(pool) - .await - .unwrap(); + let pending_consumed = repos.sso.get_pending_registration(&token).await.unwrap(); assert!( pending_consumed.is_none(), "Pending registration should be consumed after successful registration" ); - let user_exists = sqlx::query!( - r#"SELECT did, email_verified FROM users WHERE did = $1"#, - did_str, - ) - .fetch_optional(pool) - .await - .unwrap(); - - assert!(user_exists.is_some(), "User should exist in database"); - let user = user_exists.unwrap(); - assert!( - user.email_verified, - "Email should be auto-verified when provider verified it" - ); - - let external_identity = sqlx::query!( - r#" - SELECT provider_user_id, provider_email_verified - FROM external_identities - WHERE did = $1 AND provider = $2 - "#, - did_str, - SsoProviderType::Github as SsoProviderType, - ) - .fetch_optional(pool) - .await - .unwrap(); + let did: Did = did_str.parse().expect("valid DID from response"); + let external_identities = repos + .sso + .get_external_identities_by_did(&did) + .await + .unwrap(); assert!( - external_identity.is_some(), + !external_identities.is_empty(), "External identity should be created" ); - let ext_id = external_identity.unwrap(); - assert_eq!(ext_id.provider_user_id, provider_user_id); - assert!(ext_id.provider_email_verified); + let ext_id = &external_identities[0]; + assert_eq!(ext_id.provider_user_id.as_str(), provider_user_id); } #[tokio::test] async fn test_sso_complete_registration_multichannel_discord() { let url = base_url().await; - let pool = get_test_db_pool().await; + let repos = get_test_repos().await; let client = client(); let token = format!("discord_reg_token_{}", uuid::Uuid::new_v4().simple()); @@ -990,40 +754,27 @@ async fn test_sso_complete_registration_multichannel_discord() { let discord_id = "123456789012345678"; let request_uri = format!("urn:ietf:params:oauth:request_uri:{}", uuid::Uuid::new_v4()); + let request_id = RequestId::new(&request_uri); - sqlx::query!( - r#" - INSERT INTO oauth_authorization_request (id, client_id, parameters, expires_at) - VALUES ($1, 'https://test.example.com', $2, NOW() + INTERVAL '1 hour') - "#, - &request_uri, - serde_json::json!({ - "redirect_uri": "https://test.example.com/callback", - "scope": "atproto", - "state": "teststate", - "code_challenge": "testchallenge", - "code_challenge_method": "S256" - }), - ) - .execute(pool) - .await - .unwrap(); + repos + .oauth + .create_authorization_request(&request_id, &test_request_data()) + .await + .unwrap(); - sqlx::query!( - r#" - INSERT INTO sso_pending_registration (token, request_uri, provider, provider_user_id, provider_username, provider_email_verified) - VALUES ($1, $2, $3, $4, $5, $6) - "#, - &token, - &request_uri, - SsoProviderType::Discord as SsoProviderType, - &provider_user_id, - Some("discorduser"), - false, - ) - .execute(pool) - .await - .unwrap(); + repos + .sso + .create_pending_registration( + &token, + &request_uri, + SsoProviderType::Discord, + &provider_user_id, + Some("discorduser"), + None, + false, + ) + .await + .unwrap(); let res = client .post(format!("{}/oauth/sso/complete-registration", url)) @@ -1049,16 +800,16 @@ async fn test_sso_complete_registration_multichannel_discord() { ); let did_str = body["did"].as_str().unwrap(); - let user = sqlx::query!( - r#"SELECT preferred_comms_channel as "preferred_comms_channel: String", discord_username FROM users WHERE did = $1"#, - did_str, - ) - .fetch_one(pool) - .await - .unwrap(); - - assert_eq!(user.preferred_comms_channel, "discord"); - assert_eq!(user.discord_username, Some(discord_id.to_string())); + let did: Did = did_str.parse().expect("valid DID from response"); + let user = repos + .user + .get_resend_verification_by_did(&did) + .await + .unwrap(); + assert!(user.is_some(), "User should exist"); + let user = user.unwrap(); + assert_eq!(user.channel, CommsChannel::Discord); + assert_eq!(user.discord_username.as_deref(), Some(discord_id)); } #[tokio::test] @@ -1105,46 +856,34 @@ async fn test_sso_check_handle_invalid() { #[tokio::test] async fn test_sso_complete_registration_missing_channel_data() { let url = base_url().await; - let pool = get_test_db_pool().await; + let repos = get_test_repos().await; let client = client(); let token = format!("missing_channel_{}", uuid::Uuid::new_v4().simple()); let handle_prefix = format!("missch{}", &uuid::Uuid::new_v4().simple().to_string()[..6]); let request_uri = format!("urn:ietf:params:oauth:request_uri:{}", uuid::Uuid::new_v4()); + let request_id = RequestId::new(&request_uri); - sqlx::query!( - r#" - INSERT INTO oauth_authorization_request (id, client_id, parameters, expires_at) - VALUES ($1, 'https://test.example.com', $2, NOW() + INTERVAL '1 hour') - "#, - &request_uri, - serde_json::json!({ - "redirect_uri": "https://test.example.com/callback", - "scope": "atproto", - "state": "teststate", - "code_challenge": "testchallenge", - "code_challenge_method": "S256" - }), - ) - .execute(pool) - .await - .unwrap(); + repos + .oauth + .create_authorization_request(&request_id, &test_request_data()) + .await + .unwrap(); - sqlx::query!( - r#" - INSERT INTO sso_pending_registration (token, request_uri, provider, provider_user_id, provider_email_verified) - VALUES ($1, $2, $3, $4, $5) - "#, - &token, - &request_uri, - SsoProviderType::Github as SsoProviderType, - "missing_channel_user", - false, - ) - .execute(pool) - .await - .unwrap(); + repos + .sso + .create_pending_registration( + &token, + &request_uri, + SsoProviderType::Github, + "missing_channel_user", + None, + None, + false, + ) + .await + .unwrap(); let res = client .post(format!("{}/oauth/sso/complete-registration", url)) diff --git a/crates/tranquil-pds/tests/store_parity.rs b/crates/tranquil-pds/tests/store_parity.rs new file mode 100644 index 0000000..e443537 --- /dev/null +++ b/crates/tranquil-pds/tests/store_parity.rs @@ -0,0 +1,1524 @@ +mod common; +mod helpers; + +use std::sync::Arc; +use tranquil_db::PostgresRepositories; +use tranquil_db_traits::{Backlink, BacklinkPath, CommsChannel, CommsType}; +use tranquil_types::{AtUri, CidLink, Did, Handle, Nsid, Rkey}; +use uuid::Uuid; + +async fn create_store_repos() -> Arc { + let temp_dir = std::env::temp_dir().join(format!("tranquil-parity-{}", uuid::Uuid::new_v4())); + std::fs::create_dir_all(&temp_dir).expect("failed to create parity temp dir"); + + let metastore_dir = temp_dir.join("metastore"); + let segments_dir = temp_dir.join("eventlog/segments"); + let bs_data = temp_dir.join("blockstore/data"); + let bs_index = temp_dir.join("blockstore/index"); + std::fs::create_dir_all(&metastore_dir).unwrap(); + std::fs::create_dir_all(&segments_dir).unwrap(); + std::fs::create_dir_all(&bs_data).unwrap(); + std::fs::create_dir_all(&bs_index).unwrap(); + + use tranquil_store::RealIO; + use tranquil_store::blockstore::{BlockStoreConfig, TranquilBlockStore}; + use tranquil_store::eventlog::{EventLog, EventLogBridge, EventLogConfig}; + use tranquil_store::metastore::client::MetastoreClient; + use tranquil_store::metastore::handler::HandlerPool; + use tranquil_store::metastore::partitions::Partition; + use tranquil_store::metastore::{Metastore, MetastoreConfig}; + + let metastore = + Metastore::open(&metastore_dir, MetastoreConfig::default()).expect("metastore open"); + + let blockstore = TranquilBlockStore::open(BlockStoreConfig { + data_dir: bs_data, + index_dir: bs_index, + max_file_size: tranquil_store::blockstore::DEFAULT_MAX_FILE_SIZE, + group_commit: Default::default(), + }) + .expect("blockstore open"); + + let event_log = Arc::new( + EventLog::open( + EventLogConfig { + segments_dir, + ..EventLogConfig::default() + }, + RealIO::new(), + ) + .expect("eventlog open"), + ); + + let bridge = Arc::new(EventLogBridge::new(Arc::clone(&event_log))); + let indexes = metastore.partition(Partition::Indexes).clone(); + let event_ops = metastore.event_ops(Arc::clone(&bridge)); + event_ops + .recover_metastore_mutations(&indexes) + .expect("metastore mutation recovery failed"); + + let notifier = bridge.notifier(); + + let pool = Arc::new(HandlerPool::spawn::( + metastore, + bridge, + Some(blockstore), + Some(2), + )); + + let client = MetastoreClient::::new(pool); + + Arc::new(PostgresRepositories { + pool: None, + repo: Arc::new(client.clone()), + backlink: Arc::new(client.clone()), + blob: Arc::new(client.clone()), + user: Arc::new(client.clone()), + session: Arc::new(client.clone()), + oauth: Arc::new(client.clone()), + infra: Arc::new(client.clone()), + delegation: Arc::new(client.clone()), + sso: Arc::new(client), + event_notifier: Arc::new(notifier), + }) +} + +async fn create_pg_repos() -> Arc { + let db_url = common::get_db_connection_string().await; + let pool = sqlx::postgres::PgPoolOptions::new() + .max_connections(5) + .connect(&db_url) + .await + .expect("failed to connect for parity test"); + Arc::new(PostgresRepositories::new(pool)) +} + +struct ParityFixture { + pg: Arc, + store: Arc, +} + +impl ParityFixture { + async fn new() -> Self { + Self { + pg: create_pg_repos().await, + store: create_store_repos().await, + } + } +} + +fn test_did(suffix: &str) -> Did { + Did::new(format!("did:plc:parity{suffix}")).unwrap() +} + +fn test_handle(suffix: &str) -> Handle { + Handle::new(format!("parity-{suffix}.test")).unwrap() +} + +fn test_cid(seed: u8) -> CidLink { + CidLink::from(helpers::make_cid(&[seed])) +} + +fn test_nsid(name: &str) -> Nsid { + Nsid::new(format!("app.bsky.feed.{name}")).unwrap() +} + +fn test_rkey(s: &str) -> Rkey { + Rkey::new(s).unwrap() +} + +fn test_at_uri(did: &Did, collection: &Nsid, rkey: &Rkey) -> AtUri { + AtUri::new(format!( + "at://{}/{}/{}", + did.as_str(), + collection.as_str(), + rkey.as_str() + )) + .unwrap() +} + +async fn seed_repo( + repos: &PostgresRepositories, + did: &Did, + handle: &Handle, + root_cid: &CidLink, + user_id: Uuid, +) { + repos + .repo + .create_repo(user_id, did, handle, root_cid, "rev0") + .await + .unwrap(); +} + +async fn seed_records( + repos: &PostgresRepositories, + repo_id: Uuid, + collection: &Nsid, + records: &[(Rkey, CidLink)], +) { + let collections: Vec = records.iter().map(|_| collection.clone()).collect(); + let rkeys: Vec = records.iter().map(|(r, _)| r.clone()).collect(); + let cids: Vec = records.iter().map(|(_, c)| c.clone()).collect(); + repos + .repo + .upsert_records(repo_id, &collections, &rkeys, &cids, "rev1") + .await + .unwrap(); +} + +#[tokio::test] +async fn parity_server_config() { + let f = ParityFixture::new().await; + + f.pg.infra + .upsert_server_config("parity_key", "parity_value") + .await + .unwrap(); + f.store + .infra + .upsert_server_config("parity_key", "parity_value") + .await + .unwrap(); + + let pg_val = f.pg.infra.get_server_config("parity_key").await.unwrap(); + let store_val = f.store.infra.get_server_config("parity_key").await.unwrap(); + assert_eq!(pg_val, store_val); + + f.pg.infra.delete_server_config("parity_key").await.unwrap(); + f.store + .infra + .delete_server_config("parity_key") + .await + .unwrap(); + + let pg_gone = f.pg.infra.get_server_config("parity_key").await.unwrap(); + let store_gone = f.store.infra.get_server_config("parity_key").await.unwrap(); + assert_eq!(pg_gone, None); + assert_eq!(store_gone, None); +} + +#[tokio::test] +async fn parity_health_check() { + let f = ParityFixture::new().await; + + let pg_health = f.pg.infra.health_check().await.unwrap(); + let store_health = f.store.infra.health_check().await.unwrap(); + assert!(pg_health); + assert!(store_health); +} + +#[tokio::test] +async fn parity_rkey_sort_order() { + let f = ParityFixture::new().await; + let uid = Uuid::new_v4(); + let did = test_did("rkey"); + let handle = test_handle("rkey"); + let root_cid = test_cid(0); + let collection = test_nsid("post"); + + seed_repo(&f.pg, &did, &handle, &root_cid, uid).await; + seed_repo(&f.store, &did, &handle, &root_cid, uid).await; + + let records: Vec<(Rkey, CidLink)> = (0u8..10) + .map(|i| { + let rkey = test_rkey(&format!("3l{i}aaaaaaaa{i}")); + let cid = test_cid(i + 1); + (rkey, cid) + }) + .collect(); + + seed_records(&f.pg, uid, &collection, &records).await; + seed_records(&f.store, uid, &collection, &records).await; + + let pg_fwd = + f.pg.repo + .list_records(uid, &collection, None, 100, false, None, None) + .await + .unwrap(); + let store_fwd = f + .store + .repo + .list_records(uid, &collection, None, 100, false, None, None) + .await + .unwrap(); + + let pg_rkeys: Vec<&str> = pg_fwd.iter().map(|r| r.rkey.as_str()).collect(); + let store_rkeys: Vec<&str> = store_fwd.iter().map(|r| r.rkey.as_str()).collect(); + assert_eq!(pg_rkeys, store_rkeys, "forward rkey order mismatch"); + + let pg_rev = + f.pg.repo + .list_records(uid, &collection, None, 100, true, None, None) + .await + .unwrap(); + let store_rev = f + .store + .repo + .list_records(uid, &collection, None, 100, true, None, None) + .await + .unwrap(); + + let pg_rkeys_rev: Vec<&str> = pg_rev.iter().map(|r| r.rkey.as_str()).collect(); + let store_rkeys_rev: Vec<&str> = store_rev.iter().map(|r| r.rkey.as_str()).collect(); + assert_eq!(pg_rkeys_rev, store_rkeys_rev, "reverse rkey order mismatch"); + + let pg_cids: Vec<&str> = pg_fwd.iter().map(|r| r.record_cid.as_str()).collect(); + let store_cids: Vec<&str> = store_fwd.iter().map(|r| r.record_cid.as_str()).collect(); + assert_eq!(pg_cids, store_cids, "cid mapping mismatch"); +} + +#[tokio::test] +async fn parity_cursor_pagination() { + let f = ParityFixture::new().await; + let uid = Uuid::new_v4(); + let did = test_did("cursor"); + let handle = test_handle("cursor"); + let root_cid = test_cid(0); + let collection = test_nsid("post"); + + seed_repo(&f.pg, &did, &handle, &root_cid, uid).await; + seed_repo(&f.store, &did, &handle, &root_cid, uid).await; + + let records: Vec<(Rkey, CidLink)> = (0u8..20) + .map(|i| { + let rkey = test_rkey(&format!("3l{:02}aaaaaaaaa", i)); + let cid = test_cid(i + 1); + (rkey, cid) + }) + .collect(); + + seed_records(&f.pg, uid, &collection, &records).await; + seed_records(&f.store, uid, &collection, &records).await; + + let mut pg_all = Vec::new(); + let mut store_all = Vec::new(); + let mut pg_cursor: Option = None; + let mut store_cursor: Option = None; + let limit = 5i64; + let mut pages = 0; + + loop { + let pg_page = + f.pg.repo + .list_records( + uid, + &collection, + pg_cursor.as_ref(), + limit, + false, + None, + None, + ) + .await + .unwrap(); + let store_page = f + .store + .repo + .list_records( + uid, + &collection, + store_cursor.as_ref(), + limit, + false, + None, + None, + ) + .await + .unwrap(); + + assert_eq!( + pg_page.len(), + store_page.len(), + "page size mismatch at page {pages}" + ); + + let pg_rkeys: Vec<&str> = pg_page.iter().map(|r| r.rkey.as_str()).collect(); + let store_rkeys: Vec<&str> = store_page.iter().map(|r| r.rkey.as_str()).collect(); + assert_eq!( + pg_rkeys, store_rkeys, + "page content mismatch at page {pages}" + ); + + pg_all.extend(pg_page.iter().map(|r| r.rkey.clone())); + store_all.extend(store_page.iter().map(|r| r.rkey.clone())); + + if pg_page.len() < limit as usize { + break; + } + + pg_cursor = pg_page.last().map(|r| r.rkey.clone()); + store_cursor = store_page.last().map(|r| r.rkey.clone()); + pages += 1; + } + + assert_eq!(pg_all.len(), 20); + assert_eq!(store_all.len(), 20); +} + +#[tokio::test] +async fn parity_cursor_pagination_reverse() { + let f = ParityFixture::new().await; + let uid = Uuid::new_v4(); + let did = test_did("currev"); + let handle = test_handle("currev"); + let root_cid = test_cid(0); + let collection = test_nsid("post"); + + seed_repo(&f.pg, &did, &handle, &root_cid, uid).await; + seed_repo(&f.store, &did, &handle, &root_cid, uid).await; + + let records: Vec<(Rkey, CidLink)> = (0u8..15) + .map(|i| { + let rkey = test_rkey(&format!("3l{:02}aaaaaaaaa", i)); + let cid = test_cid(i + 1); + (rkey, cid) + }) + .collect(); + + seed_records(&f.pg, uid, &collection, &records).await; + seed_records(&f.store, uid, &collection, &records).await; + + let mut pg_all = Vec::new(); + let mut store_all = Vec::new(); + let mut pg_cursor: Option = None; + let mut store_cursor: Option = None; + let limit = 4i64; + + loop { + let pg_page = + f.pg.repo + .list_records( + uid, + &collection, + pg_cursor.as_ref(), + limit, + true, + None, + None, + ) + .await + .unwrap(); + let store_page = f + .store + .repo + .list_records( + uid, + &collection, + store_cursor.as_ref(), + limit, + true, + None, + None, + ) + .await + .unwrap(); + + let pg_rkeys: Vec<&str> = pg_page.iter().map(|r| r.rkey.as_str()).collect(); + let store_rkeys: Vec<&str> = store_page.iter().map(|r| r.rkey.as_str()).collect(); + assert_eq!(pg_rkeys, store_rkeys, "reverse page mismatch"); + + pg_all.extend(pg_page.iter().map(|r| r.rkey.clone())); + store_all.extend(store_page.iter().map(|r| r.rkey.clone())); + + if pg_page.len() < limit as usize { + break; + } + + pg_cursor = pg_page.last().map(|r| r.rkey.clone()); + store_cursor = store_page.last().map(|r| r.rkey.clone()); + } + + assert_eq!(pg_all.len(), 15); + assert_eq!(store_all.len(), 15); +} + +#[tokio::test] +async fn parity_rkey_range_query() { + let f = ParityFixture::new().await; + let uid = Uuid::new_v4(); + let did = test_did("range"); + let handle = test_handle("range"); + let root_cid = test_cid(0); + let collection = test_nsid("post"); + + seed_repo(&f.pg, &did, &handle, &root_cid, uid).await; + seed_repo(&f.store, &did, &handle, &root_cid, uid).await; + + let records: Vec<(Rkey, CidLink)> = (0u8..10) + .map(|i| { + let rkey = test_rkey(&format!("3l{:02}aaaaaaaaa", i)); + let cid = test_cid(i + 1); + (rkey, cid) + }) + .collect(); + + seed_records(&f.pg, uid, &collection, &records).await; + seed_records(&f.store, uid, &collection, &records).await; + + let start = test_rkey("3l03aaaaaaaaa"); + let end = test_rkey("3l07aaaaaaaaa"); + + let pg_range = + f.pg.repo + .list_records(uid, &collection, None, 100, false, Some(&start), Some(&end)) + .await + .unwrap(); + let store_range = f + .store + .repo + .list_records(uid, &collection, None, 100, false, Some(&start), Some(&end)) + .await + .unwrap(); + + let pg_rkeys: Vec<&str> = pg_range.iter().map(|r| r.rkey.as_str()).collect(); + let store_rkeys: Vec<&str> = store_range.iter().map(|r| r.rkey.as_str()).collect(); + assert_eq!(pg_rkeys, store_rkeys, "range query mismatch"); +} + +#[tokio::test] +async fn parity_collection_listing() { + let f = ParityFixture::new().await; + let uid = Uuid::new_v4(); + let did = test_did("colls"); + let handle = test_handle("colls"); + let root_cid = test_cid(0); + + seed_repo(&f.pg, &did, &handle, &root_cid, uid).await; + seed_repo(&f.store, &did, &handle, &root_cid, uid).await; + + let post_ns = test_nsid("post"); + let like_ns = Nsid::new("app.bsky.feed.like").unwrap(); + let repost_ns = Nsid::new("app.bsky.feed.repost").unwrap(); + let follow_ns = Nsid::new("app.bsky.graph.follow").unwrap(); + + let post_records = vec![(test_rkey("3laaaaaaaaa01"), test_cid(1))]; + let like_records = vec![(test_rkey("3laaaaaaaaa02"), test_cid(2))]; + let repost_records = vec![(test_rkey("3laaaaaaaaa03"), test_cid(3))]; + let follow_records = vec![(test_rkey("3laaaaaaaaa04"), test_cid(4))]; + + seed_records(&f.pg, uid, &post_ns, &post_records).await; + seed_records(&f.pg, uid, &like_ns, &like_records).await; + seed_records(&f.pg, uid, &repost_ns, &repost_records).await; + seed_records(&f.pg, uid, &follow_ns, &follow_records).await; + + seed_records(&f.store, uid, &post_ns, &post_records).await; + seed_records(&f.store, uid, &like_ns, &like_records).await; + seed_records(&f.store, uid, &repost_ns, &repost_records).await; + seed_records(&f.store, uid, &follow_ns, &follow_records).await; + + let mut pg_colls: Vec = + f.pg.repo + .list_collections(uid) + .await + .unwrap() + .into_iter() + .map(|n| n.as_str().to_owned()) + .collect(); + pg_colls.sort(); + + let mut store_colls: Vec = f + .store + .repo + .list_collections(uid) + .await + .unwrap() + .into_iter() + .map(|n| n.as_str().to_owned()) + .collect(); + store_colls.sort(); + + assert_eq!(pg_colls, store_colls, "collection listing mismatch"); + assert_eq!(pg_colls.len(), 4); + + let pg_count = f.pg.repo.count_records(uid).await.unwrap(); + let store_count = f.store.repo.count_records(uid).await.unwrap(); + assert_eq!(pg_count, store_count, "record count mismatch"); + assert_eq!(pg_count, 4); +} + +#[tokio::test] +async fn parity_record_get_and_delete() { + let f = ParityFixture::new().await; + let uid = Uuid::new_v4(); + let did = test_did("getdel"); + let handle = test_handle("getdel"); + let root_cid = test_cid(0); + let collection = test_nsid("post"); + let rkey = test_rkey("3laaaaaaaaa01"); + let cid = test_cid(1); + + seed_repo(&f.pg, &did, &handle, &root_cid, uid).await; + seed_repo(&f.store, &did, &handle, &root_cid, uid).await; + + seed_records(&f.pg, uid, &collection, &[(rkey.clone(), cid.clone())]).await; + seed_records(&f.store, uid, &collection, &[(rkey.clone(), cid.clone())]).await; + + let pg_cid = + f.pg.repo + .get_record_cid(uid, &collection, &rkey) + .await + .unwrap(); + let store_cid = f + .store + .repo + .get_record_cid(uid, &collection, &rkey) + .await + .unwrap(); + assert_eq!(pg_cid, store_cid, "get_record_cid mismatch"); + assert!(pg_cid.is_some()); + + f.pg.repo + .delete_records(uid, &[collection.clone()], &[rkey.clone()]) + .await + .unwrap(); + f.store + .repo + .delete_records(uid, &[collection.clone()], &[rkey.clone()]) + .await + .unwrap(); + + let pg_gone = + f.pg.repo + .get_record_cid(uid, &collection, &rkey) + .await + .unwrap(); + let store_gone = f + .store + .repo + .get_record_cid(uid, &collection, &rkey) + .await + .unwrap(); + assert_eq!(pg_gone, None); + assert_eq!(store_gone, None); +} + +#[tokio::test] +async fn parity_backlink_queries() { + let f = ParityFixture::new().await; + let uid = Uuid::new_v4(); + let did = test_did("blink"); + let handle = test_handle("blink"); + let root_cid = test_cid(0); + let like_ns = Nsid::new("app.bsky.feed.like").unwrap(); + let target_did = test_did("target"); + + seed_repo(&f.pg, &did, &handle, &root_cid, uid).await; + seed_repo(&f.store, &did, &handle, &root_cid, uid).await; + + let rkey1 = test_rkey("3laaaaaaaaa01"); + let rkey2 = test_rkey("3laaaaaaaaa02"); + let uri1 = test_at_uri(&did, &like_ns, &rkey1); + let uri2 = test_at_uri(&did, &like_ns, &rkey2); + let target_uri = format!( + "at://{}/app.bsky.feed.post/3laaaaaaaaa99", + target_did.as_str() + ); + + let backlinks = vec![ + Backlink { + uri: uri1.clone(), + path: BacklinkPath::Subject, + link_to: target_uri.clone(), + }, + Backlink { + uri: uri2.clone(), + path: BacklinkPath::Subject, + link_to: target_uri.clone(), + }, + ]; + + f.pg.backlink.add_backlinks(uid, &backlinks).await.unwrap(); + f.store + .backlink + .add_backlinks(uid, &backlinks) + .await + .unwrap(); + + let conflict_backlink = Backlink { + uri: uri1.clone(), + path: BacklinkPath::Subject, + link_to: target_uri.clone(), + }; + + let pg_conflicts = + f.pg.backlink + .get_backlink_conflicts(uid, &like_ns, &[conflict_backlink.clone()]) + .await + .unwrap(); + let store_conflicts = f + .store + .backlink + .get_backlink_conflicts(uid, &like_ns, &[conflict_backlink]) + .await + .unwrap(); + + assert_eq!( + pg_conflicts.len(), + store_conflicts.len(), + "backlink conflict count mismatch" + ); + + f.pg.backlink.remove_backlinks_by_uri(&uri1).await.unwrap(); + f.store + .backlink + .remove_backlinks_by_uri(&uri1) + .await + .unwrap(); + + let post_removal = Backlink { + uri: uri1.clone(), + path: BacklinkPath::Subject, + link_to: target_uri.clone(), + }; + + let pg_after = + f.pg.backlink + .get_backlink_conflicts(uid, &like_ns, &[post_removal.clone()]) + .await + .unwrap(); + let store_after = f + .store + .backlink + .get_backlink_conflicts(uid, &like_ns, &[post_removal]) + .await + .unwrap(); + + assert_eq!( + pg_after.len(), + store_after.len(), + "backlink conflicts after removal mismatch" + ); +} + +#[tokio::test] +async fn parity_backlink_remove_by_repo() { + let f = ParityFixture::new().await; + let uid = Uuid::new_v4(); + let did = test_did("blrep"); + let handle = test_handle("blrep"); + let root_cid = test_cid(0); + let like_ns = Nsid::new("app.bsky.feed.like").unwrap(); + + seed_repo(&f.pg, &did, &handle, &root_cid, uid).await; + seed_repo(&f.store, &did, &handle, &root_cid, uid).await; + + let rkey = test_rkey("3laaaaaaaaa01"); + let uri = test_at_uri(&did, &like_ns, &rkey); + let backlinks = vec![Backlink { + uri: uri.clone(), + path: BacklinkPath::Subject, + link_to: "at://did:plc:sometarget/app.bsky.feed.post/abc".to_owned(), + }]; + + f.pg.backlink.add_backlinks(uid, &backlinks).await.unwrap(); + f.store + .backlink + .add_backlinks(uid, &backlinks) + .await + .unwrap(); + + f.pg.backlink.remove_backlinks_by_repo(uid).await.unwrap(); + f.store + .backlink + .remove_backlinks_by_repo(uid) + .await + .unwrap(); + + let probe = Backlink { + uri, + path: BacklinkPath::Subject, + link_to: "at://did:plc:sometarget/app.bsky.feed.post/abc".to_owned(), + }; + let pg_after = + f.pg.backlink + .get_backlink_conflicts(uid, &like_ns, &[probe.clone()]) + .await + .unwrap(); + let store_after = f + .store + .backlink + .get_backlink_conflicts(uid, &like_ns, &[probe]) + .await + .unwrap(); + assert_eq!(pg_after.len(), 0); + assert_eq!(store_after.len(), 0); +} + +#[tokio::test] +async fn parity_blob_metadata() { + let f = ParityFixture::new().await; + let uid = Uuid::new_v4(); + let did = test_did("blob"); + let handle = test_handle("blob"); + let root_cid = test_cid(0); + + seed_repo(&f.pg, &did, &handle, &root_cid, uid).await; + seed_repo(&f.store, &did, &handle, &root_cid, uid).await; + + let blob_cid1 = test_cid(101); + let blob_cid2 = test_cid(102); + let blob_cid3 = test_cid(103); + + let blobs = [ + (&blob_cid1, "image/png", 1024i64, "blobs/a.png"), + (&blob_cid2, "image/jpeg", 2048, "blobs/b.jpg"), + (&blob_cid3, "application/pdf", 4096, "blobs/c.pdf"), + ]; + + blobs.iter().for_each(|(cid, mime, size, key)| { + let pg = Arc::clone(&f.pg); + let store = Arc::clone(&f.store); + let cid = (*cid).clone(); + let size = *size; + let mime = mime.to_string(); + let key = key.to_string(); + tokio::task::block_in_place(|| { + tokio::runtime::Handle::current().block_on(async { + pg.blob + .insert_blob(&cid, &mime, size, uid, &key) + .await + .unwrap(); + store + .blob + .insert_blob(&cid, &mime, size, uid, &key) + .await + .unwrap(); + }); + }); + }); + + let pg_meta = + f.pg.blob + .get_blob_metadata(&blob_cid1) + .await + .unwrap() + .unwrap(); + let store_meta = f + .store + .blob + .get_blob_metadata(&blob_cid1) + .await + .unwrap() + .unwrap(); + assert_eq!(pg_meta.mime_type, store_meta.mime_type); + assert_eq!(pg_meta.size_bytes, store_meta.size_bytes); + assert_eq!(pg_meta.storage_key, store_meta.storage_key); + + let pg_key = f.pg.blob.get_blob_storage_key(&blob_cid2).await.unwrap(); + let store_key = f.store.blob.get_blob_storage_key(&blob_cid2).await.unwrap(); + assert_eq!(pg_key, store_key); + + let pg_count = f.pg.blob.count_blobs_by_user(uid).await.unwrap(); + let store_count = f.store.blob.count_blobs_by_user(uid).await.unwrap(); + assert_eq!(pg_count, store_count); + assert_eq!(pg_count, 3); + + let pg_list = f.pg.blob.list_blobs_by_user(uid, None, 100).await.unwrap(); + let store_list = f + .store + .blob + .list_blobs_by_user(uid, None, 100) + .await + .unwrap(); + assert_eq!(pg_list.len(), store_list.len()); +} + +#[tokio::test] +async fn parity_blob_pagination() { + let f = ParityFixture::new().await; + let uid = Uuid::new_v4(); + let did = test_did("blobpg"); + let handle = test_handle("blobpg"); + let root_cid = test_cid(0); + + seed_repo(&f.pg, &did, &handle, &root_cid, uid).await; + seed_repo(&f.store, &did, &handle, &root_cid, uid).await; + + (0u8..8).for_each(|i| { + let cid = test_cid(200 + i); + let key = format!("blobs/pg_{i}.bin"); + let pg = Arc::clone(&f.pg); + let store = Arc::clone(&f.store); + tokio::task::block_in_place(|| { + tokio::runtime::Handle::current().block_on(async { + pg.blob + .insert_blob( + &cid, + "application/octet-stream", + 512 * (i as i64 + 1), + uid, + &key, + ) + .await + .unwrap(); + store + .blob + .insert_blob( + &cid, + "application/octet-stream", + 512 * (i as i64 + 1), + uid, + &key, + ) + .await + .unwrap(); + }); + }); + }); + + let mut pg_all = Vec::new(); + let mut store_all = Vec::new(); + let mut pg_cursor: Option = None; + let mut store_cursor: Option = None; + let limit = 3i64; + + loop { + let pg_page = + f.pg.blob + .list_blobs_by_user(uid, pg_cursor.as_deref(), limit) + .await + .unwrap(); + let store_page = f + .store + .blob + .list_blobs_by_user(uid, store_cursor.as_deref(), limit) + .await + .unwrap(); + + assert_eq!(pg_page.len(), store_page.len(), "blob page size mismatch"); + + let pg_cids: Vec<&str> = pg_page.iter().map(|c| c.as_str()).collect(); + let store_cids: Vec<&str> = store_page.iter().map(|c| c.as_str()).collect(); + assert_eq!(pg_cids, store_cids, "blob page content mismatch"); + + pg_all.extend(pg_page.iter().map(|c| c.as_str().to_owned())); + store_all.extend(store_page.iter().map(|c| c.as_str().to_owned())); + + if pg_page.len() < limit as usize { + break; + } + + pg_cursor = pg_page.last().map(|c| c.as_str().to_owned()); + store_cursor = store_page.last().map(|c| c.as_str().to_owned()); + } + + assert_eq!(pg_all.len(), 8); + assert_eq!(store_all.len(), 8); +} + +#[tokio::test] +async fn parity_blob_duplicate_insert() { + let f = ParityFixture::new().await; + let uid = Uuid::new_v4(); + let did = test_did("blobdup"); + let handle = test_handle("blobdup"); + let root_cid = test_cid(0); + + seed_repo(&f.pg, &did, &handle, &root_cid, uid).await; + seed_repo(&f.store, &did, &handle, &root_cid, uid).await; + + let cid = test_cid(150); + + let pg_first = + f.pg.blob + .insert_blob(&cid, "image/png", 1024, uid, "blobs/dup.png") + .await + .unwrap(); + let store_first = f + .store + .blob + .insert_blob(&cid, "image/png", 1024, uid, "blobs/dup.png") + .await + .unwrap(); + assert_eq!(pg_first, store_first); + + let pg_dup = + f.pg.blob + .insert_blob(&cid, "image/png", 1024, uid, "blobs/dup.png") + .await + .unwrap(); + let store_dup = f + .store + .blob + .insert_blob(&cid, "image/png", 1024, uid, "blobs/dup.png") + .await + .unwrap(); + assert_eq!(pg_dup, store_dup); +} + +#[tokio::test] +async fn parity_get_all_records() { + let f = ParityFixture::new().await; + let uid = Uuid::new_v4(); + let did = test_did("allrec"); + let handle = test_handle("allrec"); + let root_cid = test_cid(0); + + seed_repo(&f.pg, &did, &handle, &root_cid, uid).await; + seed_repo(&f.store, &did, &handle, &root_cid, uid).await; + + let post_ns = test_nsid("post"); + let like_ns = Nsid::new("app.bsky.feed.like").unwrap(); + + let posts = vec![ + (test_rkey("3laaaaaaaaa01"), test_cid(1)), + (test_rkey("3laaaaaaaaa02"), test_cid(2)), + ]; + let likes = vec![(test_rkey("3laaaaaaaaa03"), test_cid(3))]; + + seed_records(&f.pg, uid, &post_ns, &posts).await; + seed_records(&f.pg, uid, &like_ns, &likes).await; + seed_records(&f.store, uid, &post_ns, &posts).await; + seed_records(&f.store, uid, &like_ns, &likes).await; + + let mut pg_all = f.pg.repo.get_all_records(uid).await.unwrap(); + let mut store_all = f.store.repo.get_all_records(uid).await.unwrap(); + + pg_all.sort_by(|a, b| { + a.collection + .as_str() + .cmp(b.collection.as_str()) + .then(a.rkey.as_str().cmp(b.rkey.as_str())) + }); + store_all.sort_by(|a, b| { + a.collection + .as_str() + .cmp(b.collection.as_str()) + .then(a.rkey.as_str().cmp(b.rkey.as_str())) + }); + + assert_eq!(pg_all.len(), store_all.len()); + pg_all.iter().zip(store_all.iter()).for_each(|(p, s)| { + assert_eq!(p.collection.as_str(), s.collection.as_str()); + assert_eq!(p.rkey.as_str(), s.rkey.as_str()); + assert_eq!(p.record_cid.as_str(), s.record_cid.as_str()); + }); +} + +#[tokio::test] +async fn parity_comms_queue() { + let f = ParityFixture::new().await; + let uid = Uuid::new_v4(); + + let pg_id = + f.pg.infra + .enqueue_comms( + Some(uid), + CommsChannel::Email, + CommsType::Welcome, + "test@example.com", + Some("Welcome"), + "Welcome body", + None, + ) + .await + .unwrap(); + + let store_id = f + .store + .infra + .enqueue_comms( + Some(uid), + CommsChannel::Email, + CommsType::Welcome, + "test@example.com", + Some("Welcome"), + "Welcome body", + None, + ) + .await + .unwrap(); + + assert_ne!(pg_id, Uuid::nil()); + assert_ne!(store_id, Uuid::nil()); + + let pg_latest = + f.pg.infra + .get_latest_comms_for_user(uid, CommsType::Welcome, 10) + .await + .unwrap(); + let store_latest = f + .store + .infra + .get_latest_comms_for_user(uid, CommsType::Welcome, 10) + .await + .unwrap(); + + assert_eq!(pg_latest.len(), store_latest.len()); + assert_eq!(pg_latest[0].body, store_latest[0].body); + + let pg_count = + f.pg.infra + .count_comms_by_type(uid, CommsType::Welcome) + .await + .unwrap(); + let store_count = f + .store + .infra + .count_comms_by_type(uid, CommsType::Welcome) + .await + .unwrap(); + assert_eq!(pg_count, store_count); + assert_eq!(pg_count, 1); +} + +#[tokio::test] +async fn parity_invite_codes() { + let f = ParityFixture::new().await; + let code = format!("parity-invite-{}", Uuid::new_v4()); + + let pg_created = f.pg.infra.create_invite_code(&code, 5, None).await.unwrap(); + let store_created = f + .store + .infra + .create_invite_code(&code, 5, None) + .await + .unwrap(); + assert_eq!(pg_created, store_created); + + let pg_uses = + f.pg.infra + .get_invite_code_available_uses(&code) + .await + .unwrap(); + let store_uses = f + .store + .infra + .get_invite_code_available_uses(&code) + .await + .unwrap(); + assert_eq!(pg_uses, store_uses); + assert_eq!(pg_uses, Some(5)); +} + +#[tokio::test] +async fn parity_account_preferences() { + let f = ParityFixture::new().await; + let uid = Uuid::new_v4(); + let did = test_did("prefs"); + let handle = test_handle("prefs"); + let root_cid = test_cid(0); + + seed_repo(&f.pg, &did, &handle, &root_cid, uid).await; + seed_repo(&f.store, &did, &handle, &root_cid, uid).await; + + let pref_value = serde_json::json!({ + "$type": "app.bsky.actor.defs#adultContentPref", + "enabled": false + }); + + f.pg.infra + .upsert_account_preference( + uid, + "app.bsky.actor.defs#adultContentPref/0", + pref_value.clone(), + ) + .await + .unwrap(); + f.store + .infra + .upsert_account_preference(uid, "app.bsky.actor.defs#adultContentPref/0", pref_value) + .await + .unwrap(); + + let mut pg_prefs = f.pg.infra.get_account_preferences(uid).await.unwrap(); + let mut store_prefs = f.store.infra.get_account_preferences(uid).await.unwrap(); + + pg_prefs.sort_by(|a, b| a.0.cmp(&b.0)); + store_prefs.sort_by(|a, b| a.0.cmp(&b.0)); + + assert_eq!(pg_prefs.len(), store_prefs.len()); + pg_prefs.iter().zip(store_prefs.iter()).for_each(|(p, s)| { + assert_eq!(p.0, s.0); + assert_eq!(p.1, s.1); + }); +} + +#[tokio::test] +async fn parity_record_upsert_overwrites() { + let f = ParityFixture::new().await; + let uid = Uuid::new_v4(); + let did = test_did("upsert"); + let handle = test_handle("upsert"); + let root_cid = test_cid(0); + let collection = test_nsid("post"); + let rkey = test_rkey("3laaaaaaaaa01"); + + seed_repo(&f.pg, &did, &handle, &root_cid, uid).await; + seed_repo(&f.store, &did, &handle, &root_cid, uid).await; + + let cid_v1 = test_cid(1); + seed_records(&f.pg, uid, &collection, &[(rkey.clone(), cid_v1.clone())]).await; + seed_records( + &f.store, + uid, + &collection, + &[(rkey.clone(), cid_v1.clone())], + ) + .await; + + let cid_v2 = test_cid(2); + seed_records(&f.pg, uid, &collection, &[(rkey.clone(), cid_v2.clone())]).await; + seed_records( + &f.store, + uid, + &collection, + &[(rkey.clone(), cid_v2.clone())], + ) + .await; + + let pg_cid = + f.pg.repo + .get_record_cid(uid, &collection, &rkey) + .await + .unwrap(); + let store_cid = f + .store + .repo + .get_record_cid(uid, &collection, &rkey) + .await + .unwrap(); + assert_eq!(pg_cid, store_cid); + assert_eq!(pg_cid.unwrap().as_str(), cid_v2.as_str()); + + let pg_count = f.pg.repo.count_records(uid).await.unwrap(); + let store_count = f.store.repo.count_records(uid).await.unwrap(); + assert_eq!(pg_count, 1); + assert_eq!(store_count, 1); +} + +#[tokio::test] +async fn parity_empty_queries() { + let f = ParityFixture::new().await; + let uid = Uuid::new_v4(); + let did = test_did("empty"); + let handle = test_handle("empty"); + let root_cid = test_cid(0); + let collection = test_nsid("post"); + + seed_repo(&f.pg, &did, &handle, &root_cid, uid).await; + seed_repo(&f.store, &did, &handle, &root_cid, uid).await; + + let pg_records = + f.pg.repo + .list_records(uid, &collection, None, 100, false, None, None) + .await + .unwrap(); + let store_records = f + .store + .repo + .list_records(uid, &collection, None, 100, false, None, None) + .await + .unwrap(); + assert_eq!(pg_records.len(), 0); + assert_eq!(store_records.len(), 0); + + let pg_colls = f.pg.repo.list_collections(uid).await.unwrap(); + let store_colls = f.store.repo.list_collections(uid).await.unwrap(); + assert_eq!(pg_colls.len(), 0); + assert_eq!(store_colls.len(), 0); + + let pg_count = f.pg.repo.count_records(uid).await.unwrap(); + let store_count = f.store.repo.count_records(uid).await.unwrap(); + assert_eq!(pg_count, 0); + assert_eq!(store_count, 0); + + let pg_blobs = f.pg.blob.list_blobs_by_user(uid, None, 100).await.unwrap(); + let store_blobs = f + .store + .blob + .list_blobs_by_user(uid, None, 100) + .await + .unwrap(); + assert_eq!(pg_blobs.len(), 0); + assert_eq!(store_blobs.len(), 0); + + let nonexistent_cid = test_cid(255); + let pg_meta = f.pg.blob.get_blob_metadata(&nonexistent_cid).await.unwrap(); + let store_meta = f + .store + .blob + .get_blob_metadata(&nonexistent_cid) + .await + .unwrap(); + assert_eq!(pg_meta.is_none(), store_meta.is_none()); +} + +#[tokio::test] +async fn parity_deletion_requests() { + let f = ParityFixture::new().await; + let did = test_did("delreq"); + let token = format!("del-token-{}", Uuid::new_v4()); + let expires = chrono::Utc::now() + chrono::Duration::hours(24); + + f.pg.infra + .create_deletion_request(&token, &did, expires) + .await + .unwrap(); + f.store + .infra + .create_deletion_request(&token, &did, expires) + .await + .unwrap(); + + let pg_req = f.pg.infra.get_deletion_request(&token).await.unwrap(); + let store_req = f.store.infra.get_deletion_request(&token).await.unwrap(); + assert!(pg_req.is_some()); + assert!(store_req.is_some()); + assert_eq!( + pg_req.as_ref().unwrap().did, + store_req.as_ref().unwrap().did + ); + + let pg_by_did = f.pg.infra.get_deletion_request_by_did(&did).await.unwrap(); + let store_by_did = f + .store + .infra + .get_deletion_request_by_did(&did) + .await + .unwrap(); + assert!(pg_by_did.is_some()); + assert!(store_by_did.is_some()); + assert_eq!(pg_by_did.unwrap().token, store_by_did.unwrap().token); + + f.pg.infra.delete_deletion_request(&token).await.unwrap(); + f.store.infra.delete_deletion_request(&token).await.unwrap(); + + let pg_gone = f.pg.infra.get_deletion_request(&token).await.unwrap(); + let store_gone = f.store.infra.get_deletion_request(&token).await.unwrap(); + assert!(pg_gone.is_none()); + assert!(store_gone.is_none()); +} + +#[tokio::test] +async fn parity_signing_key_reservation() { + let f = ParityFixture::new().await; + let did = test_did("sigkey"); + let expires = chrono::Utc::now() + chrono::Duration::hours(1); + let pub_key = format!("did:key:z6Mk{}", Uuid::new_v4().simple()); + let priv_bytes = vec![1u8, 2, 3, 4, 5, 6, 7, 8]; + + f.pg.infra + .reserve_signing_key(Some(&did), &pub_key, &priv_bytes, expires) + .await + .unwrap(); + f.store + .infra + .reserve_signing_key(Some(&did), &pub_key, &priv_bytes, expires) + .await + .unwrap(); + + let pg_key = f.pg.infra.get_reserved_signing_key(&pub_key).await.unwrap(); + let store_key = f + .store + .infra + .get_reserved_signing_key(&pub_key) + .await + .unwrap(); + assert!(pg_key.is_some()); + assert!(store_key.is_some()); + assert_eq!( + pg_key.unwrap().private_key_bytes, + store_key.unwrap().private_key_bytes + ); + + let pg_full = + f.pg.infra + .get_reserved_signing_key_full(&pub_key) + .await + .unwrap(); + let store_full = f + .store + .infra + .get_reserved_signing_key_full(&pub_key) + .await + .unwrap(); + assert!(pg_full.is_some()); + assert!(store_full.is_some()); + let pg_f = pg_full.unwrap(); + let store_f = store_full.unwrap(); + assert_eq!(pg_f.public_key_did_key, store_f.public_key_did_key); + assert_eq!(pg_f.did, store_f.did); +} + +#[tokio::test] +async fn parity_repo_root_operations() { + let f = ParityFixture::new().await; + let uid = Uuid::new_v4(); + let did = test_did("root"); + let handle = test_handle("root"); + let root_cid = test_cid(0); + + seed_repo(&f.pg, &did, &handle, &root_cid, uid).await; + seed_repo(&f.store, &did, &handle, &root_cid, uid).await; + + let pg_root = f.pg.repo.get_repo_root_by_did(&did).await.unwrap(); + let store_root = f.store.repo.get_repo_root_by_did(&did).await.unwrap(); + assert_eq!(pg_root, store_root); + + let new_root = test_cid(99); + f.pg.repo + .update_repo_root(uid, &new_root, "rev1") + .await + .unwrap(); + f.store + .repo + .update_repo_root(uid, &new_root, "rev1") + .await + .unwrap(); + + let pg_updated = f.pg.repo.get_repo_root_by_did(&did).await.unwrap(); + let store_updated = f.store.repo.get_repo_root_by_did(&did).await.unwrap(); + assert_eq!(pg_updated, store_updated); + assert_eq!(pg_updated.unwrap().as_str(), new_root.as_str()); + + let pg_info = f.pg.repo.get_repo(uid).await.unwrap().unwrap(); + let store_info = f.store.repo.get_repo(uid).await.unwrap().unwrap(); + assert_eq!(pg_info.repo_rev, store_info.repo_rev); + assert_eq!( + pg_info.repo_root_cid.as_str(), + store_info.repo_root_cid.as_str() + ); +} + +#[tokio::test] +async fn parity_delete_all_records() { + let f = ParityFixture::new().await; + let uid = Uuid::new_v4(); + let did = test_did("delall"); + let handle = test_handle("delall"); + let root_cid = test_cid(0); + let collection = test_nsid("post"); + + seed_repo(&f.pg, &did, &handle, &root_cid, uid).await; + seed_repo(&f.store, &did, &handle, &root_cid, uid).await; + + let records: Vec<(Rkey, CidLink)> = (0u8..5) + .map(|i| (test_rkey(&format!("3l{:02}aaaaaaaaa", i)), test_cid(i + 1))) + .collect(); + + seed_records(&f.pg, uid, &collection, &records).await; + seed_records(&f.store, uid, &collection, &records).await; + + f.pg.repo.delete_all_records(uid).await.unwrap(); + f.store.repo.delete_all_records(uid).await.unwrap(); + + let pg_count = f.pg.repo.count_records(uid).await.unwrap(); + let store_count = f.store.repo.count_records(uid).await.unwrap(); + assert_eq!(pg_count, 0); + assert_eq!(store_count, 0); + + let pg_colls = f.pg.repo.list_collections(uid).await.unwrap(); + let store_colls = f.store.repo.list_collections(uid).await.unwrap(); + assert_eq!(pg_colls.len(), 0); + assert_eq!(store_colls.len(), 0); +} + +#[tokio::test] +async fn parity_plc_tokens() { + let f = ParityFixture::new().await; + let uid = Uuid::new_v4(); + let did = test_did("plctok"); + let handle = test_handle("plctok"); + let root_cid = test_cid(0); + + seed_repo(&f.pg, &did, &handle, &root_cid, uid).await; + seed_repo(&f.store, &did, &handle, &root_cid, uid).await; + + let token = format!("plc-{}", Uuid::new_v4()); + let expires = chrono::Utc::now() + chrono::Duration::hours(1); + + f.pg.infra + .insert_plc_token(uid, &token, expires) + .await + .unwrap(); + f.store + .infra + .insert_plc_token(uid, &token, expires) + .await + .unwrap(); + + let pg_expiry = f.pg.infra.get_plc_token_expiry(uid, &token).await.unwrap(); + let store_expiry = f + .store + .infra + .get_plc_token_expiry(uid, &token) + .await + .unwrap(); + assert!(pg_expiry.is_some()); + assert!(store_expiry.is_some()); + + let pg_by_did = f.pg.infra.get_plc_tokens_by_did(&did).await.unwrap(); + let store_by_did = f.store.infra.get_plc_tokens_by_did(&did).await.unwrap(); + assert_eq!(pg_by_did.len(), store_by_did.len()); + + let pg_count = f.pg.infra.count_plc_tokens_by_did(&did).await.unwrap(); + let store_count = f.store.infra.count_plc_tokens_by_did(&did).await.unwrap(); + assert_eq!(pg_count, store_count); + assert_eq!(pg_count, 1); + + f.pg.infra.delete_plc_token(uid, &token).await.unwrap(); + f.store.infra.delete_plc_token(uid, &token).await.unwrap(); + + let pg_gone = f.pg.infra.get_plc_token_expiry(uid, &token).await.unwrap(); + let store_gone = f + .store + .infra + .get_plc_token_expiry(uid, &token) + .await + .unwrap(); + assert!(pg_gone.is_none()); + assert!(store_gone.is_none()); +} + +#[tokio::test] +async fn parity_blob_delete_and_takedown() { + let f = ParityFixture::new().await; + let uid = Uuid::new_v4(); + let did = test_did("blobdel"); + let handle = test_handle("blobdel"); + let root_cid = test_cid(0); + + seed_repo(&f.pg, &did, &handle, &root_cid, uid).await; + seed_repo(&f.store, &did, &handle, &root_cid, uid).await; + + let cid = test_cid(180); + f.pg.blob + .insert_blob(&cid, "image/png", 1024, uid, "blobs/td.png") + .await + .unwrap(); + f.store + .blob + .insert_blob(&cid, "image/png", 1024, uid, "blobs/td.png") + .await + .unwrap(); + + let pg_td = + f.pg.blob + .update_blob_takedown(&cid, Some("mod-action-1")) + .await + .unwrap(); + let store_td = f + .store + .blob + .update_blob_takedown(&cid, Some("mod-action-1")) + .await + .unwrap(); + assert_eq!(pg_td, store_td); + + let pg_with_td = f.pg.blob.get_blob_with_takedown(&cid).await.unwrap(); + let store_with_td = f.store.blob.get_blob_with_takedown(&cid).await.unwrap(); + assert_eq!( + pg_with_td.as_ref().map(|b| b.takedown_ref.as_deref()), + store_with_td.as_ref().map(|b| b.takedown_ref.as_deref()) + ); + + f.pg.blob.delete_blob_by_cid(&cid).await.unwrap(); + f.store.blob.delete_blob_by_cid(&cid).await.unwrap(); + + let pg_meta = f.pg.blob.get_blob_metadata(&cid).await.unwrap(); + let store_meta = f.store.blob.get_blob_metadata(&cid).await.unwrap(); + assert!(pg_meta.is_none()); + assert!(store_meta.is_none()); +} diff --git a/crates/tranquil-pds/tests/whole_story.rs b/crates/tranquil-pds/tests/whole_story.rs index a683f69..1d2279e 100644 --- a/crates/tranquil-pds/tests/whole_story.rs +++ b/crates/tranquil-pds/tests/whole_story.rs @@ -177,31 +177,31 @@ async fn test_complete_user_journey_signup_to_deletion() { .expect("Request delete failed"); assert_eq!(request_delete_res.status(), StatusCode::OK); - let pool = get_test_db_pool().await; - let row = sqlx::query!( - "SELECT token FROM account_deletion_requests WHERE did = $1", - did - ) - .fetch_one(pool) - .await - .expect("Failed to get deletion token"); + let repos = get_test_repos().await; + let deletion_request = repos + .infra + .get_deletion_request_by_did(&tranquil_types::Did::new(did.clone()).unwrap()) + .await + .unwrap() + .unwrap(); let final_delete_res = client .post(format!("{}/xrpc/com.atproto.server.deleteAccount", base)) .json(&json!({ "did": did, "password": password, - "token": row.token + "token": deletion_request.token })) .send() .await .expect("Final delete failed"); assert_eq!(final_delete_res.status(), StatusCode::OK); - let user_gone = sqlx::query!("SELECT id FROM users WHERE did = $1", did) - .fetch_optional(pool) + let user_gone = repos + .user + .get_by_did(&tranquil_types::Did::new(did.clone()).unwrap()) .await - .expect("Failed to check user"); + .unwrap(); assert!(user_gone.is_none(), "User should be deleted"); } diff --git a/crates/tranquil-server/src/main.rs b/crates/tranquil-server/src/main.rs index e69e5a5..ae6811c 100644 --- a/crates/tranquil-server/src/main.rs +++ b/crates/tranquil-server/src/main.rs @@ -114,8 +114,8 @@ async fn run() -> Result<(), Box> { let signal_sender = if tranquil_config::get().signal.enabled { let slot = Arc::new(tranquil_signal::SignalSlot::default()); state = state.with_signal_sender(slot.clone()); - if let Some(client) = - tranquil_signal::SignalClient::from_pool(&state.repos.pool, shutdown.clone()).await + if let Some(provider) = &state.signal_store_provider + && let Some(client) = provider.load_signal_client(shutdown.clone()).await { slot.set_client(client).await; info!("Signal device already linked"); diff --git a/crates/tranquil-signal/Cargo.toml b/crates/tranquil-signal/Cargo.toml index b83246c..d36d661 100644 --- a/crates/tranquil-signal/Cargo.toml +++ b/crates/tranquil-signal/Cargo.toml @@ -4,10 +4,14 @@ version.workspace = true edition.workspace = true license.workspace = true +[features] +fjall-store = ["dep:fjall"] + [dependencies] presage = { workspace = true } async-trait = { workspace = true } chrono = { workspace = true } +fjall = { version = "3", optional = true } sqlx = { workspace = true } tracing = { workspace = true } tokio = { workspace = true } diff --git a/crates/tranquil-signal/src/client.rs b/crates/tranquil-signal/src/client.rs index 56ec96d..244df74 100644 --- a/crates/tranquil-signal/src/client.rs +++ b/crates/tranquil-signal/src/client.rs @@ -7,6 +7,7 @@ use std::time::{Duration, SystemTime, UNIX_EPOCH}; use presage::libsignal_service::configuration::SignalServers; use presage::manager::Registered; use presage::proto::DataMessage; +use presage::store::Store; use sqlx::PgPool; use tokio::sync::{RwLock, mpsc, oneshot}; use tokio_util::sync::CancellationToken; @@ -190,7 +191,7 @@ impl Drop for LinkingGuard { #[derive(Debug, thiserror::Error)] pub enum SignalError { #[error("store: {0}")] - Store(#[from] crate::store::PgStoreError), + Store(String), #[error("presage: {0}")] Presage(String), #[error("username lookup failed: {0}")] @@ -209,7 +210,11 @@ pub enum SignalError { Runtime(String), } -type Manager = presage::Manager; +impl From for SignalError { + fn from(e: crate::store::PgStoreError) -> Self { + Self::Store(e.to_string()) + } +} struct SendRequest { recipient: SignalUsername, @@ -311,7 +316,10 @@ pub struct SignalClient { } impl SignalClient { - fn from_manager(manager: Manager, shutdown: CancellationToken) -> Result { + fn from_manager( + manager: presage::Manager, + shutdown: CancellationToken, + ) -> Result { let (tx, rx) = mpsc::channel::(64); spawn_signal_thread("signal-worker", move || { @@ -322,8 +330,7 @@ impl SignalClient { Ok(Self { tx }) } - pub async fn from_pool(db: &PgPool, shutdown: CancellationToken) -> Option { - let store = PgSignalStore::new(db.clone()); + pub async fn from_store(store: S, shutdown: CancellationToken) -> Option { let (init_tx, init_rx) = oneshot::channel(); spawn_signal_thread("signal-init", move || { @@ -348,8 +355,12 @@ impl SignalClient { .ok() } - async fn worker_loop( - mut manager: Manager, + pub async fn from_pool(db: &PgPool, shutdown: CancellationToken) -> Option { + Self::from_store(PgSignalStore::new(db.clone()), shutdown).await + } + + async fn worker_loop( + mut manager: presage::Manager, mut rx: mpsc::Receiver, shutdown: CancellationToken, ) { @@ -391,8 +402,8 @@ impl SignalClient { } } - async fn handle_send( - manager: &mut Manager, + async fn handle_send( + manager: &mut presage::Manager, recipient: &SignalUsername, message: &MessageBody, ) -> Result<(), SignalError> { @@ -444,8 +455,8 @@ impl SignalClient { .map_err(|_| SignalError::Runtime("signal worker dropped request".into()))? } - pub async fn link_device( - db: &PgPool, + pub async fn link_device_with_store( + store: S, device_name: DeviceName, shutdown: CancellationToken, link_cancel: CancellationToken, @@ -457,7 +468,6 @@ impl SignalClient { )); } - let store = PgSignalStore::new(db.clone()); let (url_tx, url_rx) = oneshot::channel::>(); let (done_tx, done_rx) = oneshot::channel::>(); @@ -538,4 +548,21 @@ impl SignalClient { completion: done_rx, }) } + + pub async fn link_device( + db: &PgPool, + device_name: DeviceName, + shutdown: CancellationToken, + link_cancel: CancellationToken, + linking_flag: Arc, + ) -> Result { + Self::link_device_with_store( + PgSignalStore::new(db.clone()), + device_name, + shutdown, + link_cancel, + linking_flag, + ) + .await + } } diff --git a/crates/tranquil-signal/src/fjall_store.rs b/crates/tranquil-signal/src/fjall_store.rs new file mode 100644 index 0000000..47b3884 --- /dev/null +++ b/crates/tranquil-signal/src/fjall_store.rs @@ -0,0 +1,1157 @@ +use std::ops::RangeBounds; +use std::sync::Arc; + +use async_trait::async_trait; +use fjall::{Database, Keyspace}; +use presage::{ + AvatarBytes, + libsignal_service::{ + Profile, + pre_keys::{KyberPreKeyStoreExt, PreKeysStore}, + prelude::{Content, MasterKey, ProfileKey, SessionStoreExt, Uuid}, + protocol::{ + CiphertextMessageType, DeviceId, Direction, GenericSignedPreKey, IdentityChange, + IdentityKey, IdentityKeyPair, IdentityKeyStore, KyberPreKeyId, KyberPreKeyRecord, + KyberPreKeyStore, PreKeyId, PreKeyRecord, PreKeyStore, ProtocolAddress, ProtocolStore, + PublicKey, SenderCertificate, SenderKeyRecord, SenderKeyStore, ServiceId, + SessionRecord, SessionStore, SignalProtocolError, SignedPreKeyId, SignedPreKeyRecord, + SignedPreKeyStore, + }, + push_service::DEFAULT_DEVICE_ID, + zkgroup::GroupMasterKeyBytes, + }, + manager::RegistrationData, + model::{contacts::Contact, groups::Group}, + store::{ContentsStore, StateStore, StickerPack, Store, Thread}, +}; +use tracing::warn; + +use crate::{DeviceName, LinkResult, SignalClient, SignalError}; + +#[derive(Debug, thiserror::Error)] +pub enum FjallStoreError { + #[error("fjall: {0}")] + Fjall(#[from] fjall::Error), + #[error("json: {0}")] + Json(#[from] serde_json::Error), + #[error("protocol: {0}")] + Protocol(#[from] SignalProtocolError), + #[error("not found: {0}")] + NotFound(String), +} + +impl presage::store::StoreError for FjallStoreError {} + +#[derive(Clone)] +pub struct FjallSignalStore { + db: Database, + ks: Keyspace, +} + +#[derive(Clone)] +pub struct FjallProtocolStore { + store: FjallSignalStore, + identity: IdentityType, +} + +#[derive(Debug, Clone, Copy)] +enum IdentityType { + Aci, + Pni, +} + +impl IdentityType { + fn tag(self) -> u8 { + match self { + Self::Aci => b'a', + Self::Pni => b'p', + } + } +} + +fn kv_key(name: &[u8]) -> Vec { + let mut k = Vec::with_capacity(1 + name.len()); + k.push(b'K'); + k.extend_from_slice(name); + k +} + +fn session_key(identity: IdentityType, address: &str, device_id: u32) -> Vec { + let mut k = Vec::with_capacity(2 + address.len() + 1 + 4); + k.push(b'S'); + k.push(identity.tag()); + k.extend_from_slice(address.as_bytes()); + k.push(0); + k.extend_from_slice(&device_id.to_be_bytes()); + k +} + +fn session_prefix(identity: IdentityType, address: &str) -> Vec { + let mut k = Vec::with_capacity(2 + address.len() + 1); + k.push(b'S'); + k.push(identity.tag()); + k.extend_from_slice(address.as_bytes()); + k.push(0); + k +} + +fn identity_rec_key(identity: IdentityType, address: &str) -> Vec { + let mut k = Vec::with_capacity(2 + address.len()); + k.push(b'I'); + k.push(identity.tag()); + k.extend_from_slice(address.as_bytes()); + k +} + +fn pre_key_key(identity: IdentityType, id: u32) -> Vec { + let mut k = vec![b'P', identity.tag()]; + k.extend_from_slice(&id.to_be_bytes()); + k +} + +fn pre_key_prefix(identity: IdentityType) -> Vec { + vec![b'P', identity.tag()] +} + +fn signed_pre_key_key(identity: IdentityType, id: u32) -> Vec { + let mut k = vec![b'G', identity.tag()]; + k.extend_from_slice(&id.to_be_bytes()); + k +} + +fn signed_pre_key_prefix(identity: IdentityType) -> Vec { + vec![b'G', identity.tag()] +} + +fn kyber_key(identity: IdentityType, id: u32) -> Vec { + let mut k = vec![b'Q', identity.tag()]; + k.extend_from_slice(&id.to_be_bytes()); + k +} + +fn kyber_prefix(identity: IdentityType) -> Vec { + vec![b'Q', identity.tag()] +} + +fn sender_key_key( + identity: IdentityType, + address: &str, + device_id: u32, + distribution_id: Uuid, +) -> Vec { + let mut k = Vec::with_capacity(2 + address.len() + 1 + 4 + 16); + k.push(b'N'); + k.push(identity.tag()); + k.extend_from_slice(address.as_bytes()); + k.push(0); + k.extend_from_slice(&device_id.to_be_bytes()); + k.extend_from_slice(distribution_id.as_bytes()); + k +} + +fn base_key_seen_key( + identity: IdentityType, + kyber_id: u32, + signed_id: u32, + base_key: &[u8], +) -> Vec { + let mut k = Vec::with_capacity(2 + 4 + 4 + base_key.len()); + k.push(b'B'); + k.push(identity.tag()); + k.extend_from_slice(&kyber_id.to_be_bytes()); + k.extend_from_slice(&signed_id.to_be_bytes()); + k.extend_from_slice(base_key); + k +} + +fn profile_key_key(uuid: &Uuid) -> Vec { + let mut k = Vec::with_capacity(1 + 16); + k.push(b'R'); + k.extend_from_slice(uuid.as_bytes()); + k +} + +const KYBER_META_SIZE: usize = 1 + 8; + +fn encode_kyber_value(record: &[u8], is_last_resort: bool, stale_at_ms: Option) -> Vec { + let mut v = Vec::with_capacity(KYBER_META_SIZE + record.len()); + v.push(u8::from(is_last_resort)); + v.extend_from_slice(&stale_at_ms.unwrap_or(0).to_be_bytes()); + v.extend_from_slice(record); + v +} + +fn decode_kyber_value(data: &[u8]) -> Option<(bool, Option, &[u8])> { + (data.len() >= KYBER_META_SIZE).then_some(())?; + let is_last_resort = data[0] != 0; + let stale_ms = i64::from_be_bytes(data[1..9].try_into().ok()?); + let stale_at = (stale_ms != 0).then_some(stale_ms); + Some((is_last_resort, stale_at, &data[KYBER_META_SIZE..])) +} + +fn into_protocol_err(e: E) -> SignalProtocolError { + SignalProtocolError::InvalidState("fjall", e.to_string()) +} + +fn guard_to_kv(guard: fjall::Guard) -> Result { + guard.into_inner() +} + +fn max_id_in_prefix(ks: &Keyspace, prefix: &[u8]) -> Result, SignalProtocolError> { + let mut max: Option = None; + ks.prefix(prefix).try_for_each(|guard| { + let (key_bytes, _) = guard_to_kv(guard).map_err(into_protocol_err)?; + let id_offset = prefix.len(); + if key_bytes.len() >= id_offset + 4 { + let id = u32::from_be_bytes( + key_bytes[id_offset..id_offset + 4] + .try_into() + .map_err(into_protocol_err)?, + ); + max = Some(max.map_or(id, |m| m.max(id))); + } + Ok::<(), SignalProtocolError>(()) + })?; + Ok(max) +} + +fn next_id_from_max(max_id: Option) -> Result { + match max_id { + None => Ok(1), + Some(id) => id.checked_add(1).ok_or_else(|| { + SignalProtocolError::InvalidState( + "pre key id space exhausted", + format!("max id {id} has no successor"), + ) + }), + } +} + +impl FjallSignalStore { + pub fn new(db: Database, ks: Keyspace) -> Self { + Self { db, ks } + } + + pub fn is_linked(&self) -> Result { + Ok(self.ks.get(kv_key(b"registration"))?.is_some()) + } + + pub fn clear_all(&self) -> Result<(), FjallStoreError> { + let mut batch = self.db.batch(); + self.ks.prefix([]).try_for_each(|guard| { + let (key, _) = guard.into_inner()?; + batch.remove(&self.ks, key.as_ref()); + Ok::<(), fjall::Error>(()) + })?; + batch.commit()?; + Ok(()) + } + + fn set_identity_key_pair( + &self, + identity: IdentityType, + key_pair: IdentityKeyPair, + ) -> Result<(), FjallStoreError> { + let key_name = match identity { + IdentityType::Aci => b"identity_keypair_aci".as_slice(), + IdentityType::Pni => b"identity_keypair_pni".as_slice(), + }; + let serialized = key_pair.serialize(); + self.ks.insert(kv_key(key_name), &*serialized)?; + Ok(()) + } +} + +impl Store for FjallSignalStore { + type Error = FjallStoreError; + type AciStore = FjallProtocolStore; + type PniStore = FjallProtocolStore; + + async fn clear(&mut self) -> Result<(), FjallStoreError> { + self.clear_all() + } + + fn aci_protocol_store(&self) -> Self::AciStore { + FjallProtocolStore { + store: self.clone(), + identity: IdentityType::Aci, + } + } + + fn pni_protocol_store(&self) -> Self::PniStore { + FjallProtocolStore { + store: self.clone(), + identity: IdentityType::Pni, + } + } +} + +impl StateStore for FjallSignalStore { + type StateStoreError = FjallStoreError; + + async fn load_registration_data(&self) -> Result, FjallStoreError> { + self.ks + .get(kv_key(b"registration"))? + .map(|v| serde_json::from_slice(&v)) + .transpose() + .map_err(From::from) + } + + async fn save_registration_data( + &mut self, + state: &RegistrationData, + ) -> Result<(), FjallStoreError> { + let value = serde_json::to_vec(state)?; + self.ks.insert(kv_key(b"registration"), &value)?; + Ok(()) + } + + async fn is_registered(&self) -> bool { + self.ks + .get(kv_key(b"registration")) + .ok() + .flatten() + .is_some() + } + + async fn clear_registration(&mut self) -> Result<(), FjallStoreError> { + let protocol_tags: &[u8] = b"SIPGQNB"; + let mut batch = self.db.batch(); + batch.remove(&self.ks, kv_key(b"registration")); + self.ks.prefix([]).try_for_each(|guard| { + let (key, _) = guard.into_inner()?; + if key.first().is_some_and(|b| protocol_tags.contains(b)) { + batch.remove(&self.ks, key.as_ref()); + } + Ok::<(), fjall::Error>(()) + })?; + batch.commit()?; + Ok(()) + } + + async fn set_aci_identity_key_pair( + &self, + key_pair: IdentityKeyPair, + ) -> Result<(), FjallStoreError> { + self.set_identity_key_pair(IdentityType::Aci, key_pair) + } + + async fn set_pni_identity_key_pair( + &self, + key_pair: IdentityKeyPair, + ) -> Result<(), FjallStoreError> { + self.set_identity_key_pair(IdentityType::Pni, key_pair) + } + + async fn sender_certificate(&self) -> Result, FjallStoreError> { + self.ks + .get(kv_key(b"sender_certificate"))? + .map(|v| SenderCertificate::deserialize(&v)) + .transpose() + .map_err(From::from) + } + + async fn save_sender_certificate( + &self, + certificate: &SenderCertificate, + ) -> Result<(), FjallStoreError> { + let serialized = certificate.serialized()?; + self.ks.insert(kv_key(b"sender_certificate"), serialized)?; + Ok(()) + } + + async fn fetch_master_key(&self) -> Result, FjallStoreError> { + self.ks + .get(kv_key(b"master_key"))? + .map(|v| MasterKey::from_slice(&v)) + .transpose() + .map_err(|_| FjallStoreError::NotFound("master key has wrong length".into())) + } + + async fn store_master_key( + &self, + master_key: Option<&MasterKey>, + ) -> Result<(), FjallStoreError> { + match master_key { + Some(k) => self.ks.insert(kv_key(b"master_key"), k.inner)?, + None => self.ks.remove(kv_key(b"master_key"))?, + } + Ok(()) + } +} + +impl ProtocolStore for FjallProtocolStore {} + +#[async_trait(?Send)] +impl SessionStore for FjallProtocolStore { + async fn load_session( + &self, + address: &ProtocolAddress, + ) -> Result, SignalProtocolError> { + let key = session_key( + self.identity, + address.name(), + u32::from(address.device_id()), + ); + self.store + .ks + .get(key) + .map_err(into_protocol_err)? + .map(|v| SessionRecord::deserialize(&v)) + .transpose() + } + + async fn store_session( + &mut self, + address: &ProtocolAddress, + record: &SessionRecord, + ) -> Result<(), SignalProtocolError> { + let key = session_key( + self.identity, + address.name(), + u32::from(address.device_id()), + ); + let serialized = record.serialize()?; + self.store + .ks + .insert(key, serialized.as_slice()) + .map_err(into_protocol_err) + } +} + +#[async_trait(?Send)] +impl SessionStoreExt for FjallProtocolStore { + async fn get_sub_device_sessions( + &self, + name: &ServiceId, + ) -> Result, SignalProtocolError> { + let address = name.raw_uuid().to_string(); + let default_device = u32::from(*DEFAULT_DEVICE_ID); + let prefix = session_prefix(self.identity, &address); + let prefix_len = prefix.len(); + + let mut devices = Vec::new(); + self.store.ks.prefix(prefix).try_for_each(|guard| { + let (key, _) = guard_to_kv(guard).map_err(into_protocol_err)?; + if key.len() >= prefix_len + 4 { + let id = u32::from_be_bytes( + key[prefix_len..prefix_len + 4] + .try_into() + .map_err(into_protocol_err)?, + ); + if id != default_device + && let Ok(byte) = u8::try_from(id) + && let Ok(device_id) = DeviceId::new(byte) + { + devices.push(device_id); + } + } + Ok::<(), SignalProtocolError>(()) + })?; + Ok(devices) + } + + async fn delete_session(&self, address: &ProtocolAddress) -> Result<(), SignalProtocolError> { + let key = session_key( + self.identity, + address.name(), + u32::from(address.device_id()), + ); + self.store.ks.remove(key).map_err(into_protocol_err) + } + + async fn delete_all_sessions(&self, name: &ServiceId) -> Result { + let address = name.raw_uuid().to_string(); + let prefix = session_prefix(self.identity, &address); + let mut batch = self.store.db.batch(); + let mut count = 0usize; + self.store.ks.prefix(prefix).try_for_each(|guard| { + let (key, _) = guard_to_kv(guard).map_err(into_protocol_err)?; + batch.remove(&self.store.ks, key.as_ref()); + count = count.saturating_add(1); + Ok::<(), SignalProtocolError>(()) + })?; + batch.commit().map_err(into_protocol_err)?; + Ok(count) + } +} + +#[async_trait(?Send)] +impl PreKeyStore for FjallProtocolStore { + async fn get_pre_key(&self, prekey_id: PreKeyId) -> Result { + let key = pre_key_key(self.identity, u32::from(prekey_id)); + let data = self + .store + .ks + .get(key) + .map_err(into_protocol_err)? + .ok_or(SignalProtocolError::InvalidPreKeyId)?; + PreKeyRecord::deserialize(&data) + } + + async fn save_pre_key( + &mut self, + prekey_id: PreKeyId, + record: &PreKeyRecord, + ) -> Result<(), SignalProtocolError> { + let key = pre_key_key(self.identity, u32::from(prekey_id)); + let serialized = record.serialize()?; + self.store + .ks + .insert(key, serialized.as_slice()) + .map_err(into_protocol_err) + } + + async fn remove_pre_key(&mut self, prekey_id: PreKeyId) -> Result<(), SignalProtocolError> { + let key = pre_key_key(self.identity, u32::from(prekey_id)); + self.store.ks.remove(key).map_err(into_protocol_err) + } +} + +#[async_trait(?Send)] +impl PreKeysStore for FjallProtocolStore { + async fn next_pre_key_id(&self) -> Result { + next_id_from_max(max_id_in_prefix( + &self.store.ks, + &pre_key_prefix(self.identity), + )?) + } + + async fn next_signed_pre_key_id(&self) -> Result { + next_id_from_max(max_id_in_prefix( + &self.store.ks, + &signed_pre_key_prefix(self.identity), + )?) + } + + async fn next_pq_pre_key_id(&self) -> Result { + next_id_from_max(max_id_in_prefix( + &self.store.ks, + &kyber_prefix(self.identity), + )?) + } + + async fn signed_pre_keys_count(&self) -> Result { + let prefix = signed_pre_key_prefix(self.identity); + let mut count = 0usize; + self.store.ks.prefix(prefix).try_for_each(|guard| { + guard_to_kv(guard).map_err(into_protocol_err)?; + count = count.saturating_add(1); + Ok::<(), SignalProtocolError>(()) + })?; + Ok(count) + } + + async fn kyber_pre_keys_count(&self, last_resort: bool) -> Result { + let prefix = kyber_prefix(self.identity); + let mut count = 0usize; + self.store.ks.prefix(prefix).try_for_each(|guard| { + let (_, val) = guard_to_kv(guard).map_err(into_protocol_err)?; + if decode_kyber_value(&val).is_some_and(|(is_lr, _, _)| is_lr == last_resort) { + count = count.saturating_add(1); + } + Ok::<(), SignalProtocolError>(()) + })?; + Ok(count) + } + + async fn signed_prekey_id(&self) -> Result, SignalProtocolError> { + max_id_in_prefix(&self.store.ks, &signed_pre_key_prefix(self.identity)) + .map(|opt| opt.map(SignedPreKeyId::from)) + } + + async fn last_resort_kyber_prekey_id( + &self, + ) -> Result, SignalProtocolError> { + let prefix = kyber_prefix(self.identity); + let prefix_len = prefix.len(); + let mut max: Option = None; + self.store.ks.prefix(prefix).try_for_each(|guard| { + let (key, val) = guard_to_kv(guard).map_err(into_protocol_err)?; + if let Some((true, _, _)) = decode_kyber_value(&val) + && key.len() >= prefix_len + 4 + { + let id = u32::from_be_bytes( + key[prefix_len..prefix_len + 4] + .try_into() + .map_err(into_protocol_err)?, + ); + max = Some(max.map_or(id, |m| m.max(id))); + } + Ok::<(), SignalProtocolError>(()) + })?; + Ok(max.map(KyberPreKeyId::from)) + } +} + +#[async_trait(?Send)] +impl SignedPreKeyStore for FjallProtocolStore { + async fn get_signed_pre_key( + &self, + signed_prekey_id: SignedPreKeyId, + ) -> Result { + let key = signed_pre_key_key(self.identity, u32::from(signed_prekey_id)); + let data = self + .store + .ks + .get(key) + .map_err(into_protocol_err)? + .ok_or(SignalProtocolError::InvalidSignedPreKeyId)?; + SignedPreKeyRecord::deserialize(&data) + } + + async fn save_signed_pre_key( + &mut self, + signed_prekey_id: SignedPreKeyId, + record: &SignedPreKeyRecord, + ) -> Result<(), SignalProtocolError> { + let key = signed_pre_key_key(self.identity, u32::from(signed_prekey_id)); + let serialized = record.serialize()?; + self.store + .ks + .insert(key, serialized.as_slice()) + .map_err(into_protocol_err) + } +} + +#[async_trait(?Send)] +impl KyberPreKeyStore for FjallProtocolStore { + async fn get_kyber_pre_key( + &self, + kyber_prekey_id: KyberPreKeyId, + ) -> Result { + let key = kyber_key(self.identity, u32::from(kyber_prekey_id)); + let data = self + .store + .ks + .get(key) + .map_err(into_protocol_err)? + .ok_or(SignalProtocolError::InvalidKyberPreKeyId)?; + let (_, _, record) = decode_kyber_value(&data).ok_or_else(|| { + SignalProtocolError::InvalidState("kyber", "corrupted kyber pre key record".into()) + })?; + KyberPreKeyRecord::deserialize(record) + } + + async fn save_kyber_pre_key( + &mut self, + kyber_prekey_id: KyberPreKeyId, + record: &KyberPreKeyRecord, + ) -> Result<(), SignalProtocolError> { + let key = kyber_key(self.identity, u32::from(kyber_prekey_id)); + let serialized = record.serialize()?; + let value = encode_kyber_value(&serialized, false, None); + self.store.ks.insert(key, &value).map_err(into_protocol_err) + } + + async fn mark_kyber_pre_key_used( + &mut self, + kyber_prekey_id: KyberPreKeyId, + ec_prekey_id: SignedPreKeyId, + base_key: &PublicKey, + ) -> Result<(), SignalProtocolError> { + let key = kyber_key(self.identity, u32::from(kyber_prekey_id)); + let data = self + .store + .ks + .get(&key) + .map_err(into_protocol_err)? + .ok_or(SignalProtocolError::InvalidKyberPreKeyId)?; + let (is_last_resort, _, _) = decode_kyber_value(&data).ok_or_else(|| { + SignalProtocolError::InvalidState("kyber", "corrupted kyber pre key record".into()) + })?; + + if is_last_resort { + let base_key_bytes = base_key.serialize(); + let seen_key = base_key_seen_key( + self.identity, + u32::from(kyber_prekey_id), + u32::from(ec_prekey_id), + base_key_bytes.as_ref(), + ); + if self + .store + .ks + .get(&seen_key) + .map_err(into_protocol_err)? + .is_some() + { + return Err(SignalProtocolError::InvalidMessage( + CiphertextMessageType::PreKey, + "reused base key", + )); + } + self.store + .ks + .insert(seen_key, []) + .map_err(into_protocol_err)?; + } else { + self.store.ks.remove(key).map_err(into_protocol_err)?; + } + Ok(()) + } +} + +#[async_trait(?Send)] +impl KyberPreKeyStoreExt for FjallProtocolStore { + async fn store_last_resort_kyber_pre_key( + &mut self, + kyber_prekey_id: KyberPreKeyId, + record: &KyberPreKeyRecord, + ) -> Result<(), SignalProtocolError> { + let key = kyber_key(self.identity, u32::from(kyber_prekey_id)); + let serialized = record.serialize()?; + let value = encode_kyber_value(&serialized, true, None); + self.store.ks.insert(key, &value).map_err(into_protocol_err) + } + + async fn load_last_resort_kyber_pre_keys( + &self, + ) -> Result, SignalProtocolError> { + let prefix = kyber_prefix(self.identity); + let mut result = Vec::new(); + self.store.ks.prefix(prefix).try_for_each(|guard| { + let (_, val) = guard_to_kv(guard).map_err(into_protocol_err)?; + if let Some((true, _, record)) = decode_kyber_value(&val) { + result.push(KyberPreKeyRecord::deserialize(record)?); + } + Ok::<(), SignalProtocolError>(()) + })?; + Ok(result) + } + + async fn remove_kyber_pre_key( + &mut self, + kyber_prekey_id: KyberPreKeyId, + ) -> Result<(), SignalProtocolError> { + let key = kyber_key(self.identity, u32::from(kyber_prekey_id)); + self.store.ks.remove(key).map_err(into_protocol_err) + } + + async fn mark_all_one_time_kyber_pre_keys_stale_if_necessary( + &mut self, + stale_time: chrono::DateTime, + ) -> Result<(), SignalProtocolError> { + let stale_ms = stale_time.timestamp_millis(); + let prefix = kyber_prefix(self.identity); + let mut batch = self.store.db.batch(); + self.store.ks.prefix(&prefix).try_for_each(|guard| { + let (key, val) = guard_to_kv(guard).map_err(into_protocol_err)?; + if let Some((false, None, record)) = decode_kyber_value(&val) { + let new_val = encode_kyber_value(record, false, Some(stale_ms)); + batch.insert(&self.store.ks, key.as_ref(), &new_val); + } + Ok::<(), SignalProtocolError>(()) + })?; + batch.commit().map_err(into_protocol_err) + } + + async fn delete_all_stale_one_time_kyber_pre_keys( + &mut self, + threshold: chrono::DateTime, + min_count: usize, + ) -> Result<(), SignalProtocolError> { + let threshold_ms = threshold.timestamp_millis(); + let prefix = kyber_prefix(self.identity); + + let mut total_one_time = 0usize; + self.store.ks.prefix(&prefix).try_for_each(|guard| { + let (_, val) = guard_to_kv(guard).map_err(into_protocol_err)?; + if decode_kyber_value(&val).is_some_and(|(is_lr, _, _)| !is_lr) { + total_one_time = total_one_time.saturating_add(1); + } + Ok::<(), SignalProtocolError>(()) + })?; + + if total_one_time <= min_count { + return Ok(()); + } + + let mut batch = self.store.db.batch(); + self.store.ks.prefix(&prefix).try_for_each(|guard| { + let (key, val) = guard_to_kv(guard).map_err(into_protocol_err)?; + if let Some((false, Some(stale_at), _)) = decode_kyber_value(&val) + && stale_at < threshold_ms + { + batch.remove(&self.store.ks, key.as_ref()); + } + Ok::<(), SignalProtocolError>(()) + })?; + batch.commit().map_err(into_protocol_err) + } +} + +#[async_trait(?Send)] +impl IdentityKeyStore for FjallProtocolStore { + async fn get_identity_key_pair(&self) -> Result { + let key_name = match self.identity { + IdentityType::Aci => b"identity_keypair_aci".as_slice(), + IdentityType::Pni => b"identity_keypair_pni".as_slice(), + }; + let bytes = self + .store + .ks + .get(kv_key(key_name)) + .map_err(into_protocol_err)? + .ok_or_else(|| { + SignalProtocolError::InvalidState("identity key pair", "not found in store".into()) + })?; + IdentityKeyPair::try_from(&*bytes) + } + + async fn get_local_registration_id(&self) -> Result { + let data = self + .store + .load_registration_data() + .await + .map_err(into_protocol_err)? + .ok_or_else(|| { + SignalProtocolError::InvalidState( + "failed to load registration ID", + "no registration data".into(), + ) + })?; + Ok(data.registration_id) + } + + async fn save_identity( + &mut self, + address: &ProtocolAddress, + identity_key_val: &IdentityKey, + ) -> Result { + let existing = self.get_identity(address).await?; + + let key = identity_rec_key(self.identity, address.name()); + let serialized = identity_key_val.serialize(); + self.store + .ks + .insert(key, &*serialized) + .map_err(into_protocol_err)?; + + Ok(match existing { + Some(k) if k == *identity_key_val => IdentityChange::NewOrUnchanged, + Some(_) => IdentityChange::ReplacedExisting, + None => IdentityChange::NewOrUnchanged, + }) + } + + async fn is_trusted_identity( + &self, + address: &ProtocolAddress, + identity_key_val: &IdentityKey, + _direction: Direction, + ) -> Result { + match self.get_identity(address).await? { + Some(trusted_key) if identity_key_val == &trusted_key => Ok(true), + Some(_) => { + warn!(%address, "trusting changed identity"); + Ok(true) + } + None => { + warn!(%address, "trusting new identity"); + Ok(true) + } + } + } + + async fn get_identity( + &self, + address: &ProtocolAddress, + ) -> Result, SignalProtocolError> { + let key = identity_rec_key(self.identity, address.name()); + self.store + .ks + .get(key) + .map_err(into_protocol_err)? + .map(|bytes| IdentityKey::decode(&bytes)) + .transpose() + } +} + +#[async_trait(?Send)] +impl SenderKeyStore for FjallProtocolStore { + async fn store_sender_key( + &mut self, + sender: &ProtocolAddress, + distribution_id: Uuid, + record: &SenderKeyRecord, + ) -> Result<(), SignalProtocolError> { + let key = sender_key_key( + self.identity, + sender.name(), + u32::from(sender.device_id()), + distribution_id, + ); + let serialized = record.serialize()?; + self.store + .ks + .insert(key, serialized.as_slice()) + .map_err(into_protocol_err) + } + + async fn load_sender_key( + &mut self, + sender: &ProtocolAddress, + distribution_id: Uuid, + ) -> Result, SignalProtocolError> { + let key = sender_key_key( + self.identity, + sender.name(), + u32::from(sender.device_id()), + distribution_id, + ); + self.store + .ks + .get(key) + .map_err(into_protocol_err)? + .map(|record| SenderKeyRecord::deserialize(&record)) + .transpose() + } +} + +type EmptyIter = std::iter::Empty>; + +impl ContentsStore for FjallSignalStore { + type ContentsStoreError = FjallStoreError; + type ContactsIter = EmptyIter; + type GroupsIter = EmptyIter<(GroupMasterKeyBytes, Group)>; + type MessagesIter = EmptyIter; + type StickerPacksIter = EmptyIter; + + async fn clear_profiles(&mut self) -> Result<(), FjallStoreError> { + let mut batch = self.db.batch(); + self.ks.prefix([b'R']).try_for_each(|guard| { + let (key, _) = guard.into_inner()?; + batch.remove(&self.ks, key.as_ref()); + Ok::<(), fjall::Error>(()) + })?; + batch.commit()?; + Ok(()) + } + + async fn clear_contents(&mut self) -> Result<(), FjallStoreError> { + Ok(()) + } + + async fn clear_messages(&mut self) -> Result<(), FjallStoreError> { + Ok(()) + } + + async fn clear_thread(&mut self, _thread: &Thread) -> Result<(), FjallStoreError> { + Ok(()) + } + + async fn save_message( + &self, + _thread: &Thread, + _message: Content, + ) -> Result<(), FjallStoreError> { + Ok(()) + } + + async fn delete_message( + &mut self, + _thread: &Thread, + _timestamp: u64, + ) -> Result { + Ok(false) + } + + async fn message( + &self, + _thread: &Thread, + _timestamp: u64, + ) -> Result, FjallStoreError> { + Ok(None) + } + + async fn messages( + &self, + _thread: &Thread, + _range: impl RangeBounds, + ) -> Result { + Ok(std::iter::empty()) + } + + async fn clear_contacts(&mut self) -> Result<(), FjallStoreError> { + Ok(()) + } + + async fn save_contact(&mut self, _contact: &Contact) -> Result<(), FjallStoreError> { + Ok(()) + } + + async fn contacts(&self) -> Result { + Ok(std::iter::empty()) + } + + async fn contact_by_id(&self, _id: &ServiceId) -> Result, FjallStoreError> { + Ok(None) + } + + async fn clear_groups(&mut self) -> Result<(), FjallStoreError> { + Ok(()) + } + + async fn save_group( + &self, + _master_key: GroupMasterKeyBytes, + _group: impl Into, + ) -> Result<(), FjallStoreError> { + Ok(()) + } + + async fn groups(&self) -> Result { + Ok(std::iter::empty()) + } + + async fn group( + &self, + _master_key: GroupMasterKeyBytes, + ) -> Result, FjallStoreError> { + Ok(None) + } + + async fn save_group_avatar( + &self, + _master_key: GroupMasterKeyBytes, + _avatar: &AvatarBytes, + ) -> Result<(), FjallStoreError> { + Ok(()) + } + + async fn group_avatar( + &self, + _master_key: GroupMasterKeyBytes, + ) -> Result, FjallStoreError> { + Ok(None) + } + + async fn upsert_profile_key( + &mut self, + uuid: &Uuid, + key: ProfileKey, + ) -> Result { + let k = profile_key_key(uuid); + let existed = self.ks.get(&k)?.is_some(); + self.ks.insert(k, key.bytes)?; + Ok(!existed) + } + + async fn profile_key( + &self, + service_id: &ServiceId, + ) -> Result, FjallStoreError> { + let uuid = service_id.raw_uuid(); + let k = profile_key_key(&uuid); + Ok(self + .ks + .get(k)? + .and_then(|v| match <[u8; 32]>::try_from(v.as_ref()) { + Ok(arr) => Some(ProfileKey { bytes: arr }), + Err(_) => { + warn!(%uuid, len = v.len(), "corrupted profile key (expected 32 bytes)"); + None + } + })) + } + + async fn save_profile( + &mut self, + _uuid: Uuid, + _key: ProfileKey, + _profile: Profile, + ) -> Result<(), FjallStoreError> { + Ok(()) + } + + async fn profile( + &self, + _uuid: Uuid, + _key: ProfileKey, + ) -> Result, FjallStoreError> { + Ok(None) + } + + async fn save_profile_avatar( + &mut self, + _uuid: Uuid, + _key: ProfileKey, + _profile: &AvatarBytes, + ) -> Result<(), FjallStoreError> { + Ok(()) + } + + async fn profile_avatar( + &self, + _uuid: Uuid, + _key: ProfileKey, + ) -> Result, FjallStoreError> { + Ok(None) + } + + async fn add_sticker_pack(&mut self, _pack: &StickerPack) -> Result<(), FjallStoreError> { + Ok(()) + } + + async fn sticker_pack(&self, _id: &[u8]) -> Result, FjallStoreError> { + Ok(None) + } + + async fn remove_sticker_pack(&mut self, _id: &[u8]) -> Result { + Ok(false) + } + + async fn sticker_packs(&self) -> Result { + Ok(std::iter::empty()) + } +} + +pub struct FjallSignalStoreProvider { + store: FjallSignalStore, +} + +impl FjallSignalStoreProvider { + pub fn new(db: Database, ks: Keyspace) -> Self { + Self { + store: FjallSignalStore::new(db, ks), + } + } +} + +#[async_trait::async_trait] +impl crate::SignalStoreProvider for FjallSignalStoreProvider { + async fn is_signal_linked(&self) -> bool { + self.store.is_linked().unwrap_or(false) + } + + async fn clear_signal_data(&self) -> Result<(), SignalError> { + self.store + .clear_all() + .map_err(|e| SignalError::Store(e.to_string())) + } + + async fn link_signal_device( + &self, + device_name: DeviceName, + shutdown: tokio_util::sync::CancellationToken, + link_cancel: tokio_util::sync::CancellationToken, + linking_flag: Arc, + ) -> Result { + SignalClient::link_device_with_store( + self.store.clone(), + device_name, + shutdown, + link_cancel, + linking_flag, + ) + .await + } + + async fn load_signal_client( + &self, + shutdown: tokio_util::sync::CancellationToken, + ) -> Option { + SignalClient::from_store(self.store.clone(), shutdown).await + } +} diff --git a/crates/tranquil-signal/src/lib.rs b/crates/tranquil-signal/src/lib.rs index 205ef25..7388b98 100644 --- a/crates/tranquil-signal/src/lib.rs +++ b/crates/tranquil-signal/src/lib.rs @@ -1,6 +1,9 @@ mod client; pub mod store; +#[cfg(feature = "fjall-store")] +pub mod fjall_store; + #[cfg(test)] mod tests; @@ -10,3 +13,59 @@ pub use client::{ }; pub use presage; pub use store::PgSignalStore; + +#[async_trait::async_trait] +pub trait SignalStoreProvider: Send + Sync { + async fn is_signal_linked(&self) -> bool; + async fn clear_signal_data(&self) -> Result<(), SignalError>; + async fn link_signal_device( + &self, + device_name: DeviceName, + shutdown: tokio_util::sync::CancellationToken, + link_cancel: tokio_util::sync::CancellationToken, + linking_flag: std::sync::Arc, + ) -> Result; + async fn load_signal_client( + &self, + shutdown: tokio_util::sync::CancellationToken, + ) -> Option; +} + +pub struct PgSignalStoreProvider { + pub pool: sqlx::PgPool, +} + +#[async_trait::async_trait] +impl SignalStoreProvider for PgSignalStoreProvider { + async fn is_signal_linked(&self) -> bool { + PgSignalStore::new(self.pool.clone()) + .is_linked() + .await + .unwrap_or(false) + } + + async fn clear_signal_data(&self) -> Result<(), SignalError> { + PgSignalStore::new(self.pool.clone()) + .clear_all() + .await + .map_err(SignalError::from) + } + + async fn link_signal_device( + &self, + device_name: DeviceName, + shutdown: tokio_util::sync::CancellationToken, + link_cancel: tokio_util::sync::CancellationToken, + linking_flag: std::sync::Arc, + ) -> Result { + SignalClient::link_device(&self.pool, device_name, shutdown, link_cancel, linking_flag) + .await + } + + async fn load_signal_client( + &self, + shutdown: tokio_util::sync::CancellationToken, + ) -> Option { + SignalClient::from_pool(&self.pool, shutdown).await + } +} diff --git a/crates/tranquil-store/src/lib.rs b/crates/tranquil-store/src/lib.rs index fd9a379..d29c4c7 100644 --- a/crates/tranquil-store/src/lib.rs +++ b/crates/tranquil-store/src/lib.rs @@ -1,6 +1,7 @@ pub mod blockstore; pub mod eventlog; pub mod fsync_order; +#[cfg(any(test, feature = "test-harness"))] mod harness; mod io; pub mod metastore; @@ -19,4 +20,5 @@ pub use record::{ FILE_MAGIC, FORMAT_VERSION, HEADER_SIZE, MAX_RECORD_PAYLOAD, RECORD_OVERHEAD, ReadRecord, RecordReader, RecordWriter, }; +#[cfg(any(test, feature = "test-harness"))] pub use sim::{FaultConfig, OpRecord, SimulatedIO}; diff --git a/crates/tranquil-store/src/metastore/client.rs b/crates/tranquil-store/src/metastore/client.rs index 0367b75..504ae53 100644 --- a/crates/tranquil-store/src/metastore/client.rs +++ b/crates/tranquil-store/src/metastore/client.rs @@ -9,12 +9,13 @@ use tranquil_db_traits::{ ApplyCommitResult, Backlink, BrokenGenesisCommit, CommitEventData, CommsChannel, CommsType, CompletePasskeySetupInput, CreateAccountError, CreateDelegatedAccountInput, CreatePasskeyAccountInput, CreatePasswordAccountInput, CreatePasswordAccountResult, - CreateSsoAccountInput, DbError, DeletionRequest, DidWebOverrides, EventBlocksCids, ImportBlock, - ImportRecord, ImportRepoError, InviteCodeError, InviteCodeInfo, InviteCodeRow, - InviteCodeSortOrder, InviteCodeUse, MigrationReactivationError, MigrationReactivationInput, - NotificationHistoryRow, NotificationPrefs, OAuthTokenWithUser, PasswordResetResult, - QueuedComms, ReactivatedAccountInfo, RecoverPasskeyAccountInput, RecoverPasskeyAccountResult, - RepoAccountInfo, RepoInfo, RepoListItem, RepoWithoutRev, ReservedSigningKey, + CreateSsoAccountInput, DbError, DeletionRequest, DeletionRequestWithToken, DidWebOverrides, + EventBlocksCids, ImportBlock, ImportRecord, ImportRepoError, InviteCodeError, InviteCodeInfo, + InviteCodeRow, InviteCodeSortOrder, InviteCodeUse, MigrationReactivationError, + MigrationReactivationInput, NotificationHistoryRow, NotificationPrefs, OAuthTokenWithUser, + PasswordResetResult, PlcTokenInfo, QueuedComms, ReactivatedAccountInfo, + RecoverPasskeyAccountInput, RecoverPasskeyAccountResult, RepoAccountInfo, RepoInfo, + RepoListItem, RepoWithoutRev, ReservedSigningKey, ReservedSigningKeyFull, ScheduledDeletionAccount, ScopePreference, SequenceNumber, SequencedEvent, StoredBackupCode, StoredPasskey, TokenFamilyId, TotpRecord, TotpRecordState, User2faStatus, UserAuthInfo, UserCommsPrefs, UserConfirmSignup, UserDidWebInfo, UserEmailInfo, UserForDeletion, @@ -2459,6 +2460,114 @@ impl tranquil_db_traits::InfraRepository for MetastoreCl ))?; recv(rx).await } + + async fn get_deletion_request_by_did( + &self, + did: &Did, + ) -> Result, DbError> { + let (tx, rx) = oneshot::channel(); + self.pool.send(MetastoreRequest::Infra( + InfraRequest::GetDeletionRequestByDid { + did: did.clone(), + tx, + }, + ))?; + recv(rx).await + } + + async fn get_latest_comms_for_user( + &self, + user_id: Uuid, + comms_type: CommsType, + limit: i64, + ) -> Result, DbError> { + let (tx, rx) = oneshot::channel(); + self.pool.send(MetastoreRequest::Infra( + InfraRequest::GetLatestCommsForUser { + user_id, + comms_type, + limit, + tx, + }, + ))?; + recv(rx).await + } + + async fn count_comms_by_type( + &self, + user_id: Uuid, + comms_type: CommsType, + ) -> Result { + let (tx, rx) = oneshot::channel(); + self.pool + .send(MetastoreRequest::Infra(InfraRequest::CountCommsByType { + user_id, + comms_type, + tx, + }))?; + recv(rx).await + } + + async fn delete_comms_by_type_for_user( + &self, + user_id: Uuid, + comms_type: CommsType, + ) -> Result { + let (tx, rx) = oneshot::channel(); + self.pool.send(MetastoreRequest::Infra( + InfraRequest::DeleteCommsByTypeForUser { + user_id, + comms_type, + tx, + }, + ))?; + recv(rx).await + } + + async fn expire_deletion_request(&self, token: &str) -> Result<(), DbError> { + let (tx, rx) = oneshot::channel(); + self.pool.send(MetastoreRequest::Infra( + InfraRequest::ExpireDeletionRequest { + token: token.to_owned(), + tx, + }, + ))?; + recv(rx).await + } + + async fn get_reserved_signing_key_full( + &self, + public_key_did_key: &str, + ) -> Result, DbError> { + let (tx, rx) = oneshot::channel(); + self.pool.send(MetastoreRequest::Infra( + InfraRequest::GetReservedSigningKeyFull { + public_key_did_key: public_key_did_key.to_owned(), + tx, + }, + ))?; + recv(rx).await + } + + async fn get_plc_tokens_by_did(&self, did: &Did) -> Result, DbError> { + let (tx, rx) = oneshot::channel(); + self.pool + .send(MetastoreRequest::Infra(InfraRequest::GetPlcTokensByDid { + did: did.clone(), + tx, + }))?; + recv(rx).await + } + + async fn count_plc_tokens_by_did(&self, did: &Did) -> Result { + let (tx, rx) = oneshot::channel(); + self.pool + .send(MetastoreRequest::Infra(InfraRequest::CountPlcTokensByDid { + did: did.clone(), + tx, + }))?; + recv(rx).await + } } #[async_trait] @@ -3227,6 +3336,19 @@ impl tranquil_db_traits::OAuthRepository for MetastoreCl ))?; recv(rx).await } + + async fn get_2fa_challenge_code( + &self, + request_uri: &RequestId, + ) -> Result, DbError> { + let (tx, rx) = oneshot::channel(); + self.pool + .send(MetastoreRequest::OAuth(OAuthRequest::Get2faChallengeCode { + request_uri: request_uri.clone(), + tx, + }))?; + recv(rx).await + } } #[async_trait] @@ -3690,6 +3812,17 @@ impl tranquil_db_traits::UserRepository for MetastoreCli recv(rx).await } + async fn set_admin_status(&self, did: &Did, is_admin: bool) -> Result<(), DbError> { + let (tx, rx) = oneshot::channel(); + self.pool + .send(MetastoreRequest::User(UserRequest::SetAdminStatus { + did: did.clone(), + is_admin, + tx, + }))?; + recv(rx).await + } + async fn get_notification_prefs( &self, did: &Did, @@ -4889,4 +5022,54 @@ impl tranquil_db_traits::UserRepository for MetastoreCli }))?; recv(rx).await } + + async fn get_password_reset_info( + &self, + email: &str, + ) -> Result, DbError> { + let (tx, rx) = oneshot::channel(); + self.pool + .send(MetastoreRequest::User(UserRequest::GetPasswordResetInfo { + email: email.to_owned(), + tx, + }))?; + recv(rx).await + } + + async fn enable_totp_verified( + &self, + did: &Did, + encrypted_secret: &[u8], + ) -> Result<(), DbError> { + let (tx, rx) = oneshot::channel(); + self.pool + .send(MetastoreRequest::User(UserRequest::EnableTotpVerified { + did: did.clone(), + encrypted_secret: encrypted_secret.to_vec(), + tx, + }))?; + recv(rx).await + } + + async fn set_two_factor_enabled(&self, did: &Did, enabled: bool) -> Result<(), DbError> { + let (tx, rx) = oneshot::channel(); + self.pool + .send(MetastoreRequest::User(UserRequest::SetTwoFactorEnabled { + did: did.clone(), + enabled, + tx, + }))?; + recv(rx).await + } + + async fn expire_password_reset_code(&self, email: &str) -> Result<(), DbError> { + let (tx, rx) = oneshot::channel(); + self.pool.send(MetastoreRequest::User( + UserRequest::ExpirePasswordResetCode { + email: email.to_owned(), + tx, + }, + ))?; + recv(rx).await + } } diff --git a/crates/tranquil-store/src/metastore/handler.rs b/crates/tranquil-store/src/metastore/handler.rs index a5a7863..e7b6bed 100644 --- a/crates/tranquil-store/src/metastore/handler.rs +++ b/crates/tranquil-store/src/metastore/handler.rs @@ -10,22 +10,23 @@ use tranquil_db_traits::{ ApplyCommitResult, Backlink, BrokenGenesisCommit, CommitEventData, CommsChannel, CommsType, CompletePasskeySetupInput, CreateAccountError, CreateDelegatedAccountInput, CreatePasskeyAccountInput, CreatePasswordAccountInput, CreatePasswordAccountResult, - CreateSsoAccountInput, DbError, DelegationActionType, DeletionRequest, DidWebOverrides, - EventBlocksCids, ImportBlock, ImportRecord, ImportRepoError, InviteCodeError, InviteCodeInfo, - InviteCodeRow, InviteCodeSortOrder, InviteCodeUse, MigrationReactivationError, - MigrationReactivationInput, NotificationHistoryRow, NotificationPrefs, OAuthTokenWithUser, - PasswordResetResult, QueuedComms, ReactivatedAccountInfo, RecoverPasskeyAccountInput, - RecoverPasskeyAccountResult, RefreshSessionResult, ReservedSigningKey, - ScheduledDeletionAccount, ScopePreference, SequenceNumber, SequencedEvent, SessionId, - StoredBackupCode, StoredPasskey, TokenFamilyId, TotpRecord, TotpRecordState, User2faStatus, - UserAuthInfo, UserCommsPrefs, UserConfirmSignup, UserDidWebInfo, UserEmailInfo, - UserForDeletion, UserForDidDoc, UserForDidDocBuild, UserForPasskeyRecovery, - UserForPasskeySetup, UserForRecovery, UserForVerification, UserIdAndHandle, - UserIdAndPasswordHash, UserIdHandleEmail, UserInfoForAuth, UserKeyInfo, UserKeyWithId, - UserLegacyLoginPref, UserLoginCheck, UserLoginFull, UserLoginInfo, - UserNeedingRecordBlobsBackfill, UserPasswordInfo, UserResendVerification, UserResetCodeInfo, - UserRow, UserSessionInfo, UserStatus, UserVerificationInfo, UserWithKey, UserWithoutBlocks, - ValidatedInviteCode, WebauthnChallengeType, + CreateSsoAccountInput, DbError, DelegationActionType, DeletionRequest, + DeletionRequestWithToken, DidWebOverrides, EventBlocksCids, ImportBlock, ImportRecord, + ImportRepoError, InviteCodeError, InviteCodeInfo, InviteCodeRow, InviteCodeSortOrder, + InviteCodeUse, MigrationReactivationError, MigrationReactivationInput, NotificationHistoryRow, + NotificationPrefs, OAuthTokenWithUser, PasswordResetResult, PlcTokenInfo, QueuedComms, + ReactivatedAccountInfo, RecoverPasskeyAccountInput, RecoverPasskeyAccountResult, + RefreshSessionResult, ReservedSigningKey, ReservedSigningKeyFull, ScheduledDeletionAccount, + ScopePreference, SequenceNumber, SequencedEvent, SessionId, StoredBackupCode, StoredPasskey, + TokenFamilyId, TotpRecord, TotpRecordState, User2faStatus, UserAuthInfo, UserCommsPrefs, + UserConfirmSignup, UserDidWebInfo, UserEmailInfo, UserForDeletion, UserForDidDoc, + UserForDidDocBuild, UserForPasskeyRecovery, UserForPasskeySetup, UserForRecovery, + UserForVerification, UserIdAndHandle, UserIdAndPasswordHash, UserIdHandleEmail, + UserInfoForAuth, UserKeyInfo, UserKeyWithId, UserLegacyLoginPref, UserLoginCheck, + UserLoginFull, UserLoginInfo, UserNeedingRecordBlobsBackfill, UserPasswordInfo, + UserResendVerification, UserResetCodeInfo, UserRow, UserSessionInfo, UserStatus, + UserVerificationInfo, UserWithKey, UserWithoutBlocks, ValidatedInviteCode, + WebauthnChallengeType, }; use tranquil_oauth::{AuthorizedClientData, DeviceData, RequestData, TokenData}; use tranquil_types::{ @@ -1127,6 +1128,11 @@ pub enum UserRequest { password_hash: String, tx: Tx, }, + SetAdminStatus { + did: Did, + is_admin: bool, + tx: Tx<()>, + }, GetNotificationPrefs { did: Did, tx: Tx>, @@ -1559,6 +1565,24 @@ pub enum UserRequest { input: RecoverPasskeyAccountInput, tx: Tx, }, + GetPasswordResetInfo { + email: String, + tx: Tx>, + }, + EnableTotpVerified { + did: Did, + encrypted_secret: Vec, + tx: Tx<()>, + }, + SetTwoFactorEnabled { + did: Did, + enabled: bool, + tx: Tx<()>, + }, + ExpirePasswordResetCode { + email: String, + tx: Tx<()>, + }, } impl UserRequest { @@ -1585,6 +1609,7 @@ impl UserRequest { | Self::AdminUpdateEmail { did, .. } | Self::AdminUpdateHandle { did, .. } | Self::AdminUpdatePassword { did, .. } + | Self::SetAdminStatus { did, .. } | Self::GetNotificationPrefs { did, .. } | Self::GetIdHandleEmailByDid { did, .. } | Self::UpdatePreferredCommsChannel { did, .. } @@ -1662,7 +1687,9 @@ impl UserRequest { input: RecoverPasskeyAccountInput { did, .. }, .. } - | Self::SetRecoveryToken { did, .. } => did_to_routing(did.as_str()), + | Self::SetRecoveryToken { did, .. } + | Self::EnableTotpVerified { did, .. } + | Self::SetTwoFactorEnabled { did, .. } => did_to_routing(did.as_str()), Self::GetCommsPrefs { user_id, .. } | Self::GetUserKeyById { user_id, .. } @@ -1729,7 +1756,9 @@ impl UserRequest { | Self::GetUserForPasskeyRecovery { .. } | Self::GetAccountsScheduledForDeletion { .. } | Self::CleanupExpiredHandleReservations { .. } - | Self::CheckAndConsumeInviteCode { .. } => Routing::Global, + | Self::CheckAndConsumeInviteCode { .. } + | Self::GetPasswordResetInfo { .. } + | Self::ExpirePasswordResetCode { .. } => Routing::Global, } } } @@ -1969,6 +1998,42 @@ pub enum InfraRequest { user_ids: Vec, tx: Tx>, }, + GetDeletionRequestByDid { + did: Did, + tx: Tx>, + }, + GetLatestCommsForUser { + user_id: Uuid, + comms_type: CommsType, + limit: i64, + tx: Tx>, + }, + CountCommsByType { + user_id: Uuid, + comms_type: CommsType, + tx: Tx, + }, + DeleteCommsByTypeForUser { + user_id: Uuid, + comms_type: CommsType, + tx: Tx, + }, + ExpireDeletionRequest { + token: String, + tx: Tx<()>, + }, + GetReservedSigningKeyFull { + public_key_did_key: String, + tx: Tx>, + }, + GetPlcTokensByDid { + did: Did, + tx: Tx>, + }, + CountPlcTokensByDid { + did: Did, + tx: Tx, + }, } impl InfraRequest { @@ -1986,7 +2051,10 @@ impl InfraRequest { | Self::GetInvitesCreatedByUser { user_id, .. } | Self::GetInviteCodeUsedByUser { user_id, .. } | Self::DeleteInviteCodeUsesByUser { user_id, .. } - | Self::DeleteInviteCodesByUser { user_id, .. } => { + | Self::DeleteInviteCodesByUser { user_id, .. } + | Self::GetLatestCommsForUser { user_id, .. } + | Self::CountCommsByType { user_id, .. } + | Self::DeleteCommsByTypeForUser { user_id, .. } => { uuid_to_routing(user_hashes, user_id) } Self::CreateInviteCodesBatch { @@ -2003,7 +2071,10 @@ impl InfraRequest { } Self::DeleteDeletionRequestsByDid { did, .. } | Self::CreateDeletionRequest { did, .. } - | Self::GetAdminAccountInfoByDid { did, .. } => did_to_routing(did.as_str()), + | Self::GetAdminAccountInfoByDid { did, .. } + | Self::GetDeletionRequestByDid { did, .. } + | Self::GetPlcTokensByDid { did, .. } + | Self::CountPlcTokensByDid { did, .. } => did_to_routing(did.as_str()), Self::GetBlobStorageKeyByCid { cid, .. } | Self::DeleteBlobByCid { cid, .. } => { cid_to_routing(cid) } @@ -2279,6 +2350,10 @@ pub enum OAuthRequest { except_token_id: TokenId, tx: Tx, }, + Get2faChallengeCode { + request_uri: RequestId, + tx: Tx>, + }, } impl OAuthRequest { @@ -2342,7 +2417,8 @@ impl OAuthRequest { | Self::RevokeDeviceTrust { .. } | Self::UpdateDeviceFriendlyName { .. } | Self::TrustDevice { .. } - | Self::ExtendDeviceTrust { .. } => Routing::Global, + | Self::ExtendDeviceTrust { .. } + | Self::Get2faChallengeCode { .. } => Routing::Global, } } } @@ -4266,6 +4342,106 @@ fn dispatch_infra(state: &HandlerState, req: InfraRequest) { .map_err(metastore_to_db); let _ = tx.send(result); } + InfraRequest::GetDeletionRequestByDid { did, tx } => { + let result = state + .metastore + .infra_ops() + .get_deletion_request_by_did(&did) + .map_err(metastore_to_db); + let _ = tx.send(result); + } + InfraRequest::GetLatestCommsForUser { + user_id, + comms_type, + limit, + tx, + } => { + let result = state + .metastore + .infra_ops() + .get_latest_comms_for_user(user_id, comms_type, limit) + .map_err(metastore_to_db); + let _ = tx.send(result); + } + InfraRequest::CountCommsByType { + user_id, + comms_type, + tx, + } => { + let result = state + .metastore + .infra_ops() + .count_comms_by_type(user_id, comms_type) + .map_err(metastore_to_db); + let _ = tx.send(result); + } + InfraRequest::DeleteCommsByTypeForUser { + user_id, + comms_type, + tx, + } => { + let result = state + .metastore + .infra_ops() + .delete_comms_by_type_for_user(user_id, comms_type) + .map_err(metastore_to_db); + let _ = tx.send(result); + } + InfraRequest::ExpireDeletionRequest { token, tx } => { + let result = state + .metastore + .infra_ops() + .expire_deletion_request(&token) + .map_err(metastore_to_db); + let _ = tx.send(result); + } + InfraRequest::GetReservedSigningKeyFull { + public_key_did_key, + tx, + } => { + let result = state + .metastore + .infra_ops() + .get_reserved_signing_key_full(&public_key_did_key) + .map_err(metastore_to_db); + let _ = tx.send(result); + } + InfraRequest::GetPlcTokensByDid { did, tx } => { + let result = (|| { + let user_id = state + .metastore + .user_ops() + .get_id_by_did(&did) + .map_err(metastore_to_db)?; + match user_id { + Some(uid) => state + .metastore + .infra_ops() + .get_plc_tokens_for_user(uid) + .map_err(metastore_to_db), + None => Ok(Vec::new()), + } + })(); + let _ = tx.send(result); + } + InfraRequest::CountPlcTokensByDid { did, tx } => { + let result = (|| { + let user_id = state + .metastore + .user_ops() + .get_id_by_did(&did) + .map_err(metastore_to_db)?; + match user_id { + Some(uid) => state + .metastore + .infra_ops() + .count_plc_tokens_for_user(uid) + .map_err(metastore_to_db), + None => Ok(0), + } + })(); + let _ = tx.send(result); + } } } @@ -4818,6 +4994,14 @@ fn dispatch_oauth(state: &HandlerState, req: OAuthRequest) { .map_err(metastore_to_db); let _ = tx.send(result); } + OAuthRequest::Get2faChallengeCode { request_uri, tx } => { + let result = state + .metastore + .oauth_ops() + .get_2fa_challenge_code(&request_uri) + .map_err(metastore_to_db); + let _ = tx.send(result); + } } } @@ -5048,6 +5232,12 @@ fn dispatch_user(state: &HandlerState, req: UserRequest) { .map_err(metastore_to_db), ); } + UserRequest::SetAdminStatus { did, is_admin, tx } => { + let _ = tx.send( + user.set_admin_status(&did, is_admin) + .map_err(metastore_to_db), + ); + } UserRequest::GetNotificationPrefs { did, tx } => { let _ = tx.send(user.get_notification_prefs(&did).map_err(metastore_to_db)); } @@ -5567,16 +5757,44 @@ fn dispatch_user(state: &HandlerState, req: UserRequest) { ); } UserRequest::DeleteAccountWithFirehose { user_id, did, tx } => { - let _ = tx.send( - user.delete_account_with_firehose(user_id, &did) - .map_err(metastore_to_db), - ); + let result = user + .delete_account_complete(user_id, &did) + .map_err(metastore_to_db) + .and_then(|()| { + state + .event_ops + .insert_account_event(&did, AccountStatus::Deleted) + }); + let _ = tx.send(result.map(|seq| seq.as_i64())); } UserRequest::CreatePasswordAccount { input, tx } => { let _ = tx.send(user.create_password_account(&input)); } UserRequest::CreateDelegatedAccount { input, tx } => { - let _ = tx.send(user.create_delegated_account(&input)); + let result = user.create_delegated_account(&input).and_then(|account| { + let scope = + tranquil_db_traits::DbScope::new(&input.controller_scopes).map_err(|e| { + tranquil_db_traits::CreateAccountError::Database(format!( + "invalid delegation scope: {e}" + )) + })?; + state + .metastore + .delegation_ops() + .create_delegation( + &input.did, + &input.controller_did, + &scope, + &input.controller_did, + ) + .map_err(|e| { + tranquil_db_traits::CreateAccountError::Database(format!( + "delegation grant creation failed: {e}" + )) + })?; + Ok(account) + }); + let _ = tx.send(result); } UserRequest::CreatePasskeyAccount { input, tx } => { let _ = tx.send(user.create_passkey_account(&input)); @@ -5635,6 +5853,34 @@ fn dispatch_user(state: &HandlerState, req: UserRequest) { .map_err(metastore_to_db), ); } + UserRequest::GetPasswordResetInfo { email, tx } => { + let _ = tx.send( + user.get_password_reset_info(&email) + .map_err(metastore_to_db), + ); + } + UserRequest::EnableTotpVerified { + did, + encrypted_secret, + tx, + } => { + let _ = tx.send( + user.enable_totp_verified(&did, &encrypted_secret) + .map_err(metastore_to_db), + ); + } + UserRequest::SetTwoFactorEnabled { did, enabled, tx } => { + let _ = tx.send( + user.set_two_factor_enabled(&did, enabled) + .map_err(metastore_to_db), + ); + } + UserRequest::ExpirePasswordResetCode { email, tx } => { + let _ = tx.send( + user.expire_password_reset_code(&email) + .map_err(metastore_to_db), + ); + } } } diff --git a/crates/tranquil-store/src/metastore/infra_ops.rs b/crates/tranquil-store/src/metastore/infra_ops.rs index e9040cb..cc02617 100644 --- a/crates/tranquil-store/src/metastore/infra_ops.rs +++ b/crates/tranquil-store/src/metastore/infra_ops.rs @@ -22,9 +22,10 @@ use super::user_hash::UserHashMap; use super::users::UserValue; use tranquil_db_traits::{ - AdminAccountInfo, CommsChannel, CommsStatus, CommsType, DeletionRequest, InviteCodeError, - InviteCodeInfo, InviteCodeRow, InviteCodeSortOrder, InviteCodeState, InviteCodeUse, - NotificationHistoryRow, QueuedComms, ReservedSigningKey, ValidatedInviteCode, + AdminAccountInfo, CommsChannel, CommsStatus, CommsType, DeletionRequest, + DeletionRequestWithToken, InviteCodeError, InviteCodeInfo, InviteCodeRow, InviteCodeSortOrder, + InviteCodeState, InviteCodeUse, NotificationHistoryRow, PlcTokenInfo, QueuedComms, + ReservedSigningKey, ReservedSigningKeyFull, ValidatedInviteCode, }; use tranquil_types::{CidLink, Did, Handle}; @@ -713,6 +714,7 @@ impl InfraOps { used: false, created_at_ms: now_ms, expires_at_ms: expires_at.timestamp_millis(), + used_at_ms: None, }; let primary_key = signing_key_key(public_key_did_key); @@ -771,6 +773,7 @@ impl InfraOps { ))?; val.used = true; + val.used_at_ms = Some(Utc::now().timestamp_millis()); self.infra .insert(primary_key.as_slice(), val.serialize()) .map_err(MetastoreError::Fjall) @@ -906,9 +909,14 @@ impl InfraOps { let mut reader = super::encoding::KeyReader::new(&key_bytes); reader.tag(); reader.bytes(); - let name = reader + let raw_name = reader .string() .ok_or(MetastoreError::CorruptData("corrupt account pref key"))?; + let name = raw_name + .split('\x00') + .next() + .unwrap_or(&raw_name) + .to_owned(); let value: serde_json::Value = serde_json::from_slice(&val_bytes) .map_err(|_| MetastoreError::CorruptData("corrupt account pref json"))?; acc.push((name, value)); @@ -939,8 +947,15 @@ impl InfraOps { Ok::<(), MetastoreError>(()) })?; + let mut counts: std::collections::HashMap<&str, u32> = std::collections::HashMap::new(); preferences.iter().try_for_each(|(name, value)| { - let key = account_pref_key(user_id, name); + let idx = counts.entry(name.as_str()).or_insert(0); + let indexed_name = match *idx { + 0 => name.clone(), + n => format!("{}\x00{}", name, n), + }; + *idx += 1; + let key = account_pref_key(user_id, &indexed_name); let bytes = serde_json::to_vec(value) .map_err(|_| MetastoreError::InvalidInput("invalid json for account preference"))?; batch.insert(&self.infra, key.as_slice(), bytes); @@ -1206,4 +1221,193 @@ impl InfraOps { }) .collect() } + + pub fn get_deletion_request_by_did( + &self, + did: &Did, + ) -> Result, MetastoreError> { + let did_key = deletion_by_did_key(did.as_str()); + let token = match self + .infra + .get(did_key.as_slice()) + .map_err(MetastoreError::Fjall)? + { + Some(raw) => String::from_utf8(raw.to_vec()) + .map_err(|_| MetastoreError::CorruptData("deletion by_did not valid utf8"))?, + None => return Ok(None), + }; + + let primary_key = deletion_request_key(&token); + let val: Option = point_lookup( + &self.infra, + primary_key.as_slice(), + DeletionRequestValue::deserialize, + "corrupt deletion request", + )?; + Ok(val.map(|v| DeletionRequestWithToken { + token, + did: Did::new(v.did).expect("valid DID in database"), + expires_at: DateTime::from_timestamp_millis(v.expires_at_ms).unwrap_or_default(), + })) + } + + pub fn get_latest_comms_for_user( + &self, + user_id: Uuid, + comms_type: CommsType, + limit: i64, + ) -> Result, MetastoreError> { + let target_type = comms_type_to_u8(comms_type); + let limit = usize::try_from(limit).unwrap_or(0); + let prefix = comms_queue_prefix(); + + let mut results: Vec = self + .infra + .prefix(prefix.as_slice()) + .map(|guard| -> Result, MetastoreError> { + let (_, val_bytes) = guard.into_inner().map_err(MetastoreError::Fjall)?; + let val = QueuedCommsValue::deserialize(&val_bytes) + .ok_or(MetastoreError::CorruptData("corrupt comms queue entry"))?; + let matches_user = val.user_id == Some(user_id); + let matches_type = val.comms_type == target_type; + match matches_user && matches_type { + true => Ok(Some(self.value_to_queued_comms(&val)?)), + false => Ok(None), + } + }) + .filter_map(Result::transpose) + .collect::, _>>()?; + + results.sort_by_key(|r| std::cmp::Reverse(r.created_at)); + results.truncate(limit); + Ok(results) + } + + pub fn count_comms_by_type( + &self, + user_id: Uuid, + comms_type: CommsType, + ) -> Result { + let target_type = comms_type_to_u8(comms_type); + let prefix = comms_queue_prefix(); + + let count = self + .infra + .prefix(prefix.as_slice()) + .try_fold(0i64, |acc, guard| { + let (_, val_bytes) = guard.into_inner().map_err(MetastoreError::Fjall)?; + let val = QueuedCommsValue::deserialize(&val_bytes) + .ok_or(MetastoreError::CorruptData("corrupt comms queue entry"))?; + let matches = val.user_id == Some(user_id) && val.comms_type == target_type; + Ok::(acc + i64::from(matches)) + })?; + + Ok(count) + } + + pub fn delete_comms_by_type_for_user( + &self, + user_id: Uuid, + comms_type: CommsType, + ) -> Result { + let target_type = comms_type_to_u8(comms_type); + let prefix = comms_queue_prefix(); + + let keys_to_delete: Vec> = self + .infra + .prefix(prefix.as_slice()) + .map(|guard| -> Result>, MetastoreError> { + let (key_bytes, val_bytes) = guard.into_inner().map_err(MetastoreError::Fjall)?; + let val = QueuedCommsValue::deserialize(&val_bytes) + .ok_or(MetastoreError::CorruptData("corrupt comms queue entry"))?; + let matches = val.user_id == Some(user_id) && val.comms_type == target_type; + match matches { + true => Ok(Some(key_bytes.to_vec())), + false => Ok(None), + } + }) + .filter_map(Result::transpose) + .collect::, _>>()?; + + let count = u64::try_from(keys_to_delete.len()).unwrap_or(u64::MAX); + let mut batch = self.db.batch(); + keys_to_delete + .iter() + .for_each(|k| batch.remove(&self.infra, k.as_slice())); + batch.commit().map_err(MetastoreError::Fjall)?; + Ok(count) + } + + pub fn expire_deletion_request(&self, token: &str) -> Result<(), MetastoreError> { + let key = deletion_request_key(token); + let mut val: DeletionRequestValue = point_lookup( + &self.infra, + key.as_slice(), + DeletionRequestValue::deserialize, + "corrupt deletion request", + )? + .ok_or(MetastoreError::InvalidInput("deletion request not found"))?; + + val.expires_at_ms = Utc::now().timestamp_millis() - 3_600_000; + self.infra + .insert(key.as_slice(), val.serialize()) + .map_err(MetastoreError::Fjall) + } + + pub fn get_reserved_signing_key_full( + &self, + public_key_did_key: &str, + ) -> Result, MetastoreError> { + let key = signing_key_key(public_key_did_key); + let val: Option = point_lookup( + &self.infra, + key.as_slice(), + SigningKeyValue::deserialize, + "corrupt signing key", + )?; + Ok(val.map(|v| ReservedSigningKeyFull { + id: v.id, + did: v.did.and_then(|d| Did::new(d).ok()), + public_key_did_key: v.public_key_did_key, + private_key_bytes: v.private_key_bytes, + expires_at: DateTime::from_timestamp_millis(v.expires_at_ms).unwrap_or_default(), + used_at: v.used_at_ms.and_then(DateTime::from_timestamp_millis), + })) + } + + pub fn get_plc_tokens_for_user( + &self, + user_id: Uuid, + ) -> Result, MetastoreError> { + let prefix = plc_token_prefix(user_id); + self.infra + .prefix(prefix.as_slice()) + .map(|guard| -> Result { + let (key_bytes, val_bytes) = guard.into_inner().map_err(MetastoreError::Fjall)?; + let arr: [u8; 8] = val_bytes + .as_ref() + .try_into() + .map_err(|_| MetastoreError::CorruptData("plc token expiry not 8 bytes"))?; + let expires_at = + DateTime::from_timestamp_millis(i64::from_be_bytes(arr)).unwrap_or_default(); + let mut reader = super::encoding::KeyReader::new(&key_bytes); + let _tag = reader.tag(); + let _user_id = reader.bytes(); + let token = reader + .string() + .ok_or(MetastoreError::CorruptData("plc token key missing token"))?; + Ok(PlcTokenInfo { token, expires_at }) + }) + .collect() + } + + pub fn count_plc_tokens_for_user(&self, user_id: Uuid) -> Result { + let prefix = plc_token_prefix(user_id); + Ok(self + .infra + .prefix(prefix.as_slice()) + .count() + .try_into() + .unwrap_or(0)) + } } diff --git a/crates/tranquil-store/src/metastore/infra_schema.rs b/crates/tranquil-store/src/metastore/infra_schema.rs index ba2f2d7..47fc794 100644 --- a/crates/tranquil-store/src/metastore/infra_schema.rs +++ b/crates/tranquil-store/src/metastore/infra_schema.rs @@ -7,7 +7,8 @@ use super::keys::KeyTag; const COMMS_SCHEMA_VERSION: u8 = 1; const INVITE_CODE_SCHEMA_VERSION: u8 = 1; const INVITE_USE_SCHEMA_VERSION: u8 = 1; -const SIGNING_KEY_SCHEMA_VERSION: u8 = 1; +const SIGNING_KEY_SCHEMA_V1: u8 = 1; +const SIGNING_KEY_SCHEMA_V2: u8 = 2; const DELETION_REQUEST_SCHEMA_VERSION: u8 = 1; const REPORT_SCHEMA_VERSION: u8 = 1; const NOTIFICATION_HISTORY_SCHEMA_VERSION: u8 = 1; @@ -104,6 +105,17 @@ impl InviteCodeUseValue { } } +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +struct SigningKeyValueV1 { + id: uuid::Uuid, + did: Option, + public_key_did_key: String, + private_key_bytes: Vec, + used: bool, + created_at_ms: i64, + expires_at_ms: i64, +} + #[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] pub struct SigningKeyValue { pub id: uuid::Uuid, @@ -113,6 +125,7 @@ pub struct SigningKeyValue { pub used: bool, pub created_at_ms: i64, pub expires_at_ms: i64, + pub used_at_ms: Option, } impl SigningKeyValue { @@ -120,7 +133,7 @@ impl SigningKeyValue { let payload = postcard::to_allocvec(self).expect("SigningKeyValue serialization cannot fail"); let mut buf = Vec::with_capacity(1 + payload.len()); - buf.push(SIGNING_KEY_SCHEMA_VERSION); + buf.push(SIGNING_KEY_SCHEMA_V2); buf.extend_from_slice(&payload); buf } @@ -128,7 +141,20 @@ impl SigningKeyValue { pub fn deserialize(bytes: &[u8]) -> Option { let (&version, payload) = bytes.split_first()?; match version { - SIGNING_KEY_SCHEMA_VERSION => postcard::from_bytes(payload).ok(), + SIGNING_KEY_SCHEMA_V1 => { + let v1: SigningKeyValueV1 = postcard::from_bytes(payload).ok()?; + Some(Self { + id: v1.id, + did: v1.did, + public_key_did_key: v1.public_key_did_key, + private_key_bytes: v1.private_key_bytes, + used: v1.used, + created_at_ms: v1.created_at_ms, + expires_at_ms: v1.expires_at_ms, + used_at_ms: None, + }) + } + SIGNING_KEY_SCHEMA_V2 => postcard::from_bytes(payload).ok(), _ => None, } } @@ -518,12 +544,35 @@ mod tests { used: false, created_at_ms: 1700000000000, expires_at_ms: 1700000600000, + used_at_ms: None, }; let bytes = val.serialize(); let decoded = SigningKeyValue::deserialize(&bytes).unwrap(); assert_eq!(val, decoded); } + #[test] + fn signing_key_value_v1_migration() { + let v1 = SigningKeyValueV1 { + id: uuid::Uuid::new_v4(), + did: Some("did:plc:test".to_owned()), + public_key_did_key: "did:key:z123".to_owned(), + private_key_bytes: vec![1, 2, 3, 4], + used: true, + created_at_ms: 1700000000000, + expires_at_ms: 1700000600000, + }; + let payload = postcard::to_allocvec(&v1).unwrap(); + let mut bytes = Vec::with_capacity(1 + payload.len()); + bytes.push(SIGNING_KEY_SCHEMA_V1); + bytes.extend_from_slice(&payload); + + let decoded = SigningKeyValue::deserialize(&bytes).unwrap(); + assert_eq!(decoded.id, v1.id); + assert!(decoded.used); + assert_eq!(decoded.used_at_ms, None); + } + #[test] fn deletion_request_value_roundtrip() { let val = DeletionRequestValue { diff --git a/crates/tranquil-store/src/metastore/mod.rs b/crates/tranquil-store/src/metastore/mod.rs index 84e1fa0..c0b8d7f 100644 --- a/crates/tranquil-store/src/metastore/mod.rs +++ b/crates/tranquil-store/src/metastore/mod.rs @@ -233,6 +233,10 @@ impl Metastore { &self.partitions[p.index()] } + pub fn signal_keyspace(&self) -> Keyspace { + self.partitions[Partition::Signal.index()].clone() + } + pub fn user_hashes(&self) -> &Arc { &self.user_hashes } diff --git a/crates/tranquil-store/src/metastore/oauth_ops.rs b/crates/tranquil-store/src/metastore/oauth_ops.rs index 16e9145..95f8756 100644 --- a/crates/tranquil-store/src/metastore/oauth_ops.rs +++ b/crates/tranquil-store/src/metastore/oauth_ops.rs @@ -405,6 +405,10 @@ impl OAuthOps { family_id: token.family_id, }; + let used_val = UsedRefreshValue { + family_id: token.family_id, + }; + let mut batch = self.db.batch(); batch.remove( &self.auth, @@ -429,6 +433,11 @@ impl OAuthOps { oauth_token_by_prev_refresh_key(&old_refresh).as_slice(), index.serialize_with_ttl(token.expires_at_ms), ); + batch.insert( + &self.auth, + oauth_used_refresh_key(&old_refresh).as_slice(), + used_val.serialize_with_ttl(token.expires_at_ms), + ); batch.insert( &self.auth, oauth_token_by_id_key(&token.token_id).as_slice(), @@ -1582,6 +1591,14 @@ impl OAuthOps { } } } + + pub fn get_2fa_challenge_code( + &self, + request_uri: &RequestId, + ) -> Result, MetastoreError> { + self.get_2fa_challenge(request_uri) + .map(|opt| opt.map(|c| c.code)) + } } fn default_parameters(client_id: &str) -> tranquil_oauth::AuthorizationRequestParameters { diff --git a/crates/tranquil-store/src/metastore/partitions.rs b/crates/tranquil-store/src/metastore/partitions.rs index 4727064..24c6633 100644 --- a/crates/tranquil-store/src/metastore/partitions.rs +++ b/crates/tranquil-store/src/metastore/partitions.rs @@ -13,15 +13,17 @@ pub enum Partition { Users, Infra, Indexes, + Signal, } impl Partition { - pub const ALL: [Partition; 5] = [ + pub const ALL: [Partition; 6] = [ Partition::RepoData, Partition::Auth, Partition::Users, Partition::Infra, Partition::Indexes, + Partition::Signal, ]; pub const fn index(self) -> usize { @@ -31,6 +33,7 @@ impl Partition { Self::Users => 2, Self::Infra => 3, Self::Indexes => 4, + Self::Signal => 5, } } @@ -41,6 +44,7 @@ impl Partition { Self::Users => "users", Self::Infra => "infra", Self::Indexes => "indexes", + Self::Signal => "signal", } } @@ -52,7 +56,9 @@ impl Partition { FilterPolicyEntry::Bloom(BloomConstructionPolicy::BitsPerKey(10.0)), ])) } - Self::Auth | Self::Users | Self::Infra => KeyspaceCreateOptions::default(), + Self::Auth | Self::Users | Self::Infra | Self::Signal => { + KeyspaceCreateOptions::default() + } } } } @@ -112,7 +118,7 @@ mod tests { #[test] fn all_partitions_covered() { - assert_eq!(Partition::ALL.len(), 5); + assert_eq!(Partition::ALL.len(), 6); } #[test] diff --git a/crates/tranquil-store/src/metastore/record_ops.rs b/crates/tranquil-store/src/metastore/record_ops.rs index 31dddf9..0fd8a68 100644 --- a/crates/tranquil-store/src/metastore/record_ops.rs +++ b/crates/tranquil-store/src/metastore/record_ops.rs @@ -179,6 +179,12 @@ impl RecordOps { let effective_cursor = match query.reverse { false => { + if let Some(ck) = cursor_key.as_ref().filter(|ck| ck.as_slice() < range_hi) { + range_hi = ck.as_slice(); + } + None + } + true => { let narrowed = cursor_key.as_ref().filter(|ck| ck.as_slice() > range_lo); match narrowed { Some(ck) => { @@ -188,21 +194,15 @@ impl RecordOps { None => None, } } - true => { - if let Some(ck) = cursor_key.as_ref().filter(|ck| ck.as_slice() < range_hi) { - range_hi = ck.as_slice(); - } - None - } }; match range_lo >= range_hi { true => Ok(Vec::new()), false => match query.reverse { - false => { + false => self.list_records_reverse(range_lo, range_hi, query.limit), + true => { self.list_records_forward(range_lo, range_hi, effective_cursor, query.limit) } - true => self.list_records_reverse(range_lo, range_hi, query.limit), }, } } @@ -674,7 +674,7 @@ mod tests { } #[test] - fn list_records_forward_with_limit() { + fn list_records_default_desc_with_limit() { let (_dir, ms) = open_fresh(); let (user_id, user_hash) = setup_user(&ms); let rec_ops = ms.record_ops(); @@ -699,84 +699,13 @@ mod tests { .list_records(&lrq(user_id, &collection, None, 3, false, None, None)) .unwrap(); assert_eq!(results.len(), 3); - assert_eq!(results[0].rkey.as_str(), "rkey000"); - assert_eq!(results[1].rkey.as_str(), "rkey001"); - assert_eq!(results[2].rkey.as_str(), "rkey002"); - } - - #[test] - fn list_records_with_cursor() { - let (_dir, ms) = open_fresh(); - let (user_id, user_hash) = setup_user(&ms); - let rec_ops = ms.record_ops(); - - let collection = Nsid::from("app.bsky.feed.post".to_string()); - let rkeys: Vec = (0..5).map(|i| Rkey::from(format!("rkey{i:03}"))).collect(); - let cids: Vec = (0..5).map(|i| test_cid_link(i + 1)).collect(); - - let writes: Vec> = rkeys - .iter() - .zip(cids.iter()) - .map(|(rk, c)| rw(&collection, rk, c)) - .collect(); - - let mut batch = ms.database().batch(); - rec_ops - .upsert_records(&mut batch, user_hash, &writes) - .unwrap(); - batch.commit().unwrap(); - - let cursor = Rkey::from("rkey001".to_string()); - let results = rec_ops - .list_records(&lrq( - user_id, - &collection, - Some(&cursor), - 10, - false, - None, - None, - )) - .unwrap(); - assert_eq!(results.len(), 3); - assert_eq!(results[0].rkey.as_str(), "rkey002"); - assert_eq!(results[1].rkey.as_str(), "rkey003"); - assert_eq!(results[2].rkey.as_str(), "rkey004"); - } - - #[test] - fn list_records_reverse() { - let (_dir, ms) = open_fresh(); - let (user_id, user_hash) = setup_user(&ms); - let rec_ops = ms.record_ops(); - - let collection = Nsid::from("app.bsky.feed.post".to_string()); - let rkeys: Vec = (0..5).map(|i| Rkey::from(format!("rkey{i:03}"))).collect(); - let cids: Vec = (0..5).map(|i| test_cid_link(i + 1)).collect(); - - let writes: Vec> = rkeys - .iter() - .zip(cids.iter()) - .map(|(rk, c)| rw(&collection, rk, c)) - .collect(); - - let mut batch = ms.database().batch(); - rec_ops - .upsert_records(&mut batch, user_hash, &writes) - .unwrap(); - batch.commit().unwrap(); - - let results = rec_ops - .list_records(&lrq(user_id, &collection, None, 3, true, None, None)) - .unwrap(); - assert_eq!(results.len(), 3); assert_eq!(results[0].rkey.as_str(), "rkey004"); assert_eq!(results[1].rkey.as_str(), "rkey003"); assert_eq!(results[2].rkey.as_str(), "rkey002"); } #[test] - fn list_records_reverse_with_cursor() { + fn list_records_default_desc_with_cursor() { let (_dir, ms) = open_fresh(); let (user_id, user_hash) = setup_user(&ms); let rec_ops = ms.record_ops(); @@ -804,7 +733,7 @@ mod tests { &collection, Some(&cursor), 10, - true, + false, None, None, )) @@ -816,7 +745,78 @@ mod tests { } #[test] - fn list_records_rkey_range_bounds() { + fn list_records_reverse_asc() { + let (_dir, ms) = open_fresh(); + let (user_id, user_hash) = setup_user(&ms); + let rec_ops = ms.record_ops(); + + let collection = Nsid::from("app.bsky.feed.post".to_string()); + let rkeys: Vec = (0..5).map(|i| Rkey::from(format!("rkey{i:03}"))).collect(); + let cids: Vec = (0..5).map(|i| test_cid_link(i + 1)).collect(); + + let writes: Vec> = rkeys + .iter() + .zip(cids.iter()) + .map(|(rk, c)| rw(&collection, rk, c)) + .collect(); + + let mut batch = ms.database().batch(); + rec_ops + .upsert_records(&mut batch, user_hash, &writes) + .unwrap(); + batch.commit().unwrap(); + + let results = rec_ops + .list_records(&lrq(user_id, &collection, None, 3, true, None, None)) + .unwrap(); + assert_eq!(results.len(), 3); + assert_eq!(results[0].rkey.as_str(), "rkey000"); + assert_eq!(results[1].rkey.as_str(), "rkey001"); + assert_eq!(results[2].rkey.as_str(), "rkey002"); + } + + #[test] + fn list_records_reverse_asc_with_cursor() { + let (_dir, ms) = open_fresh(); + let (user_id, user_hash) = setup_user(&ms); + let rec_ops = ms.record_ops(); + + let collection = Nsid::from("app.bsky.feed.post".to_string()); + let rkeys: Vec = (0..5).map(|i| Rkey::from(format!("rkey{i:03}"))).collect(); + let cids: Vec = (0..5).map(|i| test_cid_link(i + 1)).collect(); + + let writes: Vec> = rkeys + .iter() + .zip(cids.iter()) + .map(|(rk, c)| rw(&collection, rk, c)) + .collect(); + + let mut batch = ms.database().batch(); + rec_ops + .upsert_records(&mut batch, user_hash, &writes) + .unwrap(); + batch.commit().unwrap(); + + let cursor = Rkey::from("rkey001".to_string()); + let results = rec_ops + .list_records(&lrq( + user_id, + &collection, + Some(&cursor), + 10, + true, + None, + None, + )) + .unwrap(); + assert_eq!(results.len(), 3); + assert_eq!(results[0].rkey.as_str(), "rkey002"); + assert_eq!(results[1].rkey.as_str(), "rkey003"); + assert_eq!(results[2].rkey.as_str(), "rkey004"); + } + + #[test] + fn list_records_default_desc_rkey_range_bounds() { let (_dir, ms) = open_fresh(); let (user_id, user_hash) = setup_user(&ms); let rec_ops = ms.record_ops(); @@ -851,8 +851,8 @@ mod tests { )) .unwrap(); assert_eq!(results.len(), 4); - assert_eq!(results[0].rkey.as_str(), "rkey003"); - assert_eq!(results[3].rkey.as_str(), "rkey006"); + assert_eq!(results[0].rkey.as_str(), "rkey006"); + assert_eq!(results[3].rkey.as_str(), "rkey003"); } #[test] @@ -1220,7 +1220,7 @@ mod tests { .list_records(&lrq(user_id, &collection, None, 100, false, None, None)) .unwrap(); let result_rkeys: Vec<&str> = results.iter().map(|r| r.rkey.as_str()).collect(); - assert_eq!(result_rkeys, ["apple", "banana", "mango", "zebra"]); + assert_eq!(result_rkeys, ["zebra", "mango", "banana", "apple"]); } #[test] @@ -1287,8 +1287,8 @@ mod tests { .list_records(&lrq(user_id, &collection, None, 100, false, None, None)) .unwrap(); assert_eq!(results.len(), 2); - assert_eq!(results[0].rkey.as_str(), "abc"); - assert_eq!(results[1].rkey.as_str(), "abc\x00def"); + assert_eq!(results[0].rkey.as_str(), "abc\x00def"); + assert_eq!(results[1].rkey.as_str(), "abc"); } #[test] @@ -1341,7 +1341,7 @@ mod tests { } #[test] - fn list_records_cursor_past_rkey_end_returns_empty() { + fn list_records_cursor_before_all_returns_empty() { let (_dir, ms) = open_fresh(); let (user_id, user_hash) = setup_user(&ms); let rec_ops = ms.record_ops(); @@ -1362,8 +1362,7 @@ mod tests { .unwrap(); batch.commit().unwrap(); - let cursor = Rkey::from("rkey010".to_string()); - let rkey_end = Rkey::from("rkey003".to_string()); + let cursor = Rkey::from("a".to_string()); let results = rec_ops .list_records(&lrq( user_id, @@ -1372,14 +1371,14 @@ mod tests { 100, false, None, - Some(&rkey_end), + None, )) .unwrap(); assert!(results.is_empty()); } #[test] - fn list_records_reverse_with_rkey_bounds() { + fn list_records_reverse_asc_with_rkey_bounds() { let (_dir, ms) = open_fresh(); let (user_id, user_hash) = setup_user(&ms); let rec_ops = ms.record_ops(); @@ -1414,12 +1413,12 @@ mod tests { )) .unwrap(); assert_eq!(results.len(), 6); - assert_eq!(results[0].rkey.as_str(), "rkey007"); - assert_eq!(results[5].rkey.as_str(), "rkey002"); + assert_eq!(results[0].rkey.as_str(), "rkey002"); + assert_eq!(results[5].rkey.as_str(), "rkey007"); } #[test] - fn list_records_reverse_cursor_narrows_range() { + fn list_records_default_desc_cursor_narrows_range() { let (_dir, ms) = open_fresh(); let (user_id, user_hash) = setup_user(&ms); let rec_ops = ms.record_ops(); @@ -1447,7 +1446,7 @@ mod tests { &collection, Some(&cursor), 3, - true, + false, None, None, )) diff --git a/crates/tranquil-store/src/metastore/user_ops.rs b/crates/tranquil-store/src/metastore/user_ops.rs index 755e785..d9eddf6 100644 --- a/crates/tranquil-store/src/metastore/user_ops.rs +++ b/crates/tranquil-store/src/metastore/user_ops.rs @@ -410,10 +410,9 @@ impl UserOps { return Ok(None); } - let email_match = email_filter.map_or(true, |f| { - val.email.as_deref().map_or(false, |e| e.contains(f)) - }); - let handle_match = handle_filter.map_or(true, |f| val.handle.contains(f)); + let email_match = email_filter + .is_none_or(|f| val.email.as_deref().is_some_and(|e| e.contains(f))); + let handle_match = handle_filter.is_none_or(|f| val.handle.contains(f)); match email_match && handle_match { true => Ok(Some(AccountSearchResult { @@ -689,13 +688,13 @@ impl UserOps { pub fn is_account_migrated(&self, did: &Did) -> Result { Ok(self .load_user_by_did(did.as_str())? - .map_or(false, |v| v.migrated_to_pds.is_some())) + .is_some_and(|v| v.migrated_to_pds.is_some())) } pub fn has_verified_comms_channel(&self, did: &Did) -> Result { Ok(self .load_user_by_did(did.as_str())? - .map_or(false, |v| v.channel_verification() != 0)) + .is_some_and(|v| v.channel_verification() != 0)) } pub fn get_id_by_handle(&self, handle: &Handle) -> Result, MetastoreError> { @@ -861,6 +860,14 @@ impl UserOps { } } + pub fn set_admin_status(&self, did: &Did, is_admin: bool) -> Result<(), MetastoreError> { + let user_hash = self.resolve_hash(did.as_str()); + self.mutate_user(user_hash, |u| { + u.is_admin = is_admin; + })?; + Ok(()) + } + pub fn get_notification_prefs( &self, did: &Did, @@ -1198,9 +1205,8 @@ impl UserOps { telegram_username: &str, ) -> Result<(), MetastoreError> { self.mutate_user_by_uuid(user_id, |u| { - match u.telegram_username.as_deref() == Some(telegram_username) { - true => u.telegram_verified = true, - false => {} + if u.telegram_username.as_deref() == Some(telegram_username) { + u.telegram_verified = true; } })?; Ok(()) @@ -1212,9 +1218,8 @@ impl UserOps { signal_username: &str, ) -> Result<(), MetastoreError> { self.mutate_user_by_uuid(user_id, |u| { - match u.signal_username.as_deref() == Some(signal_username) { - true => u.signal_verified = true, - false => {} + if u.signal_username.as_deref() == Some(signal_username) { + u.signal_verified = true; } })?; Ok(()) @@ -1251,7 +1256,7 @@ impl UserOps { pub fn has_totp_enabled(&self, did: &Did) -> Result { Ok(self .load_user_by_did(did.as_str())? - .map_or(false, |v| v.totp_enabled)) + .is_some_and(|v| v.totp_enabled)) } pub fn has_passkeys(&self, did: &Did) -> Result { @@ -2077,7 +2082,7 @@ impl UserOps { let prefix = [super::keys::KeyTag::USER_RESET_CODE.raw()]; let keys_to_remove: Vec> = self .auth - .prefix(&prefix) + .prefix(prefix) .filter_map(|guard| { let (key_bytes, val_bytes) = guard.into_inner().ok()?; let rc = ResetCodeValue::deserialize(&val_bytes)?; @@ -2566,18 +2571,7 @@ impl UserOps { .collect() } - pub fn delete_account_with_firehose( - &self, - user_id: Uuid, - did: &Did, - ) -> Result { - self.delete_account_complete(user_id, did)?; - tracing::warn!( - "delete_account_with_firehose: no firehose event emitted (not yet implemented for tranquil-store)" - ); - Ok(0) - } - + #[allow(clippy::too_many_arguments)] fn build_user_value( &self, did: &Did, @@ -2894,9 +2888,8 @@ impl UserOps { .map_err(|e| MigrationReactivationError::Database(e.to_string()))? .ok_or(MigrationReactivationError::NotFound)?; - match user.deactivated_at_ms { - None => return Err(MigrationReactivationError::NotDeactivated), - Some(_) => {} + if user.deactivated_at_ms.is_none() { + return Err(MigrationReactivationError::NotDeactivated); } let existing_handle = self @@ -3096,4 +3089,138 @@ impl UserOps { Ok(RecoverPasskeyAccountResult { passkeys_deleted }) } + + pub fn get_password_reset_info( + &self, + email: &str, + ) -> Result, MetastoreError> { + let by_email_key = user_by_email_key(email); + let user_hash = match self + .users + .get(by_email_key.as_slice()) + .map_err(MetastoreError::Fjall)? + { + Some(raw) => { + let arr: [u8; 8] = raw + .as_ref() + .try_into() + .map_err(|_| MetastoreError::CorruptData("email index not 8 bytes"))?; + UserHash::from_raw(u64::from_be_bytes(arr)) + } + None => return Ok(None), + }; + + let user = match self.load_user(user_hash)? { + Some(u) => u, + None => return Ok(None), + }; + + let prefix = [super::keys::KeyTag::USER_RESET_CODE.raw()]; + let reset_code = self + .auth + .prefix(prefix) + .filter_map(|guard| { + let (_, val_bytes) = guard.into_inner().ok()?; + let rc = ResetCodeValue::deserialize(&val_bytes)?; + match rc.user_id == user.id { + true => Some(rc), + false => None, + } + }) + .next(); + + Ok(Some(tranquil_db_traits::PasswordResetInfo { + code: reset_code.as_ref().map(|rc| rc.code.clone()), + expires_at: reset_code.and_then(|rc| DateTime::from_timestamp_millis(rc.expires_at_ms)), + })) + } + + pub fn enable_totp_verified( + &self, + did: &Did, + encrypted_secret: &[u8], + ) -> Result<(), MetastoreError> { + let user_hash = self.resolve_hash(did.as_str()); + let key = totp_key(user_hash); + + let value = TotpValue { + secret_encrypted: encrypted_secret.to_vec(), + encryption_version: 1, + verified: true, + last_used_at_ms: None, + }; + + let mut batch = self.db.batch(); + batch.insert(&self.users, key.as_slice(), value.serialize()); + + if let Some(mut user) = self.load_user(user_hash)? { + user.totp_enabled = true; + batch.insert( + &self.users, + user_primary_key(user_hash).as_slice(), + user.serialize(), + ); + } + + batch.commit().map_err(MetastoreError::Fjall) + } + + pub fn set_two_factor_enabled(&self, did: &Did, enabled: bool) -> Result<(), MetastoreError> { + let user_hash = self.resolve_hash(did.as_str()); + let mut user = self + .load_user(user_hash)? + .ok_or(MetastoreError::InvalidInput("user not found"))?; + + user.two_factor_enabled = enabled; + self.users + .insert(user_primary_key(user_hash).as_slice(), user.serialize()) + .map_err(MetastoreError::Fjall) + } + + pub fn expire_password_reset_code(&self, email: &str) -> Result<(), MetastoreError> { + let by_email_key = user_by_email_key(email); + let user_hash_bytes = match self + .users + .get(by_email_key.as_slice()) + .map_err(MetastoreError::Fjall)? + { + Some(raw) => raw, + None => return Ok(()), + }; + + let arr: [u8; 8] = user_hash_bytes + .as_ref() + .try_into() + .map_err(|_| MetastoreError::CorruptData("email index not 8 bytes"))?; + let user_hash = UserHash::from_raw(u64::from_be_bytes(arr)); + let user = match self.load_user(user_hash)? { + Some(u) => u, + None => return Ok(()), + }; + + let prefix = [super::keys::KeyTag::USER_RESET_CODE.raw()]; + let keys_to_remove: Vec> = self + .auth + .prefix(prefix) + .filter_map(|guard| { + let (key_bytes, val_bytes) = guard.into_inner().ok()?; + let rc = ResetCodeValue::deserialize(&val_bytes)?; + match rc.user_id == user.id { + true => Some(key_bytes.to_vec()), + false => None, + } + }) + .collect(); + + match keys_to_remove.is_empty() { + true => Ok(()), + false => { + let mut batch = self.db.batch(); + keys_to_remove.iter().for_each(|key| { + batch.remove(&self.auth, key); + }); + batch.commit().map_err(MetastoreError::Fjall) + } + } + } }