Skip to main content

memra_tokenizer/
json.rs

1//! Minimal recursive-descent JSON parser for HF tokenizer sidecar files
2//! (tokenizer.json / tokenizer_config.json / generation_config.json).
3//!
4//! memra-tokenizer stays serde-free (crate policy, same as memra-gguf's hand
5//! parser in safetensors.rs). This is a full-value parser (objects, arrays,
6//! strings with \uXXXX + surrogate pairs, numbers, bools, null) — the
7//! tokenizer.json vocab is ~1e5 entries so parsing is byte-based, no regex.
8
9use std::collections::HashMap;
10
11#[derive(Debug, Clone, PartialEq)]
12pub enum Value {
13    Null,
14    Bool(bool),
15    Num(f64),
16    Str(String),
17    Arr(Vec<Value>),
18    Obj(HashMap<String, Value>),
19}
20
21impl Value {
22    pub fn as_str(&self) -> Option<&str> {
23        match self {
24            Value::Str(s) => Some(s),
25            _ => None,
26        }
27    }
28    pub fn as_u64(&self) -> Option<u64> {
29        match self {
30            Value::Num(n) if *n >= 0.0 && n.fract() == 0.0 => Some(*n as u64),
31            _ => None,
32        }
33    }
34    pub fn as_bool(&self) -> Option<bool> {
35        match self {
36            Value::Bool(b) => Some(*b),
37            _ => None,
38        }
39    }
40    pub fn as_arr(&self) -> Option<&[Value]> {
41        match self {
42            Value::Arr(a) => Some(a),
43            _ => None,
44        }
45    }
46    pub fn as_obj(&self) -> Option<&HashMap<String, Value>> {
47        match self {
48            Value::Obj(o) => Some(o),
49            _ => None,
50        }
51    }
52    /// obj["key"] convenience (None on non-objects / missing keys).
53    pub fn get(&self, key: &str) -> Option<&Value> {
54        self.as_obj().and_then(|o| o.get(key))
55    }
56}
57
58pub fn parse(text: &str) -> Result<Value, String> {
59    let mut p = Parser {
60        b: text.as_bytes(),
61        i: 0,
62    };
63    p.ws();
64    let v = p.value()?;
65    p.ws();
66    if p.i != p.b.len() {
67        return Err(format!("json: trailing bytes at {}", p.i));
68    }
69    Ok(v)
70}
71
72struct Parser<'a> {
73    b: &'a [u8],
74    i: usize,
75}
76
77impl<'a> Parser<'a> {
78    fn ws(&mut self) {
79        while self.i < self.b.len() && matches!(self.b[self.i], b' ' | b'\t' | b'\n' | b'\r') {
80            self.i += 1;
81        }
82    }
83
84    fn peek(&self) -> Result<u8, String> {
85        self.b
86            .get(self.i)
87            .copied()
88            .ok_or_else(|| "json: unexpected EOF".to_string())
89    }
90
91    fn expect(&mut self, c: u8) -> Result<(), String> {
92        if self.peek()? != c {
93            return Err(format!(
94                "json: expected '{}' at byte {}, got '{}'",
95                c as char, self.i, self.b[self.i] as char
96            ));
97        }
98        self.i += 1;
99        Ok(())
100    }
101
102    fn value(&mut self) -> Result<Value, String> {
103        match self.peek()? {
104            b'{' => self.object(),
105            b'[' => self.array(),
106            b'"' => Ok(Value::Str(self.string()?)),
107            b't' => self.lit(b"true", Value::Bool(true)),
108            b'f' => self.lit(b"false", Value::Bool(false)),
109            b'n' => self.lit(b"null", Value::Null),
110            _ => self.number(),
111        }
112    }
113
114    fn lit(&mut self, word: &[u8], v: Value) -> Result<Value, String> {
115        if self.b.len() - self.i >= word.len() && &self.b[self.i..self.i + word.len()] == word {
116            self.i += word.len();
117            Ok(v)
118        } else {
119            Err(format!("json: bad literal at byte {}", self.i))
120        }
121    }
122
123    fn number(&mut self) -> Result<Value, String> {
124        let start = self.i;
125        while self.i < self.b.len()
126            && matches!(
127                self.b[self.i],
128                b'0'..=b'9' | b'-' | b'+' | b'.' | b'e' | b'E'
129            )
130        {
131            self.i += 1;
132        }
133        if self.i == start {
134            return Err(format!("json: expected value at byte {}", start));
135        }
136        std::str::from_utf8(&self.b[start..self.i])
137            .ok()
138            .and_then(|s| s.parse::<f64>().ok())
139            .map(Value::Num)
140            .ok_or_else(|| format!("json: bad number at byte {start}"))
141    }
142
143    fn hex4(&mut self) -> Result<u32, String> {
144        if self.i + 4 > self.b.len() {
145            return Err("json: truncated \\u escape".into());
146        }
147        let s = std::str::from_utf8(&self.b[self.i..self.i + 4])
148            .map_err(|_| "json: bad \\u escape".to_string())?;
149        let v = u32::from_str_radix(s, 16).map_err(|_| "json: bad \\u escape".to_string())?;
150        self.i += 4;
151        Ok(v)
152    }
153
154    fn string(&mut self) -> Result<String, String> {
155        self.expect(b'"')?;
156        let mut out = String::new();
157        loop {
158            let c = self.peek()?;
159            match c {
160                b'"' => {
161                    self.i += 1;
162                    return Ok(out);
163                }
164                b'\\' => {
165                    self.i += 1;
166                    let e = self.peek()?;
167                    self.i += 1;
168                    match e {
169                        b'"' => out.push('"'),
170                        b'\\' => out.push('\\'),
171                        b'/' => out.push('/'),
172                        b'b' => out.push('\u{8}'),
173                        b'f' => out.push('\u{c}'),
174                        b'n' => out.push('\n'),
175                        b'r' => out.push('\r'),
176                        b't' => out.push('\t'),
177                        b'u' => {
178                            let hi = self.hex4()?;
179                            let cp = if (0xD800..0xDC00).contains(&hi) {
180                                // surrogate pair: require \uXXXX low surrogate
181                                if self.i + 2 > self.b.len()
182                                    || self.b[self.i] != b'\\'
183                                    || self.b[self.i + 1] != b'u'
184                                {
185                                    return Err("json: lone high surrogate".into());
186                                }
187                                self.i += 2;
188                                let lo = self.hex4()?;
189                                if !(0xDC00..0xE000).contains(&lo) {
190                                    return Err("json: bad low surrogate".into());
191                                }
192                                0x10000 + ((hi - 0xD800) << 10) + (lo - 0xDC00)
193                            } else {
194                                hi
195                            };
196                            out.push(
197                                char::from_u32(cp)
198                                    .ok_or_else(|| "json: invalid codepoint".to_string())?,
199                            );
200                        }
201                        _ => return Err(format!("json: bad escape at byte {}", self.i - 1)),
202                    }
203                }
204                _ => {
205                    // consume one UTF-8 sequence verbatim (fast path: run of plain bytes)
206                    let start = self.i;
207                    while self.i < self.b.len() && self.b[self.i] != b'"' && self.b[self.i] != b'\\'
208                    {
209                        self.i += 1;
210                    }
211                    out.push_str(
212                        std::str::from_utf8(&self.b[start..self.i])
213                            .map_err(|_| "json: invalid utf-8 in string".to_string())?,
214                    );
215                }
216            }
217        }
218    }
219
220    fn array(&mut self) -> Result<Value, String> {
221        self.expect(b'[')?;
222        let mut out = Vec::new();
223        self.ws();
224        if self.peek()? == b']' {
225            self.i += 1;
226            return Ok(Value::Arr(out));
227        }
228        loop {
229            self.ws();
230            out.push(self.value()?);
231            self.ws();
232            match self.peek()? {
233                b',' => self.i += 1,
234                b']' => {
235                    self.i += 1;
236                    return Ok(Value::Arr(out));
237                }
238                c => {
239                    return Err(format!(
240                        "json: expected ',' or ']' at byte {}, got '{}'",
241                        self.i, c as char
242                    ));
243                }
244            }
245        }
246    }
247
248    fn object(&mut self) -> Result<Value, String> {
249        self.expect(b'{')?;
250        let mut out = HashMap::new();
251        self.ws();
252        if self.peek()? == b'}' {
253            self.i += 1;
254            return Ok(Value::Obj(out));
255        }
256        loop {
257            self.ws();
258            let k = self.string()?;
259            self.ws();
260            self.expect(b':')?;
261            self.ws();
262            let v = self.value()?;
263            out.insert(k, v);
264            self.ws();
265            match self.peek()? {
266                b',' => self.i += 1,
267                b'}' => {
268                    self.i += 1;
269                    return Ok(Value::Obj(out));
270                }
271                c => {
272                    return Err(format!(
273                        "json: expected ',' or '}}' at byte {}, got '{}'",
274                        self.i, c as char
275                    ));
276                }
277            }
278        }
279    }
280}
281
282#[cfg(test)]
283mod tests {
284    use super::*;
285
286    #[test]
287    fn scalars_and_nesting() {
288        let v = parse(r#"{"a": [1, -2.5, "xĠy", true, null], "b": {"c": "😀"}}"#).unwrap();
289        assert_eq!(v.get("a").unwrap().as_arr().unwrap()[0].as_u64(), Some(1));
290        assert_eq!(
291            v.get("a").unwrap().as_arr().unwrap()[2].as_str(),
292            Some("x\u{120}y")
293        );
294        assert_eq!(
295            v.get("b").unwrap().get("c").unwrap().as_str(),
296            Some("\u{1F600}")
297        );
298    }
299}