use alloc::{
format,
string::{String, ToString},
vec,
vec::Vec,
};
#[cfg(feature = "std")]
use crate::{Decode, Encode, Engine, EscnulError, Progress, StreamEncodeError};
#[cfg(not(feature = "std"))]
use crate::{Decode, Encode, Engine, EscnulError, Progress};
use core::fmt;
use memchr::memmem;
#[cfg(feature = "std")]
use std::io::{Read, Write};
pub const NULL_SUB: u8 = 0x01;
struct Cursor<'a> {
data: &'a mut [u8],
pos: usize,
}
impl<'a> Cursor<'a> {
fn new(data: &'a mut [u8]) -> Self {
Self { data, pos: 0 }
}
fn remaining(&self) -> usize {
self.data.len() - self.pos
}
fn write(&mut self, buf: &[u8]) -> Result<usize, usize> {
let len = buf.len().min(self.remaining());
self.data[self.pos..self.pos + len].copy_from_slice(buf);
self.pos += len;
if len < buf.len() { Err(len) } else { Ok(len) }
}
}
struct Process<'a> {
cursor: Cursor<'a>,
processed: usize,
}
impl<'a> Process<'a> {
fn new(data: &'a mut [u8]) -> Self {
Self {
cursor: Cursor::new(data),
processed: 0,
}
}
fn proceed(&mut self, nread: usize) {
self.processed += nread;
}
fn write_partial(&mut self, buf: &[u8]) -> Result<usize, EscnulError> {
let len = self.cursor.write(buf).map_err(|len| {
EscnulError::BufferTooSmall(Progress {
processed: self.processed + len,
written: len,
})
})?;
self.processed += len;
Ok(self.cursor.pos)
}
fn write_full(&mut self, buf: &[u8], nread: usize) -> Result<usize, EscnulError> {
if buf.len() > self.cursor.remaining() {
return Err(EscnulError::BufferTooSmall(Progress {
processed: self.processed,
written: self.cursor.pos,
}));
}
self.cursor.write(buf).unwrap();
self.processed += nread;
Ok(self.cursor.pos)
}
}
pub fn encode_with<T: AsRef<[u8]>>(
esc: u8,
input: T,
output: &mut [u8],
) -> Result<usize, EscnulError> {
if esc == 0 || esc == NULL_SUB {
return Err(EscnulError::BadEscapes(esc));
}
let inp = input.as_ref();
let mut w = 0;
for (processed, &b) in inp.iter().enumerate() {
let need = match b {
0x00 => 1,
x if x == NULL_SUB || x == esc => 2,
_ => 1,
};
if w + need > output.len() {
return Err(EscnulError::BufferTooSmall(Progress {
processed,
written: w,
}));
}
match b {
0x00 => output[w] = NULL_SUB,
x if x == NULL_SUB => {
output[w] = esc;
output[w + 1] = esc;
}
x if x == esc => {
output[w] = esc;
output[w + 1] = b'_';
}
_ => output[w] = b,
}
w += need;
}
Ok(w)
}
pub fn decode_with<T: AsRef<[u8]>>(
esc: u8,
input: T,
output: &mut [u8],
) -> Result<usize, EscnulError> {
if esc == 0 || esc == NULL_SUB {
return Err(EscnulError::BadEscapes(esc));
}
let inp = input.as_ref();
let mut w = 0; let mut i = 0;
while i < inp.len() {
if w >= output.len() {
return Err(EscnulError::BufferTooSmall(Progress {
processed: i,
written: w,
}));
}
match inp[i] {
NULL_SUB => {
output[w] = 0;
w += 1;
i += 1;
}
b if b == esc && i + 1 >= inp.len() => {
return Err(EscnulError::TruncatedInput(Progress {
processed: i,
written: w,
}));
}
b if b == esc && inp[i + 1] == esc => {
output[w] = NULL_SUB;
w += 1;
i += 2;
}
b if b == esc && inp[i + 1] == b'_' => {
output[w] = esc;
w += 1;
i += 2;
}
b if b == esc => {
output[w] = esc;
w += 1;
i += 1;
}
b => {
output[w] = b;
w += 1;
i += 1;
}
}
}
Ok(w)
}
#[derive(Debug, Clone)]
pub struct OnePassEscapeEngine {
esc: u8,
}
impl OnePassEscapeEngine {
pub fn new(esc: u8) -> Self {
Self { esc }
}
}
impl Encode for OnePassEscapeEngine {
fn encode_slice<T: AsRef<[u8]>>(
&self,
input: T,
output: &mut [u8],
) -> Result<usize, EscnulError> {
encode_with(self.esc, input, output)
}
}
impl Decode for OnePassEscapeEngine {
fn decode_slice<T: AsRef<[u8]>>(
&self,
input: T,
output: &mut [u8],
) -> Result<usize, EscnulError> {
decode_with(self.esc, input, output)
}
}
impl Engine for OnePassEscapeEngine {}
#[derive(Debug, Clone)]
pub struct TwoPassEscapeEngine {
allowed_escapes: [bool; 256],
}
pub const TWO_PASS_SED: TwoPassEscapeEngine = TwoPassEscapeEngine {
allowed_escapes: allowed_escapes(),
};
const fn allowed_escapes() -> [bool; 256] {
let mut arr = [false; 256];
let mut i = 0;
while i < arr.len() {
match i as u8 {
b'!' | b'@' | b'#' | b'%' => {
arr[i] = true;
}
_ => {}
}
i += 1;
}
arr
}
fn build_script(command: &str, encoded: &[u8], delimiter: &str) -> Vec<u8> {
let mut program = command.as_bytes().to_vec();
program.extend_from_slice(format!(" << {delimiter}\n").as_bytes());
program.extend_from_slice(encoded);
program.extend_from_slice(format!("\n{delimiter}\n").as_bytes());
program
}
impl TwoPassEscapeEngine {
pub fn allowed_escapes(&self) -> &[bool; 256] {
&self.allowed_escapes
}
pub fn choose_shell_decode_mode<T: AsRef<[u8]>>(&self, input: T) -> ShellDecodeMode {
let (has_null, esc) = self.has_null_and_find_escapes(input.as_ref());
if has_null {
ShellDecodeMode::TrSed { escape: esc }
} else {
ShellDecodeMode::Raw
}
}
#[cfg(feature = "std")]
pub fn choose_shell_decode_mode_reader<R: Read>(
&self,
reader: &mut R,
) -> std::io::Result<ShellDecodeMode> {
let mut has_null = false;
let mut counter = [0usize; 256];
let mut buffer = [0u8; 16 * 1024];
loop {
let read = reader.read(&mut buffer)?;
if read == 0 {
break;
}
for &byte in &buffer[..read] {
if byte == 0 {
has_null = true;
} else if self.allowed_escapes[byte as usize] {
counter[byte as usize] += 1;
}
}
}
if !has_null {
return Ok(ShellDecodeMode::Raw);
}
let escape = self
.allowed_escapes()
.iter()
.enumerate()
.filter(|(_, allowed)| **allowed)
.min_by_key(|(idx, _)| counter[*idx])
.map(|(idx, _)| idx as u8)
.expect("allowed_escapes is never empty");
Ok(ShellDecodeMode::TrSed { escape })
}
fn has_null_and_find_escapes(&self, input: &[u8]) -> (bool, u8) {
let mut has_null = false;
let mut counter = [0; 256];
for c in input {
if *c == 0 {
has_null = true;
} else if self.allowed_escapes[*c as usize] {
counter[*c as usize] += 1;
}
}
let esc = self
.allowed_escapes()
.iter()
.enumerate()
.filter(|(_, x)| **x)
.min_by_key(|(i, _)| counter[*i])
.unwrap()
.0 as u8;
(has_null, esc)
}
}
impl Encode for TwoPassEscapeEngine {
fn encode_slice<T: AsRef<[u8]>>(
&self,
input: T,
output: &mut [u8],
) -> Result<usize, EscnulError> {
let (has_null, esc) = self.has_null_and_find_escapes(input.as_ref());
if output.is_empty() {
return Err(EscnulError::BufferTooSmall(Progress {
processed: 0,
written: 0,
}));
}
if has_null {
output[0] = esc;
OnePassEscapeEngine { esc }
.encode_slice(input, &mut output[1..])
.map(|n| n + 1)
} else {
let mut output = Process::new(output);
output.write_full(&[NULL_SUB], 0)?;
output.write_partial(input.as_ref())
}
}
fn encode<T: AsRef<[u8]>>(&self, input: T) -> Vec<u8> {
let mut output = vec![0; input.as_ref().len() * 2 + 1];
let n = self.encode_slice(input, &mut output).unwrap();
output.truncate(n);
output
}
}
impl Decode for TwoPassEscapeEngine {
fn decode_slice<T: AsRef<[u8]>>(
&self,
input: T,
output: &mut [u8],
) -> Result<usize, EscnulError> {
if input.as_ref().is_empty() {
return Err(EscnulError::BufferTooSmall(Progress {
processed: 0,
written: 0,
}));
}
let esc = input.as_ref()[0];
if esc == NULL_SUB {
let inp = input.as_ref();
let mut output = Process::new(output);
let payload = &inp[1..];
output.proceed(1);
output.write_partial(payload)
} else {
OnePassEscapeEngine { esc }.decode_slice(&input.as_ref()[1..], output)
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ShellDecodeMode {
Raw,
TrSed {
escape: u8,
},
}
impl ShellDecodeMode {
pub fn write_decoder_fragment<W: fmt::Write>(self, size: u64, out: &mut W) -> fmt::Result {
match self {
Self::Raw => write!(out, "dd bs=1 count={} 2>/dev/null", size),
Self::TrSed { escape } => write!(
out,
"{{ tr '\\001' '\\000' | sed -e \"s/{0}{0}/$(printf '\\001')/g;s/{0}_/{0}/g\" | dd bs=1 count={1} 2>/dev/null; }}",
escape as char, size
),
}
}
pub fn decoder_fragment(self, size: u64) -> String {
let mut fragment = String::new();
self.write_decoder_fragment(size, &mut fragment)
.expect("writing to String cannot fail");
fragment
}
#[cfg(feature = "std")]
pub fn encode_reader<R: Read, W: Write>(
self,
reader: &mut R,
writer: &mut W,
) -> Result<(), StreamEncodeError> {
let mut input = [0u8; 16 * 1024];
let mut output = [0u8; 16 * 1024 * 2];
match self {
Self::Raw => loop {
let read = reader.read(&mut input)?;
if read == 0 {
break;
}
writer.write_all(&input[..read])?;
},
Self::TrSed { escape } => {
let engine = OnePassEscapeEngine::new(escape);
loop {
let read = reader.read(&mut input)?;
if read == 0 {
break;
}
let written = engine.encode_slice(&input[..read], &mut output)?;
writer.write_all(&output[..written])?;
}
}
}
Ok(())
}
}
#[derive(Debug, Clone)]
pub struct TrSedEngine(TwoPassEscapeEngine);
pub const TR_SED: TrSedEngine = TrSedEngine(TWO_PASS_SED);
impl TrSedEngine {
fn find_delimiter(input: &[u8]) -> String {
let mut delimiter = "__ES".to_string();
let mut pos = 0;
while let Some(p) = memmem::find(&input[pos..], delimiter.as_bytes()) {
let next = input.get(p + delimiter.len()).unwrap_or(&b'\0');
if *next == b'_' {
delimiter.push('X');
} else {
delimiter.push('_');
}
pos = p;
}
delimiter.push('_');
delimiter
}
}
impl Encode for TrSedEngine {
fn encode_slice<T: AsRef<[u8]>>(
&self,
input: T,
output: &mut [u8],
) -> Result<usize, EscnulError> {
let out = self.encode(input);
let mut output = Process::new(output);
let n = out.len();
output.write_full(&out, n)
}
fn encode<T: AsRef<[u8]>>(&self, input: T) -> Vec<u8> {
let delimiter = Self::find_delimiter(input.as_ref());
let mode = self.0.choose_shell_decode_mode(input.as_ref());
let command = mode.decoder_fragment(input.as_ref().len() as u64);
let encoded = match mode {
ShellDecodeMode::Raw => input.as_ref().to_vec(),
ShellDecodeMode::TrSed { escape } => OnePassEscapeEngine::new(escape).encode(input),
};
build_script(
&command,
encoded.as_slice(),
&delimiter,
)
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::process::{Command, Stdio};
#[test]
fn encode_with_fixed_buffer() {
let mut out = [0u8; 16];
let written = encode_with(b'@', b"\0@\x01X", &mut out).unwrap();
assert_eq!(&out[..written], b"\x01@_@@X");
}
#[test]
fn encode_with_buffer_too_small() {
let mut out = [0u8; 2];
let err = encode_with(b'@', b"\0@", &mut out).unwrap_err();
assert_eq!(
err,
EscnulError::BufferTooSmall(Progress {
processed: 1,
written: 1,
})
);
}
#[test]
fn decode_simple() {
let data = b"X\0Y@Z";
let mut enc = [0u8; 16];
let n = encode_with(b'@', &data[..], &mut enc).unwrap();
let mut dec = [0u8; 16];
let m = decode_with(b'@', &enc[..n], &mut dec).unwrap();
assert_eq!(&dec[..m], data);
assert_eq!(m, data.len());
assert_eq!(n, data.len() + 1);
assert_eq!(&enc[..n], b"X\x01Y@_Z");
}
#[test]
fn decode_truncated() {
let err = decode_with(b'@', b"a@", &mut [0u8; 8]).unwrap_err();
assert!(matches!(
err,
EscnulError::TruncatedInput(Progress {
processed: 1,
written: 1,
})
));
}
#[test]
fn decode_buffer_too_small_errors() {
let mut enc = [0u8; 1];
let _ = encode_with(b'@', b"\0", &mut enc).unwrap();
match decode_with(b'@', &enc, &mut []) {
Err(EscnulError::BufferTooSmall(progress)) => {
assert_eq!(
progress,
Progress {
processed: 0,
written: 0
}
);
}
other => panic!("expected BufferTooSmall, got {:?}", other),
}
}
#[test]
fn one_pass_engine_roundtrip() {
let engine = OnePassEscapeEngine::new(b'%');
let src = b"Hello\0Rust%World";
let enc = engine.encode(src);
assert!(!enc.is_empty(), "encoded output should not be empty");
assert_ne!(enc[0], NULL_SUB);
let dec = engine.decode(enc.clone());
assert_eq!(&dec[..], src, "decoded must match original");
}
#[test]
fn two_pass_engine_raw_roundtrip() {
let src = b"JustPlainAsciiNoNULs";
let enc = TWO_PASS_SED.encode(src);
assert_eq!(enc[0], NULL_SUB,);
assert_eq!(&enc[1..], src);
let dec = TWO_PASS_SED.decode(enc);
assert_eq!(&dec[..], src);
}
#[test]
fn two_pass_engine_escape_roundtrip() {
let src = b"Has\0Some\0NULs";
let enc = TWO_PASS_SED.encode(src);
assert_ne!(enc[0], NULL_SUB);
let dec = TWO_PASS_SED.decode(enc.clone());
assert_eq!(&dec[..], src);
}
#[test]
fn encode_slice_and_vec_methods_agree() {
let engine = OnePassEscapeEngine::new(b'@');
let src = b"foo\0bar@baz";
let mut buf = vec![0u8; src.len() * 2];
let n = engine.encode_slice(src, &mut buf).unwrap();
let slice_res = buf[..n].to_vec();
let vec_res = engine.encode(src);
assert_eq!(slice_res, vec_res);
}
#[test]
fn decode_slice_and_vec_methods_agree() {
let engine = OnePassEscapeEngine::new(b'#');
let src = b"\0###abc";
let enc = engine.encode(src);
let mut buf = vec![0u8; src.len()];
let n = engine.decode_slice(&enc, &mut buf).unwrap();
let slice_res = buf[..n].to_vec();
let vec_res = engine.decode(&enc);
assert_eq!(slice_res, vec_res);
assert_eq!(vec_res, src);
}
fn run_script(command: &str, input: &[u8]) -> Vec<u8> {
use std::io::Write;
let command = command.split(' ').collect::<Vec<&str>>();
let mut c = Command::new(command[0])
.args(&command[1..])
.stdout(Stdio::piped())
.stdin(Stdio::piped())
.spawn()
.unwrap();
c.stdin.take().unwrap().write_all(input).unwrap();
let output = c.wait_with_output().unwrap();
assert!(output.status.success());
output.stdout
}
#[test]
fn build_command_fragment_run() {
for esc in TWO_PASS_SED
.allowed_escapes()
.iter()
.enumerate()
.filter(|(_, x)| **x)
.map(|(i, _)| i as u8)
{
let mut expected = b"Hello\0World\x01".to_vec();
expected.push(esc);
let fragment = ShellDecodeMode::TrSed { escape: esc }.decoder_fragment(expected.len() as u64);
let encoded = format!("Hello\x01World{esc}{esc}{esc}_", esc = esc as char)
.as_bytes()
.to_vec();
let script = build_script(&fragment, encoded.as_slice(), "EOF");
let dec = run_script("/bin/sh", script.as_slice());
assert_eq!(dec, expected, "esc: {esc}");
}
}
#[test]
fn tr_sed_engine_raw_roundtrip() {
for src in &[
b"Hello Rust%World\n",
b" Hello Rust%World",
b"\nHello Rust%World",
] {
let enc = TR_SED.encode(src);
assert!(!enc.is_empty(), "encoded output should not be empty");
assert_eq!(run_script("/bin/sh", &enc), &src[..], "src: {:?}", src);
}
}
#[test]
fn tr_sed_engine_roundtrip() {
let src = b"Hello\0Rust%World";
let enc = TR_SED.encode(src);
assert!(!enc.is_empty(), "encoded output should not be empty");
assert_eq!(run_script("/bin/sh", &enc), src);
}
#[test]
fn choose_shell_decode_mode_reports_raw() {
assert_eq!(TWO_PASS_SED.choose_shell_decode_mode(b"hello"), ShellDecodeMode::Raw);
}
#[test]
fn choose_shell_decode_mode_reports_escape() {
assert!(matches!(
TWO_PASS_SED.choose_shell_decode_mode(b"he\0llo"),
ShellDecodeMode::TrSed { .. }
));
}
#[test]
fn encode_reader_roundtrip() {
let src = b"Hello\0Rust%World";
let mode = TWO_PASS_SED.choose_shell_decode_mode(src);
let mut encoded = Vec::new();
mode.encode_reader(&mut &src[..], &mut encoded).unwrap();
let script = build_script(&mode.decoder_fragment(src.len() as u64), &encoded, "EOF");
assert_eq!(run_script("/bin/sh", &script), src);
}
#[test]
fn write_decoder_fragment_matches_string_api() {
let mode = ShellDecodeMode::TrSed { escape: b'!' };
let mut out = String::new();
mode.write_decoder_fragment(42, &mut out).unwrap();
assert_eq!(out, mode.decoder_fragment(42));
}
#[test]
fn find_delimiter() {
let input = b"Hello__ES__World__ES__";
let delimiter = TrSedEngine::find_delimiter(input);
assert_eq!(delimiter, "__ESX_");
}
}