Skip to main content

weavatrix_refactor_plan/
parser.rs

1use crate::{PlanError, PlanErrorCode, RefactorPlan, RefactorPlanLimits};
2use serde::de::{DeserializeSeed, MapAccess, SeqAccess, Visitor};
3use std::{collections::BTreeSet, fmt};
4
5/// Parses, duplicate-checks, bounds, and validates one refactor-plan document.
6pub fn parse_refactor_plan(
7    bytes: &[u8],
8    limits: RefactorPlanLimits,
9) -> Result<RefactorPlan, PlanError> {
10    limits.validate()?;
11    let byte_limit = parse_byte_limit(limits);
12    if bytes.len() > byte_limit {
13        return Err(PlanError::new(
14            PlanErrorCode::PlanTooLarge,
15            format!("serialized plan exceeds the {byte_limit}-byte parse limit"),
16        ));
17    }
18    let mut budget = ParseBudget {
19        nodes: 0,
20        max_nodes: parse_node_limit(limits),
21        max_depth: limits.max_extension_depth.saturating_add(16),
22        max_key_bytes: limits.max_path_bytes,
23    };
24    let mut deserializer = serde_json::Deserializer::from_slice(bytes);
25    let value = CheckedValueSeed {
26        budget: &mut budget,
27        depth: 1,
28    }
29    .deserialize(&mut deserializer)
30    .map_err(|error| parse_error(&error))?;
31    deserializer.end().map_err(|error| parse_error(&error))?;
32    let plan = serde_json::from_value(value).map_err(|error| parse_error(&error))?;
33    crate::validate_consumer_plan(&plan, limits)?;
34    Ok(plan)
35}
36
37fn parse_byte_limit(limits: RefactorPlanLimits) -> usize {
38    limits
39        .max_total_create_bytes
40        .saturating_add(limits.max_total_text_bytes)
41        .saturating_add(limits.max_extension_bytes)
42        .saturating_add(limits.max_evidence_text_bytes)
43        .saturating_add(limits.max_paths.saturating_mul(limits.max_path_bytes))
44        .saturating_add(1024 * 1024)
45}
46
47fn parse_node_limit(limits: RefactorPlanLimits) -> usize {
48    limits
49        .max_extension_nodes
50        .saturating_add(limits.max_operations.saturating_mul(16))
51        .saturating_add(limits.max_total_edits.saturating_mul(16))
52        .saturating_add(limits.max_evidence_entries.saturating_mul(16))
53        .saturating_add(64)
54}
55
56struct ParseBudget {
57    nodes: usize,
58    max_nodes: usize,
59    max_depth: usize,
60    max_key_bytes: usize,
61}
62
63struct CheckedValueSeed<'a> {
64    budget: &'a mut ParseBudget,
65    depth: usize,
66}
67
68impl<'de> DeserializeSeed<'de> for CheckedValueSeed<'_> {
69    type Value = serde_json::Value;
70
71    fn deserialize<D>(self, deserializer: D) -> Result<Self::Value, D::Error>
72    where
73        D: serde::Deserializer<'de>,
74    {
75        if self.depth > self.budget.max_depth {
76            return Err(serde::de::Error::custom(
77                "JSON nesting exceeds parse depth limit",
78            ));
79        }
80        if self.budget.nodes >= self.budget.max_nodes {
81            return Err(serde::de::Error::custom(
82                "JSON value count exceeds parse node limit",
83            ));
84        }
85        self.budget.nodes += 1;
86        deserializer.deserialize_any(CheckedValueVisitor {
87            budget: self.budget,
88            depth: self.depth,
89        })
90    }
91}
92
93struct CheckedValueVisitor<'a> {
94    budget: &'a mut ParseBudget,
95    depth: usize,
96}
97
98impl<'de> Visitor<'de> for CheckedValueVisitor<'_> {
99    type Value = serde_json::Value;
100
101    fn expecting(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
102        formatter.write_str("a duplicate-free JSON value")
103    }
104
105    fn visit_bool<E>(self, value: bool) -> Result<Self::Value, E> {
106        Ok(serde_json::Value::Bool(value))
107    }
108
109    fn visit_i64<E>(self, value: i64) -> Result<Self::Value, E> {
110        Ok(serde_json::Value::Number(value.into()))
111    }
112
113    fn visit_u64<E>(self, value: u64) -> Result<Self::Value, E> {
114        Ok(serde_json::Value::Number(value.into()))
115    }
116
117    fn visit_f64<E>(self, value: f64) -> Result<Self::Value, E>
118    where
119        E: serde::de::Error,
120    {
121        serde_json::Number::from_f64(value)
122            .map(serde_json::Value::Number)
123            .ok_or_else(|| E::custom("non-finite JSON number"))
124    }
125
126    fn visit_str<E>(self, value: &str) -> Result<Self::Value, E> {
127        Ok(serde_json::Value::String(value.to_owned()))
128    }
129
130    fn visit_string<E>(self, value: String) -> Result<Self::Value, E> {
131        Ok(serde_json::Value::String(value))
132    }
133
134    fn visit_none<E>(self) -> Result<Self::Value, E> {
135        Ok(serde_json::Value::Null)
136    }
137
138    fn visit_unit<E>(self) -> Result<Self::Value, E> {
139        Ok(serde_json::Value::Null)
140    }
141
142    fn visit_seq<A>(self, mut sequence: A) -> Result<Self::Value, A::Error>
143    where
144        A: SeqAccess<'de>,
145    {
146        let mut values = Vec::with_capacity(sequence.size_hint().unwrap_or(0).min(1024));
147        while let Some(value) = sequence.next_element_seed(CheckedValueSeed {
148            budget: self.budget,
149            depth: self.depth + 1,
150        })? {
151            values.push(value);
152        }
153        Ok(serde_json::Value::Array(values))
154    }
155
156    fn visit_map<A>(self, mut object: A) -> Result<Self::Value, A::Error>
157    where
158        A: MapAccess<'de>,
159    {
160        let mut values = serde_json::Map::new();
161        let mut keys = BTreeSet::new();
162        while let Some(key) = object.next_key_seed(BoundedKeySeed {
163            max_bytes: self.budget.max_key_bytes,
164        })? {
165            if !keys.insert(key.clone()) {
166                return Err(serde::de::Error::custom("duplicate JSON object member"));
167            }
168            let value = object.next_value_seed(CheckedValueSeed {
169                budget: self.budget,
170                depth: self.depth + 1,
171            })?;
172            values.insert(key, value);
173        }
174        Ok(serde_json::Value::Object(values))
175    }
176}
177
178struct BoundedKeySeed {
179    max_bytes: usize,
180}
181
182impl<'de> DeserializeSeed<'de> for BoundedKeySeed {
183    type Value = String;
184
185    fn deserialize<D>(self, deserializer: D) -> Result<Self::Value, D::Error>
186    where
187        D: serde::Deserializer<'de>,
188    {
189        deserializer.deserialize_str(BoundedKeyVisitor {
190            max_bytes: self.max_bytes,
191        })
192    }
193}
194
195struct BoundedKeyVisitor {
196    max_bytes: usize,
197}
198
199impl Visitor<'_> for BoundedKeyVisitor {
200    type Value = String;
201
202    fn expecting(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
203        formatter.write_str("a bounded JSON object member name")
204    }
205
206    fn visit_borrowed_str<E>(self, value: &str) -> Result<Self::Value, E>
207    where
208        E: serde::de::Error,
209    {
210        self.check(value)
211    }
212
213    fn visit_str<E>(self, value: &str) -> Result<Self::Value, E>
214    where
215        E: serde::de::Error,
216    {
217        self.check(value)
218    }
219
220    fn visit_string<E>(self, value: String) -> Result<Self::Value, E>
221    where
222        E: serde::de::Error,
223    {
224        if value.len() > self.max_bytes {
225            return Err(E::custom("JSON object member name exceeds byte limit"));
226        }
227        Ok(value)
228    }
229}
230
231impl BoundedKeyVisitor {
232    fn check<E>(self, value: &str) -> Result<String, E>
233    where
234        E: serde::de::Error,
235    {
236        if value.len() > self.max_bytes {
237            return Err(E::custom("JSON object member name exceeds byte limit"));
238        }
239        Ok(value.to_owned())
240    }
241}
242
243fn parse_error(error: &serde_json::Error) -> PlanError {
244    PlanError::new(
245        PlanErrorCode::EvidenceMalformed,
246        format!("could not parse refactor plan: {error}"),
247    )
248}
249
250#[cfg(test)]
251mod tests {
252    use super::parse_refactor_plan;
253    use crate::RefactorPlanLimits;
254
255    #[test]
256    fn rejects_duplicate_members_at_any_depth() {
257        for source in [
258            br#"{"schemaVersion":"a","schemaVersion":"b"}"#.as_slice(),
259            br#"{"outer":{"same":1,"same":2}}"#.as_slice(),
260            br#"[{"same":1,"same":2}]"#.as_slice(),
261        ] {
262            assert!(parse_refactor_plan(source, RefactorPlanLimits::default()).is_err());
263        }
264    }
265}