Skip to main content

bitwarden_crypto/safe/
helpers.rs

1use std::fmt::DebugStruct;
2
3use ciborium::Value;
4
5use crate::cose::{
6    ContentNamespace, SAFE_CONTENT_NAMESPACE, SAFE_OBJECT_NAMESPACE, SafeObjectNamespace,
7    extract_integer, symmetric::CoseContentEncryptionAlgorithm,
8};
9
10#[derive(Debug)]
11pub(super) enum ExtractionError {
12    MissingNamespace,
13    InvalidNamespace,
14}
15
16pub(super) fn extract_safe_object_namespace(
17    header: &coset::Header,
18) -> Result<SafeObjectNamespace, ExtractionError> {
19    match extract_integer(header, SAFE_OBJECT_NAMESPACE, "safe object namespace") {
20        Ok(value) => value
21            .try_into()
22            .map_err(|_| ExtractionError::InvalidNamespace),
23        Err(_) => Err(ExtractionError::MissingNamespace),
24    }
25}
26
27pub(super) fn extract_safe_content_namespace<T: ContentNamespace>(
28    header: &coset::Header,
29) -> Result<T, ExtractionError> {
30    match extract_integer(header, SAFE_CONTENT_NAMESPACE, "safe content namespace") {
31        Ok(value) => value
32            .try_into()
33            .map_err(|_| ExtractionError::InvalidNamespace),
34        Err(_) => Err(ExtractionError::MissingNamespace),
35    }
36}
37
38pub(super) fn debug_fmt<C: ContentNamespace>(
39    debug_struct: &mut DebugStruct,
40    header: &coset::Header,
41) {
42    if let Ok(object_namespace) = extract_safe_object_namespace(header) {
43        debug_struct.field("object_namespace", &object_namespace);
44    }
45    if let Ok(content_namespace) = extract_safe_content_namespace::<C>(header) {
46        debug_struct.field("content_namespace", &content_namespace);
47    }
48    if let Some(algorithm) = header.alg.as_ref()
49        && let Ok(content_encryption_algorithm) =
50            CoseContentEncryptionAlgorithm::try_from(algorithm)
51    {
52        let label = match content_encryption_algorithm {
53            CoseContentEncryptionAlgorithm::Aes256Gcm => "AES-256-GCM",
54            CoseContentEncryptionAlgorithm::XAes256Gcm => "XAES-256-GCM",
55            CoseContentEncryptionAlgorithm::XChaCha20Poly1305 => "XChaCha20-Poly1305",
56            CoseContentEncryptionAlgorithm::Aes256CbcHmacSha256 => "AES-256-CBC-HMAC-SHA256-AEAD",
57        };
58        debug_struct.field("content_encryption_algorithm", &label);
59    }
60}
61
62pub(super) fn set_header_value(header: &mut coset::Header, label: i64, value: Value) {
63    if let Some((_, existing_value)) =
64        header
65            .rest
66            .iter_mut()
67            .find(|(existing_label, _)| matches!(existing_label, coset::Label::Int(existing) if *existing == label))
68    {
69        *existing_value = value;
70    } else {
71        header.rest.push((coset::Label::Int(label), value));
72    }
73}
74
75pub(super) fn set_safe_namespaces<T: ContentNamespace>(
76    header: &mut coset::Header,
77    object_namespace: SafeObjectNamespace,
78    content_namespace: T,
79) {
80    set_header_value(
81        header,
82        SAFE_OBJECT_NAMESPACE,
83        Value::from(i128::from(object_namespace)),
84    );
85    set_header_value(
86        header,
87        SAFE_CONTENT_NAMESPACE,
88        Value::from(content_namespace.into()),
89    );
90}
91
92/// Validates the provided header contains the expected object and content namespace.
93/// For backward compatibility, missing values are OK, but incorrect values are not.
94/// The validation happens individually for both namespace layers, and either one
95/// missing with the other being present is OK.
96pub(super) fn validate_safe_namespaces<T: ContentNamespace>(
97    header: &coset::Header,
98    expected_object_namespace: SafeObjectNamespace,
99    expected_content_namespace: T,
100) -> Result<(), ExtractionError> {
101    match extract_safe_object_namespace(header) {
102        Ok(ns) if ns == expected_object_namespace => (),
103        // If the namespace is present but doesn't match, return an error immediately.
104        Ok(_) => return Err(ExtractionError::InvalidNamespace),
105        // If the namespace is missing, do not validate for backward compatibility
106        Err(ExtractionError::MissingNamespace) => (),
107        // If the namespace is present but invalid (e.g., not an integer or out of range), return an
108        // error.
109        Err(ExtractionError::InvalidNamespace) => return Err(ExtractionError::InvalidNamespace),
110    }
111
112    match extract_safe_content_namespace::<T>(header) {
113        Ok(ns) if ns == expected_content_namespace => Ok(()),
114        // If the namespace is present but doesn't match, return an error immediately.
115        Ok(_) => Err(ExtractionError::InvalidNamespace),
116        // If the namespace is missing, do not validate for backward compatibility
117        Err(ExtractionError::MissingNamespace) => Ok(()),
118        // If the namespace is present but invalid (e.g., not an integer or out of range), return an
119        // error.
120        Err(ExtractionError::InvalidNamespace) => Err(ExtractionError::InvalidNamespace),
121    }
122}
123
124#[cfg(test)]
125mod tests {
126    use ciborium::Value;
127
128    use super::*;
129    use crate::{cose::SAFE_OBJECT_NAMESPACE, safe::DataEnvelopeNamespace};
130
131    fn count_label(header: &coset::Header, label: i64) -> usize {
132        header
133            .rest
134            .iter()
135            .filter(
136                |(existing_label, _)| {
137                    matches!(existing_label, coset::Label::Int(existing) if *existing == label)
138                },
139            )
140            .count()
141    }
142
143    fn extract_safe_namespaces<T: ContentNamespace>(
144        header: &coset::Header,
145    ) -> Result<(SafeObjectNamespace, T), ExtractionError> {
146        let object_namespace = extract_safe_object_namespace(header)?;
147        let content_namespace = extract_safe_content_namespace(header)?;
148
149        Ok((object_namespace, content_namespace))
150    }
151
152    #[test]
153    fn set_safe_namespaces_sets_both_namespace_labels() {
154        let mut header = coset::HeaderBuilder::new().build();
155
156        set_safe_namespaces(
157            &mut header,
158            SafeObjectNamespace::DataEnvelope,
159            DataEnvelopeNamespace::ExampleNamespace,
160        );
161
162        let extracted = extract_safe_namespaces::<DataEnvelopeNamespace>(&header);
163        assert!(matches!(
164            extracted,
165            Ok((
166                SafeObjectNamespace::DataEnvelope,
167                DataEnvelopeNamespace::ExampleNamespace
168            ))
169        ));
170    }
171
172    #[test]
173    fn set_safe_namespaces_overwrites_existing_namespace_values() {
174        let mut header = coset::HeaderBuilder::new()
175            .value(SAFE_OBJECT_NAMESPACE, Value::from(999_i64))
176            .value(SAFE_CONTENT_NAMESPACE, Value::from(999_i64))
177            .build();
178
179        set_safe_namespaces(
180            &mut header,
181            SafeObjectNamespace::DataEnvelope,
182            DataEnvelopeNamespace::ExampleNamespace,
183        );
184
185        assert_eq!(count_label(&header, SAFE_OBJECT_NAMESPACE), 1);
186        assert_eq!(count_label(&header, SAFE_CONTENT_NAMESPACE), 1);
187        assert!(matches!(
188            extract_safe_namespaces::<DataEnvelopeNamespace>(&header),
189            Ok((
190                SafeObjectNamespace::DataEnvelope,
191                DataEnvelopeNamespace::ExampleNamespace
192            ))
193        ));
194    }
195
196    #[test]
197    fn extract_safe_namespaces_fails_when_namespace_missing() {
198        let header = coset::HeaderBuilder::new().build();
199
200        assert!(matches!(
201            extract_safe_namespaces::<DataEnvelopeNamespace>(&header),
202            Err(ExtractionError::MissingNamespace)
203        ));
204    }
205
206    #[test]
207    fn extract_safe_namespaces_fails_when_namespace_invalid() {
208        let header = coset::HeaderBuilder::new()
209            .value(
210                SAFE_OBJECT_NAMESPACE,
211                Value::from(SafeObjectNamespace::DataEnvelope as i64),
212            )
213            .value(SAFE_CONTENT_NAMESPACE, Value::from(999_i64))
214            .build();
215
216        assert!(matches!(
217            extract_safe_namespaces::<DataEnvelopeNamespace>(&header),
218            Err(ExtractionError::InvalidNamespace)
219        ));
220    }
221
222    #[test]
223    fn validate_safe_namespaces_allows_missing_labels_for_backwards_compat() {
224        let header = coset::HeaderBuilder::new().build();
225
226        let result = validate_safe_namespaces(
227            &header,
228            SafeObjectNamespace::DataEnvelope,
229            DataEnvelopeNamespace::ExampleNamespace,
230        );
231        assert!(result.is_ok());
232    }
233
234    #[test]
235    fn validate_safe_namespaces_rejects_namespace_mismatch() {
236        let mut header = coset::HeaderBuilder::new().build();
237        set_safe_namespaces(
238            &mut header,
239            SafeObjectNamespace::DataEnvelope,
240            DataEnvelopeNamespace::ExampleNamespace,
241        );
242
243        let result = validate_safe_namespaces(
244            &header,
245            SafeObjectNamespace::DataEnvelope,
246            DataEnvelopeNamespace::ExampleNamespace2,
247        );
248        assert!(matches!(result, Err(ExtractionError::InvalidNamespace)));
249    }
250}