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#[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 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 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 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 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 let request = {
132 let mut ctx = self.key_store.context();
133
134 let invite_key =
136 invite.unseal_invite_key_with_invite_secret(&invite_secret, &mut ctx)?;
137
138 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 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
206enum 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 {
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 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 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 *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 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 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 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}