Skip to main content

libsignal_service/
utils.rs

1mod phonenumber;
2use libsignal_core::{Aci, Pni, ServiceId};
3pub use phonenumber::*;
4use uuid::Uuid;
5
6// Signal sometimes adds padding, sometimes it does not.
7// This requires a custom decoding engine.
8// This engine is as general as possible.
9pub const BASE64_RELAXED: base64::engine::GeneralPurpose =
10    base64::engine::GeneralPurpose::new(
11        &base64::alphabet::STANDARD,
12        base64::engine::GeneralPurposeConfig::new()
13            .with_encode_padding(true)
14            .with_decode_padding_mode(
15                base64::engine::DecodePaddingMode::Indifferent,
16            ),
17    );
18
19pub fn parse_aci_with_fallback(
20    bytes: Option<&[u8]>,
21    utf8: Option<&str>,
22) -> Option<Aci> {
23    let binary = bytes.and_then(|bytes| {
24        let bytes = bytes
25            .try_into()
26            .inspect_err(|_e| tracing::warn!("binary ACI not 16 bytes"))
27            .ok()?;
28        Some(Aci::from_uuid_bytes(bytes))
29    });
30
31    binary.or_else(|| {
32        let utf8 = utf8?;
33        match Aci::parse_from_service_id_string(utf8) {
34            Some(sid) => Some(sid),
35            None => {
36                tracing::warn!("unparseable utf8 ACI");
37                None
38            },
39        }
40    })
41}
42
43pub fn parse_pni_with_fallback(
44    bytes: Option<&[u8]>,
45    utf8: Option<&str>,
46    pni_is_uuid: bool,
47) -> Option<Pni> {
48    let binary = bytes.and_then(|bytes| {
49        let bytes = bytes
50            .try_into()
51            .inspect_err(|_e| tracing::warn!("binary PNI not 16 bytes"))
52            .ok()?;
53        Some(Pni::from_uuid_bytes(bytes))
54    });
55
56    binary.or_else(|| {
57        let utf8 = utf8?;
58        if pni_is_uuid {
59            let uuid: uuid::Uuid = utf8
60                .parse()
61                .inspect_err(|e| {
62                    tracing::warn!(error = %e, "unparseable UUID");
63                })
64                .ok()?;
65            Some(Pni::from_uuid_bytes(*uuid.as_bytes()))
66        } else {
67            match Pni::parse_from_service_id_string(utf8) {
68                Some(sid) => Some(sid),
69                None => {
70                    tracing::warn!("unparseable utf8 PNI");
71                    None
72                },
73            }
74        }
75    })
76}
77
78pub fn parse_service_id_with_fallback(
79    bytes: Option<&[u8]>,
80    utf8: Option<&str>,
81) -> Option<ServiceId> {
82    let binary = bytes.and_then(|bytes| {
83        match ServiceId::parse_from_service_id_binary(bytes) {
84            Some(sid) => Some(sid),
85            None => {
86                tracing::warn!("unparseable binary ServiceId");
87                None
88            },
89        }
90    });
91
92    binary.or_else(|| {
93        let utf8 = utf8?;
94        match ServiceId::parse_from_service_id_string(utf8) {
95            Some(sid) => Some(sid),
96            None => {
97                tracing::warn!("unparseable utf8 ServiceId");
98                None
99            },
100        }
101    })
102}
103
104/// Parse protobuf UUIDs specified in both binary and utf8 formats
105///
106/// Prefers the binary format
107pub fn parse_uuid_with_fallback(
108    binary: Option<&[u8]>,
109    utf8: Option<&str>,
110) -> Option<Uuid> {
111    let binary = binary
112        .map(<[u8; 16]>::try_from)
113        .transpose()
114        .inspect_err(|_e| tracing::warn!("invalid binary UUID length"))
115        .ok()
116        .flatten()
117        .map(Uuid::from_bytes);
118
119    binary.or_else(|| {
120        let utf8 = utf8?;
121        utf8.parse()
122            .inspect_err(|e| tracing::warn!(error=%e, "unparseable UUID"))
123            .ok()
124    })
125}
126
127pub fn random_length_padding<R: rand::Rng + rand::CryptoRng>(
128    csprng: &mut R,
129    max_len: usize,
130) -> Vec<u8> {
131    let length = csprng.random_range(0..max_len);
132    let mut padding = vec![0u8; length];
133    csprng.fill_bytes(&mut padding);
134    padding
135}
136
137pub mod serde_base64 {
138    use super::BASE64_RELAXED;
139    use base64::prelude::*;
140    use serde::{Deserialize, Deserializer, Serializer};
141
142    pub fn serialize<T, S>(bytes: &T, serializer: S) -> Result<S::Ok, S::Error>
143    where
144        T: AsRef<[u8]>,
145        S: Serializer,
146    {
147        serializer.serialize_str(&BASE64_RELAXED.encode(bytes.as_ref()))
148    }
149
150    pub fn deserialize<'de, D>(deserializer: D) -> Result<Vec<u8>, D::Error>
151    where
152        D: Deserializer<'de>,
153    {
154        use serde::de::Error;
155        <&str>::deserialize(deserializer).and_then(|string| {
156            BASE64_RELAXED
157                .decode(string)
158                .map_err(|err| Error::custom(err.to_string()))
159        })
160    }
161}
162
163pub mod serde_optional_base64 {
164    use super::BASE64_RELAXED;
165    use base64::prelude::*;
166    use serde::{Deserialize, Deserializer, Serializer};
167
168    use super::serde_base64;
169
170    pub fn serialize<T, S>(
171        bytes: &Option<T>,
172        serializer: S,
173    ) -> Result<S::Ok, S::Error>
174    where
175        T: AsRef<[u8]>,
176        S: Serializer,
177    {
178        match bytes {
179            Some(bytes) => serde_base64::serialize(bytes, serializer),
180            None => serializer.serialize_none(),
181        }
182    }
183
184    pub fn deserialize<'de, D>(
185        deserializer: D,
186    ) -> Result<Option<Vec<u8>>, D::Error>
187    where
188        D: Deserializer<'de>,
189    {
190        use serde::de::Error;
191        match Option::<String>::deserialize(deserializer)? {
192            Some(s) => BASE64_RELAXED
193                .decode(s)
194                .map_err(|err| Error::custom(err.to_string()))
195                .map(Some),
196            None => Ok(None),
197        }
198    }
199}
200
201pub mod serde_optional_base64_url_safe_no_pad {
202    use base64::{prelude::BASE64_URL_SAFE_NO_PAD, Engine};
203    use serde::{Deserialize, Deserializer, Serializer};
204
205    use super::serde_base64_url_safe_no_pad;
206
207    pub fn serialize<T, S>(
208        bytes: &Option<T>,
209        serializer: S,
210    ) -> Result<S::Ok, S::Error>
211    where
212        T: AsRef<[u8]>,
213        S: Serializer,
214    {
215        match bytes {
216            Some(bytes) => {
217                serde_base64_url_safe_no_pad::serialize(bytes, serializer)
218            },
219            None => serializer.serialize_none(),
220        }
221    }
222
223    pub fn deserialize<'de, D>(
224        deserializer: D,
225    ) -> Result<Option<Vec<u8>>, D::Error>
226    where
227        D: Deserializer<'de>,
228    {
229        use serde::de::Error;
230        match Option::<String>::deserialize(deserializer)? {
231            Some(s) => BASE64_URL_SAFE_NO_PAD
232                .decode(s)
233                .map_err(|err| Error::custom(err.to_string()))
234                .map(Some),
235            None => Ok(None),
236        }
237    }
238}
239
240pub mod serde_base64_url_safe_no_pad {
241    use base64::{prelude::BASE64_URL_SAFE_NO_PAD, Engine};
242    use serde::{Deserialize, Deserializer, Serializer};
243
244    pub fn serialize<T, S>(bytes: &T, serializer: S) -> Result<S::Ok, S::Error>
245    where
246        T: AsRef<[u8]>,
247        S: Serializer,
248    {
249        serializer.serialize_str(&BASE64_URL_SAFE_NO_PAD.encode(bytes.as_ref()))
250    }
251
252    pub fn deserialize<'de, D>(deserializer: D) -> Result<Vec<u8>, D::Error>
253    where
254        D: Deserializer<'de>,
255    {
256        use serde::de::Error;
257        <&str>::deserialize(deserializer).and_then(|string| {
258            BASE64_URL_SAFE_NO_PAD
259                .decode(string)
260                .map_err(|err| Error::custom(err.to_string()))
261        })
262    }
263}
264
265pub mod serde_identity_key {
266    use super::BASE64_RELAXED;
267    use base64::prelude::*;
268    use libsignal_protocol::IdentityKey;
269    use serde::{Deserialize, Deserializer, Serializer};
270
271    pub fn serialize<S>(
272        public_key: &IdentityKey,
273        serializer: S,
274    ) -> Result<S::Ok, S::Error>
275    where
276        S: Serializer,
277    {
278        let public_key = public_key.serialize();
279        serializer.serialize_str(&BASE64_RELAXED.encode(&public_key))
280    }
281
282    pub fn deserialize<'de, D>(deserializer: D) -> Result<IdentityKey, D::Error>
283    where
284        D: Deserializer<'de>,
285    {
286        IdentityKey::decode(
287            &BASE64_RELAXED
288                .decode(<&str>::deserialize(deserializer)?)
289                .map_err(serde::de::Error::custom)?,
290        )
291        .map_err(serde::de::Error::custom)
292    }
293}
294
295pub mod serde_optional_identity_key {
296    use super::BASE64_RELAXED;
297    use base64::prelude::*;
298    use libsignal_protocol::IdentityKey;
299    use serde::{Deserialize, Deserializer, Serializer};
300
301    use super::serde_identity_key;
302
303    pub fn serialize<S>(
304        public_key: &Option<IdentityKey>,
305        serializer: S,
306    ) -> Result<S::Ok, S::Error>
307    where
308        S: Serializer,
309    {
310        match public_key {
311            Some(public_key) => {
312                serde_identity_key::serialize(public_key, serializer)
313            },
314            None => serializer.serialize_none(),
315        }
316    }
317
318    pub fn deserialize<'de, D>(
319        deserializer: D,
320    ) -> Result<Option<IdentityKey>, D::Error>
321    where
322        D: Deserializer<'de>,
323    {
324        match Option::<String>::deserialize(deserializer)? {
325            Some(public_key) => Ok(Some(
326                IdentityKey::decode(
327                    &BASE64_RELAXED
328                        .decode(public_key)
329                        .map_err(serde::de::Error::custom)?,
330                )
331                .map_err(serde::de::Error::custom)?,
332            )),
333            None => Ok(None),
334        }
335    }
336}
337
338pub mod serde_private_key {
339    use super::BASE64_RELAXED;
340    use base64::prelude::*;
341    use libsignal_protocol::PrivateKey;
342    use serde::{Deserialize, Deserializer, Serializer};
343
344    pub fn serialize<S>(
345        public_key: &PrivateKey,
346        serializer: S,
347    ) -> Result<S::Ok, S::Error>
348    where
349        S: Serializer,
350    {
351        let public_key = public_key.serialize();
352        serializer.serialize_str(&BASE64_RELAXED.encode(public_key))
353    }
354
355    pub fn deserialize<'de, D>(deserializer: D) -> Result<PrivateKey, D::Error>
356    where
357        D: Deserializer<'de>,
358    {
359        PrivateKey::deserialize(
360            &BASE64_RELAXED
361                .decode(<&str>::deserialize(deserializer)?)
362                .map_err(serde::de::Error::custom)?,
363        )
364        .map_err(serde::de::Error::custom)
365    }
366}
367
368pub mod serde_optional_private_key {
369    use super::BASE64_RELAXED;
370    use base64::prelude::*;
371    use libsignal_protocol::PrivateKey;
372    use serde::{Deserialize, Deserializer, Serializer};
373
374    use super::serde_private_key;
375
376    pub fn serialize<S>(
377        private_key: &Option<PrivateKey>,
378        serializer: S,
379    ) -> Result<S::Ok, S::Error>
380    where
381        S: Serializer,
382    {
383        match private_key {
384            Some(private_key) => {
385                serde_private_key::serialize(private_key, serializer)
386            },
387            None => serializer.serialize_none(),
388        }
389    }
390
391    pub fn deserialize<'de, D>(
392        deserializer: D,
393    ) -> Result<Option<PrivateKey>, D::Error>
394    where
395        D: Deserializer<'de>,
396    {
397        match Option::<String>::deserialize(deserializer)? {
398            Some(private_key) => Ok(Some(
399                PrivateKey::deserialize(
400                    &BASE64_RELAXED
401                        .decode(private_key)
402                        .map_err(serde::de::Error::custom)?,
403                )
404                .map_err(serde::de::Error::custom)?,
405            )),
406            None => Ok(None),
407        }
408    }
409}
410
411pub mod serde_optional_e164 {
412    use libsignal_core::E164;
413    use serde::{Deserialize, Deserializer, Serializer};
414
415    pub fn serialize<S>(
416        phone_number: &Option<E164>,
417        serializer: S,
418    ) -> Result<S::Ok, S::Error>
419    where
420        S: Serializer,
421    {
422        match phone_number {
423            Some(p) => serializer.serialize_str(&p.to_string()),
424            None => serializer.serialize_none(),
425        }
426    }
427
428    pub fn deserialize<'de, D>(
429        deserializer: D,
430    ) -> Result<Option<E164>, D::Error>
431    where
432        D: Deserializer<'de>,
433    {
434        match Option::<String>::deserialize(deserializer)? {
435            Some(s) => s.parse().map_err(serde::de::Error::custom).map(Some),
436            None => Ok(None),
437        }
438    }
439}
440
441pub mod serde_e164 {
442    use libsignal_core::E164;
443    use serde::{Deserialize, Deserializer, Serializer};
444
445    pub fn serialize<S>(
446        phone_number: &E164,
447        serializer: S,
448    ) -> Result<S::Ok, S::Error>
449    where
450        S: Serializer,
451    {
452        serializer.serialize_str(&phone_number.to_string())
453    }
454
455    pub fn deserialize<'de, D>(deserializer: D) -> Result<E164, D::Error>
456    where
457        D: Deserializer<'de>,
458    {
459        <&str>::deserialize(deserializer)?
460            .parse()
461            .map_err(serde::de::Error::custom)
462    }
463}
464
465#[cfg(feature = "phonenumber")]
466pub mod serde_phone_number {
467    use phonenumber::PhoneNumber;
468    use serde::{Deserialize, Deserializer, Serializer};
469
470    pub fn serialize<S>(
471        phone_number: &PhoneNumber,
472        serializer: S,
473    ) -> Result<S::Ok, S::Error>
474    where
475        S: Serializer,
476    {
477        serializer.serialize_str(&phone_number.to_string())
478    }
479
480    pub fn deserialize<'de, D>(deserializer: D) -> Result<PhoneNumber, D::Error>
481    where
482        D: Deserializer<'de>,
483    {
484        phonenumber::parse(None, <&str>::deserialize(deserializer)?)
485            .map_err(serde::de::Error::custom)
486    }
487}
488
489pub mod serde_service_id {
490    use libsignal_protocol::ServiceId;
491    use serde::{Deserialize, Deserializer, Serializer};
492
493    pub fn serialize<S>(
494        service_id: &ServiceId,
495        serializer: S,
496    ) -> Result<S::Ok, S::Error>
497    where
498        S: Serializer,
499    {
500        serializer.serialize_str(&service_id.service_id_string())
501    }
502
503    pub fn deserialize<'de, D>(deserializer: D) -> Result<ServiceId, D::Error>
504    where
505        D: Deserializer<'de>,
506    {
507        ServiceId::parse_from_service_id_string(<&str>::deserialize(
508            deserializer,
509        )?)
510        .ok_or_else(|| serde::de::Error::custom("invalid service ID string"))
511    }
512}
513
514pub mod serde_aci {
515    use libsignal_core::Aci;
516    use serde::{Deserialize, Deserializer, Serializer};
517
518    pub fn serialize<S>(aci: &Aci, serializer: S) -> Result<S::Ok, S::Error>
519    where
520        S: Serializer,
521    {
522        serializer.serialize_str(&aci.service_id_string())
523    }
524
525    pub fn deserialize<'de, D>(deserializer: D) -> Result<Aci, D::Error>
526    where
527        D: Deserializer<'de>,
528    {
529        Aci::parse_from_service_id_string(<&str>::deserialize(deserializer)?)
530            .ok_or_else(|| serde::de::Error::custom("invalid ACI string"))
531    }
532}
533
534pub mod serde_device_id {
535    use libsignal_core::DeviceId;
536    use serde::{Deserialize, Deserializer, Serializer};
537
538    pub fn serialize<S>(id: &DeviceId, serializer: S) -> Result<S::Ok, S::Error>
539    where
540        S: Serializer,
541    {
542        serializer.serialize_u8(u8::from(*id))
543    }
544
545    pub fn deserialize<'de, D>(deserializer: D) -> Result<DeviceId, D::Error>
546    where
547        D: Deserializer<'de>,
548    {
549        DeviceId::try_from(u8::deserialize(deserializer)?)
550            .map_err(|_| serde::de::Error::custom("invalid device id"))
551    }
552}
553
554pub mod serde_device_id_vec {
555    use libsignal_core::DeviceId;
556    use serde::{ser::SerializeSeq, Deserialize, Deserializer, Serializer};
557
558    pub fn serialize<S>(
559        ids: &Vec<DeviceId>,
560        serializer: S,
561    ) -> Result<S::Ok, S::Error>
562    where
563        S: Serializer,
564    {
565        let mut seq = serializer.serialize_seq(Some(ids.len()))?;
566        for id in ids {
567            seq.serialize_element(&u8::from(*id))?;
568        }
569        seq.end()
570    }
571
572    pub fn deserialize<'de, D>(
573        deserializer: D,
574    ) -> Result<Vec<DeviceId>, D::Error>
575    where
576        D: Deserializer<'de>,
577    {
578        Vec::<u8>::deserialize(deserializer)?
579            .into_iter()
580            .map(DeviceId::try_from)
581            .collect::<Result<Vec<_>, _>>()
582            .map_err(|_| serde::de::Error::custom("invalid device id"))
583    }
584}
585
586pub mod serde_prost_base64 {
587    use super::BASE64_RELAXED;
588    use base64::Engine;
589    use prost::Message;
590    use serde::{Deserialize, Deserializer, Serializer};
591
592    // Serializes a Prost message into a Base64 string
593    pub fn serialize<T, S>(value: &T, serializer: S) -> Result<S::Ok, S::Error>
594    where
595        T: Message,
596        S: Serializer,
597    {
598        let b64 = BASE64_RELAXED.encode(value.encode_to_vec());
599        serializer.serialize_str(&b64)
600    }
601
602    // Deserializes a Base64 string back into a Prost message
603    pub fn deserialize<'de, T, D>(deserializer: D) -> Result<T, D::Error>
604    where
605        T: Message + Default,
606        D: Deserializer<'de>,
607    {
608        let bytes = BASE64_RELAXED
609            .decode(<&str>::deserialize(deserializer)?)
610            .map_err(serde::de::Error::custom)?;
611
612        T::decode(bytes.as_slice()).map_err(serde::de::Error::custom)
613    }
614}
615
616pub mod serde_optional_prost_base64 {
617    use base64::Engine;
618    use prost::Message;
619    use serde::{Deserialize, Deserializer, Serializer};
620
621    use super::{serde_prost_base64, BASE64_RELAXED};
622
623    pub fn serialize<T, S>(
624        value: &Option<T>,
625        serializer: S,
626    ) -> Result<S::Ok, S::Error>
627    where
628        T: Message,
629        S: Serializer,
630    {
631        match value {
632            Some(msg) => serde_prost_base64::serialize(msg, serializer),
633            None => serializer.serialize_none(),
634        }
635    }
636
637    pub fn deserialize<'de, T, D>(
638        deserializer: D,
639    ) -> Result<Option<T>, D::Error>
640    where
641        T: Message + Default,
642        D: Deserializer<'de>,
643    {
644        match Option::<String>::deserialize(deserializer)? {
645            Some(s) => {
646                let bytes = BASE64_RELAXED
647                    .decode(s)
648                    .map_err(serde::de::Error::custom)?;
649                let msg = T::decode(bytes.as_slice())
650                    .map_err(serde::de::Error::custom)?;
651                Ok(Some(msg))
652            },
653            None => Ok(None),
654        }
655    }
656}