bitwarden_vault/cipher/cipher_client/
bulk_update_collections.rs1use std::collections::HashSet;
2
3use bitwarden_api_api::models::CipherBulkUpdateCollectionsRequestModel;
4use bitwarden_collections::collection::CollectionId;
5use bitwarden_core::{ApiError, OrganizationId};
6use bitwarden_error::bitwarden_error;
7use bitwarden_state::repository::{RepositoryError, RepositoryOption};
8use thiserror::Error;
9#[cfg(feature = "wasm")]
10use wasm_bindgen::prelude::wasm_bindgen;
11
12use crate::{CipherId, CiphersClient};
13
14#[allow(missing_docs)]
15#[bitwarden_error(flat)]
16#[derive(Debug, Error)]
17pub enum BulkUpdateCollectionsCipherError {
18 #[error(transparent)]
19 Api(#[from] ApiError),
20 #[error(transparent)]
21 Repository(#[from] RepositoryError),
22}
23
24#[cfg_attr(feature = "wasm", wasm_bindgen)]
25impl CiphersClient {
26 pub async fn bulk_update_collections(
31 &self,
32 organization_id: OrganizationId,
33 cipher_ids: Vec<CipherId>,
34 collection_ids: Vec<CollectionId>,
35 remove_collections: bool,
36 ) -> Result<(), BulkUpdateCollectionsCipherError> {
37 self.api_configurations
38 .api_client
39 .ciphers_api()
40 .post_bulk_collections(Some(CipherBulkUpdateCollectionsRequestModel {
41 organization_id: Some(organization_id.into()),
42 cipher_ids: Some(cipher_ids.iter().map(|id| (*id).into()).collect()),
43 collection_ids: Some(collection_ids.iter().map(|id| (*id).into()).collect()),
44 remove_collections: Some(remove_collections),
45 }))
46 .await?;
47
48 let repository = self.repository.require()?;
49 let mut updated_ciphers = Vec::new();
50 let collection_ids = collection_ids.iter().copied().collect::<HashSet<_>>();
51 for cipher_id in cipher_ids {
52 if let Some(mut cipher) = repository.get(cipher_id).await? {
53 if remove_collections {
54 cipher
55 .collection_ids
56 .retain(|id| !collection_ids.contains(id));
57 } else {
58 let existing = cipher
59 .collection_ids
60 .iter()
61 .copied()
62 .collect::<HashSet<_>>();
63 cipher.collection_ids = cipher
64 .collection_ids
65 .into_iter()
66 .chain(
67 collection_ids
68 .clone()
69 .into_iter()
70 .filter(|id| !existing.contains(id)),
71 )
72 .collect();
73 }
74 updated_ciphers.push((cipher_id, cipher));
75 }
76 }
77 repository.set_bulk(updated_ciphers).await?;
78
79 Ok(())
80 }
81}
82
83#[cfg(test)]
84mod tests {
85 use std::sync::Arc;
86
87 use bitwarden_api_api::apis::ApiClient;
88 use bitwarden_collections::collection::CollectionId;
89 use bitwarden_core::{
90 OrganizationId, client::ApiConfigurations, key_management::create_test_crypto_with_user_key,
91 };
92 use bitwarden_crypto::{SymmetricCryptoKey, SymmetricKeyAlgorithm};
93 use bitwarden_state::repository::Repository;
94 use bitwarden_test::MemoryRepository;
95
96 use crate::{Cipher, CipherId, CiphersClient};
97
98 const TEST_CIPHER_ID: &str = "5faa9684-c793-4a2d-8a12-b33900187097";
99 const TEST_ORG_ID: &str = "7faa9684-c793-4a2d-8a12-b33900187099";
100 const TEST_COLLECTION_ID_1: &str = "8faa9684-c793-4a2d-8a12-b33900187100";
101
102 fn generate_test_cipher() -> Cipher {
103 Cipher {
104 partial_data: None,
105 id: TEST_CIPHER_ID.parse().ok(),
106 name: Some("2.pMS6/icTQABtulw52pq2lg==|XXbxKxDTh+mWiN1HjH2N1w==|Q6PkuT+KX/axrgN9ubD5Ajk2YNwxQkgs3WJM0S0wtG8=".parse().unwrap()),
107 r#type: crate::CipherType::Login,
108 notes: Default::default(),
109 organization_id: Default::default(),
110 folder_id: Default::default(),
111 favorite: Default::default(),
112 reprompt: Default::default(),
113 fields: Default::default(),
114 collection_ids: Default::default(),
115 key: Default::default(),
116 login: Default::default(),
117 identity: Default::default(),
118 card: Default::default(),
119 secure_note: Default::default(),
120 ssh_key: Default::default(),
121 bank_account: Default::default(),
122 drivers_license: Default::default(),
123 passport: Default::default(),
124 organization_use_totp: Default::default(),
125 edit: Default::default(),
126 permissions: Default::default(),
127 view_password: Default::default(),
128 local_data: Default::default(),
129 attachments: Default::default(),
130 password_history: Default::default(),
131 creation_date: Default::default(),
132 deleted_date: Default::default(),
133 revision_date: Default::default(),
134 archived_date: Default::default(),
135 data: Default::default(),
136 }
137 }
138
139 fn create_test_client(api_client: ApiClient) -> (CiphersClient, Arc<MemoryRepository<Cipher>>) {
140 let repository = Arc::new(MemoryRepository::<Cipher>::default());
141 #[allow(deprecated)]
142 let client = CiphersClient {
143 key_store: create_test_crypto_with_user_key(SymmetricCryptoKey::make(
144 SymmetricKeyAlgorithm::Aes256CbcHmac,
145 )),
146 api_configurations: Arc::new(ApiConfigurations::from_api_client(api_client)),
147 repository: Some(repository.clone() as Arc<dyn Repository<Cipher>>),
148 client: bitwarden_core::Client::new_test(None),
149 };
150 (client, repository)
151 }
152
153 fn make_api_client() -> ApiClient {
154 ApiClient::new_mocked(|mock| {
155 mock.ciphers_api
156 .expect_post_bulk_collections()
157 .returning(|_| Ok(()));
158 })
159 }
160
161 #[tokio::test]
162 async fn test_bulk_update_adds_collections() {
163 let (client, repository) = create_test_client(make_api_client());
164
165 let cipher_id: CipherId = TEST_CIPHER_ID.parse().unwrap();
166 let org_id: OrganizationId = TEST_ORG_ID.parse().unwrap();
167 let collection_id: CollectionId = TEST_COLLECTION_ID_1.parse().unwrap();
168
169 repository
170 .set(cipher_id, generate_test_cipher())
171 .await
172 .unwrap();
173
174 client
175 .bulk_update_collections(org_id, vec![cipher_id], vec![collection_id], false)
176 .await
177 .unwrap();
178
179 let c: Cipher = repository.get(cipher_id).await.unwrap().unwrap();
180 assert!(c.collection_ids.contains(&collection_id));
181 }
182
183 #[tokio::test]
184 async fn test_bulk_update_removes_collections() {
185 let (client, repository) = create_test_client(make_api_client());
186
187 let cipher_id: CipherId = TEST_CIPHER_ID.parse().unwrap();
188 let org_id: OrganizationId = TEST_ORG_ID.parse().unwrap();
189 let collection_id: CollectionId = TEST_COLLECTION_ID_1.parse().unwrap();
190
191 let mut cipher = generate_test_cipher();
192 cipher.collection_ids = vec![collection_id];
193 repository.set(cipher_id, cipher).await.unwrap();
194
195 client
196 .bulk_update_collections(org_id, vec![cipher_id], vec![collection_id], true)
197 .await
198 .unwrap();
199
200 let c: Cipher = repository.get(cipher_id).await.unwrap().unwrap();
201 assert!(!c.collection_ids.contains(&collection_id));
202 }
203
204 #[tokio::test]
205 async fn test_bulk_update_no_duplicates_when_adding() {
206 let (client, repository) = create_test_client(make_api_client());
207
208 let cipher_id: CipherId = TEST_CIPHER_ID.parse().unwrap();
209 let org_id: OrganizationId = TEST_ORG_ID.parse().unwrap();
210 let collection_id: CollectionId = TEST_COLLECTION_ID_1.parse().unwrap();
211
212 let mut cipher = generate_test_cipher();
213 cipher.collection_ids = vec![collection_id];
214 repository.set(cipher_id, cipher).await.unwrap();
215
216 client
217 .bulk_update_collections(org_id, vec![cipher_id], vec![collection_id], false)
218 .await
219 .unwrap();
220
221 let c: Cipher = repository.get(cipher_id).await.unwrap().unwrap();
222 assert_eq!(
223 c.collection_ids.len(),
224 1,
225 "no duplicates introduced when collection already present"
226 );
227 }
228
229 #[tokio::test]
230 async fn test_bulk_update_skips_missing_ciphers() {
231 let (client, _repository) = create_test_client(make_api_client());
232
233 let cipher_id: CipherId = TEST_CIPHER_ID.parse().unwrap();
234 let org_id: OrganizationId = TEST_ORG_ID.parse().unwrap();
235 let collection_id: CollectionId = TEST_COLLECTION_ID_1.parse().unwrap();
236
237 let result = client
238 .bulk_update_collections(org_id, vec![cipher_id], vec![collection_id], false)
239 .await;
240 assert!(result.is_ok());
241 }
242}