use std::io::{self, Read, Write};
use std::panic::{self, AssertUnwindSafe};
use std::path::PathBuf;
use std::process::{Command, Stdio};
use std::sync::OnceLock;
use std::time::{Duration, Instant};
use preflate_rs::{PreflateConfig, preflate_whole_deflate_stream, recreate_whole_deflate_stream};
use crate::error::{Error, Result};
use crate::limits::Limits;
pub const ZLIB_HEADER_LEN: usize = 2;
pub const ZLIB_TRAILER_LEN: usize = 4;
pub const MIN_ZLIB_LEN: usize = ZLIB_HEADER_LEN + 1 + ZLIB_TRAILER_LEN;
pub const MAX_CHAIN_LENGTH: u32 = 4096;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ReplayPlan {
pub header: [u8; ZLIB_HEADER_LEN],
pub plaintext: Vec<u8>,
pub corrections: Vec<u8>,
pub adler: [u8; ZLIB_TRAILER_LEN],
pub raw_len: u32,
}
pub fn zlib_header_valid(bytes: &[u8]) -> bool {
if bytes.len() < ZLIB_HEADER_LEN {
return false;
}
let cmf = bytes[0];
let flg = bytes[1];
let method = cmf & 0x0f;
let cinfo = cmf >> 4;
method == 8 && cinfo <= 7 && (u16::from(cmf) * 256 + u16::from(flg)).is_multiple_of(31)
}
fn config(limits: Limits) -> PreflateConfig {
let cap = limits
.max_output_bytes
.min(u64::from(limits.max_record_len));
PreflateConfig {
max_chain_length: MAX_CHAIN_LENGTH,
plain_text_limit: cap.min(usize::MAX as u64) as usize,
verify_compression: true,
}
}
pub fn replay_raw(plaintext: &[u8], corrections: &[u8]) -> Result<Vec<u8>> {
let outcome = panic::catch_unwind(AssertUnwindSafe(|| {
recreate_whole_deflate_stream(plaintext, corrections)
}));
match outcome {
Ok(Ok(bytes)) => Ok(bytes),
Ok(Err(e)) => Err(Error::codec_replay(format!("deflate replay failed: {e}"))),
Err(_) => Err(Error::codec_replay(
"deflate replay panicked on malformed correction state",
)),
}
}
pub const REPLAY_WORKER_SUBCOMMAND: &str = "__replay-worker";
pub const REPLAY_WORKER_MAX_FIELD: u32 = 1 << 30;
pub const REPLAY_DEFAULT_TIMEOUT_MS: u64 = 30_000;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct WorkerRequest {
pub plaintext: Vec<u8>,
pub corrections: Vec<u8>,
pub declared_len: u32,
}
fn read_field<R: Read>(r: &mut R, max: u32) -> io::Result<Vec<u8>> {
let mut len = [0u8; 4];
r.read_exact(&mut len)?;
let n = u32::from_le_bytes(len);
if n > max {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
format!("field length {n} exceeds worker bound {max}"),
));
}
let mut buf = vec![0u8; n as usize];
r.read_exact(&mut buf)?;
Ok(buf)
}
pub fn read_worker_request<R: Read>(r: &mut R) -> io::Result<WorkerRequest> {
let plaintext = read_field(r, REPLAY_WORKER_MAX_FIELD)?;
let corrections = read_field(r, REPLAY_WORKER_MAX_FIELD)?;
let mut declared = [0u8; 4];
r.read_exact(&mut declared)?;
Ok(WorkerRequest {
plaintext,
corrections,
declared_len: u32::from_le_bytes(declared),
})
}
pub fn encode_worker_request(
plaintext: &[u8],
corrections: &[u8],
declared_len: u32,
) -> Result<Vec<u8>> {
let p = u32::try_from(plaintext.len())
.map_err(|_| Error::codec_replay("replay plaintext exceeds u32 framing"))?;
let c = u32::try_from(corrections.len())
.map_err(|_| Error::codec_replay("replay corrections exceed u32 framing"))?;
let mut buf = Vec::with_capacity(12 + plaintext.len() + corrections.len());
buf.extend_from_slice(&p.to_le_bytes());
buf.extend_from_slice(plaintext);
buf.extend_from_slice(&c.to_le_bytes());
buf.extend_from_slice(corrections);
buf.extend_from_slice(&declared_len.to_le_bytes());
Ok(buf)
}
pub fn write_worker_reply<W: Write>(w: &mut W, status: u8, payload: &[u8]) -> io::Result<()> {
let n = u32::try_from(payload.len())
.map_err(|_| io::Error::new(io::ErrorKind::InvalidData, "worker reply too large"))?;
w.write_all(&[status])?;
w.write_all(&n.to_le_bytes())?;
w.write_all(payload)?;
w.flush()
}
pub fn read_worker_reply<R: Read>(
r: &mut R,
max_payload: u64,
) -> io::Result<Option<(u8, Vec<u8>)>> {
let mut status = [0u8; 1];
match r.read(&mut status) {
Ok(0) => return Ok(None),
Ok(_) => {}
Err(e) if e.kind() == io::ErrorKind::UnexpectedEof => return Ok(None),
Err(e) => return Err(e),
}
let mut len = [0u8; 4];
r.read_exact(&mut len)?;
let n = u32::from_le_bytes(len);
if u64::from(n) > max_payload {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
format!("worker reply length {n} exceeds bound {max_payload}"),
));
}
let mut payload = vec![0u8; n as usize];
r.read_exact(&mut payload)?;
Ok(Some((status[0], payload)))
}
pub fn run_worker_stdio() -> ! {
std::panic::set_hook(Box::new(|info| {
eprintln!("replay worker panic (isolated): {info}");
}));
let stdin = io::stdin();
let stdout = io::stdout();
let mut r = stdin.lock();
let mut w = stdout.lock();
let code = match read_worker_request(&mut r) {
Ok(req) => {
let outcome = panic::catch_unwind(AssertUnwindSafe(|| {
recreate_whole_deflate_stream(&req.plaintext, &req.corrections)
}));
match outcome {
Ok(Ok(bytes)) => {
let _ = write_worker_reply(&mut w, 0, &bytes);
0
}
Ok(Err(e)) => {
let _ = write_worker_reply(
&mut w,
1,
format!("deflate replay failed: {e}").as_bytes(),
);
1
}
Err(_) => {
let _ = write_worker_reply(
&mut w,
1,
b"deflate replay panicked on malformed correction state",
);
1
}
}
}
Err(e) => {
let _ = write_worker_reply(
&mut w,
1,
format!("malformed worker request: {e}").as_bytes(),
);
2
}
};
std::process::exit(code);
}
static WORKER_PATH: OnceLock<Option<PathBuf>> = OnceLock::new();
static INSTALLED_WORKER: OnceLock<PathBuf> = OnceLock::new();
pub fn install_default_replay_worker(path: PathBuf) {
let _ = INSTALLED_WORKER.set(path);
}
fn resolve_worker_path() -> Option<PathBuf> {
WORKER_PATH
.get_or_init(|| {
std::env::var_os("VOLE_REPLAY_WORKER")
.filter(|v| !v.is_empty())
.map(PathBuf::from)
.or_else(|| INSTALLED_WORKER.get().cloned())
})
.clone()
}
fn env_u64(name: &str) -> Option<u64> {
std::env::var(name).ok().and_then(|v| v.trim().parse().ok())
}
fn replay_timeout() -> Duration {
Duration::from_millis(env_u64("VOLE_REPLAY_TIMEOUT_MS").unwrap_or(REPLAY_DEFAULT_TIMEOUT_MS))
}
fn replay_memory_cap_kb(
plaintext: &[u8],
corrections: &[u8],
declared_len: u32,
limits: Limits,
) -> u64 {
if let Some(mb) = env_u64("VOLE_REPLAY_MEM_MB") {
return mb.saturating_mul(1024);
}
let total = plaintext.len() as u64 + corrections.len() as u64 + u64::from(declared_len);
let bytes = total.saturating_mul(8).saturating_add(64 << 20);
let lo = 256u64 << 20;
let hi = limits.max_replay_bytes.clamp(1, 1u64 << 31);
let cap = if hi >= lo { bytes.clamp(lo, hi) } else { hi };
cap / 1024
}
fn describe_worker_exit(status: std::process::ExitStatus) -> String {
#[cfg(unix)]
{
use std::os::unix::process::ExitStatusExt;
if let Some(sig) = status.signal() {
if sig == 6 {
return "replay worker aborted (SIGABRT); likely exceeded the address-space cap"
.to_string();
}
return format!("replay worker killed by signal {sig}");
}
}
match status.code() {
Some(0) => "replay worker exited 0 without a reply".to_string(),
Some(c) => format!("replay worker exited with status {c}"),
None => "replay worker terminated abnormally".to_string(),
}
}
pub fn replay_bounded(
plaintext: &[u8],
corrections: &[u8],
declared_len: u32,
limits: Limits,
) -> Result<Vec<u8>> {
let Some(worker) = resolve_worker_path() else {
return replay_raw(plaintext, corrections);
};
let cap_kb = replay_memory_cap_kb(plaintext, corrections, declared_len, limits);
let timeout = replay_timeout();
let script = format!(
"ulimit -v {cap_kb} 2>/dev/null; ulimit -t {} 2>/dev/null; exec \"$0\" {REPLAY_WORKER_SUBCOMMAND}",
timeout.as_secs().max(1)
);
let mut child = Command::new("sh")
.arg("-c")
.arg(&script)
.arg(&worker)
.stdin(Stdio::piped())
.stdout(Stdio::piped())
.stderr(Stdio::piped())
.spawn()
.map_err(|e| {
Error::codec_replay(format!(
"failed to spawn replay worker {}: {e}",
worker.display()
))
})?;
let mut stdin = child
.stdin
.take()
.ok_or_else(|| Error::codec_replay("replay worker stdin unavailable"))?;
let mut stdout = child
.stdout
.take()
.ok_or_else(|| Error::codec_replay("replay worker stdout unavailable"))?;
let stderr = child
.stderr
.take()
.ok_or_else(|| Error::codec_replay("replay worker stderr unavailable"))?;
let request = encode_worker_request(plaintext, corrections, declared_len)?;
let stderr_thread = std::thread::spawn(move || {
let mut buf = Vec::new();
let _ = stderr.take(8192).read_to_end(&mut buf);
String::from_utf8_lossy(&buf).into_owned()
});
let max_reply = cap_kb.saturating_mul(1024).max(1);
let reply_thread = std::thread::spawn(move || read_worker_reply(&mut stdout, max_reply));
let write_err = stdin.write_all(&request).err();
drop(stdin);
let deadline = Instant::now() + timeout;
let mut timed_out = false;
let status = loop {
match child.try_wait() {
Ok(Some(status)) => break Some(status),
Ok(None) => {
if Instant::now() >= deadline {
timed_out = true;
let _ = child.kill();
break None;
}
std::thread::sleep(Duration::from_millis(2));
}
Err(_) => {
let _ = child.kill();
break None;
}
}
};
let status = match status {
Some(status) => Some(status),
None => child.wait().ok(),
};
let reply = reply_thread
.join()
.unwrap_or_else(|_| Err(io::Error::other("replay reply reader panicked")));
let stderr_text = stderr_thread.join().unwrap_or_default();
if timed_out {
return Err(Error::codec_replay(format!(
"replay worker exceeded the {timeout:?} time bound and was killed"
)));
}
match reply {
Ok(Some((0, payload))) => Ok(payload),
Ok(Some((1, payload))) => Err(Error::codec_replay(
String::from_utf8_lossy(&payload).into_owned(),
)),
Ok(Some((other, _))) => Err(Error::codec_replay(format!(
"replay worker returned unknown status {other}"
))),
Ok(None) | Err(_) => {
let mut msg = String::from("replay worker produced no reply");
if let Some(status) = status {
msg.push_str("; ");
msg.push_str(&describe_worker_exit(status));
}
if let Some(e) = write_err {
msg.push_str(&format!("; request write failed: {e}"));
}
if !stderr_text.trim().is_empty() {
msg.push_str("; stderr: ");
msg.push_str(stderr_text.trim());
}
Err(Error::codec_replay(msg))
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ReplayDecline {
TooShort,
TooLarge,
NotZlib,
EmptyPayload,
AnalyzerPanic,
AnalyzerError,
NotFullyConsumed,
DictionaryPrefix,
PlaintextTooLarge,
CorrectionsTooLarge,
NotReproducible,
}
impl ReplayDecline {
pub fn name(self) -> &'static str {
match self {
ReplayDecline::TooShort => "too_short",
ReplayDecline::TooLarge => "too_large",
ReplayDecline::NotZlib => "not_zlib",
ReplayDecline::EmptyPayload => "empty_payload",
ReplayDecline::AnalyzerPanic => "analyzer_panic",
ReplayDecline::AnalyzerError => "analyzer_error",
ReplayDecline::NotFullyConsumed => "not_fully_consumed",
ReplayDecline::DictionaryPrefix => "dictionary_prefix",
ReplayDecline::PlaintextTooLarge => "plaintext_too_large",
ReplayDecline::CorrectionsTooLarge => "corrections_too_large",
ReplayDecline::NotReproducible => "not_reproducible",
}
}
}
pub fn try_replay(bytes: &[u8], limits: Limits) -> Option<ReplayPlan> {
try_replay_detailed(bytes, limits).ok()
}
pub fn try_replay_detailed(
bytes: &[u8],
limits: Limits,
) -> std::result::Result<ReplayPlan, ReplayDecline> {
let total = bytes.len();
if total > limits.max_input_bytes as usize {
return Err(ReplayDecline::TooLarge);
}
if total < MIN_ZLIB_LEN {
return Err(ReplayDecline::TooShort);
}
if !zlib_header_valid(bytes) {
return Err(ReplayDecline::NotZlib);
}
let raw = &bytes[ZLIB_HEADER_LEN..total - ZLIB_TRAILER_LEN];
if raw.is_empty() {
return Err(ReplayDecline::EmptyPayload);
}
let analyzed = match panic::catch_unwind(AssertUnwindSafe(|| {
preflate_whole_deflate_stream(raw, &config(limits))
})) {
Ok(Ok(analyzed)) => analyzed,
Ok(Err(_)) => return Err(ReplayDecline::AnalyzerError),
Err(_) => return Err(ReplayDecline::AnalyzerPanic),
};
let (chunk, plain) = analyzed;
if chunk.compressed_size != raw.len() {
return Err(ReplayDecline::NotFullyConsumed);
}
if !plain.prefix().is_empty() {
return Err(ReplayDecline::DictionaryPrefix);
}
if plaintext_len_exceeds(plain.text(), limits) {
return Err(ReplayDecline::PlaintextTooLarge);
}
if chunk.corrections.len() as u64 > u64::from(limits.max_record_len) {
return Err(ReplayDecline::CorrectionsTooLarge);
}
let plaintext = plain.text().to_vec();
if replay_raw(&plaintext, &chunk.corrections).ok().as_deref() != Some(raw) {
return Err(ReplayDecline::NotReproducible);
}
let raw_len = match u32::try_from(raw.len()) {
Ok(n) => n,
Err(_) => return Err(ReplayDecline::TooLarge),
};
let mut header = [0u8; ZLIB_HEADER_LEN];
header.copy_from_slice(&bytes[..ZLIB_HEADER_LEN]);
let mut adler = [0u8; ZLIB_TRAILER_LEN];
adler.copy_from_slice(&bytes[total - ZLIB_TRAILER_LEN..]);
Ok(ReplayPlan {
header,
plaintext,
corrections: chunk.corrections,
adler,
raw_len,
})
}
fn plaintext_len_exceeds(plaintext: &[u8], limits: Limits) -> bool {
plaintext.len() as u64 > u64::from(limits.max_record_len)
}
#[cfg(test)]
mod tests {
use super::*;
use flate2::Compression;
use flate2::write::ZlibEncoder;
use std::io::Write;
fn zlib(data: &[u8], level: u32) -> Vec<u8> {
let mut e = ZlibEncoder::new(Vec::new(), Compression::new(level));
e.write_all(data).unwrap();
e.finish().unwrap()
}
fn content_like() -> Vec<u8> {
let mut v = Vec::new();
for i in 0..400 {
v.extend_from_slice(
format!(
"BT /F1 12 Tf 72 {} Td (Invoice line {i:05} amount 456.78) Tj ET\n",
700 - (i % 40)
)
.as_bytes(),
);
}
v
}
#[test]
fn zlib_header_shape_is_checked() {
assert!(zlib_header_valid(&[0x78, 0x9c, 0x00, 0x00, 0x00, 0x00]));
assert!(zlib_header_valid(&[0x78, 0x01, 0x00, 0x00, 0x00, 0x00]));
assert!(!zlib_header_valid(&[0x00, 0x00]));
assert!(!zlib_header_valid(&[0x3b, 0x00, 0x00, 0x00, 0x00, 0x00]));
assert!(!zlib_header_valid(&[0x78, 0x9d, 0x00, 0x00, 0x00, 0x00]));
assert!(!zlib_header_valid(&[0x78]));
}
#[test]
fn replay_plan_round_trips_byte_exactly() {
for (name, data, level) in [
("content", content_like(), 6),
(
"text",
b"the quick brown fox jumps over the lazy dog. ".repeat(200),
9,
),
("small", b"hello hello hello".to_vec(), 6),
(
"incompressible",
(0..4096u32)
.map(|i| (i.wrapping_mul(2654435761) >> 13) as u8)
.collect(),
6,
),
] {
let z = zlib(&data, level);
let plan = try_replay(&z, Limits::DEFAULT)
.unwrap_or_else(|| panic!("{name} must produce a replay plan"));
let mut rebuilt = Vec::new();
rebuilt.extend_from_slice(&plan.header);
rebuilt.extend_from_slice(&replay_raw(&plan.plaintext, &plan.corrections).unwrap());
rebuilt.extend_from_slice(&plan.adler);
assert_eq!(rebuilt, z, "{name} must replay byte-for-byte");
}
}
#[test]
fn replay_declines_non_zlib_and_short() {
assert!(try_replay(b"not zlib at all", Limits::DEFAULT).is_none());
assert!(try_replay(&[0x78, 0x9c, 0x00, 0x00, 0x00], Limits::DEFAULT).is_none());
assert!(try_replay(&[], Limits::DEFAULT).is_none());
assert!(try_replay(&[0x78, 0x9c], Limits::DEFAULT).is_none());
}
#[test]
fn hostile_corrections_never_panic() {
let plain = content_like();
let mut state = 0x1234_5678_9abc_def0u64;
let mut next = || {
state ^= state << 13;
state ^= state >> 7;
state ^= state << 17;
state
};
for _ in 0..48 {
let mut blob = vec![0u8; (next() % 96) as usize];
for b in blob.iter_mut() {
*b = next() as u8;
}
let _ = replay_raw(&plain, &blob);
let _ = replay_raw(&blob, &plain);
let _ = replay_raw(&[], &blob);
}
assert!(replay_raw(&plain, &[]).is_err() || replay_raw(&plain, &[]).is_ok());
}
#[test]
fn tiny_plaintext_limit_declines_instead_of_truncating() {
let data = content_like();
let z = zlib(&data, 6);
let limits = Limits {
max_record_len: 8,
..Limits::DEFAULT
};
assert!(try_replay(&z, limits).is_none());
}
#[test]
fn worker_framing_round_trips() {
let req = encode_worker_request(b"plain", b"corrections", 7).unwrap();
let mut cur = std::io::Cursor::new(req);
assert_eq!(
read_worker_request(&mut cur).unwrap(),
WorkerRequest {
plaintext: b"plain".to_vec(),
corrections: b"corrections".to_vec(),
declared_len: 7,
}
);
let mut reply = Vec::new();
write_worker_reply(&mut reply, 0, b"raw deflate").unwrap();
let mut rc = std::io::Cursor::new(reply);
assert_eq!(
read_worker_reply(&mut rc, 1024).unwrap(),
Some((0, b"raw deflate".to_vec()))
);
let mut bad = Vec::new();
bad.extend_from_slice(&u32::MAX.to_le_bytes());
let mut bc = std::io::Cursor::new(bad);
assert!(read_worker_request(&mut bc).is_err());
let mut big_reply = vec![0u8; 5];
big_reply[1..].copy_from_slice(&u32::MAX.to_le_bytes());
let mut brc = std::io::Cursor::new(big_reply);
assert!(read_worker_reply(&mut brc, 16).is_err());
let mut empty = std::io::Cursor::new(Vec::new());
assert_eq!(read_worker_reply(&mut empty, 1024).unwrap(), None);
}
}