Skip to main content

ahri_tre_secrets/
acquisition.rs

1//! Acquisition-only material. It cannot be cloned or serialized, and owns no store.
2use crate::{MAX_SECRET_MATERIAL_BYTES, SecretMaterial};
3use ahri_tre_types::SourceAuthentication;
4use secrecy::zeroize::{Zeroize, Zeroizing};
5use std::io::Read;
6
7pub struct SourceCredentials {
8    kind: SourceAuthentication,
9    material: Option<SecretMaterial>,
10}
11
12impl std::fmt::Debug for SourceCredentials {
13    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
14        f.debug_struct("SourceCredentials")
15            .field("kind", &self.kind)
16            .finish_non_exhaustive()
17    }
18}
19
20#[derive(Debug, Clone, Copy, thiserror::Error)]
21#[error("Source authentication material is invalid or unavailable")]
22pub struct SourceCredentialError;
23
24impl SourceCredentials {
25    /// Call only after reserving acquisition capacity. Always bounded, including rejection.
26    pub fn read(
27        kind: SourceAuthentication,
28        reader: impl Read,
29    ) -> Result<Self, SourceCredentialError> {
30        let mut bytes = Zeroizing::new(Vec::new());
31        reader
32            .take(MAX_SECRET_MATERIAL_BYTES as u64 + 1)
33            .read_to_end(&mut bytes)
34            .map_err(|_| SourceCredentialError)?;
35        Self::new(kind, std::mem::take(&mut *bytes))
36    }
37    pub fn new(
38        kind: SourceAuthentication,
39        mut bytes: Vec<u8>,
40    ) -> Result<Self, SourceCredentialError> {
41        if kind == SourceAuthentication::None {
42            let empty = bytes.is_empty();
43            bytes.zeroize();
44            return if empty {
45                Ok(Self {
46                    kind,
47                    material: None,
48                })
49            } else {
50                Err(SourceCredentialError)
51            };
52        }
53        let material = SecretMaterial::try_from(bytes).map_err(|_| SourceCredentialError)?;
54        let valid = material.expose(|bytes| match kind {
55            SourceAuthentication::Basic
56            | SourceAuthentication::PostgreSql
57            | SourceAuthentication::MsSql => {
58                serde_json::from_slice::<Password>(bytes).is_ok_and(|value| {
59                    !value.username.0.is_empty()
60                        && !value.password.0.is_empty()
61                        && !value.username.0.contains(['\r', '\n'])
62                        && (kind != SourceAuthentication::Basic || !value.username.0.contains(':'))
63                        && !value.password.0.contains(['\r', '\n'])
64                })
65            }
66            _ => std::str::from_utf8(bytes)
67                .is_ok_and(|value| !value.is_empty() && !value.chars().any(char::is_control)),
68        });
69        if !valid {
70            return Err(SourceCredentialError);
71        }
72        Ok(Self {
73            kind,
74            material: Some(material),
75        })
76    }
77    pub fn kind(&self) -> SourceAuthentication {
78        self.kind
79    }
80    pub fn expose<R>(&self, use_material: impl FnOnce(&[u8]) -> R) -> R {
81        match &self.material {
82            Some(material) => material.expose(use_material),
83            None => use_material(&[]),
84        }
85    }
86    pub fn password<R>(
87        &self,
88        use_material: impl FnOnce(&str, &str) -> R,
89    ) -> Result<R, SourceCredentialError> {
90        self.expose(|bytes| {
91            let value: Password =
92                serde_json::from_slice(bytes).map_err(|_| SourceCredentialError)?;
93            Ok(use_material(&value.username.0, &value.password.0))
94        })
95    }
96}
97#[derive(serde::Deserialize)]
98#[serde(deny_unknown_fields)]
99struct Password {
100    username: PasswordText,
101    password: PasswordText,
102}
103// Each decoded field clears itself even if a later field is invalid and the
104// containing Password is never constructed.
105struct PasswordText(Zeroizing<String>);
106impl<'de> serde::Deserialize<'de> for PasswordText {
107    fn deserialize<D: serde::Deserializer<'de>>(d: D) -> Result<Self, D::Error> {
108        let value = <String as serde::Deserialize>::deserialize(d)?;
109        Ok(Self(Zeroizing::new(value)))
110    }
111}