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
92pub(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 Ok(_) => return Err(ExtractionError::InvalidNamespace),
105 Err(ExtractionError::MissingNamespace) => (),
107 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 Ok(_) => Err(ExtractionError::InvalidNamespace),
116 Err(ExtractionError::MissingNamespace) => Ok(()),
118 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}