use mlua::{Lua, Table, Value};
const CEILING: i64 = 2000;
struct Config {
sparse_convert: bool,
sparse_ratio: i64,
sparse_safe: i64,
encode_max_depth: i64,
decode_max_depth: i64,
invalid_encode: i64,
invalid_decode: bool,
precision: usize,
array_mt: bool,
}
impl Config {
fn read(t: &Table) -> mlua::Result<Config> {
let n = |name: &str| -> mlua::Result<i64> { t.raw_get(name) };
Ok(Config {
sparse_convert: n("encode_sparse_convert")? != 0,
sparse_ratio: n("encode_sparse_ratio")?,
sparse_safe: n("encode_sparse_safe")?,
encode_max_depth: n("encode_max_depth")?,
decode_max_depth: n("decode_max_depth")?,
invalid_encode: n("encode_invalid_numbers")?,
invalid_decode: n("decode_invalid_numbers")? != 0,
precision: n("encode_number_precision")?.clamp(1, 14) as usize,
array_mt: n("decode_array_with_array_mt")? != 0,
})
}
}
pub(super) fn typename(v: &Value) -> &'static str {
match v {
Value::Nil => "nil",
Value::Boolean(_) => "boolean",
Value::Integer(_) | Value::Number(_) => "number",
Value::String(_) => "string",
Value::Table(_) => "table",
Value::Function(_) => "function",
Value::Thread(_) => "thread",
_ => "userdata",
}
}
fn g_fmt(x: f64, precision: usize) -> String {
if x.is_nan() {
return "nan".to_string();
}
if x.is_infinite() {
return if x < 0.0 { "-inf" } else { "inf" }.to_string();
}
let precision = precision.max(1);
let scientific = format!("{:.*e}", precision - 1, x);
let (mantissa, exponent) = match scientific.split_once('e') {
Some(parts) => parts,
None => return scientific,
};
let exponent: i32 = exponent.parse().unwrap_or(0);
if exponent < -4 || exponent >= precision as i32 {
let sign = if exponent < 0 { '-' } else { '+' };
return format!("{}e{}{:02}", trim(mantissa), sign, exponent.abs());
}
let places = (precision as i32 - 1 - exponent).max(0) as usize;
trim(&format!("{x:.places$}"))
}
fn trim(text: &str) -> String {
if !text.contains('.') {
return text.to_string();
}
text.trim_end_matches('0').trim_end_matches('.').to_string()
}
fn escape(byte: u8) -> Option<&'static str> {
Some(match byte {
0x08 => "\\b",
0x09 => "\\t",
0x0a => "\\n",
0x0c => "\\f",
0x0d => "\\r",
b'"' => "\\\"",
b'/' => "\\/",
b'\\' => "\\\\",
0x00..=0x1f | 0x7f => {
const LONG: [&str; 33] = [
"\\u0000", "\\u0001", "\\u0002", "\\u0003", "\\u0004", "\\u0005", "\\u0006",
"\\u0007", "\\u0008", "\\u0009", "\\u000a", "\\u000b", "\\u000c", "\\u000d",
"\\u000e", "\\u000f", "\\u0010", "\\u0011", "\\u0012", "\\u0013", "\\u0014",
"\\u0015", "\\u0016", "\\u0017", "\\u0018", "\\u0019", "\\u001a", "\\u001b",
"\\u001c", "\\u001d", "\\u001e", "\\u001f", "\\u007f",
];
LONG[if byte == 0x7f { 32 } else { byte as usize }]
}
_ => return None,
})
}
struct Encoder<'a> {
cfg: &'a Config,
out: Vec<u8>,
}
impl Encoder<'_> {
fn data(&mut self, value: &Value, depth: i64) -> Result<(), String> {
match value {
Value::String(s) => {
self.string(&s.as_bytes());
Ok(())
}
Value::Integer(n) => self.number(*n as f64),
Value::Number(n) => self.number(*n),
Value::Boolean(yes) => {
self.out
.extend_from_slice(if *yes { b"true" } else { b"false" });
Ok(())
}
Value::Nil => {
self.out.extend_from_slice(b"null");
Ok(())
}
Value::LightUserData(u) if u.0.is_null() => {
self.out.extend_from_slice(b"null");
Ok(())
}
Value::Table(t) => self.table(t, depth + 1),
other => Err(format!(
"Cannot serialise {}: type not supported",
typename(other)
)),
}
}
fn table(&mut self, t: &Table, depth: i64) -> Result<(), String> {
if depth > self.cfg.encode_max_depth || depth > CEILING {
return Err(format!("Cannot serialise, excessive nesting ({depth})"));
}
if let Some(real) = super::argue::behind(t) {
return self.table(&real, depth);
}
if marked(t) {
let len = t.raw_len() as i64;
return self.array(t, len, depth);
}
match self.length(t)? {
len if len > 0 => self.array(t, len, depth),
_ => self.object(t, depth),
}
}
fn length(&self, t: &Table) -> Result<i64, String> {
let mut max = 0i64;
let mut items = 0i64;
for pair in t.pairs::<Value, Value>() {
let (key, _) = pair.map_err(|e| e.to_string())?;
let k = match key {
Value::Integer(n) => n as f64,
Value::Number(n) => n,
_ => return Ok(-1),
};
if k == 0.0 || k.floor() != k || k < 1.0 {
return Ok(-1);
}
let k = k.min(i64::MAX as f64) as i64;
max = max.max(k);
items += 1;
}
if self.cfg.sparse_ratio > 0
&& max > items.saturating_mul(self.cfg.sparse_ratio)
&& max > self.cfg.sparse_safe
{
if !self.cfg.sparse_convert {
return Err("Cannot serialise table: excessively sparse array".to_string());
}
return Ok(-1);
}
Ok(max)
}
fn array(&mut self, t: &Table, len: i64, depth: i64) -> Result<(), String> {
self.out.push(b'[');
for i in 1..=len {
if i > 1 {
self.out.push(b',');
}
let value: Value = t.raw_get(i).unwrap_or(Value::Nil);
self.data(&value, depth)?;
}
self.out.push(b']');
Ok(())
}
fn object(&mut self, t: &Table, depth: i64) -> Result<(), String> {
self.out.push(b'{');
let mut first = true;
for pair in t.pairs::<Value, Value>() {
let (key, value) = pair.map_err(|e| e.to_string())?;
if !first {
self.out.push(b',');
}
first = false;
match &key {
Value::Integer(n) => self.number_key(*n as f64)?,
Value::Number(n) => self.number_key(*n)?,
Value::String(s) => {
self.string(&s.as_bytes());
self.out.push(b':');
}
other => {
return Err(format!(
"Cannot serialise {}: table key must be a number or string",
typename(other)
));
}
}
self.data(&value, depth)?;
}
self.out.push(b'}');
Ok(())
}
fn number_key(&mut self, x: f64) -> Result<(), String> {
self.out.push(b'"');
self.number(x)?;
self.out.extend_from_slice(b"\":");
Ok(())
}
fn number(&mut self, x: f64) -> Result<(), String> {
match self.cfg.invalid_encode {
0 if !x.is_finite() => {
return Err("Cannot serialise number: must not be NaN or Inf".to_string());
}
1 if x.is_nan() => {
self.out.extend_from_slice(b"nan");
return Ok(());
}
2 if !x.is_finite() => {
self.out.extend_from_slice(b"null");
return Ok(());
}
_ => {}
}
self.out
.extend_from_slice(g_fmt(x, self.cfg.precision).as_bytes());
Ok(())
}
fn string(&mut self, bytes: &[u8]) {
self.out.push(b'"');
for &byte in bytes {
match escape(byte) {
Some(text) => self.out.extend_from_slice(text.as_bytes()),
None => self.out.push(byte),
}
}
self.out.push(b'"');
}
}
fn marked(t: &Table) -> bool {
t.metatable()
.and_then(|mt| mt.raw_get::<Value>("__is_cjson_array").ok())
.is_some_and(|v| !matches!(v, Value::Nil | Value::Boolean(false)))
}
#[derive(Debug, PartialEq)]
enum Tok {
ObjBegin,
ObjEnd,
ArrBegin,
ArrEnd,
Str(Vec<u8>),
Num(f64),
Bool(bool),
Null,
Colon,
Comma,
End,
Bad(&'static str),
}
impl Tok {
fn found(&self) -> &str {
match self {
Tok::ObjBegin => "T_OBJ_BEGIN",
Tok::ObjEnd => "T_OBJ_END",
Tok::ArrBegin => "T_ARR_BEGIN",
Tok::ArrEnd => "T_ARR_END",
Tok::Str(_) => "T_STRING",
Tok::Num(_) => "T_NUMBER",
Tok::Bool(_) => "T_BOOLEAN",
Tok::Null => "T_NULL",
Tok::Colon => "T_COLON",
Tok::Comma => "T_COMMA",
Tok::End => "T_END",
Tok::Bad(why) => why,
}
}
}
struct Found {
tok: Tok,
at: usize,
}
struct Parser<'a> {
data: &'a [u8],
at: usize,
cfg: &'a Config,
depth: i64,
}
impl<'a> Parser<'a> {
fn byte(&self, at: usize) -> u8 {
self.data.get(at).copied().unwrap_or(0)
}
fn next(&mut self) -> Found {
while matches!(self.byte(self.at), b' ' | b'\t' | b'\n' | b'\r') {
self.at += 1;
}
let at = self.at;
let ch = self.byte(at);
let tok = match ch {
b'{' => {
self.at += 1;
Tok::ObjBegin
}
b'}' => {
self.at += 1;
Tok::ObjEnd
}
b'[' => {
self.at += 1;
Tok::ArrBegin
}
b']' => {
self.at += 1;
Tok::ArrEnd
}
b',' => {
self.at += 1;
Tok::Comma
}
b':' => {
self.at += 1;
Tok::Colon
}
0 => Tok::End,
b'"' => return self.text(at),
b'-' | b'0'..=b'9' => {
if !self.cfg.invalid_decode && self.loose() {
return Found {
tok: Tok::Bad("invalid number"),
at,
};
}
return self.number();
}
b't' if self.word(b"true") => {
self.at += 4;
Tok::Bool(true)
}
b'f' if self.word(b"false") => {
self.at += 5;
Tok::Bool(false)
}
b'n' if self.word(b"null") => {
self.at += 4;
Tok::Null
}
b'f' | b'i' | b'I' | b'n' | b'N' | b't' | b'+' => {
if self.cfg.invalid_decode && self.loose() {
return self.number();
}
Tok::Bad("invalid token")
}
_ => Tok::Bad("invalid token"),
};
Found { tok, at }
}
fn word(&self, want: &[u8]) -> bool {
want.iter()
.enumerate()
.all(|(i, &b)| self.byte(self.at + i) == b)
}
fn loose(&self) -> bool {
let mut at = self.at;
if self.byte(at) == b'+' {
return true;
}
if self.byte(at) == b'-' {
at += 1;
}
if self.byte(at) == b'0' {
let next = self.byte(at + 1);
return next | 0x20 == b'x' || next.is_ascii_digit();
}
if self.byte(at) <= b'9' {
return false;
}
let rest = &self.data[at.min(self.data.len())..];
starts_ci(rest, b"inf") || starts_ci(rest, b"nan")
}
fn number(&mut self) -> Found {
let at = self.at;
match strtod(self.data, at) {
Some((value, end)) => {
self.at = end;
Found {
tok: Tok::Num(value),
at,
}
}
None => Found {
tok: Tok::Bad("invalid number"),
at,
},
}
}
fn text(&mut self, from: usize) -> Found {
self.at += 1;
let mut out = Vec::new();
loop {
let ch = self.byte(self.at);
if ch == b'"' {
self.at += 1;
return Found {
tok: Tok::Str(out),
at: from,
};
}
if ch == 0 {
return Found {
tok: Tok::Bad("unexpected end of string"),
at: self.at,
};
}
if ch != b'\\' {
out.push(ch);
self.at += 1;
continue;
}
let what = self.byte(self.at + 1);
if what == b'u' {
match self.unicode(&mut out) {
true => continue,
false => {
return Found {
tok: Tok::Bad("invalid unicode escape code"),
at: self.at,
};
}
}
}
let plain = match what {
b'"' => b'"',
b'\\' => b'\\',
b'/' => b'/',
b'b' => 0x08,
b't' => 0x09,
b'n' => 0x0a,
b'f' => 0x0c,
b'r' => 0x0d,
_ => {
return Found {
tok: Tok::Bad("invalid escape code"),
at: self.at,
};
}
};
out.push(plain);
self.at += 2;
}
}
fn unicode(&mut self, out: &mut Vec<u8>) -> bool {
let Some(mut point) = self.hex4(self.at + 2) else {
return false;
};
let mut len = 6;
if point & 0xf800 == 0xd800 {
if point & 0x400 != 0 {
return false;
}
if self.byte(self.at + len) != b'\\' || self.byte(self.at + len + 1) != b'u' {
return false;
}
let Some(low) = self.hex4(self.at + 2 + len) else {
return false;
};
if low & 0xfc00 != 0xdc00 {
return false;
}
point = (((point & 0x3ff) << 10) | (low & 0x3ff)) + 0x10000;
len = 12;
}
let Some(ch) = char::from_u32(point) else {
return false;
};
let mut buf = [0u8; 4];
out.extend_from_slice(ch.encode_utf8(&mut buf).as_bytes());
self.at += len;
true
}
fn hex4(&self, at: usize) -> Option<u32> {
let mut value = 0u32;
for i in 0..4 {
value = value * 16 + (self.byte(at + i) as char).to_digit(16)?;
}
Some(value)
}
}
fn starts_ci(data: &[u8], want: &[u8]) -> bool {
data.len() >= want.len() && data[..want.len()].eq_ignore_ascii_case(want)
}
fn strtod(data: &[u8], from: usize) -> Option<(f64, usize)> {
let byte = |i: usize| data.get(i).copied().unwrap_or(0);
let mut at = from;
while matches!(byte(at), b' ' | b'\t' | b'\n' | 0x0b | 0x0c | b'\r') {
at += 1;
}
let start = at;
let negative = byte(at) == b'-';
if negative || byte(at) == b'+' {
at += 1;
}
let sign = if negative { -1.0 } else { 1.0 };
if byte(at) == b'0' && byte(at + 1) | 0x20 == b'x' {
return Some(hex(data, at, sign));
}
if starts_ci(&data[at.min(data.len())..], b"infinity") {
return Some((sign * f64::INFINITY, at + 8));
}
if starts_ci(&data[at.min(data.len())..], b"inf") {
return Some((sign * f64::INFINITY, at + 3));
}
if starts_ci(&data[at.min(data.len())..], b"nan") {
let mut end = at + 3;
if byte(end) == b'(' {
let mut scan = end + 1;
while byte(scan).is_ascii_alphanumeric() || byte(scan) == b'_' {
scan += 1;
}
if byte(scan) == b')' {
end = scan + 1;
}
}
return Some((sign * f64::NAN, end));
}
let mut digits = 0;
while byte(at).is_ascii_digit() {
at += 1;
digits += 1;
}
if byte(at) == b'.' {
at += 1;
while byte(at).is_ascii_digit() {
at += 1;
digits += 1;
}
}
if digits == 0 {
return None;
}
let mut end = at;
if byte(at) | 0x20 == b'e' {
let mut scan = at + 1;
if matches!(byte(scan), b'+' | b'-') {
scan += 1;
}
if byte(scan).is_ascii_digit() {
while byte(scan).is_ascii_digit() {
scan += 1;
}
end = scan;
}
}
let text = std::str::from_utf8(&data[start..end]).ok()?;
text.parse::<f64>().ok().map(|value| (value, end))
}
fn hex(data: &[u8], at: usize, sign: f64) -> (f64, usize) {
let byte = |i: usize| data.get(i).copied().unwrap_or(0);
let mut scan = at + 2;
let mut mantissa = 0u128;
let mut shift = 0i32;
let mut digits = 0;
let mut fraction = false;
loop {
let ch = byte(scan);
if ch == b'.' && !fraction {
fraction = true;
scan += 1;
continue;
}
let Some(value) = (ch as char).to_digit(16) else {
break;
};
if mantissa <= (u128::MAX - 15) / 16 {
mantissa = mantissa * 16 + u128::from(value);
if fraction {
shift -= 4;
}
} else if !fraction {
shift += 4;
}
scan += 1;
digits += 1;
}
if digits == 0 {
return (sign * 0.0, at + 1);
}
if byte(scan) | 0x20 == b'p' {
let mut walk = scan + 1;
let negative = matches!(byte(walk), b'+' | b'-');
let minus = byte(walk) == b'-';
if negative {
walk += 1;
}
let mut power = 0i32;
let mut counted = 0;
while let Some(value) = (byte(walk) as char).to_digit(10) {
power = (power * 10 + value as i32).min(100_000);
walk += 1;
counted += 1;
}
if counted > 0 {
shift += if minus { -power } else { power };
scan = walk;
}
}
(sign * (mantissa as f64) * 2f64.powi(shift), scan)
}
impl<'a> Parser<'a> {
fn value(&mut self, lua: &Lua, found: Found) -> Result<Value, String> {
match found.tok {
Tok::Str(bytes) => Ok(Value::String(
lua.create_string(&bytes).map_err(|e| e.to_string())?,
)),
Tok::Num(n) => Ok(Value::Number(n)),
Tok::Bool(b) => Ok(Value::Boolean(b)),
Tok::Null => Ok(null()),
Tok::ObjBegin => self.object(lua),
Tok::ArrBegin => self.array(lua),
other => Err(complain(
"value",
&Found {
tok: other,
at: found.at,
},
)),
}
}
fn descend(&mut self) -> Result<(), String> {
self.depth += 1;
if self.depth <= self.cfg.decode_max_depth && self.depth <= CEILING {
return Ok(());
}
Err(format!(
"Found too many nested data structures ({}) at character {}",
self.depth, self.at
))
}
fn object(&mut self, lua: &Lua) -> Result<Value, String> {
self.descend()?;
let t = lua.create_table().map_err(|e| e.to_string())?;
let mut found = self.next();
if found.tok == Tok::ObjEnd {
self.depth -= 1;
return Ok(Value::Table(t));
}
loop {
let Tok::Str(name) = found.tok else {
return Err(complain("object key string", &found));
};
let key = lua.create_string(&name).map_err(|e| e.to_string())?;
let colon = self.next();
if colon.tok != Tok::Colon {
return Err(complain("colon", &colon));
}
let next = self.next();
let value = self.value(lua, next)?;
t.raw_set(key, value).map_err(|e| e.to_string())?;
found = self.next();
if found.tok == Tok::ObjEnd {
self.depth -= 1;
return Ok(Value::Table(t));
}
if found.tok != Tok::Comma {
return Err(complain("comma or object end", &found));
}
found = self.next();
}
}
fn array(&mut self, lua: &Lua) -> Result<Value, String> {
self.descend()?;
let t = lua.create_table().map_err(|e| e.to_string())?;
if self.cfg.array_mt {
let mt = lua.create_table().map_err(|e| e.to_string())?;
mt.raw_set("__is_cjson_array", true)
.map_err(|e| e.to_string())?;
t.set_metatable(Some(mt)).map_err(|e| e.to_string())?;
}
let mut found = self.next();
if found.tok == Tok::ArrEnd {
self.depth -= 1;
return Ok(Value::Table(t));
}
let mut i = 1i64;
loop {
let value = self.value(lua, found)?;
t.raw_set(i, value).map_err(|e| e.to_string())?;
i += 1;
found = self.next();
if found.tok == Tok::ArrEnd {
self.depth -= 1;
return Ok(Value::Table(t));
}
if found.tok != Tok::Comma {
return Err(complain("comma or array end", &found));
}
found = self.next();
}
}
}
fn complain(wanted: &str, found: &Found) -> String {
format!(
"Expected {} but found {} at character {}",
wanted,
found.tok.found(),
found.at + 1
)
}
pub(super) fn null() -> Value {
Value::LightUserData(mlua::LightUserData(std::ptr::null_mut()))
}
pub(super) fn statics(lua: &Lua, raw: &Table) -> mlua::Result<()> {
raw.raw_set(
"cjson_encode",
lua.create_function(|lua, (value, settings): (Value, Table)| {
let cfg = Config::read(&settings)?;
let mut encoder = Encoder {
cfg: &cfg,
out: Vec::with_capacity(64),
};
match encoder.data(&value, 0) {
Ok(()) => Ok((true, Value::String(lua.create_string(&encoder.out)?))),
Err(why) => Ok((false, Value::String(lua.create_string(&why)?))),
}
})?,
)?;
raw.raw_set(
"cjson_decode",
lua.create_function(|lua, (text, settings): (mlua::LuaString, Table)| {
let cfg = Config::read(&settings)?;
let bytes = text.as_bytes();
if bytes.len() >= 2 && (bytes[0] == 0 || bytes[1] == 0) {
let why = "JSON parser does not support UTF-16 or UTF-32";
return Ok((false, Value::String(lua.create_string(why)?)));
}
let mut parser = Parser {
data: &bytes,
at: 0,
cfg: &cfg,
depth: 0,
};
let first = parser.next();
let answer = parser.value(lua, first).and_then(|value| {
let rest = parser.next();
match rest.tok {
Tok::End => Ok(value),
_ => Err(complain("the end", &rest)),
}
});
match answer {
Ok(value) => Ok((true, value)),
Err(why) => Ok((false, Value::String(lua.create_string(&why)?))),
}
})?,
)?;
Ok(())
}
#[cfg(test)]
mod tests {
use super::{g_fmt, strtod};
#[test]
fn a_number_is_written_the_way_the_c_library_writes_one() {
for (given, want) in [
(100.0, "100"),
(0.0, "0"),
(-0.0, "-0"),
(1e300, "1e+300"),
(1e-7, "1e-07"),
(1.0 / 3.0, "0.33333333333333"),
(std::f64::consts::PI, "3.1415926535898"),
(9_007_199_254_740_992.0, "9.007199254741e+15"),
(1e14, "1e+14"),
(1e13, "10000000000000"),
(-1.5, "-1.5"),
(f64::INFINITY, "inf"),
(f64::NEG_INFINITY, "-inf"),
] {
assert_eq!(g_fmt(given, 14), want, "{given}");
}
assert_eq!(g_fmt(std::f64::consts::PI, 3), "3.14");
assert_eq!(g_fmt(1234.0, 3), "1.23e+03");
}
#[test]
fn a_number_is_read_the_way_the_c_library_reads_one() {
for (given, value, end) in [
("1", 1.0, 1),
("01", 1.0, 2),
("+1", 1.0, 2),
("1.", 1.0, 2),
("-2.5e3", -2500.0, 6),
("1e", 1.0, 1),
("1e+", 1.0, 1),
("0x10", 16.0, 4),
("0X1f", 31.0, 4),
("-0x10,", -16.0, 5),
("0xz", 0.0, 1),
("1e999", f64::INFINITY, 5),
("inf", f64::INFINITY, 3),
("-Infinity", f64::NEG_INFINITY, 9),
] {
let (got, at) = strtod(given.as_bytes(), 0).expect(given);
assert_eq!(got, value, "{given}");
assert_eq!(at, end, "{given}");
}
let (nan, at) = strtod(b"nan", 0).expect("nan");
assert!(nan.is_nan());
assert_eq!(at, 3);
assert!(strtod(b"abc", 0).is_none());
}
}