use super::argue::{Answer, Stop, answer, number, text};
use mlua::{Lua, MultiValue, Table, Value};
struct Fmt<'a> {
text: &'a [u8],
at: usize,
}
struct Header {
big: bool,
align: usize,
}
const LONG: usize = 8;
const INT: usize = 4;
const FLOAT: usize = 4;
const DOUBLE: usize = 8;
const MAXINT: usize = 32;
const MAXALIGN: usize = 8;
impl Fmt<'_> {
fn byte(&self) -> u8 {
self.text.get(self.at).copied().unwrap_or(0)
}
fn num(&mut self, default: usize) -> Answer<usize> {
if !self.byte().is_ascii_digit() {
return Ok(default);
}
let mut a: i32 = 0;
while self.byte().is_ascii_digit() {
let digit = i32::from(self.byte() - b'0');
if a > i32::MAX / 10 || a * 10 > i32::MAX - digit {
return Err(Stop::Plain(b"integral size overflow".to_vec()));
}
a = a * 10 + digit;
self.at += 1;
}
Ok(a as usize)
}
}
fn optsize(opt: u8, fmt: &mut Fmt<'_>) -> Answer<usize> {
match opt {
b'b' | b'B' | b'x' => Ok(1),
b'h' | b'H' => Ok(2),
b'l' | b'L' | b'T' => Ok(LONG),
b'f' => Ok(FLOAT),
b'd' => Ok(DOUBLE),
b'c' => fmt.num(1),
b'i' | b'I' => {
let size = fmt.num(INT)?;
if size > MAXINT {
return Err(Stop::Plain(
format!("integral size {size} is larger than limit of {MAXINT}").into_bytes(),
));
}
Ok(size)
}
_ => Ok(0),
}
}
fn to_align(len: usize, h: &Header, opt: u8, size: usize) -> usize {
if size == 0 || opt == b'c' {
return 0;
}
let size = size.min(h.align);
(size - (len & (size - 1))) & (size - 1)
}
fn control(opt: u8, fmt: &mut Fmt<'_>, h: &mut Header) -> Answer<()> {
match opt {
b' ' => Ok(()),
b'>' => {
h.big = true;
Ok(())
}
b'<' => {
h.big = false;
Ok(())
}
b'!' => {
let want = fmt.num(MAXALIGN)?;
if want == 0 || want & (want - 1) != 0 {
return Err(Stop::Plain(
format!("alignment {want} is not a power of 2").into_bytes(),
));
}
h.align = want;
Ok(())
}
_ => {
let mut why = b"invalid format option '".to_vec();
why.push(opt);
why.push(b'\'');
Err(Stop::Arg(1, why))
}
}
}
fn until_zero(text: &[u8]) -> &[u8] {
match text.iter().position(|&b| b == 0) {
Some(end) => &text[..end],
None => text,
}
}
fn put_integer(out: &mut Vec<u8>, n: f64, big: bool, size: usize) {
let mut value = if n < 0.0 { (n as i64) as u64 } else { n as u64 };
let at = out.len();
out.resize(at + size, 0);
for i in 0..size {
let byte = (value & 0xff) as u8;
out[if big { at + size - 1 - i } else { at + i }] = byte;
value >>= 8;
}
}
fn get_integer(data: &[u8], big: bool, signed: bool, size: usize) -> f64 {
let mut value = 0u64;
for i in 0..size {
let byte = data[if big { i } else { size - 1 - i }];
value = (value << 8) | u64::from(byte);
}
if !signed {
return value as f64;
}
let mask = (!0u64) << ((size * 8 - 1) % 64);
if value & mask != 0 {
value |= mask;
}
(value as i64) as f64
}
fn pack(lua: &Lua, args: &[Value]) -> Answer<Vec<u8>> {
let held = text(lua, args, 1, "no value")?;
let mut fmt = Fmt {
text: until_zero(&held),
at: 0,
};
let mut h = Header {
big: false,
align: 1,
};
let mut arg = 2;
let mut total = 0usize;
let mut out = Vec::new();
while fmt.at < fmt.text.len() {
let opt = fmt.text[fmt.at];
fmt.at += 1;
let mut size = optsize(opt, &mut fmt)?;
let pad = to_align(total, &h, opt, size);
total += pad;
out.resize(out.len() + pad, 0);
match opt {
b'b' | b'B' | b'h' | b'H' | b'l' | b'L' | b'T' | b'i' | b'I' => {
let n = number(lua, args, arg, "nil")?;
arg += 1;
put_integer(&mut out, n, h.big, size);
}
b'x' => out.push(0),
b'f' => {
let n = number(lua, args, arg, "nil")? as f32;
arg += 1;
let bytes = if h.big {
n.to_be_bytes()
} else {
n.to_le_bytes()
};
out.extend_from_slice(&bytes);
}
b'd' => {
let n = number(lua, args, arg, "nil")?;
arg += 1;
let bytes = if h.big {
n.to_be_bytes()
} else {
n.to_le_bytes()
};
out.extend_from_slice(&bytes);
}
b'c' | b's' => {
let s = text(lua, args, arg, "nil")?;
arg += 1;
if size == 0 {
size = s.len();
}
if s.len() < size {
return Err(Stop::Arg(arg, b"string too short".to_vec()));
}
out.extend_from_slice(&s[..size]);
if opt == b's' {
out.push(0);
size += 1;
}
}
_ => control(opt, &mut fmt, &mut h)?,
}
total += size;
}
Ok(out)
}
fn unpack(lua: &Lua, args: &[Value]) -> Answer<(Vec<Value>, f64)> {
let held = text(lua, args, 1, "no value")?;
let data = text(lua, args, 2, "no value")?;
let start = match args.get(2) {
None | Some(Value::Nil) => 1.0,
Some(_) => number(lua, args, 3, "no value")?,
};
let counted = start as i64;
if counted == 0 {
return Err(Stop::Arg(3, b"offset must be 1 or greater".to_vec()));
}
let long = data.len();
let mut pos = (counted as u64).wrapping_sub(1).min(long as u64 + 1) as usize;
let mut fmt = Fmt {
text: until_zero(&held),
at: 0,
};
let mut h = Header {
big: false,
align: 1,
};
let mut found: Vec<Value> = Vec::new();
let short = || Stop::Arg(2, b"data string too short".to_vec());
while fmt.at < fmt.text.len() {
let opt = fmt.text[fmt.at];
fmt.at += 1;
let mut size = optsize(opt, &mut fmt)?;
pos = pos.saturating_add(to_align(pos, &h, opt, size));
if size > long || pos > long - size {
return Err(short());
}
match opt {
b'b' | b'B' | b'h' | b'H' | b'l' | b'L' | b'T' | b'i' | b'I' => {
let signed = opt.is_ascii_lowercase();
let n = get_integer(&data[pos..], h.big, signed, size);
found.push(Value::Number(n));
}
b'x' => {}
b'f' => {
let raw = [data[pos], data[pos + 1], data[pos + 2], data[pos + 3]];
let n = if h.big {
f32::from_be_bytes(raw)
} else {
f32::from_le_bytes(raw)
};
found.push(Value::Number(f64::from(n)));
}
b'd' => {
let mut raw = [0u8; 8];
raw.copy_from_slice(&data[pos..pos + 8]);
let n = if h.big {
f64::from_be_bytes(raw)
} else {
f64::from_le_bytes(raw)
};
found.push(Value::Number(n));
}
b'c' => {
if size == 0 {
let last = found
.last()
.and_then(|v| lua.coerce_number(v.clone()).ok().flatten());
let Some(n) = last else {
return Err(Stop::Plain(b"format 'c0' needs a previous size".to_vec()));
};
found.pop();
size = if n < 0.0 { usize::MAX } else { n as usize };
if size > long || pos > long - size {
return Err(short());
}
}
found.push(string(lua, &data[pos..pos + size])?);
}
b's' => {
let Some(end) = data[pos..].iter().position(|&b| b == 0) else {
return Err(Stop::Plain(b"unfinished string in data".to_vec()));
};
found.push(string(lua, &data[pos..pos + end])?);
size = end + 1;
}
_ => control(opt, &mut fmt, &mut h)?,
}
pos += size;
}
Ok((found, (pos + 1) as f64))
}
fn sizeof(lua: &Lua, args: &[Value]) -> Answer<f64> {
let held = text(lua, args, 1, "no value")?;
let mut fmt = Fmt {
text: until_zero(&held),
at: 0,
};
let mut h = Header {
big: false,
align: 1,
};
let mut pos = 0usize;
while fmt.at < fmt.text.len() {
let opt = fmt.text[fmt.at];
fmt.at += 1;
let size = optsize(opt, &mut fmt)?;
pos += to_align(pos, &h, opt, size);
if opt == b's' {
return Err(Stop::Arg(1, b"option 's' has no fixed size".to_vec()));
}
if opt == b'c' && size == 0 {
return Err(Stop::Arg(1, b"option 'c0' has no fixed size".to_vec()));
}
if !opt.is_ascii_alphanumeric() {
control(opt, &mut fmt, &mut h)?;
}
pos += size;
}
Ok(pos as f64)
}
fn string(lua: &Lua, bytes: &[u8]) -> Answer<Value> {
lua.create_string(bytes)
.map(Value::String)
.map_err(|e| Stop::Plain(e.to_string().into_bytes()))
}
pub(super) fn statics(lua: &Lua, raw: &Table) -> mlua::Result<()> {
raw.raw_set(
"struct_pack",
lua.create_function(|lua, args: MultiValue| {
let held: Vec<Value> = args.into_iter().collect();
let packed = pack(lua, &held).and_then(|out| string(lua, &out));
answer(lua, packed)
})?,
)?;
raw.raw_set(
"struct_unpack",
lua.create_function(|lua, args: MultiValue| {
let held: Vec<Value> = args.into_iter().collect();
let read = unpack(lua, &held).and_then(|(found, at)| {
let build = || -> mlua::Result<Value> {
let out = lua.create_table()?;
let count = found.len() + 1;
for (i, value) in found.into_iter().enumerate() {
out.raw_set(i + 1, value)?;
}
out.raw_set(count, at)?;
out.raw_set("n", count)?;
Ok(Value::Table(out))
};
build().map_err(|e| Stop::Plain(e.to_string().into_bytes()))
});
answer(lua, read)
})?,
)?;
raw.raw_set(
"struct_size",
lua.create_function(|lua, args: MultiValue| {
let held: Vec<Value> = args.into_iter().collect();
answer(lua, sizeof(lua, &held).map(Value::Number))
})?,
)?;
Ok(())
}
#[cfg(test)]
mod tests {
use super::{get_integer, put_integer};
#[test]
fn an_integer_goes_out_and_comes_back_the_way_the_c_library_moves_one() {
for (n, size, big, want) in [
(1.0, 1, false, vec![1]),
(255.0, 1, false, vec![255]),
(-1.0, 1, false, vec![255]),
(258.0, 2, false, vec![2, 1]),
(258.0, 2, true, vec![1, 2]),
(-1.0, 4, false, vec![255, 255, 255, 255]),
(
-1.0,
9,
false,
vec![255, 255, 255, 255, 255, 255, 255, 255, 0],
),
(1.0, 3, true, vec![0, 0, 1]),
] {
let mut out = Vec::new();
put_integer(&mut out, n, big, size);
assert_eq!(out, want, "{n} in {size} bytes");
}
for (bytes, size, big, signed, want) in [
(vec![255u8], 1, false, true, -1.0),
(vec![255u8], 1, false, false, 255.0),
(vec![2u8, 1], 2, false, true, 258.0),
(vec![2u8, 1], 2, true, true, 513.0),
(vec![0u8, 0, 0, 128], 4, false, true, -2147483648.0),
(vec![0u8, 0, 0, 128], 4, false, false, 2147483648.0),
] {
assert_eq!(
get_integer(&bytes, big, signed, size),
want,
"{bytes:?} {size} {big} {signed}",
);
}
}
}