1
use std::str::FromStr;
2

            
3
use serde::{Deserialize, Serialize};
4
use zbus::zvariant;
5
use zeroize::{Zeroize, ZeroizeOnDrop, Zeroizing};
6

            
7
#[derive(Default, PartialEq, Eq, Copy, Clone, Debug, zvariant::Type)]
8
#[zvariant(signature = "s")]
9
pub enum ContentType {
10
    Text,
11
    #[default]
12
    Blob,
13
}
14

            
15
impl Serialize for ContentType {
16
28
    fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
17
    where
18
        S: serde::Serializer,
19
    {
20
27
        self.as_str().serialize(serializer)
21
    }
22
}
23

            
24
impl<'de> Deserialize<'de> for ContentType {
25
23
    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
26
    where
27
        D: serde::Deserializer<'de>,
28
    {
29
23
        let s = String::deserialize(deserializer)?;
30
46
        Self::from_str(&s).map_err(serde::de::Error::custom)
31
    }
32
}
33

            
34
impl FromStr for ContentType {
35
    type Err = String;
36

            
37
26
    fn from_str(s: &str) -> Result<Self, Self::Err> {
38
        // MIME types may include parameters, which are irrelevant to
39
        // ContentType.
40
53
        let media_type = s
41
            .split_once(';')
42
29
            .map_or(s, |(media_type, _)| media_type)
43
            .trim();
44

            
45
52
        if media_type.eq_ignore_ascii_case("text/plain")
46
13
            || media_type.eq_ignore_ascii_case("text/utf8")
47
13
            || media_type.eq_ignore_ascii_case("application/json")
48
        {
49
25
            Ok(Self::Text)
50
26
        } else if media_type.eq_ignore_ascii_case("application/octet-stream")
51
2
            || media_type.eq_ignore_ascii_case("application/binary")
52
        {
53
13
            Ok(Self::Blob)
54
        } else {
55
2
            Err(format!("Invalid content type: {s}"))
56
        }
57
    }
58
}
59

            
60
impl ContentType {
61
25
    pub const fn as_str(&self) -> &'static str {
62
25
        match self {
63
23
            Self::Text => "text/plain",
64
13
            Self::Blob => "application/octet-stream",
65
        }
66
    }
67
}
68

            
69
/// A wrapper around a combination of (secret, content-type).
70
#[derive(Clone, PartialEq, Eq, Zeroize, ZeroizeOnDrop)]
71
pub enum Secret {
72
    /// Corresponds to [`ContentType::Text`]
73
    Text(String),
74
    /// Corresponds to [`ContentType::Blob`]
75
    Blob(Vec<u8>),
76
}
77

            
78
impl std::fmt::Debug for Secret {
79
2
    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
80
2
        match self {
81
2
            Self::Text(_) => write!(f, "Secret::Text([REDACTED])"),
82
2
            Self::Blob(_) => write!(f, "Secret::Blob([REDACTED])"),
83
        }
84
    }
85
}
86

            
87
impl Secret {
88
    /// Generate a random secret, used when creating a session collection.
89
41
    pub fn random() -> Result<Self, getrandom::Error> {
90
34
        let mut secret = [0; 64];
91
        // Equivalent of `ring::rand::SecureRandom`
92
37
        getrandom::fill(&mut secret)?;
93

            
94
31
        Ok(Self::blob(secret))
95
    }
96

            
97
    /// Get the sandboxed app secret if the app is sandboxed using
98
    /// org.freedesktop.portal.Secret portal.
99
    pub async fn sandboxed() -> Result<Self, crate::file::Error> {
100
        Ok(Self::blob(
101
            ashpd::desktop::secret::retrieve()
102
                .await
103
                .map_err(crate::file::Error::from)?,
104
        ))
105
    }
106

            
107
    /// Create a text secret, stored with `text/plain` content type.
108
65
    pub fn text(value: impl AsRef<str>) -> Self {
109
85
        Self::Text(value.as_ref().to_owned())
110
    }
111

            
112
    /// Create a blob secret, stored with `application/octet-stream` content
113
    /// type.
114
135
    pub fn blob(value: impl AsRef<[u8]>) -> Self {
115
165
        Self::Blob(value.as_ref().to_owned())
116
    }
117

            
118
22
    pub const fn content_type(&self) -> ContentType {
119
23
        match self {
120
23
            Self::Text(_) => ContentType::Text,
121
13
            Self::Blob(_) => ContentType::Blob,
122
        }
123
    }
124

            
125
    /// Returns the secret as a string slice, or `None` if it is not valid
126
    /// UTF-8.
127
    pub fn as_str(&self) -> Option<&str> {
128
        match self {
129
            Self::Text(text) => Some(text.as_str()),
130
            Self::Blob(bytes) => std::str::from_utf8(bytes).ok(),
131
        }
132
    }
133

            
134
37
    pub fn as_bytes(&self) -> &[u8] {
135
78
        match self {
136
34
            Self::Text(text) => text.as_bytes(),
137
25
            Self::Blob(bytes) => bytes.as_ref(),
138
        }
139
    }
140

            
141
23
    pub fn with_content_type(content_type: ContentType, secret: impl AsRef<[u8]>) -> Self {
142
23
        match content_type {
143
23
            ContentType::Text => match String::from_utf8(secret.as_ref().to_owned()) {
144
23
                Ok(text) => Secret::text(text),
145
2
                Err(_e) => {
146
                    #[cfg(feature = "tracing")]
147
                    tracing::warn!(
148
                        "Failed to decode secret as UTF-8: {}, falling back to blob",
149
                        _e
150
                    );
151

            
152
2
                    Secret::blob(secret)
153
                }
154
            },
155
26
            _ => Secret::blob(secret),
156
        }
157
    }
158
}
159

            
160
impl From<&[u8]> for Secret {
161
8
    fn from(value: &[u8]) -> Self {
162
8
        Self::blob(value)
163
    }
164
}
165

            
166
impl From<Zeroizing<Vec<u8>>> for Secret {
167
23
    fn from(value: Zeroizing<Vec<u8>>) -> Self {
168
25
        Self::blob(value)
169
    }
170
}
171

            
172
impl From<Vec<u8>> for Secret {
173
17
    fn from(value: Vec<u8>) -> Self {
174
16
        Self::blob(value)
175
    }
176
}
177

            
178
impl From<&Vec<u8>> for Secret {
179
6
    fn from(value: &Vec<u8>) -> Self {
180
6
        Self::blob(value)
181
    }
182
}
183

            
184
impl<const N: usize> From<&[u8; N]> for Secret {
185
    fn from(value: &[u8; N]) -> Self {
186
        Self::blob(value)
187
    }
188
}
189

            
190
impl From<String> for Secret {
191
6
    fn from(value: String) -> Self {
192
6
        Self::text(value)
193
    }
194
}
195

            
196
impl From<&str> for Secret {
197
35
    fn from(value: &str) -> Self {
198
36
        Self::text(value)
199
    }
200
}
201

            
202
impl std::ops::Deref for Secret {
203
    type Target = [u8];
204

            
205
28
    fn deref(&self) -> &Self::Target {
206
38
        self.as_bytes()
207
    }
208
}
209

            
210
impl AsRef<[u8]> for Secret {
211
12
    fn as_ref(&self) -> &[u8] {
212
12
        self.as_bytes()
213
    }
214
}
215

            
216
#[cfg(test)]
217
mod tests {
218
    use zbus::zvariant::{Endian, serialized::Context, to_bytes};
219

            
220
    use super::*;
221

            
222
    #[test]
223
    fn secret_debug_is_redacted() {
224
        let text_secret = Secret::text("password");
225
        let blob_secret = Secret::blob([1, 2, 3]);
226

            
227
        assert_eq!(format!("{:?}", text_secret), "Secret::Text([REDACTED])");
228
        assert_eq!(format!("{:?}", blob_secret), "Secret::Blob([REDACTED])");
229
    }
230

            
231
    #[test]
232
    fn content_type_serialization() {
233
        let ctxt = Context::new_dbus(Endian::Little, 0);
234

            
235
        // Test Text serialization
236
        let encoded = to_bytes(ctxt, &ContentType::Text).unwrap();
237
        let value: String = encoded.deserialize().unwrap().0;
238
        assert_eq!(value, "text/plain");
239

            
240
        // Test Blob serialization
241
        let encoded = to_bytes(ctxt, &ContentType::Blob).unwrap();
242
        let value: String = encoded.deserialize().unwrap().0;
243
        assert_eq!(value, "application/octet-stream");
244

            
245
        // Test Text deserialization
246
        let encoded = to_bytes(ctxt, &"text/plain").unwrap();
247
        let content_type: ContentType = encoded.deserialize().unwrap().0;
248
        assert_eq!(content_type, ContentType::Text);
249

            
250
        // Test Text deserialization with MIME parameters
251
        let encoded = to_bytes(ctxt, &"text/plain; charset=utf8").unwrap();
252
        let content_type: ContentType = encoded.deserialize().unwrap().0;
253
        assert_eq!(content_type, ContentType::Text);
254

            
255
        // Test Blob deserialization
256
        let encoded = to_bytes(ctxt, &"application/octet-stream").unwrap();
257
        let content_type: ContentType = encoded.deserialize().unwrap().0;
258
        assert_eq!(content_type, ContentType::Blob);
259

            
260
        // application/json deserializes as Text
261
        let encoded = to_bytes(ctxt, &"application/json").unwrap();
262
        let content_type: ContentType = encoded.deserialize().unwrap().0;
263
        assert_eq!(content_type, ContentType::Text);
264

            
265
        // Test invalid content type deserialization
266
        let encoded = to_bytes(ctxt, &"invalid/type").unwrap();
267
        let result: Result<(ContentType, _), _> = encoded.deserialize();
268
        assert!(result.is_err());
269
        assert!(
270
            result
271
                .unwrap_err()
272
                .to_string()
273
                .contains("Invalid content type")
274
        );
275
    }
276

            
277
    #[test]
278
    fn content_type_from_str() {
279
        for content_type in [
280
            "text/plain",
281
            "text/plain; charset=utf8",
282
            "TEXT/PLAIN; CHARSET=UTF-8",
283
            "application/json",
284
            "application/json; charset=utf8",
285
        ] {
286
            assert_eq!(
287
                ContentType::from_str(content_type).unwrap(),
288
                ContentType::Text
289
            );
290
        }
291

            
292
        for content_type in [
293
            "application/octet-stream",
294
            "application/octet-stream; version=1",
295
        ] {
296
            assert_eq!(
297
                ContentType::from_str(content_type).unwrap(),
298
                ContentType::Blob
299
            );
300
        }
301

            
302
        // Test error case
303
        let result = ContentType::from_str("text/html; charset=utf8");
304
        assert!(result.is_err());
305
        let error = result.unwrap_err();
306
        assert!(error.contains("Invalid content type"));
307
        assert!(error.contains("text/html; charset=utf8"));
308
    }
309

            
310
    #[test]
311
    fn invalid_utf8() {
312
        // Test with invalid UTF-8 bytes
313
        let invalid_utf8 = vec![0xFF, 0xFE, 0xFD];
314

            
315
        // Should fall back to blob when UTF-8 decoding fails
316
        let secret = Secret::with_content_type(ContentType::Text, &invalid_utf8);
317
        assert_eq!(secret.content_type(), ContentType::Blob);
318
        assert_eq!(&*secret, &[0xFF, 0xFE, 0xFD]);
319

            
320
        // Test with valid UTF-8
321
        let valid_utf8 = "Hello, World!";
322
        let secret = Secret::with_content_type(ContentType::Text, valid_utf8.as_bytes());
323
        assert_eq!(secret.content_type(), ContentType::Text);
324
        assert_eq!(&*secret, valid_utf8.as_bytes());
325

            
326
        // Test with blob content type
327
        let data = vec![1, 2, 3, 4];
328
        let secret = Secret::with_content_type(ContentType::Blob, &data);
329
        assert_eq!(secret.content_type(), ContentType::Blob);
330
        assert_eq!(&*secret, &[1, 2, 3, 4]);
331
    }
332

            
333
    #[test]
334
    fn random() {
335
        let secret1 = Secret::random().unwrap();
336
        let secret2 = Secret::random().unwrap();
337

            
338
        // Random secrets should be blobs
339
        assert_eq!(secret1.content_type(), ContentType::Blob);
340
        assert_eq!(secret2.content_type(), ContentType::Blob);
341

            
342
        // Should be 64 bytes
343
        assert_eq!(secret1.as_bytes().len(), 64);
344
        assert_eq!(secret2.as_bytes().len(), 64);
345

            
346
        // Should be different
347
        assert_ne!(secret1.as_bytes(), secret2.as_bytes());
348
    }
349
}