use std::env;
use std::io::Read;
use std::io::Write as _;
use std::net::TcpListener;
use std::net::TcpStream;
use std::panic::catch_unwind;
use std::panic::AssertUnwindSafe;
use std::panic::UnwindSafe;
use std::process;
use std::process::Child;
use std::process::Command;
use std::process::ExitCode;
use std::process::Stdio;
use std::process::Termination;
use crate::cmdline;
use crate::error::Result;
const OCCURS_ENV: &str = "TEST_FORK_OCCURS";
const OCCURS_TERM_LENGTH: usize = 17;
fn supervise_child(child: Child) -> ExitCode {
let output = child.wait_with_output().expect("failed to wait for child");
if !output.stdout.is_empty() {
let s = String::from_utf8_lossy(&output.stdout);
print!("{s}");
}
if !output.stderr.is_empty() {
let s = String::from_utf8_lossy(&output.stderr);
eprint!("{s}");
}
if output.status.success() {
ExitCode::SUCCESS
} else {
ExitCode::FAILURE
}
}
pub fn run_should_panic<F, T>(test: F, expected: Option<&str>) -> ExitCode
where
F: FnOnce() -> T + UnwindSafe,
{
let payload = match catch_unwind(test) {
Ok(_) => {
eprintln!("note: test did not panic as expected");
return ExitCode::FAILURE
}
Err(payload) => payload,
};
let expected = match expected {
Some(expected) => expected,
None => return ExitCode::SUCCESS,
};
let message = payload
.downcast_ref::<&str>()
.copied()
.or_else(|| payload.downcast_ref::<String>().map(String::as_str));
match message {
Some(message) if message.contains(expected) => ExitCode::SUCCESS,
_ => {
eprintln!("note: panic did not contain the expected string '{expected}'");
ExitCode::FAILURE
}
}
}
pub fn fork<F, T>(fork_id: &str, test_name: &str, test: F) -> Result<ExitCode>
where
F: Fn() -> T,
T: Termination,
{
fn no_configure_child(_child: &mut Command) {}
fork_int(
test_name,
fork_id,
no_configure_child,
supervise_child,
test,
)
}
#[expect(clippy::panic_in_result_fn, clippy::unwrap_in_result)]
pub fn fork_in_out<F, T>(
fork_id: &str,
test_name: &str,
test: F,
data: &mut [u8],
) -> Result<ExitCode>
where
F: Fn(&mut [u8]) -> T,
T: Termination,
{
let listener = TcpListener::bind("127.0.0.1:0").expect("failed to bind TCP socket");
let addr = listener.local_addr().unwrap();
let data_len = data.len();
fork_int(
test_name,
fork_id,
|cmd| {
cmd.env(fork_id, addr.to_string());
},
|child| {
let (mut stream, _addr) = listener
.accept()
.expect("failed to listen for child connection");
let () = stream
.write_all(data)
.expect("failed to send data to child");
let () = stream
.read_exact(data)
.expect("failed to receive data from child");
supervise_child(child)
},
|| {
let addr = env::var(fork_id).unwrap_or_else(|err| {
panic!("failed to retrieve {fork_id} environment variable: {err}")
});
let mut stream =
TcpStream::connect(addr).expect("failed to establish connection with parent");
let mut data = Vec::with_capacity(data_len);
let () = unsafe { data.set_len(data_len) };
let () = stream
.read_exact(&mut data)
.expect("failed to receive data from parent");
let status = test(&mut data);
let () = stream
.write_all(&data)
.expect("failed to send data to parent");
status
},
)
}
pub(crate) fn fork_int<M, P, C, R, T>(
test_name: &str,
fork_id: &str,
process_modifier: M,
in_parent: P,
in_child: C,
) -> Result<R>
where
M: FnOnce(&mut process::Command),
P: FnOnce(Child) -> R,
T: Termination,
C: FnOnce() -> T,
{
let mut process_modifier = Some(process_modifier);
let mut in_parent = Some(in_parent);
let mut in_child = Some(in_child);
fork_impl(
test_name,
fork_id,
&mut |cmd| process_modifier.take().unwrap()(cmd),
&mut |child| in_parent.take().unwrap()(child),
&mut || in_child.take().unwrap()(),
)
}
#[expect(clippy::panic_in_result_fn, clippy::unwrap_in_result)]
fn fork_impl<T: Termination, R>(
test_name: &str,
fork_id: &str,
process_modifier: &mut dyn FnMut(&mut process::Command),
in_parent: &mut dyn FnMut(Child) -> R,
in_child: &mut dyn FnMut() -> T,
) -> Result<R> {
let mut occurs = env::var(OCCURS_ENV).unwrap_or_else(|_| String::new());
if occurs.contains(fork_id) {
match catch_unwind(AssertUnwindSafe(in_child)) {
Ok(test_result) => {
let rc = if test_result.report() == ExitCode::SUCCESS {
0
} else {
70
};
process::exit(rc)
}
Err(_) => process::exit(70 ),
}
} else {
if occurs.len() > 16 * OCCURS_TERM_LENGTH {
panic!("test-fork: Not forking due to >=16 levels of recursion");
}
occurs.push_str(fork_id);
let mut command =
process::Command::new(env::current_exe().expect("current_exe() failed, cannot fork"));
command
.args(cmdline::strip_cmdline(env::args())?)
.args(cmdline::RUN_TEST_ARGS)
.arg(test_name)
.env(OCCURS_ENV, &occurs)
.stdin(Stdio::null())
.stdout(Stdio::piped())
.stderr(Stdio::piped());
process_modifier(&mut command);
let child = command.spawn()?;
let result = in_parent(child);
Ok(result)
}
}
#[cfg(test)]
mod test {
use super::*;
use std::io;
use std::process::abort;
fn wait_for_child_stdout(child: Child) -> String {
let output = child.wait_with_output().expect("failed to wait for child");
assert!(output.status.success());
let stdout = String::from_utf8(output.stdout).unwrap();
stdout
}
fn wait_for_child_failure_stderr(child: Child) -> String {
let output = child.wait_with_output().expect("failed to wait for child");
assert!(!output.status.success());
let stderr = String::from_utf8(output.stderr).unwrap();
stderr
}
#[test]
fn fork_basically_works() {
let status = fork_int(
"fork::test::fork_basically_works",
fork_id!(),
|_| (),
supervise_child,
|| println!("hello from child"),
)
.unwrap();
assert_eq!(status, ExitCode::SUCCESS);
}
#[test]
fn child_output_captured_and_repeated() {
let output = fork_int(
"fork::test::child_output_captured_and_repeated",
fork_id!(),
|_| (),
wait_for_child_stdout,
|| {
fork_int(
"fork::test::child_output_captured_and_repeated",
fork_id!(),
|_| (),
supervise_child,
|| println!("hello from child"),
)
.unwrap()
},
)
.unwrap();
assert!(output.contains("hello from child"));
}
#[test]
fn child_error_output() {
let output = fork_int(
"fork::test::child_error_output",
fork_id!(),
|_| (),
wait_for_child_failure_stderr,
|| {
fork_int(
"fork::test::child_error_output",
fork_id!(),
|_| (),
supervise_child,
|| io::Result::<()>::Err(io::Error::other("induced error")),
)
.unwrap()
},
)
.unwrap();
assert!(output.contains("induced error"));
}
#[test]
fn child_aborted_if_panics() {
let status = fork_int::<_, _, _, _, ()>(
"fork::test::child_aborted_if_panics",
fork_id!(),
|_| (),
|mut child| child.wait().unwrap(),
|| panic!("testing a panic, nothing to see here"),
)
.unwrap();
assert_eq!(70, status.code().unwrap());
}
#[test]
fn child_failure_reported() {
let status = fork_int::<_, _, _, _, ()>(
"fork::test::child_failure_reported",
fork_id!(),
|_| (),
supervise_child,
|| abort(),
)
.unwrap();
assert_eq!(status, ExitCode::FAILURE);
}
#[test]
fn data_exchange() {
let mut data = [1, 2, 3, 4, 5];
let status = fork_in_out(
fork_id!(),
"fork::test::data_exchange",
|data| {
assert_eq!(data.len(), 5);
let () = data.iter_mut().for_each(|x| *x += 1);
},
data.as_mut_slice(),
)
.unwrap();
assert_eq!(status, ExitCode::SUCCESS);
assert_eq!(data, [2, 3, 4, 5, 6]);
}
#[test]
fn run_should_panic_accepts_panic() {
let code = run_should_panic(|| panic!("boom"), None);
assert_eq!(code, ExitCode::SUCCESS);
}
#[test]
fn run_should_panic_rejects_missing_panic() {
let code = run_should_panic(|| {}, None);
assert_eq!(code, ExitCode::FAILURE);
}
#[test]
fn run_should_panic_accepts_expected_message() {
let code = run_should_panic(|| panic!("a boom occurred"), Some("boom"));
assert_eq!(code, ExitCode::SUCCESS);
}
#[test]
fn run_should_panic_rejects_unexpected_message() {
let code = run_should_panic(|| panic!("something else"), Some("boom"));
assert_eq!(code, ExitCode::FAILURE);
}
}