Skip to main content

isb_server/auth/
cbor.rs

1//! A minimal CBOR (RFC 8949) decoder for WebAuthn: attestation objects and
2//! COSE keys. Definite lengths only (CTAP2 canonical encoding never uses
3//! indefinite ones), nesting capped, every length checked against the input
4//! before anything is allocated. Floats are skipped over, not interpreted.
5
6/// A decoded item.
7#[derive(Debug, Clone, PartialEq)]
8pub enum Value {
9    Uint(u64),
10    /// A negative integer: the value is `-1 - n`.
11    Nint(u64),
12    Bytes(Vec<u8>),
13    Text(String),
14    Array(Vec<Value>),
15    Map(Vec<(Value, Value)>),
16    Bool(bool),
17    Null,
18    Undefined,
19    Float,
20}
21
22impl Value {
23    /// This integer as an i64, if it fits.
24    pub fn as_int(&self) -> Option<i64> {
25        match *self {
26            Value::Uint(n) => i64::try_from(n).ok(),
27            Value::Nint(n) => i64::try_from(n).ok().map(|n| -1 - n),
28            _ => None,
29        }
30    }
31
32    pub fn as_bytes(&self) -> Option<&[u8]> {
33        match self {
34            Value::Bytes(b) => Some(b),
35            _ => None,
36        }
37    }
38
39    pub fn as_text(&self) -> Option<&str> {
40        match self {
41            Value::Text(s) => Some(s),
42            _ => None,
43        }
44    }
45
46    /// The value under integer key `k` in a map.
47    pub fn get_int(&self, k: i64) -> Option<&Value> {
48        match self {
49            Value::Map(m) => m
50                .iter()
51                .find(|(key, _)| key.as_int() == Some(k))
52                .map(|(_, v)| v),
53            _ => None,
54        }
55    }
56
57    /// The value under text key `k` in a map.
58    pub fn get_text(&self, k: &str) -> Option<&Value> {
59        match self {
60            Value::Map(m) => m
61                .iter()
62                .find(|(key, _)| key.as_text() == Some(k))
63                .map(|(_, v)| v),
64            _ => None,
65        }
66    }
67}
68
69const MAX_DEPTH: usize = 16;
70
71/// Decode one item from the front of `input`; returns it and how many bytes
72/// it took. Trailing bytes are the caller's business (authenticator data
73/// carries extensions after the credential key).
74pub fn decode(input: &[u8]) -> Result<(Value, usize), String> {
75    let mut d = Decoder { b: input, at: 0 };
76    let v = d.item(0)?;
77    Ok((v, d.at))
78}
79
80/// Decode exactly one item spanning all of `input`.
81pub fn decode_all(input: &[u8]) -> Result<Value, String> {
82    let (v, n) = decode(input)?;
83    if n != input.len() {
84        return Err(format!(
85            "{} trailing bytes after the CBOR item",
86            input.len() - n
87        ));
88    }
89    Ok(v)
90}
91
92struct Decoder<'a> {
93    b: &'a [u8],
94    at: usize,
95}
96
97impl Decoder<'_> {
98    fn byte(&mut self) -> Result<u8, String> {
99        let c = *self.b.get(self.at).ok_or("CBOR: unexpected end of input")?;
100        self.at += 1;
101        Ok(c)
102    }
103
104    fn take(&mut self, n: u64) -> Result<&[u8], String> {
105        let n = usize::try_from(n).map_err(|_| "CBOR: length overflows")?;
106        let end = self
107            .at
108            .checked_add(n)
109            .filter(|e| *e <= self.b.len())
110            .ok_or("CBOR: length runs past the end of input")?;
111        let s = &self.b[self.at..end];
112        self.at = end;
113        Ok(s)
114    }
115
116    /// The argument of an initial byte with additional info `ai`.
117    fn arg(&mut self, ai: u8) -> Result<u64, String> {
118        Ok(match ai {
119            0..=23 => u64::from(ai),
120            24 => u64::from(self.byte()?),
121            25 => u64::from(u16::from_be_bytes(self.take(2)?.try_into().unwrap())),
122            26 => u64::from(u32::from_be_bytes(self.take(4)?.try_into().unwrap())),
123            27 => u64::from_be_bytes(self.take(8)?.try_into().unwrap()),
124            31 => return Err("CBOR: indefinite lengths are not supported".into()),
125            _ => return Err(format!("CBOR: reserved additional info {ai}")),
126        })
127    }
128
129    /// A count of items that must each take at least one byte of input.
130    fn count(&self, n: u64) -> Result<usize, String> {
131        let left = (self.b.len() - self.at) as u64;
132        if n > left {
133            return Err("CBOR: item count runs past the end of input".into());
134        }
135        Ok(n as usize)
136    }
137
138    fn item(&mut self, depth: usize) -> Result<Value, String> {
139        if depth > MAX_DEPTH {
140            return Err("CBOR: nested too deeply".into());
141        }
142        let ib = self.byte()?;
143        let (major, ai) = (ib >> 5, ib & 0x1f);
144        if major == 7 {
145            return match ai {
146                20 => Ok(Value::Bool(false)),
147                21 => Ok(Value::Bool(true)),
148                22 => Ok(Value::Null),
149                23 => Ok(Value::Undefined),
150                25 => self.take(2).map(|_| Value::Float),
151                26 => self.take(4).map(|_| Value::Float),
152                27 => self.take(8).map(|_| Value::Float),
153                _ => Err(format!("CBOR: unsupported simple value {ai}")),
154            };
155        }
156        let n = self.arg(ai)?;
157        Ok(match major {
158            0 => Value::Uint(n),
159            1 => Value::Nint(n),
160            2 => Value::Bytes(self.take(n)?.to_vec()),
161            3 => Value::Text(
162                String::from_utf8(self.take(n)?.to_vec()).map_err(|_| "CBOR: text is not UTF-8")?,
163            ),
164            4 => {
165                let n = self.count(n)?;
166                let mut v = Vec::with_capacity(n);
167                for _ in 0..n {
168                    v.push(self.item(depth + 1)?);
169                }
170                Value::Array(v)
171            }
172            5 => {
173                let n = self.count(n.saturating_mul(2))? / 2;
174                let mut v = Vec::with_capacity(n);
175                for _ in 0..n {
176                    let k = self.item(depth + 1)?;
177                    let val = self.item(depth + 1)?;
178                    v.push((k, val));
179                }
180                Value::Map(v)
181            }
182            // A tag: keep the tagged item, drop the tag.
183            6 => self.item(depth + 1)?,
184            _ => unreachable!("major type is 3 bits"),
185        })
186    }
187}
188
189/// Encode (tests build attestation objects and COSE keys with it).
190#[cfg(test)]
191pub fn encode(v: &Value) -> Vec<u8> {
192    fn head(out: &mut Vec<u8>, major: u8, n: u64) {
193        let m = major << 5;
194        match n {
195            0..=23 => out.push(m | n as u8),
196            24..=0xff => out.extend([m | 24, n as u8]),
197            0x100..=0xffff => {
198                out.push(m | 25);
199                out.extend((n as u16).to_be_bytes());
200            }
201            0x1_0000..=0xffff_ffff => {
202                out.push(m | 26);
203                out.extend((n as u32).to_be_bytes());
204            }
205            _ => {
206                out.push(m | 27);
207                out.extend(n.to_be_bytes());
208            }
209        }
210    }
211    fn go(out: &mut Vec<u8>, v: &Value) {
212        match v {
213            Value::Uint(n) => head(out, 0, *n),
214            Value::Nint(n) => head(out, 1, *n),
215            Value::Bytes(b) => {
216                head(out, 2, b.len() as u64);
217                out.extend(b);
218            }
219            Value::Text(s) => {
220                head(out, 3, s.len() as u64);
221                out.extend(s.as_bytes());
222            }
223            Value::Array(a) => {
224                head(out, 4, a.len() as u64);
225                a.iter().for_each(|x| go(out, x));
226            }
227            Value::Map(m) => {
228                head(out, 5, m.len() as u64);
229                for (k, x) in m {
230                    go(out, k);
231                    go(out, x);
232                }
233            }
234            Value::Bool(b) => out.push(if *b { 0xf5 } else { 0xf4 }),
235            Value::Null => out.push(0xf6),
236            Value::Undefined => out.push(0xf7),
237            Value::Float => out.extend([0xf9, 0, 0]),
238        }
239    }
240    let mut out = Vec::new();
241    go(&mut out, v);
242    out
243}
244
245/// An integer as a CBOR value (tests).
246#[cfg(test)]
247pub fn int(i: i64) -> Value {
248    if i >= 0 {
249        Value::Uint(i as u64)
250    } else {
251        Value::Nint((-1 - i) as u64)
252    }
253}
254
255#[cfg(test)]
256mod tests {
257    use super::*;
258
259    #[test]
260    fn round_trips_and_reports_length() {
261        let v = Value::Map(vec![
262            (int(1), int(2)),
263            (int(3), int(-7)),
264            (int(-1), int(1)),
265            (int(-2), Value::Bytes(vec![7; 32])),
266            (Value::Text("fmt".into()), Value::Text("none".into())),
267            (
268                Value::Text("a".into()),
269                Value::Array(vec![Value::Bool(true), Value::Null, int(1000), int(70000)]),
270            ),
271        ]);
272        let mut b = encode(&v);
273        let n = b.len();
274        b.extend([0xa0, 0xff]);
275        let (d, used) = decode(&b).unwrap();
276        assert_eq!(used, n);
277        assert_eq!(d, v);
278        assert_eq!(d.get_int(3).and_then(Value::as_int), Some(-7));
279        assert_eq!(d.get_text("fmt").and_then(Value::as_text), Some("none"));
280        assert!(decode_all(&b).is_err());
281    }
282
283    #[test]
284    fn known_encodings() {
285        // RFC 8949 appendix A.
286        assert_eq!(decode_all(&[0x19, 0x03, 0xe8]).unwrap(), Value::Uint(1000));
287        assert_eq!(decode_all(&[0x38, 0x63]).unwrap().as_int(), Some(-100));
288        assert_eq!(
289            decode_all(&[0x82, 0x01, 0x82, 0x02, 0x03]).unwrap(),
290            Value::Array(vec![int(1), Value::Array(vec![int(2), int(3)])])
291        );
292        // A tagged item decodes to the item.
293        assert_eq!(decode_all(&[0xc1, 0x01]).unwrap(), Value::Uint(1));
294    }
295
296    #[test]
297    fn refuses_hostile_input() {
298        // Indefinite length, a byte string longer than the input, a huge
299        // array count, deep nesting, truncation, bad UTF-8.
300        assert!(decode(&[0x5f]).is_err());
301        assert!(decode(&[0x5a, 0xff, 0xff, 0xff, 0xff, 0x00]).is_err());
302        assert!(decode(&[0x9b, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff]).is_err());
303        assert!(decode(&[0xbb, 0x80, 0, 0, 0, 0, 0, 0, 0]).is_err());
304        assert!(decode(&[0x81; 40]).is_err());
305        assert!(decode(&[0x19, 0x03]).is_err());
306        assert!(decode(&[0x62, 0xff, 0xfe]).is_err());
307        assert!(decode(&[]).is_err());
308    }
309}