use serde::{Deserialize, Serialize};
use std::fmt;
use zeroize::{Zeroize, ZeroizeOnDrop};
#[derive(Clone, Zeroize, ZeroizeOnDrop)]
pub struct SecretString {
inner: String,
}
impl<'de> Deserialize<'de> for SecretString {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: serde::Deserializer<'de>,
{
struct SecretStringVisitor;
impl<'de> serde::de::Visitor<'de> for SecretStringVisitor {
type Value = SecretString;
fn expecting(&self, formatter: &mut std::fmt::Formatter) -> std::fmt::Result {
formatter.write_str("a string or a map with a 'secret', 'env', or 'file' key")
}
fn visit_str<E>(self, value: &str) -> Result<SecretString, E>
where
E: serde::de::Error,
{
Ok(SecretString {
inner: value.to_owned(),
})
}
fn visit_map<M>(self, mut map: M) -> Result<SecretString, M::Error>
where
M: serde::de::MapAccess<'de>,
{
let mut secret = None;
while let Some(key) = map.next_key::<String>()? {
match key.as_str() {
"secret" => {
set_secret_source(&mut secret, map.next_value::<String>()?)?;
}
"env" => {
let name = map.next_value::<String>()?;
ensure_no_secret_source(&secret)?;
let value = std::env::var(&name).map_err(|error| {
serde::de::Error::custom(format!(
"failed to read secret from environment variable `{name}`: {error}"
))
})?;
set_secret_source(&mut secret, value)?;
}
"file" => {
let path = map.next_value::<String>()?;
ensure_no_secret_source(&secret)?;
let value = std::fs::read_to_string(&path).map_err(|error| {
serde::de::Error::custom(format!(
"failed to read secret from file `{path}`: {error}"
))
})?;
set_secret_source(&mut secret, trim_secret_file_newline(value))?;
}
_ => {
let _ = map.next_value::<serde::de::IgnoredAny>()?;
}
}
}
secret
.map(|s| SecretString { inner: s })
.ok_or_else(|| serde::de::Error::missing_field("secret, env, or file"))
}
}
deserializer.deserialize_any(SecretStringVisitor)
}
}
fn ensure_no_secret_source<E>(secret: &Option<String>) -> Result<(), E>
where
E: serde::de::Error,
{
if secret.is_some() {
return Err(duplicate_secret_source_error());
}
Ok(())
}
fn set_secret_source<E>(secret: &mut Option<String>, value: String) -> Result<(), E>
where
E: serde::de::Error,
{
if secret.replace(value).is_some() {
return Err(duplicate_secret_source_error());
}
Ok(())
}
fn duplicate_secret_source_error<E>() -> E
where
E: serde::de::Error,
{
E::custom("secret source maps must contain only one of `secret`, `env`, or `file`")
}
fn trim_secret_file_newline(mut value: String) -> String {
while value.ends_with('\n') || value.ends_with('\r') {
value.pop();
}
value
}
impl Serialize for SecretString {
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
where
S: serde::Serializer,
{
serializer.serialize_str("[REDACTED]")
}
}
impl SecretString {
pub fn new(secret: impl Into<String>) -> Self {
Self {
inner: secret.into(),
}
}
pub fn expose_secret(&self) -> &str {
&self.inner
}
pub fn into_inner(self) -> String {
let this = std::mem::ManuallyDrop::new(self);
unsafe { std::ptr::read(&this.inner) }
}
pub fn len(&self) -> usize {
self.inner.len()
}
pub fn is_empty(&self) -> bool {
self.inner.is_empty()
}
}
impl fmt::Debug for SecretString {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("SecretString([REDACTED])")
}
}
impl fmt::Display for SecretString {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("[REDACTED]")
}
}
impl From<String> for SecretString {
fn from(s: String) -> Self {
Self::new(s)
}
}
impl From<&str> for SecretString {
fn from(s: &str) -> Self {
Self::new(s.to_string())
}
}
impl PartialEq for SecretString {
fn eq(&self, other: &Self) -> bool {
use subtle::ConstantTimeEq;
self.inner.as_bytes().ct_eq(other.inner.as_bytes()).into()
}
}
impl Eq for SecretString {}
#[derive(Clone, Deserialize, Zeroize, ZeroizeOnDrop)]
pub struct SecretValue<T: Zeroize> {
#[serde(bound(deserialize = "T: Deserialize<'de>"))]
inner: T,
}
impl<T: Zeroize> Serialize for SecretValue<T> {
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
where
S: serde::Serializer,
{
serializer.serialize_str("[REDACTED]")
}
}
impl<T: Zeroize> SecretValue<T> {
pub fn new(value: T) -> Self {
Self { inner: value }
}
pub fn expose_secret(&self) -> &T {
&self.inner
}
pub fn into_inner(self) -> T {
let this = std::mem::ManuallyDrop::new(self);
unsafe { std::ptr::read(&this.inner) }
}
}
impl<T: Zeroize> fmt::Debug for SecretValue<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("SecretValue([REDACTED])")
}
}
impl<T: Zeroize> fmt::Display for SecretValue<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("[REDACTED]")
}
}
impl<T: Zeroize> From<T> for SecretValue<T> {
fn from(value: T) -> Self {
Self::new(value)
}
}
impl<T: Zeroize + AsRef<[u8]>> PartialEq for SecretValue<T> {
fn eq(&self, other: &Self) -> bool {
use subtle::ConstantTimeEq;
self.inner.as_ref().ct_eq(other.inner.as_ref()).into()
}
}
impl<T: Zeroize + AsRef<[u8]> + Eq> Eq for SecretValue<T> {}
#[cfg(test)]
mod tests {
use super::*;
use rstest::rstest;
use serial_test::serial;
#[rstest]
fn test_secret_string_debug() {
let secret = SecretString::new("my-secret-password");
let debug_output = format!("{:?}", secret);
assert!(!debug_output.contains("my-secret-password"));
assert!(debug_output.contains("REDACTED"));
}
#[rstest]
fn test_secret_string_display() {
let secret = SecretString::new("my-secret-password");
let display_output = format!("{}", secret);
assert!(!display_output.contains("my-secret-password"));
assert!(display_output.contains("REDACTED"));
}
#[rstest]
fn test_secret_string_expose() {
let secret = SecretString::new("my-secret-password");
assert_eq!(secret.expose_secret(), "my-secret-password");
}
#[rstest]
fn test_secret_string_len() {
let secret = SecretString::new("password");
assert_eq!(secret.len(), 8);
assert!(!secret.is_empty());
let empty = SecretString::new("");
assert_eq!(empty.len(), 0);
assert!(empty.is_empty());
}
#[rstest]
fn test_secret_string_equality() {
let secret1 = SecretString::new("password");
let secret2 = SecretString::new("password");
let secret3 = SecretString::new("different");
assert_eq!(secret1, secret2);
assert_ne!(secret1, secret3);
}
#[rstest]
fn test_secret_value_constant_time_equality() {
let val1 = SecretValue::new(vec![1u8, 2, 3, 4]);
let val2 = SecretValue::new(vec![1u8, 2, 3, 4]);
let val3 = SecretValue::new(vec![5u8, 6, 7, 8]);
assert_eq!(val1, val2);
assert_ne!(val1, val3);
}
#[rstest]
fn test_secret_value_constant_time_equality_strings() {
let val1 = SecretValue::new("secret_token".to_string());
let val2 = SecretValue::new("secret_token".to_string());
let val3 = SecretValue::new("different_token".to_string());
assert_eq!(val1, val2);
assert_ne!(val1, val3);
}
#[rstest]
fn test_secret_value_debug() {
let secret = SecretValue::new(12345);
let debug_output = format!("{:?}", secret);
assert!(!debug_output.contains("12345"));
assert!(debug_output.contains("REDACTED"));
}
#[rstest]
fn test_secret_value_expose() {
let secret = SecretValue::new(vec![1, 2, 3, 4, 5]);
assert_eq!(secret.expose_secret(), &vec![1, 2, 3, 4, 5]);
}
#[rstest]
fn test_secret_string_serialization_redacts_value() {
let secret = SecretString::new("my-super-secret-password");
let json = serde_json::to_string(&secret).unwrap();
assert!(!json.contains("my-super-secret-password"));
assert!(json.contains("[REDACTED]"));
assert_eq!(json, "\"[REDACTED]\"");
}
#[rstest]
fn test_secret_string_deserialization() {
let json = r#"{"secret":"test-secret"}"#;
let deserialized: SecretString = serde_json::from_str(json).unwrap();
assert_eq!(deserialized.expose_secret(), "test-secret");
}
#[rstest]
#[serial(env)]
fn test_secret_string_deserializes_from_env_source() {
unsafe { std::env::set_var("REINHARDT_TEST_SECRET_SOURCE", "env-secret") };
let json = r#"{"env":"REINHARDT_TEST_SECRET_SOURCE"}"#;
let deserialized: SecretString = serde_json::from_str(json).unwrap();
assert_eq!(deserialized.expose_secret(), "env-secret");
unsafe { std::env::remove_var("REINHARDT_TEST_SECRET_SOURCE") };
}
#[rstest]
fn test_secret_string_deserializes_from_file_source() {
let temp_file = tempfile::NamedTempFile::new().unwrap();
std::fs::write(temp_file.path(), "file-secret\n").unwrap();
let json = serde_json::json!({
"file": temp_file.path().to_string_lossy(),
})
.to_string();
let deserialized: SecretString = serde_json::from_str(&json).unwrap();
assert_eq!(deserialized.expose_secret(), "file-secret");
}
#[rstest]
fn test_secret_string_rejects_multiple_sources() {
let json = r#"{"secret":"inline","env":"IGNORED"}"#;
let error = serde_json::from_str::<SecretString>(json).unwrap_err();
assert!(
error
.to_string()
.contains("must contain only one of `secret`, `env`, or `file`")
);
}
#[rstest]
fn test_secret_value_serialization_redacts_value() {
let secret = SecretValue::new(42);
let json = serde_json::to_string(&secret).unwrap();
assert!(!json.contains("42"));
assert!(json.contains("[REDACTED]"));
assert_eq!(json, "\"[REDACTED]\"");
}
#[rstest]
fn test_secret_value_deserialization() {
let json = r#"{"inner":42}"#;
let deserialized: SecretValue<i32> = serde_json::from_str(json).unwrap();
assert_eq!(*deserialized.expose_secret(), 42);
}
#[rstest]
fn test_secret_string_into_inner() {
let secret = SecretString::new("my-secret-value");
let inner = secret.into_inner();
assert_eq!(inner, "my-secret-value");
}
#[rstest]
fn test_secret_value_into_inner() {
let secret = SecretValue::new(vec![1, 2, 3, 4, 5]);
let inner = secret.into_inner();
assert_eq!(inner, vec![1, 2, 3, 4, 5]);
}
#[rstest]
fn test_secret_value_into_inner_non_clone() {
struct NonClone {
inner: String,
}
impl Zeroize for NonClone {
fn zeroize(&mut self) {
self.inner.zeroize();
}
}
let secret = SecretValue::new(NonClone {
inner: "secret".to_string(),
});
let inner = secret.into_inner();
assert_eq!(inner.inner, "secret");
}
}