1use crate::cipher::{
14 decrypt, decrypt_owned_in_place, encrypt, payload_digest, random_nonce, verify_payload_digest,
15};
16use crate::compression::{decode_payload, encode_payload};
17use crate::flags::validate_flags;
18use crate::header::{encode_header, HeaderParts};
19use crate::model_types::{DntContext, DntOpenOptions, DntSealOptions, OpenedDnt, VerifiedDnt};
20use crate::plaintext::{decode_plaintext_ranges, encode_plaintext};
21use crate::{
22 inspect_header, DntAlgorithm, DntCodec, DntError, DntHeader, DntKeyProvider, DntResult,
23};
24use std::ops::Range;
25use zeroize::Zeroize;
26
27pub fn seal<P, C>(
29 payload: &[u8],
30 key_provider: &P,
31 codec: &C,
32 options: DntSealOptions,
33) -> DntResult<Vec<u8>>
34where
35 P: DntKeyProvider,
36 C: DntCodec,
37{
38 enforce_max(payload.len() as u64, options.max_payload_bytes)?;
39 validate_flags(options.flags)?;
40 let encoded = codec.encode(payload)?;
41 enforce_max(encoded.len() as u64, options.max_payload_bytes)?;
42 let mut stored_payload = encode_payload(options.flags, encoded, options.max_payload_bytes)?;
43 enforce_max(stored_payload.len() as u64, options.max_payload_bytes)?;
44 let encrypted_metadata_length =
45 u32::try_from(options.encrypted_metadata.len()).map_err(|_| DntError::PayloadTooLarge)?;
46 let codec_id = codec.codec_id();
47 let context = DntContext {
48 application_id: options.application_id.clone(),
49 tenant_id: options.tenant_id.clone(),
50 content_type: options.content_type.clone(),
51 codec_id: codec_id.clone(),
52 schema_version: options.schema_version,
53 };
54 let key = key_provider.resolve_key(&options.key_id, &context)?;
55 let nonce = random_nonce()?;
56 let payload_hash = match payload_digest(&key, &stored_payload) {
57 Ok(payload_hash) => payload_hash,
58 Err(error) => {
59 stored_payload.zeroize();
60 return Err(error);
61 }
62 };
63 let payload_length = match u64::try_from(stored_payload.len()) {
64 Ok(payload_length) => payload_length,
65 Err(_) => {
66 stored_payload.zeroize();
67 return Err(DntError::PayloadTooLarge);
68 }
69 };
70 let header = match encode_header(HeaderParts {
71 flags: options.flags,
72 algorithm: DntAlgorithm::XChaCha20Poly1305,
73 application_id: options.application_id,
74 tenant_id: options.tenant_id,
75 content_type: options.content_type,
76 codec_id,
77 key_id: options.key_id,
78 schema_version: options.schema_version,
79 created_at_ms: options.created_at_ms,
80 payload_length,
81 nonce,
82 payload_hash,
83 public_metadata: options.public_metadata,
84 encrypted_metadata_length,
85 }) {
86 Ok(header) => header,
87 Err(error) => {
88 stored_payload.zeroize();
89 return Err(error);
90 }
91 };
92 let plaintext_result = encode_plaintext(&options.encrypted_metadata, &stored_payload);
93 stored_payload.zeroize();
94 let mut plaintext = plaintext_result?;
95 let ciphertext_result = encrypt(&key, &nonce, &header, &plaintext);
96 plaintext.zeroize();
97 let ciphertext = ciphertext_result?;
98 let mut output = Vec::with_capacity(header.len() + ciphertext.len());
99 output.extend_from_slice(&header);
100 output.extend_from_slice(&ciphertext);
101 Ok(output)
102}
103
104pub fn open<P, C>(
106 input: &[u8],
107 key_provider: &P,
108 codec: &C,
109 options: &DntOpenOptions,
110) -> DntResult<OpenedDnt>
111where
112 P: DntKeyProvider,
113 C: DntCodec,
114{
115 open_internal(input, key_provider, codec, options)
116}
117
118pub fn open_owned<P, C>(
124 input: Vec<u8>,
125 key_provider: &P,
126 codec: &C,
127 options: &DntOpenOptions,
128) -> DntResult<OpenedDnt>
129where
130 P: DntKeyProvider,
131 C: DntCodec,
132{
133 open_owned_internal(input, key_provider, codec, options)
134}
135
136pub fn verify<P, C>(
138 input: &[u8],
139 key_provider: &P,
140 codec: &C,
141 options: &DntOpenOptions,
142) -> DntResult<VerifiedDnt>
143where
144 P: DntKeyProvider,
145 C: DntCodec,
146{
147 let mut authenticated = authenticate_internal(input, key_provider, codec, options)?;
148 if authenticated.header.compression().is_compacted() {
149 let encoded_result = decode_payload(
150 authenticated.header.flags,
151 authenticated.stored_payload(),
152 options.max_payload_bytes,
153 );
154 authenticated.zeroize_buffer();
155 match encoded_result {
156 Ok(mut encoded_payload) => encoded_payload.zeroize(),
157 Err(error) => {
158 authenticated.encrypted_metadata.zeroize();
159 return Err(error);
160 }
161 }
162 } else {
163 authenticated.zeroize_buffer();
164 }
165 authenticated.encrypted_metadata.zeroize();
166 Ok(VerifiedDnt {
167 header: authenticated.header,
168 })
169}
170
171struct AuthenticatedDnt {
172 header: DntHeader,
173 encrypted_metadata: Vec<u8>,
174 buffer: Vec<u8>,
175 payload_range: Range<usize>,
176}
177
178impl AuthenticatedDnt {
179 fn stored_payload(&self) -> &[u8] {
180 &self.buffer[self.payload_range.clone()]
181 }
182
183 fn take_payload_buffer(&mut self) -> Vec<u8> {
184 let payload_start = self.payload_range.start;
185 let payload_end = self.payload_range.end;
186 let payload_len = payload_end - payload_start;
187 if payload_start != 0 {
188 self.buffer.copy_within(payload_start..payload_end, 0);
189 }
190 self.buffer[payload_len..payload_end].zeroize();
191 self.buffer.truncate(payload_len);
192 std::mem::take(&mut self.buffer)
193 }
194
195 fn zeroize_buffer(&mut self) {
196 self.buffer.zeroize();
197 }
198}
199
200fn open_internal<P, C>(
201 input: &[u8],
202 key_provider: &P,
203 codec: &C,
204 options: &DntOpenOptions,
205) -> DntResult<OpenedDnt>
206where
207 P: DntKeyProvider,
208 C: DntCodec,
209{
210 let authenticated = authenticate_internal(input, key_provider, codec, options)?;
211 open_authenticated(authenticated, codec, options)
212}
213
214fn open_owned_internal<P, C>(
215 input: Vec<u8>,
216 key_provider: &P,
217 codec: &C,
218 options: &DntOpenOptions,
219) -> DntResult<OpenedDnt>
220where
221 P: DntKeyProvider,
222 C: DntCodec,
223{
224 let authenticated = authenticate_owned_internal(input, key_provider, codec, options)?;
225 open_authenticated(authenticated, codec, options)
226}
227
228fn open_authenticated<C>(
229 mut authenticated: AuthenticatedDnt,
230 codec: &C,
231 options: &DntOpenOptions,
232) -> DntResult<OpenedDnt>
233where
234 C: DntCodec,
235{
236 let payload_result = if authenticated.header.compression().is_compacted() {
237 let encoded_result = decode_payload(
238 authenticated.header.flags,
239 authenticated.stored_payload(),
240 options.max_payload_bytes,
241 );
242 authenticated.zeroize_buffer();
243 let encoded_payload = match encoded_result {
244 Ok(encoded_payload) => encoded_payload,
245 Err(error) => {
246 authenticated.encrypted_metadata.zeroize();
247 return Err(error);
248 }
249 };
250 codec.decode_owned(encoded_payload)
251 } else {
252 codec.decode_owned(authenticated.take_payload_buffer())
253 };
254 let mut payload = match payload_result {
255 Ok(payload) => payload,
256 Err(error) => {
257 authenticated.encrypted_metadata.zeroize();
258 return Err(error.into());
259 }
260 };
261 if let Err(error) = enforce_max(payload.len() as u64, options.max_payload_bytes) {
262 payload.zeroize();
263 authenticated.encrypted_metadata.zeroize();
264 return Err(error);
265 }
266 let opened = OpenedDnt {
267 header: authenticated.header,
268 payload,
269 encrypted_metadata: authenticated.encrypted_metadata,
270 };
271 Ok(opened)
272}
273
274fn authenticate_internal<P, C>(
275 input: &[u8],
276 key_provider: &P,
277 codec: &C,
278 options: &DntOpenOptions,
279) -> DntResult<AuthenticatedDnt>
280where
281 P: DntKeyProvider,
282 C: DntCodec,
283{
284 let header = inspect_header(input)?;
285 validate_header_context(&header, codec, options)?;
286 enforce_max(header.payload_length, options.max_payload_bytes)?;
287 let header_len = header.header_length as usize;
288 let ciphertext = input.get(header_len..).ok_or(DntError::InvalidFormat)?;
289 if ciphertext.is_empty() {
290 return Err(DntError::InvalidFormat);
291 }
292 let context = DntContext::from_header(&header);
293 let key = key_provider.resolve_key(&header.key_id, &context)?;
294 let mut plaintext = decrypt(&key, header.nonce(), &input[..header_len], ciphertext)?;
295 let (encrypted_metadata_range, payload_start) =
296 match decode_plaintext_ranges(&plaintext, &header) {
297 Ok(ranges) => ranges,
298 Err(error) => {
299 plaintext.zeroize();
300 return Err(error);
301 }
302 };
303 let digest_matches =
304 match verify_payload_digest(&key, &plaintext[payload_start..], &header.payload_hash) {
305 Ok(matches) => matches,
306 Err(error) => {
307 plaintext.zeroize();
308 return Err(error);
309 }
310 };
311 if !digest_matches {
312 plaintext.zeroize();
313 return Err(DntError::AuthenticationFailed);
314 }
315 let encrypted_metadata = plaintext[encrypted_metadata_range].to_vec();
316 let payload_range = payload_start..plaintext.len();
317 let authenticated = AuthenticatedDnt {
318 header,
319 encrypted_metadata,
320 buffer: plaintext,
321 payload_range,
322 };
323 Ok(authenticated)
324}
325
326fn authenticate_owned_internal<P, C>(
327 mut input: Vec<u8>,
328 key_provider: &P,
329 codec: &C,
330 options: &DntOpenOptions,
331) -> DntResult<AuthenticatedDnt>
332where
333 P: DntKeyProvider,
334 C: DntCodec,
335{
336 let header = inspect_header(&input)?;
337 validate_header_context(&header, codec, options)?;
338 enforce_max(header.payload_length, options.max_payload_bytes)?;
339 let header_len = header.header_length as usize;
340 if input.len() <= header_len {
341 return Err(DntError::InvalidFormat);
342 }
343 let context = DntContext::from_header(&header);
344 let key = key_provider.resolve_key(&header.key_id, &context)?;
345 let plaintext_range = match decrypt_owned_in_place(&key, header.nonce(), header_len, &mut input)
346 {
347 Ok(range) => range,
348 Err(error) => {
349 input.zeroize();
350 return Err(error);
351 }
352 };
353 let (encrypted_metadata_range, payload_start) =
354 match decode_plaintext_ranges(&input[plaintext_range.clone()], &header) {
355 Ok(ranges) => ranges,
356 Err(error) => {
357 input.zeroize();
358 return Err(error);
359 }
360 };
361 let payload_start = plaintext_range.start + payload_start;
362 let encrypted_metadata_range = plaintext_range.start + encrypted_metadata_range.start
363 ..plaintext_range.start + encrypted_metadata_range.end;
364 let digest_matches = match verify_payload_digest(
365 &key,
366 &input[payload_start..plaintext_range.end],
367 &header.payload_hash,
368 ) {
369 Ok(matches) => matches,
370 Err(error) => {
371 input.zeroize();
372 return Err(error);
373 }
374 };
375 if !digest_matches {
376 input.zeroize();
377 return Err(DntError::AuthenticationFailed);
378 }
379 let encrypted_metadata = input[encrypted_metadata_range].to_vec();
380 let authenticated = AuthenticatedDnt {
381 header,
382 encrypted_metadata,
383 buffer: input,
384 payload_range: payload_start..plaintext_range.end,
385 };
386 Ok(authenticated)
387}
388
389fn validate_header_context<C>(
390 header: &DntHeader,
391 codec: &C,
392 options: &DntOpenOptions,
393) -> DntResult<()>
394where
395 C: DntCodec,
396{
397 if header.application_id != options.application_id
398 || header.tenant_id != options.tenant_id
399 || header.content_type != options.content_type
400 {
401 return Err(DntError::ContextMismatch);
402 }
403 if !codec.matches_codec_id(&header.codec_id) {
404 return Err(DntError::CodecUnavailable);
405 }
406 Ok(())
407}
408
409fn enforce_max(actual: u64, max: Option<u64>) -> DntResult<()> {
410 if max.is_some_and(|max| actual > max) {
411 return Err(DntError::PayloadTooLarge);
412 }
413 Ok(())
414}