diff --git a/.sqlx/query-2ff22a8c39914689d6cf215ba201fa4ced50b7a003ce01bf7603a7f125113447.json b/.sqlx/query-2ff22a8c39914689d6cf215ba201fa4ced50b7a003ce01bf7603a7f125113447.json deleted file mode 100644 index 41b7fb2..0000000 --- a/.sqlx/query-2ff22a8c39914689d6cf215ba201fa4ced50b7a003ce01bf7603a7f125113447.json +++ /dev/null @@ -1,40 +0,0 @@ -{ - "db_name": "PostgreSQL", - "query": "\n SELECT code, available_uses, created_at, disabled\n FROM invite_codes\n WHERE created_by_user = $1\n ORDER BY created_at DESC\n ", - "describe": { - "columns": [ - { - "ordinal": 0, - "name": "code", - "type_info": "Text" - }, - { - "ordinal": 1, - "name": "available_uses", - "type_info": "Int4" - }, - { - "ordinal": 2, - "name": "created_at", - "type_info": "Timestamptz" - }, - { - "ordinal": 3, - "name": "disabled", - "type_info": "Bool" - } - ], - "parameters": { - "Left": [ - "Uuid" - ] - }, - "nullable": [ - false, - false, - false, - true - ] - }, - "hash": "2ff22a8c39914689d6cf215ba201fa4ced50b7a003ce01bf7603a7f125113447" -} diff --git a/.sqlx/query-bbe639bb24cc1bb3cc144baae263e7e3411e185bf7c91751ee1046c64a81df52.json b/.sqlx/query-59678fbb756d46bb5f51c9a52800a8d203ed52129b1fae65145df92d145d18de.json similarity index 59% rename from .sqlx/query-bbe639bb24cc1bb3cc144baae263e7e3411e185bf7c91751ee1046c64a81df52.json rename to .sqlx/query-59678fbb756d46bb5f51c9a52800a8d203ed52129b1fae65145df92d145d18de.json index 988c010..d2cee94 100644 --- a/.sqlx/query-bbe639bb24cc1bb3cc144baae263e7e3411e185bf7c91751ee1046c64a81df52.json +++ b/.sqlx/query-59678fbb756d46bb5f51c9a52800a8d203ed52129b1fae65145df92d145d18de.json @@ -1,16 +1,17 @@ { "db_name": "PostgreSQL", - "query": "INSERT INTO invite_codes (code, available_uses, created_by_user) VALUES ($1, $2, $3)", + "query": "INSERT INTO invite_codes (code, available_uses, created_by_user, for_account) VALUES ($1, $2, $3, $4)", "describe": { "columns": [], "parameters": { "Left": [ "Text", "Int4", - "Uuid" + "Uuid", + "Text" ] }, "nullable": [] }, - "hash": "bbe639bb24cc1bb3cc144baae263e7e3411e185bf7c91751ee1046c64a81df52" + "hash": "59678fbb756d46bb5f51c9a52800a8d203ed52129b1fae65145df92d145d18de" } diff --git a/.sqlx/query-704b32d9ae2234ae12dad87f5f86230e16acaa1c0c229c66b39024bf9662f1e5.json b/.sqlx/query-704b32d9ae2234ae12dad87f5f86230e16acaa1c0c229c66b39024bf9662f1e5.json new file mode 100644 index 0000000..1dd1a43 --- /dev/null +++ b/.sqlx/query-704b32d9ae2234ae12dad87f5f86230e16acaa1c0c229c66b39024bf9662f1e5.json @@ -0,0 +1,52 @@ +{ + "db_name": "PostgreSQL", + "query": "\n SELECT\n ic.code,\n ic.available_uses,\n ic.created_at,\n ic.disabled,\n ic.for_account,\n (SELECT COUNT(*) FROM invite_code_uses icu WHERE icu.code = ic.code)::int as \"use_count!\"\n FROM invite_codes ic\n WHERE ic.for_account = $1\n ORDER BY ic.created_at DESC\n ", + "describe": { + "columns": [ + { + "ordinal": 0, + "name": "code", + "type_info": "Text" + }, + { + "ordinal": 1, + "name": "available_uses", + "type_info": "Int4" + }, + { + "ordinal": 2, + "name": "created_at", + "type_info": "Timestamptz" + }, + { + "ordinal": 3, + "name": "disabled", + "type_info": "Bool" + }, + { + "ordinal": 4, + "name": "for_account", + "type_info": "Text" + }, + { + "ordinal": 5, + "name": "use_count!", + "type_info": "Int4" + } + ], + "parameters": { + "Left": [ + "Text" + ] + }, + "nullable": [ + false, + false, + false, + true, + false, + null + ] + }, + "hash": "704b32d9ae2234ae12dad87f5f86230e16acaa1c0c229c66b39024bf9662f1e5" +} diff --git a/.sqlx/query-b3d44806b6351d788048e6afe7a6623882fac70b466bf09596cad8eae1fc9dac.json b/.sqlx/query-b3d44806b6351d788048e6afe7a6623882fac70b466bf09596cad8eae1fc9dac.json new file mode 100644 index 0000000..6d04728 --- /dev/null +++ b/.sqlx/query-b3d44806b6351d788048e6afe7a6623882fac70b466bf09596cad8eae1fc9dac.json @@ -0,0 +1,16 @@ +{ + "db_name": "PostgreSQL", + "query": "INSERT INTO invite_codes (code, available_uses, created_by_user, for_account)\n SELECT $1, $2, id, $3 FROM users WHERE is_admin = true LIMIT 1", + "describe": { + "columns": [], + "parameters": { + "Left": [ + "Text", + "Int4", + "Text" + ] + }, + "nullable": [] + }, + "hash": "b3d44806b6351d788048e6afe7a6623882fac70b466bf09596cad8eae1fc9dac" +} diff --git a/.sqlx/query-ce50221e621d89f7f7d315b0ccc7893b2c344e3612b56116a785248dda296424.json b/.sqlx/query-ce50221e621d89f7f7d315b0ccc7893b2c344e3612b56116a785248dda296424.json new file mode 100644 index 0000000..9bf4f08 --- /dev/null +++ b/.sqlx/query-ce50221e621d89f7f7d315b0ccc7893b2c344e3612b56116a785248dda296424.json @@ -0,0 +1,20 @@ +{ + "db_name": "PostgreSQL", + "query": "SELECT id FROM users WHERE is_admin = true LIMIT 1", + "describe": { + "columns": [ + { + "ordinal": 0, + "name": "id", + "type_info": "Uuid" + } + ], + "parameters": { + "Left": [] + }, + "nullable": [ + false + ] + }, + "hash": "ce50221e621d89f7f7d315b0ccc7893b2c344e3612b56116a785248dda296424" +} diff --git a/.sqlx/query-da0e9a9edad3895ed5015b52335f5a0256e7bdc6c79e6faa927414d68800404c.json b/.sqlx/query-da0e9a9edad3895ed5015b52335f5a0256e7bdc6c79e6faa927414d68800404c.json deleted file mode 100644 index 0c6f6b4..0000000 --- a/.sqlx/query-da0e9a9edad3895ed5015b52335f5a0256e7bdc6c79e6faa927414d68800404c.json +++ /dev/null @@ -1,22 +0,0 @@ -{ - "db_name": "PostgreSQL", - "query": "SELECT invites_disabled FROM users WHERE did = $1", - "describe": { - "columns": [ - { - "ordinal": 0, - "name": "invites_disabled", - "type_info": "Bool" - } - ], - "parameters": { - "Left": [ - "Text" - ] - }, - "nullable": [ - true - ] - }, - "hash": "da0e9a9edad3895ed5015b52335f5a0256e7bdc6c79e6faa927414d68800404c" -} diff --git a/frontend/src/routes/Dashboard.svelte b/frontend/src/routes/Dashboard.svelte index 8f86589..eaad876 100644 --- a/frontend/src/routes/Dashboard.svelte +++ b/frontend/src/routes/Dashboard.svelte @@ -164,7 +164,7 @@

{$_('dashboard.navSessions')}

{$_('dashboard.navSessionsDesc')}

- {#if inviteCodesEnabled} + {#if inviteCodesEnabled && auth.session.isAdmin}

{$_('dashboard.navInviteCodes')}

{$_('dashboard.navInviteCodesDesc')}

diff --git a/frontend/src/routes/InviteCodes.svelte b/frontend/src/routes/InviteCodes.svelte index e5ddfc8..f7086d7 100644 --- a/frontend/src/routes/InviteCodes.svelte +++ b/frontend/src/routes/InviteCodes.svelte @@ -91,11 +91,13 @@ {/if} -
- -
+ {#if auth.session?.isAdmin} +
+ +
+ {/if}

{$_('inviteCodes.yourCodes')}

{#if loading} diff --git a/migrations/20251242_invite_code_for_account.sql b/migrations/20251242_invite_code_for_account.sql new file mode 100644 index 0000000..7197a88 --- /dev/null +++ b/migrations/20251242_invite_code_for_account.sql @@ -0,0 +1,2 @@ +ALTER TABLE invite_codes ADD COLUMN IF NOT EXISTS for_account TEXT NOT NULL DEFAULT 'admin'; +CREATE INDEX IF NOT EXISTS idx_invite_codes_for_account ON invite_codes(for_account); diff --git a/src/api/server/invite.rs b/src/api/server/invite.rs index 52d4b54..031604c 100644 --- a/src/api/server/invite.rs +++ b/src/api/server/invite.rs @@ -1,15 +1,36 @@ use crate::api::ApiError; +use crate::auth::extractor::BearerAuthAdmin; use crate::auth::BearerAuth; use crate::state::AppState; -use crate::util::get_user_id_by_did; use axum::{ Json, extract::State, response::{IntoResponse, Response}, }; +use rand::Rng; use serde::{Deserialize, Serialize}; use tracing::error; -use uuid::Uuid; + +const BASE32_ALPHABET: &[u8] = b"abcdefghijklmnopqrstuvwxyz234567"; + +fn gen_random_token() -> String { + let mut rng = rand::thread_rng(); + let mut token = String::with_capacity(11); + for i in 0..10 { + if i == 5 { + token.push('-'); + } + let idx = rng.gen_range(0..32); + token.push(BASE32_ALPHABET[idx] as char); + } + token +} + +fn gen_invite_code() -> String { + let hostname = std::env::var("PDS_HOSTNAME").unwrap_or_else(|_| "localhost".to_string()); + let hostname_prefix = hostname.replace('.', "-"); + format!("{}-{}", hostname_prefix, gen_random_token()) +} #[derive(Deserialize)] #[serde(rename_all = "camelCase")] @@ -25,59 +46,33 @@ pub struct CreateInviteCodeOutput { pub async fn create_invite_code( State(state): State, - BearerAuth(auth_user): BearerAuth, + BearerAuthAdmin(_auth_user): BearerAuthAdmin, Json(input): Json, ) -> Response { if input.use_count < 1 { return ApiError::InvalidRequest("useCount must be at least 1".into()).into_response(); } - let user_id = match get_user_id_by_did(&state.db, &auth_user.did).await { - Ok(id) => id, - Err(e) => return ApiError::from(e).into_response(), - }; - let creator_user_id = if let Some(for_account) = &input.for_account { - match sqlx::query!("SELECT id FROM users WHERE did = $1", for_account) - .fetch_optional(&state.db) - .await - { - Ok(Some(row)) => row.id, - Ok(None) => return ApiError::AccountNotFound.into_response(), - Err(e) => { - error!("DB error looking up target account: {:?}", e); - return ApiError::InternalError.into_response(); - } - } - } else { - user_id - }; - let user_invites_disabled = sqlx::query_scalar!( - "SELECT invites_disabled FROM users WHERE did = $1", - auth_user.did - ) - .fetch_optional(&state.db) - .await - .map_err(|e| { - error!("DB error checking invites_disabled: {:?}", e); - ApiError::InternalError - }) - .ok() - .flatten() - .flatten() - .unwrap_or(false); - if user_invites_disabled { - return ApiError::InvitesDisabled.into_response(); - } - let code = Uuid::new_v4().to_string(); + + let for_account = input.for_account.unwrap_or_else(|| "admin".to_string()); + let code = gen_invite_code(); + match sqlx::query!( - "INSERT INTO invite_codes (code, available_uses, created_by_user) VALUES ($1, $2, $3)", + "INSERT INTO invite_codes (code, available_uses, created_by_user, for_account) + SELECT $1, $2, id, $3 FROM users WHERE is_admin = true LIMIT 1", code, input.use_count, - creator_user_id + for_account ) .execute(&state.db) .await { - Ok(_) => Json(CreateInviteCodeOutput { code }).into_response(), + Ok(result) => { + if result.rows_affected() == 0 { + error!("No admin user found to create invite code"); + return ApiError::InternalError.into_response(); + } + Json(CreateInviteCodeOutput { code }).into_response() + } Err(e) => { error!("DB error creating invite code: {:?}", e); ApiError::InternalError.into_response() @@ -106,28 +101,48 @@ pub struct AccountCodes { pub async fn create_invite_codes( State(state): State, - BearerAuth(auth_user): BearerAuth, + BearerAuthAdmin(_auth_user): BearerAuthAdmin, Json(input): Json, ) -> Response { if input.use_count < 1 { return ApiError::InvalidRequest("useCount must be at least 1".into()).into_response(); } - let user_id = match get_user_id_by_did(&state.db, &auth_user.did).await { - Ok(id) => id, - Err(e) => return ApiError::from(e).into_response(), - }; + let code_count = input.code_count.unwrap_or(1).max(1); - let for_accounts = input.for_accounts.unwrap_or_default(); + let for_accounts = input + .for_accounts + .filter(|v| !v.is_empty()) + .unwrap_or_else(|| vec!["admin".to_string()]); + + let admin_user_id = match sqlx::query_scalar!( + "SELECT id FROM users WHERE is_admin = true LIMIT 1" + ) + .fetch_optional(&state.db) + .await + { + Ok(Some(id)) => id, + Ok(None) => { + error!("No admin user found to create invite codes"); + return ApiError::InternalError.into_response(); + } + Err(e) => { + error!("DB error looking up admin user: {:?}", e); + return ApiError::InternalError.into_response(); + } + }; + let mut result_codes = Vec::new(); - if for_accounts.is_empty() { + + for account in for_accounts { let mut codes = Vec::new(); for _ in 0..code_count { - let code = Uuid::new_v4().to_string(); + let code = gen_invite_code(); if let Err(e) = sqlx::query!( - "INSERT INTO invite_codes (code, available_uses, created_by_user) VALUES ($1, $2, $3)", + "INSERT INTO invite_codes (code, available_uses, created_by_user, for_account) VALUES ($1, $2, $3, $4)", code, input.use_count, - user_id + admin_user_id, + account ) .execute(&state.db) .await @@ -137,47 +152,9 @@ pub async fn create_invite_codes( } codes.push(code); } - result_codes.push(AccountCodes { - account: "admin".to_string(), - codes, - }); - } else { - for account_did in for_accounts { - let target_user_id = - match sqlx::query!("SELECT id FROM users WHERE did = $1", account_did) - .fetch_optional(&state.db) - .await - { - Ok(Some(row)) => row.id, - Ok(None) => continue, - Err(e) => { - error!("DB error looking up target account: {:?}", e); - return ApiError::InternalError.into_response(); - } - }; - let mut codes = Vec::new(); - for _ in 0..code_count { - let code = Uuid::new_v4().to_string(); - if let Err(e) = sqlx::query!( - "INSERT INTO invite_codes (code, available_uses, created_by_user) VALUES ($1, $2, $3)", - code, - input.use_count, - target_user_id - ) - .execute(&state.db) - .await - { - error!("DB error creating invite code: {:?}", e); - return ApiError::InternalError.into_response(); - } - codes.push(code); - } - result_codes.push(AccountCodes { - account: account_did, - codes, - }); - } + result_codes.push(AccountCodes { account, codes }); } + Json(CreateInviteCodesOutput { codes: result_codes, }) @@ -220,37 +197,45 @@ pub async fn get_account_invite_codes( BearerAuth(auth_user): BearerAuth, axum::extract::Query(params): axum::extract::Query, ) -> Response { - let user_id = match get_user_id_by_did(&state.db, &auth_user.did).await { - Ok(id) => id, - Err(e) => return ApiError::from(e).into_response(), - }; let include_used = params.include_used.unwrap_or(true); + let codes_rows = match sqlx::query!( r#" - SELECT code, available_uses, created_at, disabled - FROM invite_codes - WHERE created_by_user = $1 - ORDER BY created_at DESC + SELECT + ic.code, + ic.available_uses, + ic.created_at, + ic.disabled, + ic.for_account, + (SELECT COUNT(*) FROM invite_code_uses icu WHERE icu.code = ic.code)::int as "use_count!" + FROM invite_codes ic + WHERE ic.for_account = $1 + ORDER BY ic.created_at DESC "#, - user_id + auth_user.did ) .fetch_all(&state.db) .await { - Ok(rows) => { - if include_used { - rows - } else { - rows.into_iter().filter(|r| r.available_uses > 0).collect() - } - } + Ok(rows) => rows, Err(e) => { error!("DB error fetching invite codes: {:?}", e); return ApiError::InternalError.into_response(); } }; + let mut codes = Vec::new(); for row in codes_rows { + let disabled = row.disabled.unwrap_or(false); + if disabled { + continue; + } + + let use_count = row.use_count; + if !include_used && use_count >= row.available_uses { + continue; + } + let uses = sqlx::query!( r#" SELECT u.did, icu.used_at @@ -273,15 +258,17 @@ pub async fn get_account_invite_codes( .collect() }) .unwrap_or_default(); + codes.push(InviteCode { code: row.code, available: row.available_uses, - disabled: row.disabled.unwrap_or(false), - for_account: auth_user.did.clone(), - created_by: auth_user.did.clone(), + disabled, + for_account: row.for_account, + created_by: "admin".to_string(), created_at: row.created_at.to_rfc3339(), uses, }); } + Json(GetAccountInviteCodesOutput { codes }).into_response() } diff --git a/tests/admin_invite.rs b/tests/admin_invite.rs index 52d6f43..5440257 100644 --- a/tests/admin_invite.rs +++ b/tests/admin_invite.rs @@ -83,88 +83,6 @@ async fn test_admin_get_invite_codes_no_auth() { assert_eq!(res.status(), StatusCode::UNAUTHORIZED); } -#[tokio::test] -async fn test_disable_account_invites_success() { - let client = client(); - let (access_jwt, did) = create_admin_account_and_login(&client).await; - let payload = json!({ - "account": did - }); - let res = client - .post(format!( - "{}/xrpc/com.atproto.admin.disableAccountInvites", - base_url().await - )) - .bearer_auth(&access_jwt) - .json(&payload) - .send() - .await - .expect("Failed to send request"); - assert_eq!(res.status(), StatusCode::OK); - let create_payload = json!({ - "useCount": 1 - }); - let res = client - .post(format!( - "{}/xrpc/com.atproto.server.createInviteCode", - base_url().await - )) - .bearer_auth(&access_jwt) - .json(&create_payload) - .send() - .await - .expect("Failed to send request"); - assert_eq!(res.status(), StatusCode::FORBIDDEN); - let body: Value = res.json().await.expect("Response was not valid JSON"); - assert_eq!(body["error"], "InvitesDisabled"); -} - -#[tokio::test] -async fn test_enable_account_invites_success() { - let client = client(); - let (access_jwt, did) = create_admin_account_and_login(&client).await; - let disable_payload = json!({ - "account": did - }); - let _ = client - .post(format!( - "{}/xrpc/com.atproto.admin.disableAccountInvites", - base_url().await - )) - .bearer_auth(&access_jwt) - .json(&disable_payload) - .send() - .await; - let enable_payload = json!({ - "account": did - }); - let res = client - .post(format!( - "{}/xrpc/com.atproto.admin.enableAccountInvites", - base_url().await - )) - .bearer_auth(&access_jwt) - .json(&enable_payload) - .send() - .await - .expect("Failed to send request"); - assert_eq!(res.status(), StatusCode::OK); - let create_payload = json!({ - "useCount": 1 - }); - let res = client - .post(format!( - "{}/xrpc/com.atproto.server.createInviteCode", - base_url().await - )) - .bearer_auth(&access_jwt) - .json(&create_payload) - .send() - .await - .expect("Failed to send request"); - assert_eq!(res.status(), StatusCode::OK); -} - #[tokio::test] async fn test_disable_account_invites_no_auth() { let client = client(); @@ -206,9 +124,10 @@ async fn test_disable_account_invites_not_found() { #[tokio::test] async fn test_disable_invite_codes_by_code() { let client = client(); - let (access_jwt, _did) = create_admin_account_and_login(&client).await; + let (access_jwt, admin_did) = create_admin_account_and_login(&client).await; let create_payload = json!({ - "useCount": 5 + "useCount": 5, + "forAccount": admin_did }); let create_res = client .post(format!( @@ -236,9 +155,10 @@ async fn test_disable_invite_codes_by_code() { .await .expect("Failed to send request"); assert_eq!(res.status(), StatusCode::OK); + let list_res = client .get(format!( - "{}/xrpc/com.atproto.server.getAccountInviteCodes", + "{}/xrpc/com.atproto.admin.getInviteCodes", base_url().await )) .bearer_auth(&access_jwt) @@ -258,7 +178,8 @@ async fn test_disable_invite_codes_by_account() { let (access_jwt, did) = create_admin_account_and_login(&client).await; for _ in 0..3 { let create_payload = json!({ - "useCount": 1 + "useCount": 1, + "forAccount": did }); let _ = client .post(format!( @@ -284,9 +205,10 @@ async fn test_disable_invite_codes_by_account() { .await .expect("Failed to send request"); assert_eq!(res.status(), StatusCode::OK); + let list_res = client .get(format!( - "{}/xrpc/com.atproto.server.getAccountInviteCodes", + "{}/xrpc/com.atproto.admin.getInviteCodes", base_url().await )) .bearer_auth(&access_jwt) @@ -295,7 +217,8 @@ async fn test_disable_invite_codes_by_account() { .expect("Failed to get invite codes"); let list_body: Value = list_res.json().await.unwrap(); let codes = list_body["codes"].as_array().unwrap(); - for code in codes { + let admin_codes: Vec<_> = codes.iter().filter(|c| c["forAccount"].as_str() == Some(&did)).collect(); + for code in admin_codes { assert_eq!(code["disabled"], true); } } diff --git a/tests/invite.rs b/tests/invite.rs index 3a1b7d9..7a6dc8b 100644 --- a/tests/invite.rs +++ b/tests/invite.rs @@ -6,7 +6,7 @@ use serde_json::{Value, json}; #[tokio::test] async fn test_create_invite_code_success() { let client = client(); - let (access_jwt, _did) = create_account_and_login(&client).await; + let (access_jwt, _did) = create_admin_account_and_login(&client).await; let payload = json!({ "useCount": 5 }); @@ -25,7 +25,9 @@ async fn test_create_invite_code_success() { assert!(body["code"].is_string()); let code = body["code"].as_str().unwrap(); assert!(!code.is_empty()); - assert!(code.contains('-'), "Code should be a UUID format"); + assert!(code.contains('-'), "Code should be in hostname-xxxxx-xxxxx format"); + let parts: Vec<&str> = code.split('-').collect(); + assert!(parts.len() >= 3, "Code should have at least 3 parts (hostname + 2 random parts)"); } #[tokio::test] @@ -49,9 +51,31 @@ async fn test_create_invite_code_no_auth() { } #[tokio::test] -async fn test_create_invite_code_invalid_use_count() { +async fn test_create_invite_code_non_admin() { let client = client(); let (access_jwt, _did) = create_account_and_login(&client).await; + let payload = json!({ + "useCount": 5 + }); + let res = client + .post(format!( + "{}/xrpc/com.atproto.server.createInviteCode", + base_url().await + )) + .bearer_auth(&access_jwt) + .json(&payload) + .send() + .await + .expect("Failed to send request"); + assert_eq!(res.status(), StatusCode::FORBIDDEN); + let body: Value = res.json().await.expect("Response was not valid JSON"); + assert_eq!(body["error"], "AdminRequired"); +} + +#[tokio::test] +async fn test_create_invite_code_invalid_use_count() { + let client = client(); + let (access_jwt, _did) = create_admin_account_and_login(&client).await; let payload = json!({ "useCount": 0 }); @@ -73,7 +97,7 @@ async fn test_create_invite_code_invalid_use_count() { #[tokio::test] async fn test_create_invite_code_for_another_account() { let client = client(); - let (access_jwt1, _did1) = create_account_and_login(&client).await; + let (access_jwt1, _did1) = create_admin_account_and_login(&client).await; let (_access_jwt2, did2) = create_account_and_login(&client).await; let payload = json!({ "useCount": 3, @@ -97,7 +121,7 @@ async fn test_create_invite_code_for_another_account() { #[tokio::test] async fn test_create_invite_codes_success() { let client = client(); - let (access_jwt, _did) = create_account_and_login(&client).await; + let (access_jwt, _did) = create_admin_account_and_login(&client).await; let payload = json!({ "useCount": 2, "codeCount": 3 @@ -117,13 +141,14 @@ async fn test_create_invite_codes_success() { assert!(body["codes"].is_array()); let codes = body["codes"].as_array().unwrap(); assert_eq!(codes.len(), 1); + assert_eq!(codes[0]["account"], "admin"); assert_eq!(codes[0]["codes"].as_array().unwrap().len(), 3); } #[tokio::test] async fn test_create_invite_codes_for_multiple_accounts() { let client = client(); - let (access_jwt1, did1) = create_account_and_login(&client).await; + let (access_jwt1, did1) = create_admin_account_and_login(&client).await; let (_access_jwt2, did2) = create_account_and_login(&client).await; let payload = json!({ "useCount": 1, @@ -169,28 +194,54 @@ async fn test_create_invite_codes_no_auth() { } #[tokio::test] -async fn test_get_account_invite_codes_success() { +async fn test_create_invite_codes_non_admin() { let client = client(); let (access_jwt, _did) = create_account_and_login(&client).await; + let payload = json!({ + "useCount": 2 + }); + let res = client + .post(format!( + "{}/xrpc/com.atproto.server.createInviteCodes", + base_url().await + )) + .bearer_auth(&access_jwt) + .json(&payload) + .send() + .await + .expect("Failed to send request"); + assert_eq!(res.status(), StatusCode::FORBIDDEN); + let body: Value = res.json().await.expect("Response was not valid JSON"); + assert_eq!(body["error"], "AdminRequired"); +} + +#[tokio::test] +async fn test_get_account_invite_codes_success() { + let client = client(); + let (admin_jwt, _admin_did) = create_admin_account_and_login(&client).await; + let (user_jwt, user_did) = create_account_and_login(&client).await; + let create_payload = json!({ - "useCount": 5 + "useCount": 5, + "forAccount": user_did }); let _ = client .post(format!( "{}/xrpc/com.atproto.server.createInviteCode", base_url().await )) - .bearer_auth(&access_jwt) + .bearer_auth(&admin_jwt) .json(&create_payload) .send() .await .expect("Failed to create invite code"); + let res = client .get(format!( "{}/xrpc/com.atproto.server.getAccountInviteCodes", base_url().await )) - .bearer_auth(&access_jwt) + .bearer_auth(&user_jwt) .send() .await .expect("Failed to send request"); @@ -205,6 +256,8 @@ async fn test_get_account_invite_codes_success() { assert!(code["disabled"].is_boolean()); assert!(code["createdAt"].is_string()); assert!(code["uses"].is_array()); + assert_eq!(code["forAccount"], user_did); + assert_eq!(code["createdBy"], "admin"); } #[tokio::test] @@ -224,26 +277,30 @@ async fn test_get_account_invite_codes_no_auth() { #[tokio::test] async fn test_get_account_invite_codes_include_used_filter() { let client = client(); - let (access_jwt, _did) = create_account_and_login(&client).await; + let (admin_jwt, _admin_did) = create_admin_account_and_login(&client).await; + let (user_jwt, user_did) = create_account_and_login(&client).await; + let create_payload = json!({ - "useCount": 5 + "useCount": 5, + "forAccount": user_did }); let _ = client .post(format!( "{}/xrpc/com.atproto.server.createInviteCode", base_url().await )) - .bearer_auth(&access_jwt) + .bearer_auth(&admin_jwt) .json(&create_payload) .send() .await .expect("Failed to create invite code"); + let res = client .get(format!( "{}/xrpc/com.atproto.server.getAccountInviteCodes", base_url().await )) - .bearer_auth(&access_jwt) + .bearer_auth(&user_jwt) .query(&[("includeUsed", "false")]) .send() .await @@ -255,3 +312,56 @@ async fn test_get_account_invite_codes_include_used_filter() { assert!(code["available"].as_i64().unwrap() > 0); } } + +#[tokio::test] +async fn test_get_account_invite_codes_filters_disabled() { + let client = client(); + let (admin_jwt, admin_did) = create_admin_account_and_login(&client).await; + + let create_payload = json!({ + "useCount": 5, + "forAccount": admin_did + }); + let create_res = client + .post(format!( + "{}/xrpc/com.atproto.server.createInviteCode", + base_url().await + )) + .bearer_auth(&admin_jwt) + .json(&create_payload) + .send() + .await + .expect("Failed to create invite code"); + let create_body: Value = create_res.json().await.unwrap(); + let code = create_body["code"].as_str().unwrap(); + + let disable_payload = json!({ + "codes": [code] + }); + let _ = client + .post(format!( + "{}/xrpc/com.atproto.admin.disableInviteCodes", + base_url().await + )) + .bearer_auth(&admin_jwt) + .json(&disable_payload) + .send() + .await + .expect("Failed to disable invite code"); + + let res = client + .get(format!( + "{}/xrpc/com.atproto.server.getAccountInviteCodes", + base_url().await + )) + .bearer_auth(&admin_jwt) + .send() + .await + .expect("Failed to send request"); + assert_eq!(res.status(), StatusCode::OK); + let body: Value = res.json().await.expect("Response was not valid JSON"); + let codes = body["codes"].as_array().unwrap(); + for c in codes { + assert_ne!(c["code"].as_str().unwrap(), code, "Disabled code should be filtered out"); + } +}