Skip to main content

macula_rust/frame/
check_payload.rs

1//! The checks a sender runs so that nothing the decoding rule refuses on
2//! arrival leaves this node, as macula_frame's check_payload/1 and
3//! check_frame/1: a map key that is not text or an integer, two keys of one
4//! map that encode alike, an integer outside -2^63 to 2^63-1, a float that is
5//! NaN or infinite, too deep a nesting, and too many items. Text is valid
6//! UTF-8 by construction here.
7
8use crate::cbor::{self, Value, MAX_ELEMENTS, MAX_NESTING_DEPTH};
9
10use super::{FrameError, MAX_FRAME_BYTES};
11
12/// How many lists and maps a payload may nest, the outermost counted: a
13/// payload travels inside a frame's map, which takes one level.
14pub const MAX_PAYLOAD_NESTING: usize = MAX_NESTING_DEPTH - 1;
15
16/// How many of the decoding rule's items a payload leaves for the frame
17/// around it.
18pub const FRAME_RESERVED_ELEMENTS: usize = 64;
19
20/// How many CBOR items a payload may hold, itself included.
21pub const MAX_PAYLOAD_ELEMENTS: usize = MAX_ELEMENTS - FRAME_RESERVED_ELEMENTS;
22
23/// Whether `payload` is admissible as a frame payload. It also refuses a
24/// payload whose own encoding is over the frame cap.
25pub fn check_payload(payload: &Value) -> Result<(), FrameError> {
26    let mut check = RuleCheck {
27        subject: "payload",
28        max_items: MAX_PAYLOAD_ELEMENTS,
29        max_nesting: MAX_PAYLOAD_NESTING,
30        items: 0,
31    };
32    check
33        .value(payload, &mut Vec::new())
34        .map_err(FrameError::Payload)?;
35    let encoded = cbor::encode(payload).map_err(|e| FrameError::Payload(e.to_string()))?;
36    if encoded.len() > MAX_FRAME_BYTES {
37        return Err(FrameError::Payload(format!(
38            "the payload encodes to {} bytes, over the {MAX_FRAME_BYTES}-byte frame cap",
39            encoded.len()
40        )));
41    }
42    Ok(())
43}
44
45/// Whether the whole `frame` is one the decoding rule accepts where it
46/// arrives, the check macula runs on every frame before it is sent.
47pub fn check_frame(frame: &Value) -> Result<(), FrameError> {
48    let mut check = RuleCheck {
49        subject: "frame",
50        max_items: MAX_ELEMENTS,
51        max_nesting: MAX_NESTING_DEPTH,
52        items: 0,
53    };
54    check
55        .value(frame, &mut Vec::new())
56        .map_err(FrameError::BreaksDecodingRule)
57}
58
59/// A walk of a payload or a whole frame, its subject, under the decoding
60/// rule's limits for it, counting its items.
61struct RuleCheck {
62    subject: &'static str,
63    max_items: usize,
64    max_nesting: usize,
65    items: usize,
66}
67
68impl RuleCheck {
69    fn value(&mut self, v: &Value, path: &mut Vec<String>) -> Result<(), String> {
70        self.items += 1;
71        if self.items > self.max_items {
72            return Err(format!(
73                "the {} holds more than {} items, at {}",
74                self.subject,
75                self.max_items,
76                self.at(path)
77            ));
78        }
79        match v {
80            Value::Float(f) if !f.is_finite() => {
81                Err(format!("a float that is not finite at {}", self.at(path)))
82            }
83            Value::Int(n) if i64::try_from(*n).is_err() => Err(format!(
84                "an integer outside -2^63 to 2^63-1 at {}",
85                self.at(path)
86            )),
87            Value::List(items) => self.list(items, path),
88            Value::Map(pairs) => self.map(pairs, path),
89            _ => Ok(()),
90        }
91    }
92
93    /// Walks a list at `path`, each item under its index.
94    fn list(&mut self, items: &[Value], path: &mut Vec<String>) -> Result<(), String> {
95        self.nesting(path)?;
96        for (i, item) in items.iter().enumerate() {
97            path.push(i.to_string());
98            self.value(item, path)?;
99            path.pop();
100        }
101        Ok(())
102    }
103
104    /// Walks a map at `path`, refusing a key that is not text or an integer
105    /// and two keys that encode alike.
106    fn map(&mut self, pairs: &[(Value, Value)], path: &mut Vec<String>) -> Result<(), String> {
107        self.nesting(path)?;
108        let mut seen = std::collections::HashSet::with_capacity(pairs.len());
109        for (key, value) in pairs {
110            self.map_entry(key, value, &mut seen, path)?;
111        }
112        Ok(())
113    }
114
115    /// Checks one key and its value of the map at `path`, `seen` holding the
116    /// encodings of the keys before it.
117    fn map_entry(
118        &mut self,
119        key: &Value,
120        value: &Value,
121        seen: &mut std::collections::HashSet<Vec<u8>>,
122        path: &mut Vec<String>,
123    ) -> Result<(), String> {
124        if !matches!(key, Value::Text(_) | Value::Int(_)) {
125            return Err(format!(
126                "a map key that is not text or an integer at {}",
127                self.at(path)
128            ));
129        }
130        self.value(key, path)?;
131        let encoded = cbor::encode(key).map_err(|e| e.to_string())?;
132        if !seen.insert(encoded) {
133            return Err(format!(
134                "two keys of the map at {} encode alike",
135                self.at(path)
136            ));
137        }
138        path.push(match key {
139            Value::Text(t) => t.clone(),
140            other => format!("{other:?}"),
141        });
142        self.value(value, path)?;
143        path.pop();
144        Ok(())
145    }
146
147    /// Refuses a list or map at `path` that would nest more than the limit.
148    fn nesting(&self, path: &[String]) -> Result<(), String> {
149        if path.len() >= self.max_nesting {
150            return Err(format!(
151                "lists and maps at {} nest more than {} levels",
152                self.at(path),
153                self.max_nesting
154            ));
155        }
156        Ok(())
157    }
158
159    fn at(&self, path: &[String]) -> String {
160        if path.is_empty() {
161            format!("the {} root", self.subject)
162        } else {
163            path.join(".")
164        }
165    }
166}