Skip to main content

micromegas_transit/
parser.rs

1use crate::{
2    UserDefinedType, advance_window, read_advance_string_in, try_read_consume_pod, try_read_pod_at,
3    value::{Object, Value},
4};
5use anyhow::{Context, Result, bail};
6use bumpalo::Bump;
7use std::{collections::HashMap, sync::Arc};
8
9/// A reader for a custom (dynamically-sized) transit type.
10///
11/// Two higher-ranked lifetimes keep the returned `Value<'a>` (which borrows the
12/// arena and source buffer) independent of the `&'dep` borrow of the dependency
13/// map, so the caller can insert the result back into the map afterwards.
14pub type CustomReader = Arc<
15    dyn for<'a, 'dep> Fn(
16        &'a Bump,
17        &'a UserDefinedType,
18        &'a [UserDefinedType],
19        &'dep HashMap<u64, Value<'a>>,
20        &'a [u8],
21    ) -> Result<Value<'a>>,
22>;
23pub type CustomReaderMap = HashMap<String, CustomReader>;
24
25pub fn read_dependencies<'a>(
26    bump: &'a Bump,
27    custom_readers: &CustomReaderMap,
28    udts: &'a [UserDefinedType],
29    buffer: &'a [u8],
30) -> Result<HashMap<u64, Value<'a>>> {
31    let mut hash = HashMap::new();
32    let mut offset = 0;
33    while offset < buffer.len() {
34        let type_index = buffer[offset] as usize;
35        if type_index >= udts.len() {
36            bail!(
37                "Invalid type index parsing transit dependencies: {}",
38                type_index
39            );
40        }
41        offset += 1;
42        let udt = &udts[type_index];
43        let object_size = match udt.size {
44            0 => {
45                //dynamic size
46                let obj_size = try_read_pod_at::<u32>(buffer, offset)?;
47                offset += std::mem::size_of::<u32>();
48                obj_size as usize
49            }
50            static_size => static_size,
51        };
52        // Single guard: bounds every slice/advance derived from this object below
53        // (offset..object_end) against the actual buffer length.
54        let object_end = match offset.checked_add(object_size) {
55            Some(end) if end <= buffer.len() => end,
56            _ => bail!(
57                "corrupt block: object at offset {offset} with size {object_size} exceeds {}-byte buffer",
58                buffer.len()
59            ),
60        };
61
62        match udt.name.as_str() {
63            "StaticString" => {
64                // zero-copy: the string borrows the (whole-block) source buffer.
65                let string_id = try_read_pod_at::<u64>(buffer, offset)?;
66                let nb_utf8_bytes = object_size
67                    .checked_sub(std::mem::size_of::<usize>())
68                    .with_context(|| {
69                        format!("StaticString object_size {object_size} smaller than header")
70                    })?;
71                let str_start = offset + std::mem::size_of::<usize>();
72                let s = std::str::from_utf8(&buffer[str_start..str_start + nb_utf8_bytes])
73                    .with_context(|| "str::from_utf8")?;
74                let insert_res = hash.insert(string_id, Value::String(s));
75                if insert_res.is_some() {
76                    bail!("duplicate dependency id {string_id}");
77                }
78            }
79            "StaticStringDependency" => {
80                // This window covers the entire remaining buffer, not just
81                // object_size bytes, so the guard above doesn't bound what's
82                // read here — the checked reads below do that.
83                let mut window = advance_window(buffer, offset);
84                let string_id: u64 = try_read_consume_pod(&mut window)?;
85                let s =
86                    read_advance_string_in(bump, &mut window).with_context(|| "parsing string")?;
87                let insert_res = hash.insert(string_id, Value::String(s));
88                if insert_res.is_some() {
89                    bail!("duplicate dependency id {string_id}");
90                }
91            }
92            type_name => {
93                if let Some(reader) = custom_readers.get(type_name) {
94                    // Same unbounded-window property as StaticStringDependency
95                    // above: the reader validates every field it consumes.
96                    let window = advance_window(buffer, offset);
97                    if let Value::Object(obj) = (*reader)(bump, udt, udts, &hash, window)
98                        .with_context(|| "parsing custom dependency")?
99                    {
100                        let id: u64 = obj
101                            .get("id")
102                            .with_context(|| "reading id of custom dependency")?;
103                        let value = *obj
104                            .get_ref("value")
105                            .with_context(|| "reading value of custom dependency")?;
106                        let insert_res = hash.insert(id, value);
107                        if insert_res.is_some() {
108                            bail!("duplicate dependency id {id}");
109                        }
110                    } else {
111                        anyhow::bail!("custom dependency is not an object");
112                    }
113                } else {
114                    if udt.size == 0 {
115                        anyhow::bail!("invalid user-defined type {:?}", udt);
116                    }
117                    let instance =
118                        parse_pod_instance(bump, udt, udts, &hash, &buffer[offset..object_end])
119                            .with_context(|| "parse_pod_instance")?;
120                    if let Value::Object(obj) = instance {
121                        let id = obj.get::<u64>("id")?;
122                        let insert_res = hash.insert(id, Value::Object(obj));
123                        if insert_res.is_some() {
124                            bail!("duplicate dependency id {id}");
125                        }
126                    }
127                }
128            }
129        }
130        offset = object_end;
131    }
132
133    Ok(hash)
134}
135
136fn parse_custom_instance<'a>(
137    bump: &'a Bump,
138    custom_readers: &CustomReaderMap,
139    udt: &'a UserDefinedType,
140    udts: &'a [UserDefinedType],
141    dependencies: &HashMap<u64, Value<'a>>,
142    object_window: &'a [u8],
143) -> Result<Value<'a>> {
144    if let Some(reader) = custom_readers.get(&*udt.name) {
145        (*reader)(bump, udt, udts, dependencies, object_window)
146    } else {
147        log::warn!("unknown custom object {}", udt.name);
148        Ok(Value::Object(bump.alloc(Object {
149            type_name: udt.name.as_str(),
150            members: &[],
151        })))
152    }
153}
154
155pub fn parse_pod_instance<'a>(
156    bump: &'a Bump,
157    udt: &'a UserDefinedType,
158    udts: &'a [UserDefinedType],
159    dependencies: &HashMap<u64, Value<'a>>,
160    object_window: &'a [u8],
161) -> Result<Value<'a>> {
162    let mut members = bumpalo::collections::Vec::with_capacity_in(udt.members.len(), bump);
163    for member_meta in &udt.members {
164        let name: &'a str = member_meta.name.as_str();
165        // Bounds the reference-key read and every intrinsic read below against
166        // the object window (member metadata is untrusted, same as the payload).
167        match member_meta.offset.checked_add(member_meta.size) {
168            Some(end) if end <= object_window.len() => {}
169            _ => bail!(
170                "corrupt block: member {member_meta:?} exceeds {}-byte object window",
171                object_window.len()
172            ),
173        }
174        let value = if member_meta.is_reference {
175            if member_meta.size < std::mem::size_of::<u64>() {
176                bail!(
177                    "member references have to be at least 8 bytes {:?}",
178                    member_meta
179                );
180            }
181            let key = try_read_pod_at::<u64>(object_window, member_meta.offset)?;
182            if let Some(v) = dependencies.get(&key) {
183                *v
184            } else {
185                bail!("dependency not found: member={member_meta:?} key={key}");
186            }
187        } else {
188            match member_meta.type_name.as_str() {
189                "u8" | "uint8" => {
190                    if std::mem::size_of::<u8>() != member_meta.size {
191                        bail!("type size mismatch for member {member_meta:?}");
192                    }
193                    Value::U8(try_read_pod_at::<u8>(object_window, member_meta.offset)?)
194                }
195                "u32" | "uint32" => {
196                    if std::mem::size_of::<u32>() != member_meta.size {
197                        bail!("type size mismatch for member {member_meta:?}");
198                    }
199                    Value::U32(try_read_pod_at::<u32>(object_window, member_meta.offset)?)
200                }
201                "u64" | "uint64" => {
202                    if std::mem::size_of::<u64>() != member_meta.size {
203                        bail!("type size mismatch for member {member_meta:?}");
204                    }
205                    Value::U64(try_read_pod_at::<u64>(object_window, member_meta.offset)?)
206                }
207                "i64" | "int64" => {
208                    if std::mem::size_of::<i64>() != member_meta.size {
209                        bail!("type size mismatch for member {member_meta:?}");
210                    }
211                    Value::I64(try_read_pod_at::<i64>(object_window, member_meta.offset)?)
212                }
213                "f64" => {
214                    if std::mem::size_of::<f64>() != member_meta.size {
215                        bail!("type size mismatch for member {member_meta:?}");
216                    }
217                    Value::F64(try_read_pod_at::<f64>(object_window, member_meta.offset)?)
218                }
219                non_intrinsic_member_type_name => {
220                    if let Some(index) = udts
221                        .iter()
222                        .position(|t| *t.name == non_intrinsic_member_type_name)
223                    {
224                        let member_udt = &udts[index];
225                        // member_udt.size may differ from member_meta.size, so it
226                        // needs its own guard before slicing.
227                        let udt_end = match member_meta.offset.checked_add(member_udt.size) {
228                            Some(end) if end <= object_window.len() => end,
229                            _ => bail!(
230                                "corrupt block: nested member {member_meta:?} exceeds {}-byte object window",
231                                object_window.len()
232                            ),
233                        };
234                        parse_pod_instance(
235                            bump,
236                            member_udt,
237                            udts,
238                            dependencies,
239                            &object_window[member_meta.offset..udt_end],
240                        )
241                        .with_context(|| "parse_pod_instance")?
242                    } else {
243                        bail!("unknown member type {}", non_intrinsic_member_type_name);
244                    }
245                }
246            }
247        };
248        members.push((name, value));
249    }
250
251    if udt.is_reference {
252        // reference objects need a member called 'id' which is the key to the dependency
253        if let Some(id_index) = members.iter().position(|m| m.0 == "id") {
254            return Ok(members[id_index].1);
255        }
256        bail!("reference object has no 'id' member");
257    }
258
259    Ok(Value::Object(bump.alloc(Object {
260        type_name: udt.name.as_str(),
261        members: members.into_bump_slice(),
262    })))
263}
264
265// parse_object_buffer calls fun for each object in the buffer until fun returns
266// `false`
267pub fn parse_object_buffer<'a, F>(
268    bump: &'a Bump,
269    custom_readers: &CustomReaderMap,
270    dependencies: &HashMap<u64, Value<'a>>,
271    udts: &'a [UserDefinedType],
272    buffer: &'a [u8],
273    mut fun: F,
274) -> Result<bool>
275where
276    F: FnMut(Value<'a>) -> Result<bool>,
277{
278    let mut offset = 0;
279    while offset < buffer.len() {
280        let type_index = buffer[offset] as usize;
281        if type_index >= udts.len() {
282            bail!("Invalid type index parsing transit objects: {}", type_index);
283        }
284        offset += 1;
285        let udt = &udts[type_index];
286        let (object_size, is_size_dynamic) = match udt.size {
287            0 => {
288                //dynamic size
289                let obj_size = try_read_pod_at::<u32>(buffer, offset)?;
290                offset += std::mem::size_of::<u32>();
291                (obj_size as usize, true)
292            }
293            static_size => (static_size, false),
294        };
295        let object_end = match offset.checked_add(object_size) {
296            Some(end) if end <= buffer.len() => end,
297            _ => bail!(
298                "corrupt block: object at offset {offset} with size {object_size} exceeds {}-byte buffer",
299                buffer.len()
300            ),
301        };
302        let instance = if is_size_dynamic {
303            parse_custom_instance(
304                bump,
305                custom_readers,
306                udt,
307                udts,
308                dependencies,
309                &buffer[offset..object_end],
310            )
311            .with_context(|| "parse_custom_instance")?
312        } else {
313            parse_pod_instance(bump, udt, udts, dependencies, &buffer[offset..object_end])
314                .with_context(|| "parse_pod_instance")?
315        };
316        if !fun(instance)? {
317            return Ok(false);
318        }
319        offset = object_end;
320    }
321    Ok(true)
322}