Skip to main content

bitwarden_organization_invite_link/
invite_link_user_client.rs

1use std::sync::Arc;
2
3use bitwarden_api_api::models::{
4    AcceptOrganizationInviteLinkRequestModel, ConfirmOrganizationInviteLinkRequestModel,
5    GetOrganizationInviteLinkStatusRequestModel, GetOrganizationInviteRequestModel,
6    OrganizationInviteLinkValidateEmailDomainRequestModel,
7};
8use bitwarden_core::{
9    FromClient, OrganizationId,
10    client::ApiConfigurations,
11    key_management::{KeySlotIds, PrivateKeySlotId, SymmetricKeySlotId},
12    require,
13};
14use bitwarden_crypto::{
15    CoseKeyThumbprintExt, KeyStore, PrimitiveEncryptable, PublicKey, SpkiPublicKeyBytes,
16    UnsignedSharedKey,
17};
18use bitwarden_encoding::B64;
19use bitwarden_organization_crypto::invite::{Invite, InviteSecret};
20#[cfg(feature = "wasm")]
21use wasm_bindgen::prelude::wasm_bindgen;
22
23use crate::{InviteLinkError, OrganizationInviteLinkStatusView};
24
25/// Client for organization invite link invitee (user) operations: checking link status, validating
26/// email eligibility, and accepting or self-confirming an invite.
27#[cfg_attr(feature = "wasm", wasm_bindgen)]
28#[derive(FromClient)]
29pub struct InviteLinkUserClient {
30    pub(crate) key_store: KeyStore<KeySlotIds>,
31    pub(crate) api_configurations: Arc<ApiConfigurations>,
32}
33
34#[cfg_attr(feature = "wasm", wasm_bindgen)]
35impl InviteLinkUserClient {
36    /// Retrieves the status of an invite link.
37    /// Used to verify basic availability before attempting to accept.
38    pub async fn get_status(
39        &self,
40        organization_id: OrganizationId,
41        code: String,
42    ) -> Result<OrganizationInviteLinkStatusView, InviteLinkError> {
43        let code =
44            uuid::Uuid::parse_str(&code).map_err(|_| InviteLinkError::ParseFailure("code"))?;
45
46        let response = self
47            .api_configurations
48            .api_client
49            .organization_invite_links_api()
50            .get_status(Some(GetOrganizationInviteLinkStatusRequestModel {
51                organization_id: organization_id.into(),
52                code,
53            }))
54            .await?;
55
56        OrganizationInviteLinkStatusView::try_from(response)
57    }
58
59    /// Returns whether the given email address is in the allowed domains for an invite link.
60    pub async fn is_email_allowed(
61        &self,
62        organization_id: OrganizationId,
63        code: String,
64        email: String,
65    ) -> Result<bool, InviteLinkError> {
66        let code =
67            uuid::Uuid::parse_str(&code).map_err(|_| InviteLinkError::ParseFailure("code"))?;
68
69        let response = self
70            .api_configurations
71            .api_client
72            .organization_invite_links_api()
73            .validate_email_domain(Some(
74                OrganizationInviteLinkValidateEmailDomainRequestModel {
75                    organization_id: organization_id.into(),
76                    code,
77                    email,
78                },
79            ))
80            .await?;
81
82        Ok(require!(response.is_allowed))
83    }
84
85    /// Accepts an organization invite for the current user, optionally enrolling into account
86    /// recovery (when `enroll_into_account_recovery` is set) and — when the invite supports
87    /// confirmation — self-confirming.
88    pub async fn accept_and_optionally_confirm(
89        &self,
90        organization_id: OrganizationId,
91        code: String,
92        invite_secret: InviteSecret,
93        default_collection_name: String,
94        enroll_into_account_recovery: bool,
95    ) -> Result<(), InviteLinkError> {
96        let code =
97            uuid::Uuid::parse_str(&code).map_err(|_| InviteLinkError::ParseFailure("code"))?;
98
99        // When enrolling into account recovery, fetch the organization's public key (which is the
100        // account-recovery public key) from the server.
101        let recovery_public_key = if enroll_into_account_recovery {
102            let response = self
103                .api_configurations
104                .api_client
105                .organizations_api()
106                .get_public_key(&organization_id.to_string())
107                .await?;
108            Some(
109                require!(response.public_key)
110                    .parse::<B64>()
111                    .map_err(|_| InviteLinkError::ParseFailure("public_key"))?,
112            )
113        } else {
114            None
115        };
116
117        let invite_response = self
118            .api_configurations
119            .api_client
120            .organization_users_api()
121            .get_invite(Some(GetOrganizationInviteRequestModel {
122                organization_id: organization_id.into(),
123                code,
124            }))
125            .await?;
126
127        let invite: Invite = require!(invite_response.invite).parse()?;
128
129        // Confine the (non-Send) key store context to a synchronous scope; it produces the owned
130        // request payload consumed after the `.await`s below.
131        let request = {
132            let mut ctx = self.key_store.context();
133
134            // Recover the invite key from the invite secret the invitee holds.
135            let invite_key =
136                invite.unseal_invite_key_with_invite_secret(&invite_secret, &mut ctx)?;
137
138            // Enroll into account recovery when requested. Verify the account-recovery public key
139            // against the organization public-key thumbprint bound into the invite before
140            // enrolling: a substituted recovery key would not match, so the organization key cannot
141            // be captured by an attacker-supplied key. Then encapsulate the user key to it.
142            let reset_password_key = match &recovery_public_key {
143                Some(recovery_public_key) => {
144                    let recovery_public_key =
145                        PublicKey::from_der(&SpkiPublicKeyBytes::from(recovery_public_key))?;
146                    let bound_thumbprint =
147                        invite.get_public_key_thumbprint(invite_key, &mut ctx)?;
148                    if bound_thumbprint != recovery_public_key.thumbprint()? {
149                        return Err(InviteLinkError::RecoveryKeyMismatch);
150                    }
151                    Some(
152                        UnsignedSharedKey::encapsulate(
153                            SymmetricKeySlotId::User,
154                            &recovery_public_key,
155                            &ctx,
156                        )?
157                        .to_string(),
158                    )
159                }
160                None => None,
161            };
162
163            if invite.supports_confirmation() {
164                // Self-confirm: recover the organization key and encapsulate it to the user.
165                let org_key = invite.unseal_organization_key(invite_key, &mut ctx)?;
166                let user_public_key = ctx.get_public_key(PrivateKeySlotId::UserPrivateKey)?;
167                let org_user_key =
168                    UnsignedSharedKey::encapsulate(org_key, &user_public_key, &ctx)?.to_string();
169                let default_user_collection_name = default_collection_name
170                    .encrypt(&mut ctx, org_key)?
171                    .to_string();
172                PendingPost::Confirm(ConfirmOrganizationInviteLinkRequestModel {
173                    organization_id: organization_id.into(),
174                    code,
175                    org_user_key,
176                    reset_password_key,
177                    default_user_collection_name,
178                })
179            } else {
180                PendingPost::Accept(AcceptOrganizationInviteLinkRequestModel {
181                    organization_id: organization_id.into(),
182                    code,
183                    reset_password_key,
184                })
185            }
186        };
187
188        let organization_users_api = self.api_configurations.api_client.organization_users_api();
189        match request {
190            PendingPost::Confirm(model) => {
191                organization_users_api
192                    .confirm_invite_link(Some(model))
193                    .await?
194            }
195            PendingPost::Accept(model) => {
196                organization_users_api
197                    .accept_invite_link(Some(model))
198                    .await?
199            }
200        }
201
202        Ok(())
203    }
204}
205
206/// A prepared invite acceptance request, built while the key store context is held and posted once
207/// it has been dropped.
208enum PendingPost {
209    Confirm(ConfirmOrganizationInviteLinkRequestModel),
210    Accept(AcceptOrganizationInviteLinkRequestModel),
211}
212
213#[cfg(test)]
214mod tests {
215    use bitwarden_api_api::{
216        apis::ApiClient,
217        models::{
218            OrganizationInviteLinkSsoResponseModel, OrganizationInviteLinkStatusResponseModel,
219            OrganizationInviteLinkValidateEmailDomainResponseModel,
220            OrganizationInviteResponseModel, OrganizationPublicKeyResponseModel,
221        },
222    };
223    use bitwarden_core::{
224        client::ApiConfigurations, key_management::create_test_crypto_with_user_and_org_key,
225    };
226    use bitwarden_crypto::{
227        PublicKeyEncryptionAlgorithm, SymmetricCryptoKey, SymmetricKeyAlgorithm,
228    };
229
230    use super::*;
231
232    fn make_client(org_id: OrganizationId, api_client: ApiClient) -> InviteLinkUserClient {
233        let user_key = SymmetricCryptoKey::make(SymmetricKeyAlgorithm::Aes256CbcHmac);
234        let org_key = SymmetricCryptoKey::make(SymmetricKeyAlgorithm::Aes256CbcHmac);
235        let key_store = create_test_crypto_with_user_and_org_key(user_key, org_id, org_key);
236        // Give the store a user private key so the confirmation branch can derive a user public
237        // key.
238        {
239            let mut ctx = key_store.context_mut();
240            let local = ctx.make_private_key(PublicKeyEncryptionAlgorithm::RsaOaepSha1);
241            ctx.persist_private_key(local, PrivateKeySlotId::UserPrivateKey)
242                .expect("persisting the user private key should work");
243        }
244        InviteLinkUserClient {
245            key_store,
246            api_configurations: Arc::new(ApiConfigurations::from_api_client(api_client)),
247        }
248    }
249
250    /// Builds an invite + its secret and the organization public key it binds, all consistent with
251    /// the client's org key.
252    fn build_invite(
253        client: &InviteLinkUserClient,
254        org_id: OrganizationId,
255    ) -> (InviteSecret, Invite, B64) {
256        let mut ctx = client.key_store.context();
257        let org_key = SymmetricKeySlotId::Organization(org_id);
258        let private_key = ctx.make_private_key(PublicKeyEncryptionAlgorithm::RsaOaepSha1);
259        let org_public_key = B64::from(
260            ctx.get_public_key(private_key)
261                .unwrap()
262                .to_der()
263                .unwrap()
264                .as_ref(),
265        );
266        let wrapped = ctx.wrap_private_key(org_key, private_key).unwrap();
267        let (secret, invite) = Invite::make_for_private_key(org_key, &wrapped, &mut ctx).unwrap();
268        (secret, invite, org_public_key)
269    }
270
271    #[tokio::test]
272    async fn accept_and_confirm_succeeds_for_confirmable_invite() {
273        let org_id = OrganizationId::new_v4();
274        // `get_public_key` returns the base64 key held in this cell, and `get_invite` returns the
275        // serialized invite; both are filled after the invite is generated below.
276        let recovery = Arc::new(std::sync::Mutex::new(None::<String>));
277        let invite_cell = Arc::new(std::sync::Mutex::new(None::<String>));
278        let recovery_mock = recovery.clone();
279        let invite_mock = invite_cell.clone();
280        let client = make_client(
281            org_id,
282            ApiClient::new_mocked(move |mock| {
283                mock.organizations_api
284                    .expect_get_public_key()
285                    .returning(move |_id| {
286                        Ok(OrganizationPublicKeyResponseModel {
287                            object: None,
288                            public_key: recovery_mock.lock().unwrap().clone(),
289                        })
290                    })
291                    .once();
292                mock.organization_users_api
293                    .expect_get_invite()
294                    .returning(move |_model| {
295                        Ok(OrganizationInviteResponseModel {
296                            invite: invite_mock.lock().unwrap().clone(),
297                        })
298                    })
299                    .once();
300                mock.organization_users_api
301                    .expect_confirm_invite_link()
302                    .returning(|_model| Ok(()))
303                    .once();
304            }),
305        );
306
307        let (secret, invite, org_public_key) = build_invite(&client, org_id);
308        assert!(invite.supports_confirmation());
309        // The recovery public key returned by the "server" matches the invite's bound org key.
310        *recovery.lock().unwrap() = Some(String::from(&org_public_key));
311        *invite_cell.lock().unwrap() = Some(String::from(&invite));
312
313        client
314            .accept_and_optionally_confirm(
315                org_id,
316                uuid::Uuid::new_v4().to_string(),
317                secret,
318                "Default".to_string(),
319                true,
320            )
321            .await
322            .unwrap();
323    }
324
325    #[tokio::test]
326    async fn accept_without_enrollment_confirms_without_recovery_key() {
327        let org_id = OrganizationId::new_v4();
328        // Without enrollment the recovery key is never fetched, so only `get_invite` and
329        // `confirm_invite_link` run.
330        let invite_cell = Arc::new(std::sync::Mutex::new(None::<String>));
331        let invite_mock = invite_cell.clone();
332        let client = make_client(
333            org_id,
334            ApiClient::new_mocked(move |mock| {
335                mock.organization_users_api
336                    .expect_get_invite()
337                    .returning(move |_model| {
338                        Ok(OrganizationInviteResponseModel {
339                            invite: invite_mock.lock().unwrap().clone(),
340                        })
341                    })
342                    .once();
343                mock.organization_users_api
344                    .expect_confirm_invite_link()
345                    .returning(|_model| Ok(()))
346                    .once();
347            }),
348        );
349
350        let (secret, invite, _org_public_key) = build_invite(&client, org_id);
351        *invite_cell.lock().unwrap() = Some(String::from(&invite));
352        client
353            .accept_and_optionally_confirm(
354                org_id,
355                uuid::Uuid::new_v4().to_string(),
356                secret,
357                "Default".to_string(),
358                false,
359            )
360            .await
361            .unwrap();
362    }
363
364    #[tokio::test]
365    async fn accept_without_confirmation_posts_acceptance() {
366        let org_id = OrganizationId::new_v4();
367        let invite_cell = Arc::new(std::sync::Mutex::new(None::<String>));
368        let invite_mock = invite_cell.clone();
369        let client = make_client(
370            org_id,
371            ApiClient::new_mocked(move |mock| {
372                mock.organization_users_api
373                    .expect_get_invite()
374                    .returning(move |_model| {
375                        Ok(OrganizationInviteResponseModel {
376                            invite: invite_mock.lock().unwrap().clone(),
377                        })
378                    })
379                    .once();
380                mock.organization_users_api
381                    .expect_accept_invite_link()
382                    .returning(|_model| Ok(()))
383                    .once();
384            }),
385        );
386
387        // An invite with confirmation disabled routes to the acceptance branch.
388        let (secret, mut invite, _org_public_key) = build_invite(&client, org_id);
389        invite.disable_confirmation();
390        assert!(!invite.supports_confirmation());
391        *invite_cell.lock().unwrap() = Some(String::from(&invite));
392
393        client
394            .accept_and_optionally_confirm(
395                org_id,
396                uuid::Uuid::new_v4().to_string(),
397                secret,
398                "Default".to_string(),
399                false,
400            )
401            .await
402            .unwrap();
403    }
404
405    #[tokio::test]
406    async fn accept_with_mismatched_recovery_key_fails() {
407        let org_id = OrganizationId::new_v4();
408        let recovery = Arc::new(std::sync::Mutex::new(None::<String>));
409        let invite_cell = Arc::new(std::sync::Mutex::new(None::<String>));
410        let recovery_mock = recovery.clone();
411        let invite_mock = invite_cell.clone();
412        let client = make_client(
413            org_id,
414            ApiClient::new_mocked(move |mock| {
415                mock.organizations_api
416                    .expect_get_public_key()
417                    .returning(move |_id| {
418                        Ok(OrganizationPublicKeyResponseModel {
419                            object: None,
420                            public_key: recovery_mock.lock().unwrap().clone(),
421                        })
422                    })
423                    .once();
424                mock.organization_users_api
425                    .expect_get_invite()
426                    .returning(move |_model| {
427                        Ok(OrganizationInviteResponseModel {
428                            invite: invite_mock.lock().unwrap().clone(),
429                        })
430                    })
431                    .once();
432            }),
433        );
434
435        let (secret, invite, _org_public_key) = build_invite(&client, org_id);
436        *invite_cell.lock().unwrap() = Some(String::from(&invite));
437        // The "server" returns an unrelated public key that must not match the invite's bound
438        // thumbprint.
439        let unrelated_public_key = {
440            let mut ctx = client.key_store.context();
441            let private_key = ctx.make_private_key(PublicKeyEncryptionAlgorithm::RsaOaepSha1);
442            B64::from(
443                ctx.get_public_key(private_key)
444                    .unwrap()
445                    .to_der()
446                    .unwrap()
447                    .as_ref(),
448            )
449        };
450        *recovery.lock().unwrap() = Some(String::from(&unrelated_public_key));
451
452        let result = client
453            .accept_and_optionally_confirm(
454                org_id,
455                uuid::Uuid::new_v4().to_string(),
456                secret,
457                "Default".to_string(),
458                true,
459            )
460            .await;
461
462        assert!(matches!(result, Err(InviteLinkError::RecoveryKeyMismatch)));
463    }
464
465    #[tokio::test]
466    async fn get_status_returns_mapped_view() {
467        let org_id = OrganizationId::new_v4();
468        let client = make_client(
469            org_id,
470            ApiClient::new_mocked(|mock| {
471                mock.organization_invite_links_api
472                    .expect_get_status()
473                    .returning(|_model| {
474                        Ok(OrganizationInviteLinkStatusResponseModel {
475                            object: None,
476                            organization_name: Some("Test Org".to_string()),
477                            links_enabled: Some(true),
478                            seats_available: Some(true),
479                            supports_confirmation: Some(true),
480                            sso: Some(Box::new(OrganizationInviteLinkSsoResponseModel {
481                                object: None,
482                                org_sso_id: Some("sso-id".to_string()),
483                                required: Some(true),
484                            })),
485                        })
486                    })
487                    .once();
488            }),
489        );
490
491        let status = client
492            .get_status(org_id, uuid::Uuid::new_v4().to_string())
493            .await
494            .unwrap();
495
496        assert_eq!(status.organization_name, "Test Org");
497        assert!(status.links_enabled);
498        assert!(status.seats_available);
499        assert!(status.supports_confirmation);
500        let sso = status.sso.expect("sso should be present");
501        assert_eq!(sso.org_sso_id.as_deref(), Some("sso-id"));
502        assert!(sso.required);
503    }
504
505    #[tokio::test]
506    async fn is_email_allowed_returns_is_allowed() {
507        let org_id = OrganizationId::new_v4();
508        let client = make_client(
509            org_id,
510            ApiClient::new_mocked(|mock| {
511                mock.organization_invite_links_api
512                    .expect_validate_email_domain()
513                    .returning(|_model| {
514                        Ok(OrganizationInviteLinkValidateEmailDomainResponseModel {
515                            is_allowed: Some(true),
516                        })
517                    })
518                    .once();
519            }),
520        );
521
522        let allowed = client
523            .is_email_allowed(
524                org_id,
525                uuid::Uuid::new_v4().to_string(),
526                "[email protected]".to_string(),
527            )
528            .await
529            .unwrap();
530
531        assert!(allowed);
532    }
533}