1use std::collections::HashSet;
4
5use archive_trait::{
6 Archive as ArchiveTrait, Member, MemberMetadata, MemberPayload as MemberPayloadTrait,
7 SpecialKind,
8};
9use tar_framing::{
10 ArchiveFormat, FrameError, PaxKeyword, PaxKind, PaxRecord, PaxValue, StreamPolicy, UstarKind,
11 logical::{MemberExtensions, MemberFrame, MemberPayload as FramingMemberPayload, TarReader},
12};
13use thiserror::Error;
14use tokio::io::AsyncRead;
15
16pub use tar_framing::{
17 DEFAULT_MAX_GLOBAL_PAX_EXTENSIONS_SIZE, DEFAULT_MAX_GNU_EXTENSION_SIZE,
18 DEFAULT_MAX_PAX_EXTENSION_SIZE,
19};
20
21pub struct TarArchive<R> {
26 reader: TarReader<R>,
27 policy: DecodePolicy,
28 fused: bool,
29}
30
31impl<R> TarArchive<R> {
32 pub fn new(reader: R) -> Self {
34 Self {
35 reader: TarReader::new(reader),
36 policy: DecodePolicy::default(),
37 fused: false,
38 }
39 }
40
41 pub fn with_policy(mut self, policy: DecodePolicy) -> Self {
45 let stream_policy = StreamPolicy::default()
46 .max_pax_extension_size(policy.pax_policy.max_extension_size)
47 .max_global_pax_extensions_size(policy.pax_policy.max_global_extensions_size)
48 .allow_all_nul_numeric_fields(policy.allow_all_nul_numeric_fields)
49 .max_gnu_extension_size(policy.max_gnu_extension_size);
50 self.reader = self.reader.with_policy(stream_policy);
51 self.policy = policy;
52 self
53 }
54}
55
56#[derive(Clone, Debug)]
60pub struct DecodePolicy {
61 allow_gnu: bool,
62 allow_all_nul_numeric_fields: bool,
63 max_gnu_extension_size: u64,
64 pax_policy: PaxDecodePolicy,
65}
66
67#[derive(Clone, Debug, Eq, PartialEq)]
71pub struct PaxDecodePolicy {
72 max_extension_size: u64,
73 max_global_extensions_size: u64,
74 allow_non_utf8_pax_vendor_values: bool,
75 allow_global_pax_extensions: bool,
76 vendor_extension_policy: PaxVendorExtensionPolicy,
77 allow_duplicate_pax_records: bool,
78 allow_global_pax_member_metadata: bool,
79}
80
81#[derive(Clone, Debug, Default, Eq, PartialEq)]
83pub enum PaxVendorExtensionPolicy {
84 #[default]
86 RejectUnknown,
87 Ignore(PaxVendorExtensionAllowlist),
91 AllowUnknown,
95}
96
97impl PaxVendorExtensionPolicy {
98 pub fn ignore(keywords: impl IntoIterator<Item = &'static str>) -> Self {
102 Self::Ignore(PaxVendorExtensionAllowlist {
103 keywords: keywords.into_iter().collect(),
104 })
105 }
106}
107
108#[derive(Clone, Debug, Eq, PartialEq)]
112pub struct PaxVendorExtensionAllowlist {
113 keywords: HashSet<&'static str>,
114}
115
116impl Default for PaxDecodePolicy {
117 fn default() -> Self {
118 Self {
119 max_extension_size: DEFAULT_MAX_PAX_EXTENSION_SIZE,
120 max_global_extensions_size: DEFAULT_MAX_GLOBAL_PAX_EXTENSIONS_SIZE,
121 allow_non_utf8_pax_vendor_values: true,
122 allow_global_pax_extensions: true,
123 vendor_extension_policy: PaxVendorExtensionPolicy::default(),
124 allow_duplicate_pax_records: false,
125 allow_global_pax_member_metadata: false,
126 }
127 }
128}
129
130impl Default for DecodePolicy {
131 fn default() -> Self {
132 Self {
133 allow_gnu: true,
134 allow_all_nul_numeric_fields: true,
135 max_gnu_extension_size: DEFAULT_MAX_GNU_EXTENSION_SIZE,
136 pax_policy: PaxDecodePolicy::default(),
137 }
138 }
139}
140
141impl DecodePolicy {
142 pub fn allow_gnu(mut self, allow: bool) -> Self {
149 self.allow_gnu = allow;
150 self
151 }
152
153 pub fn allow_all_nul_numeric_fields(mut self, allow: bool) -> Self {
161 self.allow_all_nul_numeric_fields = allow;
162 self
163 }
164
165 pub fn max_gnu_extension_size(mut self, max_gnu_extension_size: u64) -> Self {
173 self.max_gnu_extension_size = max_gnu_extension_size;
174 self
175 }
176
177 pub fn pax_policy(mut self, policy: PaxDecodePolicy) -> Self {
179 self.pax_policy = policy;
180 self
181 }
182
183 fn check_format(&self, position: u64, format: ArchiveFormat) -> Result<(), DecodeError> {
184 if format == ArchiveFormat::Gnu && !self.allow_gnu {
185 return Err(DecodeError::policy_violation(
186 position,
187 DecodePolicyViolation::GnuArchive,
188 ));
189 }
190 Ok(())
191 }
192
193 fn check_global_pax(&self, position: u64, records: &[PaxRecord]) -> Result<(), DecodeError> {
194 self.pax_policy.check_global_pax_extension(position)?;
195 self.pax_policy
196 .check_pax_records(position, PaxKind::Global, records)
197 }
198
199 fn check_member<R>(&self, frame: &MemberFrame<'_, R>) -> Result<(), DecodeError> {
200 if let MemberExtensions::Pax(state) = &frame.extensions {
201 for extension in state
202 .extensions()
203 .filter(|extension| extension.kind == PaxKind::Global)
204 {
205 self.check_global_pax(extension.position, extension.records())?;
206 }
207 }
208 let format_position = match &frame.extensions {
209 MemberExtensions::Pax(_) => frame.header.position,
210 MemberExtensions::Gnu {
211 long_name,
212 long_link,
213 } => long_name
214 .iter()
215 .chain(long_link.iter())
216 .map(|header| header.position)
217 .min()
218 .unwrap_or(frame.header.position),
219 };
220 self.check_format(format_position, frame.header.format)?;
221 if let MemberExtensions::Pax(state) = &frame.extensions {
222 for extension in state
223 .extensions()
224 .filter(|extension| extension.kind == PaxKind::Local)
225 {
226 self.pax_policy.check_pax_records(
227 extension.position,
228 PaxKind::Local,
229 extension.records(),
230 )?;
231 }
232 }
233 Ok(())
234 }
235}
236
237impl PaxDecodePolicy {
238 pub fn max_extension_size(mut self, max_extension_size: u64) -> Self {
249 self.max_extension_size = max_extension_size;
250 self
251 }
252
253 pub fn max_global_extensions_size(mut self, max_global_extensions_size: u64) -> Self {
263 self.max_global_extensions_size = max_global_extensions_size;
264 self
265 }
266
267 pub fn allow_non_utf8_pax_vendor_values(mut self, allow: bool) -> Self {
277 self.allow_non_utf8_pax_vendor_values = allow;
278 self
279 }
280
281 pub fn allow_global_pax_extensions(mut self, allow: bool) -> Self {
290 self.allow_global_pax_extensions = allow;
291 self
292 }
293
294 pub fn vendor_extension_policy(mut self, policy: PaxVendorExtensionPolicy) -> Self {
310 self.vendor_extension_policy = policy;
311 self
312 }
313
314 pub fn allow_duplicate_pax_records(mut self, allow: bool) -> Self {
321 self.allow_duplicate_pax_records = allow;
322 self
323 }
324
325 pub fn allow_global_pax_member_metadata(mut self, allow: bool) -> Self {
333 self.allow_global_pax_member_metadata = allow;
334 self
335 }
336
337 fn check_global_pax_extension(&self, position: u64) -> Result<(), DecodeError> {
338 if !self.allow_global_pax_extensions {
339 return Err(DecodeError::policy_violation(
340 position,
341 DecodePolicyViolation::GlobalPaxExtension,
342 ));
343 }
344 Ok(())
345 }
346
347 fn check_pax_records(
348 &self,
349 position: u64,
350 kind: PaxKind,
351 records: &[PaxRecord],
352 ) -> Result<(), DecodeError> {
353 for record in records {
354 if let PaxRecord::Vendor {
355 vendor,
356 name,
357 value,
358 } = record
359 {
360 let allowed = match &self.vendor_extension_policy {
361 PaxVendorExtensionPolicy::RejectUnknown => false,
362 PaxVendorExtensionPolicy::Ignore(allowed) => allowed
363 .keywords
364 .contains(format!("{vendor}.{name}").as_str()),
365 PaxVendorExtensionPolicy::AllowUnknown => true,
366 };
367 if !allowed {
368 return Err(DecodeError::policy_violation(
369 position,
370 DecodePolicyViolation::PaxVendorExtension {
371 vendor: vendor.to_string(),
372 name: name.to_string(),
373 },
374 ));
375 }
376
377 if !self.allow_non_utf8_pax_vendor_values
378 && let PaxValue::Value(value) = value
379 && std::str::from_utf8(value).is_err()
380 {
381 return Err(DecodeError::policy_violation(
382 position,
383 DecodePolicyViolation::NonUtf8PaxVendorValue {
384 vendor: vendor.to_string(),
385 name: name.to_string(),
386 },
387 ));
388 }
389 }
390 }
391
392 if kind == PaxKind::Global && !self.allow_global_pax_member_metadata {
393 for record in records {
394 let keyword = match record.keyword() {
395 PaxKeyword::Path => Some("path"),
396 PaxKeyword::LinkPath => Some("linkpath"),
397 PaxKeyword::Size => Some("size"),
398 _ => None,
399 };
400 if let Some(keyword) = keyword {
401 return Err(DecodeError::policy_violation(
402 position,
403 DecodePolicyViolation::GlobalPaxMemberMetadata { keyword },
404 ));
405 }
406 }
407 }
408
409 if !self.allow_duplicate_pax_records {
410 let mut keywords = HashSet::new();
411 for record in records {
412 let keyword = record.keyword();
413 if !keywords.insert(keyword.clone()) {
414 return Err(DecodeError::policy_violation(
415 position,
416 DecodePolicyViolation::DuplicatePaxRecord {
417 keyword: keyword.to_string(),
418 },
419 ));
420 }
421 }
422 }
423
424 Ok(())
425 }
426}
427
428#[derive(Clone, Debug, Eq, PartialEq, Error)]
430pub enum DecodePolicyViolation {
431 #[error("GNU archives are not allowed")]
433 GnuArchive,
434 #[error("global pax extended headers are not allowed")]
436 GlobalPaxExtension,
437 #[error("pax vendor extension {vendor}.{name} is not allowed")]
439 PaxVendorExtension {
440 vendor: String,
442 name: String,
444 },
445 #[error("pax vendor extension {vendor}.{name} contains a non-UTF-8 value")]
447 NonUtf8PaxVendorValue {
448 vendor: String,
450 name: String,
452 },
453 #[error("pax extended header contains duplicate record {keyword}")]
455 DuplicatePaxRecord {
456 keyword: String,
458 },
459 #[error("global pax extended header contains restricted member metadata {keyword}")]
461 GlobalPaxMemberMetadata {
462 keyword: &'static str,
464 },
465}
466
467#[derive(Debug, Error)]
469pub enum DecodeError {
470 #[error(transparent)]
472 Framing(#[from] FrameError),
473 #[error("at byte {position}: {field} is not valid UTF-8")]
475 InvalidUtf8 {
476 position: u64,
478 field: &'static str,
480 },
481 #[error("at byte {position}: decode policy rejected input: {violation}")]
483 PolicyViolation {
484 position: u64,
486 violation: DecodePolicyViolation,
488 },
489}
490
491impl DecodeError {
492 fn policy_violation(position: u64, violation: DecodePolicyViolation) -> Self {
493 Self::PolicyViolation {
494 position,
495 violation,
496 }
497 }
498}
499
500pub struct TarMemberPayload<'a, R> {
502 payload: FramingMemberPayload<'a, R>,
503}
504
505impl<R: AsyncRead + Unpin> MemberPayloadTrait for TarMemberPayload<'_, R> {
506 type Error = DecodeError;
507
508 async fn next_chunk(
509 &mut self,
510 buffer: &mut Vec<u8>,
511 target_len: usize,
512 ) -> Result<bool, Self::Error> {
513 self.payload
514 .next_chunk(buffer, target_len)
515 .await
516 .map_err(Into::into)
517 }
518
519 async fn skip(self) -> Result<(), Self::Error> {
520 self.payload.skip().await.map_err(Into::into)
521 }
522}
523
524impl<R: AsyncRead + Unpin> ArchiveTrait for TarArchive<R> {
525 type Error = DecodeError;
526 type Payload<'a>
527 = TarMemberPayload<'a, R>
528 where
529 Self: 'a;
530
531 async fn next_member<'a>(
532 &'a mut self,
533 ) -> Result<Option<Member<Self::Payload<'a>>>, Self::Error> {
534 if self.fused {
535 return Ok(None);
536 }
537
538 let frame = match self.reader.next_frame().await {
539 Ok(Some(frame)) => frame,
540 Ok(None) => {
541 self.fused = true;
542 return Ok(None);
543 }
544 Err(error) => {
545 self.fused = true;
546 return Err(error.into());
547 }
548 };
549
550 if let Err(error) = self.policy.check_member(&frame) {
551 self.fused = true;
552 return Err(error);
553 }
554
555 match project_member(frame) {
556 Ok(member) => Ok(Some(member)),
557 Err(error) => {
558 self.fused = true;
559 Err(error)
560 }
561 }
562 }
563}
564
565fn project_member<'a, R>(
566 frame: MemberFrame<'a, R>,
567) -> Result<Member<TarMemberPayload<'a, R>>, DecodeError> {
568 let position = frame.header.position;
569 let kind = frame.header.kind;
570 let size = frame.header.effective_size;
571 let executable = frame.header.mode.unwrap_or_default() & 0o111 != 0;
572 let path = std::str::from_utf8(frame.effective_path()?.as_ref())
573 .map(str::to_owned)
574 .map_err(|_| DecodeError::InvalidUtf8 {
575 position,
576 field: "path",
577 })?;
578 let target = if matches!(kind, UstarKind::HardLink | UstarKind::SymbolicLink) {
579 std::str::from_utf8(frame.effective_link_path()?.as_ref())
580 .map(str::to_owned)
581 .map_err(|_| DecodeError::InvalidUtf8 {
582 position,
583 field: "linkpath",
584 })?
585 } else {
586 String::new()
587 };
588 let metadata = MemberMetadata { path, position };
589
590 Ok(match kind {
591 UstarKind::Regular | UstarKind::Contiguous => Member::File {
592 metadata,
593 size,
594 executable,
595 payload: TarMemberPayload {
596 payload: frame.payload,
597 },
598 },
599 UstarKind::Directory => Member::Directory { metadata },
600 UstarKind::SymbolicLink => Member::SymbolicLink { metadata, target },
601 UstarKind::HardLink => Member::HardLink {
602 metadata,
603 target,
604 size,
605 payload: TarMemberPayload {
606 payload: frame.payload,
607 },
608 },
609 UstarKind::CharacterDevice => Member::Special {
610 metadata,
611 kind: SpecialKind::CharacterDevice,
612 },
613 UstarKind::BlockDevice => Member::Special {
614 metadata,
615 kind: SpecialKind::BlockDevice,
616 },
617 UstarKind::Fifo => Member::Special {
618 metadata,
619 kind: SpecialKind::Fifo,
620 },
621 })
622}