use super::argue::{Answer, Stop, answer, behind, number, text};
use mlua::{Lua, MultiValue, Table, Value};
const MAX_NESTING: u32 = 16;
const TWO_63: f64 = 9_223_372_036_854_775_808.0;
fn bytes(out: &mut Vec<u8>, s: &[u8]) {
let len = s.len();
if len < 32 {
out.push(0xa0 | (len as u8 & 0x1f));
} else if len <= 0xff {
out.push(0xd9);
out.push(len as u8);
} else if len <= 0xffff {
out.push(0xda);
out.extend_from_slice(&(len as u16).to_be_bytes());
} else {
out.push(0xdb);
out.extend_from_slice(&(len as u32).to_be_bytes());
}
out.extend_from_slice(s);
}
fn real(out: &mut Vec<u8>, d: f64) {
let narrow = d as f32;
if d == f64::from(narrow) {
out.push(0xca);
out.extend_from_slice(&narrow.to_be_bytes());
} else {
out.push(0xcb);
out.extend_from_slice(&d.to_be_bytes());
}
}
fn integer(out: &mut Vec<u8>, n: i64) {
if n >= 0 {
if n <= 127 {
out.push(n as u8 & 0x7f);
} else if n <= 0xff {
out.push(0xcc);
out.push(n as u8);
} else if n <= 0xffff {
out.push(0xcd);
out.extend_from_slice(&(n as u16).to_be_bytes());
} else if n <= 0xffff_ffff {
out.push(0xce);
out.extend_from_slice(&(n as u32).to_be_bytes());
} else {
out.push(0xcf);
out.extend_from_slice(&(n as u64).to_be_bytes());
}
} else if n >= -32 {
out.push(n as u8);
} else if n >= -128 {
out.push(0xd0);
out.push(n as u8);
} else if n >= -32768 {
out.push(0xd1);
out.extend_from_slice(&(n as i16).to_be_bytes());
} else if n >= -2_147_483_648 {
out.push(0xd2);
out.extend_from_slice(&(n as i32).to_be_bytes());
} else {
out.push(0xd3);
out.extend_from_slice(&n.to_be_bytes());
}
}
fn count(out: &mut Vec<u8>, n: usize, fix: u8, wide: u8, wider: u8) {
if n <= 15 {
out.push(fix | (n as u8 & 0xf));
} else if n <= 65535 {
out.push(wide);
out.extend_from_slice(&(n as u16).to_be_bytes());
} else {
out.push(wider);
out.extend_from_slice(&(n as u32).to_be_bytes());
}
}
fn whole(n: f64) -> Option<i64> {
if !n.is_finite() || !(-TWO_63..TWO_63).contains(&n) {
return None;
}
let held = n as i64;
if held as f64 == n { Some(held) } else { None }
}
fn small(n: f64) -> bool {
n.is_finite() && f64::from(n as i32) == n
}
fn listed(t: &Table) -> mlua::Result<bool> {
let mut seen: i64 = 0;
let mut top: i64 = 0;
for pair in t.pairs::<Value, Value>() {
let (key, _) = pair?;
let n = match key {
Value::Integer(i) => i as f64,
Value::Number(f) => f,
_ => return Ok(false),
};
if n <= 0.0 || !small(n) {
return Ok(false);
}
top = top.max(n as i64);
seen += 1;
}
Ok(top == seen)
}
fn encode(out: &mut Vec<u8>, v: &Value, level: u32) -> mlua::Result<()> {
match v {
Value::String(s) => bytes(out, &s.as_bytes()),
Value::Boolean(b) => out.push(if *b { 0xc3 } else { 0xc2 }),
Value::Integer(i) => integer(out, *i),
Value::Number(f) => match whole(*f) {
Some(held) => integer(out, held),
None => real(out, *f),
},
Value::Table(t) if level < MAX_NESTING => {
let held;
let t = match behind(t) {
Some(real) => {
held = real;
&held
}
None => t,
};
if listed(t)? {
let len = t.raw_len();
count(out, len, 0x90, 0xdc, 0xdd);
for j in 1..=len {
let item: Value = t.get(j)?;
encode(out, &item, level + 1)?;
}
} else {
let mut len = 0usize;
for pair in t.pairs::<Value, Value>() {
pair?;
len += 1;
}
count(out, len, 0x80, 0xde, 0xdf);
for pair in t.pairs::<Value, Value>() {
let (key, held) = pair?;
encode(out, &key, level + 1)?;
encode(out, &held, level + 1)?;
}
}
}
_ => out.push(0xc0),
}
Ok(())
}
struct Cur<'a> {
data: &'a [u8],
at: usize,
}
fn short() -> Stop {
Stop::Plain(b"Missing bytes in input.".to_vec())
}
fn garbled() -> Stop {
Stop::Plain(b"Bad data format in input.".to_vec())
}
impl Cur<'_> {
fn left(&self) -> usize {
self.data.len() - self.at
}
fn need(&self, len: usize) -> Answer<()> {
if self.left() < len {
Err(short())
} else {
Ok(())
}
}
fn byte(&self, i: usize) -> u8 {
self.data[self.at + i]
}
fn take(&mut self, len: usize) {
self.at += len;
}
fn wide(&self, from: usize, len: usize) -> u64 {
let mut held = 0u64;
for i in 0..len {
held = (held << 8) | u64::from(self.byte(from + i));
}
held
}
}
fn string(lua: &Lua, c: &mut Cur, from: usize, len: usize) -> Answer<Value> {
c.need(from + len)?;
let held = lua
.create_string(&c.data[c.at + from..c.at + from + len])
.map_err(|e| Stop::Plain(e.to_string().into_bytes()))?;
c.take(from + len);
Ok(Value::String(held))
}
fn array(lua: &Lua, c: &mut Cur, len: usize) -> Answer<Value> {
let out = lua
.create_table()
.map_err(|e| Stop::Plain(e.to_string().into_bytes()))?;
for i in 1..=len {
let held = value(lua, c)?;
out.raw_set(i, held)
.map_err(|e| Stop::Plain(e.to_string().into_bytes()))?;
}
Ok(Value::Table(out))
}
fn hash(lua: &Lua, c: &mut Cur, len: usize) -> Answer<Value> {
let out = lua
.create_table()
.map_err(|e| Stop::Plain(e.to_string().into_bytes()))?;
for _ in 0..len {
let key = value(lua, c)?;
let held = value(lua, c)?;
match key {
Value::Nil => return Err(Stop::Bare(b"table index is nil".to_vec())),
Value::Number(n) if n.is_nan() => {
return Err(Stop::Bare(b"table index is NaN".to_vec()));
}
_ => {}
}
out.raw_set(key, held)
.map_err(|e| Stop::Plain(e.to_string().into_bytes()))?;
}
Ok(Value::Table(out))
}
fn value(lua: &Lua, c: &mut Cur) -> Answer<Value> {
c.need(1)?;
let tag = c.byte(0);
let number = |c: &mut Cur, len: usize, held: f64| -> Answer<Value> {
c.take(len);
Ok(Value::Number(held))
};
match tag {
0xcc => {
c.need(2)?;
let held = c.wide(1, 1) as f64;
number(c, 2, held)
}
0xcd => {
c.need(3)?;
let held = c.wide(1, 2) as f64;
number(c, 3, held)
}
0xce => {
c.need(5)?;
let held = c.wide(1, 4) as f64;
number(c, 5, held)
}
0xcf => {
c.need(9)?;
let held = c.wide(1, 8) as i64 as f64;
number(c, 9, held)
}
0xd0 => {
c.need(2)?;
let held = f64::from(c.byte(1) as i8);
number(c, 2, held)
}
0xd1 => {
c.need(3)?;
let held = f64::from(c.wide(1, 2) as u16 as i16);
number(c, 3, held)
}
0xd2 => {
c.need(5)?;
let held = f64::from(c.wide(1, 4) as u32 as i32);
number(c, 5, held)
}
0xd3 => {
c.need(9)?;
let held = c.wide(1, 8) as i64 as f64;
number(c, 9, held)
}
0xc0 => {
c.take(1);
Ok(Value::Nil)
}
0xc2 | 0xc3 => {
let held = tag == 0xc3;
c.take(1);
Ok(Value::Boolean(held))
}
0xca => {
c.need(5)?;
let held = f64::from(f32::from_bits(c.wide(1, 4) as u32));
number(c, 5, held)
}
0xcb => {
c.need(9)?;
let held = f64::from_bits(c.wide(1, 8));
number(c, 9, held)
}
0xd9 => {
c.need(2)?;
let len = c.byte(1) as usize;
string(lua, c, 2, len)
}
0xda => {
c.need(3)?;
let len = c.wide(1, 2) as usize;
string(lua, c, 3, len)
}
0xdb => {
c.need(5)?;
let len = c.wide(1, 4) as usize;
string(lua, c, 5, len)
}
0xdc => {
c.need(3)?;
let len = c.wide(1, 2) as usize;
c.take(3);
array(lua, c, len)
}
0xdd => {
c.need(5)?;
let len = c.wide(1, 4) as usize;
c.take(5);
array(lua, c, len)
}
0xde => {
c.need(3)?;
let len = c.wide(1, 2) as usize;
c.take(3);
hash(lua, c, len)
}
0xdf => {
c.need(5)?;
let len = c.wide(1, 4) as usize;
c.take(5);
hash(lua, c, len)
}
_ => {
if tag & 0x80 == 0 {
let held = f64::from(tag);
number(c, 1, held)
} else if tag & 0xe0 == 0xe0 {
let held = f64::from(tag as i8);
number(c, 1, held)
} else if tag & 0xe0 == 0xa0 {
let len = (tag & 0x1f) as usize;
string(lua, c, 1, len)
} else if tag & 0xf0 == 0x90 {
let len = (tag & 0xf) as usize;
c.take(1);
array(lua, c, len)
} else if tag & 0xf0 == 0x80 {
let len = (tag & 0xf) as usize;
c.take(1);
hash(lua, c, len)
} else {
Err(garbled())
}
}
}
}
fn pack(lua: &Lua, args: &[Value]) -> Answer<Value> {
if args.is_empty() {
return Err(Stop::Arg(0, b"MessagePack pack needs input.".to_vec()));
}
let mut out = Vec::new();
for held in args {
encode(&mut out, held, 0).map_err(|e| Stop::Plain(e.to_string().into_bytes()))?;
}
let held = lua
.create_string(&out)
.map_err(|e| Stop::Plain(e.to_string().into_bytes()))?;
Ok(Value::String(held))
}
fn unpack(lua: &Lua, args: &[Value], limit: i64, offset: i64) -> Answer<(Vec<Value>, Option<i64>)> {
let data = text(lua, args, 1, "no value")?;
let len = data.len() as i64;
let everything = limit == 0 && offset == 0;
if offset < 0 || limit < 0 {
return Err(Stop::Plain(
format!(
"Invalid request to unpack with offset of {} and limit of {}.",
offset as i32, len as i32
)
.into_bytes(),
));
}
if offset > len {
return Err(Stop::Plain(
format!(
"Start offset {} greater than input length {}.",
offset as i32, len as i32
)
.into_bytes(),
));
}
let limit = if everything {
i64::from(i32::MAX)
} else {
limit
};
let mut c = Cur {
data: &data[offset as usize..],
at: 0,
};
let mut found = Vec::new();
while c.left() > 0 && (found.len() as i64) < limit {
found.push(value(lua, &mut c)?);
}
if everything {
return Ok((found, None));
}
let stopped = if c.left() == 0 {
-1
} else {
len - c.left() as i64
};
Ok((found, Some(stopped)))
}
fn listing(lua: &Lua, given: Answer<(Vec<Value>, Option<i64>)>) -> Answer<Value> {
let (found, stopped) = given?;
let build = || -> mlua::Result<Value> {
let out = lua.create_table()?;
let mut at = 0;
if let Some(stopped) = stopped {
at += 1;
out.raw_set(at, stopped)?;
}
for held in found {
at += 1;
out.raw_set(at, held)?;
}
out.raw_set("n", at)?;
Ok(Value::Table(out))
};
build().map_err(|e| Stop::Plain(e.to_string().into_bytes()))
}
fn maybe(lua: &Lua, args: &[Value], i: usize, default: i64) -> Answer<i64> {
match args.get(i - 1) {
None | Some(Value::Nil) => Ok(default),
Some(_) => Ok(number(lua, args, i, "no value")? as i64),
}
}
pub(super) fn statics(lua: &Lua, raw: &Table) -> mlua::Result<()> {
raw.raw_set(
"cmsgpack_pack",
lua.create_function(|lua, args: MultiValue| {
let held: Vec<Value> = args.into_iter().collect();
answer(lua, pack(lua, &held))
})?,
)?;
raw.raw_set(
"cmsgpack_unpack",
lua.create_function(|lua, args: MultiValue| {
let held: Vec<Value> = args.into_iter().collect();
answer(lua, listing(lua, unpack(lua, &held, 0, 0)))
})?,
)?;
raw.raw_set(
"cmsgpack_unpack_one",
lua.create_function(|lua, args: MultiValue| {
let held: Vec<Value> = args.into_iter().collect();
let read = maybe(lua, &held, 2, 0).and_then(|at| unpack(lua, &held, 1, at));
answer(lua, listing(lua, read))
})?,
)?;
raw.raw_set(
"cmsgpack_unpack_limit",
lua.create_function(|lua, args: MultiValue| {
let read = (|| {
let held: Vec<Value> = args.into_iter().collect();
let limit = number(lua, &held, 2, "no value")? as i64;
let at = maybe(lua, &held, 3, 0)?;
unpack(lua, &held, limit, at)
})();
answer(lua, listing(lua, read))
})?,
)?;
Ok(())
}
#[cfg(test)]
mod tests {
use super::{integer, real, whole};
#[test]
fn a_whole_number_takes_the_shortest_form_that_holds_it() {
for (n, want) in [
(0i64, vec![0x00]),
(127, vec![0x7f]),
(128, vec![0xcc, 0x80]),
(255, vec![0xcc, 0xff]),
(256, vec![0xcd, 0x01, 0x00]),
(65535, vec![0xcd, 0xff, 0xff]),
(65536, vec![0xce, 0x00, 0x01, 0x00, 0x00]),
(-1, vec![0xff]),
(-32, vec![0xe0]),
(-33, vec![0xd0, 0xdf]),
(-128, vec![0xd0, 0x80]),
(-129, vec![0xd1, 0xff, 0x7f]),
(-32768, vec![0xd1, 0x80, 0x00]),
(-32769, vec![0xd2, 0xff, 0xff, 0x7f, 0xff]),
] {
let mut out = Vec::new();
integer(&mut out, n);
assert_eq!(out, want, "{n}");
}
}
#[test]
fn a_fraction_goes_out_narrow_when_that_loses_nothing() {
let mut out = Vec::new();
real(&mut out, 1.5);
assert_eq!(out, vec![0xca, 0x3f, 0xc0, 0x00, 0x00]);
out.clear();
real(&mut out, 0.1);
assert_eq!(out[0], 0xcb);
assert_eq!(out.len(), 9);
}
#[test]
fn a_number_is_whole_only_when_an_integer_holds_it_and_gives_it_back() {
assert_eq!(whole(1.0), Some(1));
assert_eq!(whole(-1.0), Some(-1));
assert_eq!(whole(0.5), None);
assert_eq!(whole(f64::INFINITY), None);
assert_eq!(whole(f64::NAN), None);
assert_eq!(whole(-9_223_372_036_854_775_808.0), Some(i64::MIN));
assert_eq!(whole(9_223_372_036_854_775_808.0), None);
assert_eq!(whole(1e30), None);
}
}