Skip to main content

bitwarden_vault/cipher/cipher_client/
bulk_update_collections.rs

1use 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    /// Updates collection membership for multiple [`Cipher`](crate::Cipher) objects.
27    ///
28    /// When `remove_collections` is `true`, the given collection IDs are removed from each cipher.
29    /// When `false`, they are added without introducing duplicates.
30    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}