use shimforge::{Session, mock};
use std::sync::{Mutex, MutexGuard};
static TEST_LOCK: Mutex<()> = Mutex::new(());
fn serial() -> MutexGuard<'static, ()> {
TEST_LOCK.lock().unwrap_or_else(|error| error.into_inner())
}
extern "C" fn native_sum(a: i64, b: i64) -> i64 {
a + b
}
unsafe fn unchecked_sum(a: i64, b: i64) -> i64 {
a + b
}
extern "system" fn system_sum(a: i64, b: i64) -> i64 {
a + b
}
extern "C-unwind" fn unwind_sum(a: i64, b: i64) -> i64 {
a + b
}
#[test]
fn native_calls_match_arguments_and_count_captured_responses() {
let _serial = serial();
let mut session = Session::new_global();
let sum = mock!(session, native_sum, extern "C" fn(i64, i64) -> i64);
let mut results = vec![42, 24].into_iter();
let count = sum
.expect()
.with(|a, b| *a == 6 && *b == 7)
.times(2)
.returning(move |_, _| results.next().unwrap());
assert_eq!(native_sum(6, 7), 42);
assert_eq!(native_sum(6, 7), 24);
assert_eq!(count.calls(), 2);
session.restore();
assert_eq!(native_sum(6, 7), 13);
}
#[test]
fn unsafe_and_system_functions_accept_expectations() {
let _serial = serial();
let mut session = Session::new_global();
let sum = mock!(session, unchecked_sum, unsafe fn(i64, i64) -> i64);
sum.expect().once().returns(42);
assert_eq!(unsafe { unchecked_sum(6, 7) }, 42);
let system = mock!(session, system_sum, extern "system" fn(i64, i64) -> i64);
system.expect().once().returning(|a, b| a * b);
assert_eq!(system_sum(6, 7), 42);
session.verify();
}
#[test]
fn unwind_abi_keeps_rust_panic_behavior() {
let _serial = serial();
let mut session = Session::new_global();
let sum = mock!(session, unwind_sum, extern "C-unwind" fn(i64, i64) -> i64);
sum.expect().once().panics("chosen failure");
assert!(std::panic::catch_unwind(|| unwind_sum(6, 7)).is_err());
session.restore();
assert_eq!(unwind_sum(6, 7), 13);
}
#[test]
fn native_panic_does_not_cross_the_abi_boundary() {
const CHILD: &str = "SHIMFORGE_NATIVE_PANIC_CHILD";
if std::env::var_os(CHILD).is_some() {
let mut session = Session::new_global();
let sum = mock!(session, native_sum, extern "C" fn(i64, i64) -> i64);
sum.expect().panics("native callback failed");
native_sum(1, 2);
std::process::exit(99);
}
let output = std::process::Command::new(std::env::current_exe().unwrap())
.args([
"--exact",
"native_panic_does_not_cross_the_abi_boundary",
"--nocapture",
])
.env(CHILD, "1")
.output()
.unwrap();
assert!(!output.status.success());
assert_ne!(output.status.code(), Some(99));
let error = String::from_utf8_lossy(&output.stderr);
assert!(error.contains("native callback failed"), "{error}");
assert!(error.contains("cannot unwind"), "{error}");
}
#[cfg(any(target_os = "linux", target_os = "macos"))]
#[test]
fn imported_system_function_is_mocked_without_a_wrapper() {
let _serial = serial();
for constructor in [Session::new_global, Session::new] {
let mut session = constructor();
let hostname = mock!(
session,
libc::gethostname,
unsafe extern "C" fn(*mut libc::c_char, usize) -> libc::c_int
);
hostname
.expect()
.with(|_, size| *size == 64)
.once()
.returning(|buffer, size| {
let name = b"test-host\0";
assert!(size >= name.len());
unsafe { std::ptr::copy_nonoverlapping(name.as_ptr().cast(), buffer, name.len()) };
0
});
let mut buffer = [0u8; 64];
let result = unsafe { libc::gethostname(buffer.as_mut_ptr().cast(), buffer.len()) };
assert_eq!(result, 0);
assert_eq!(&buffer[..10], b"test-host\0");
session.restore();
}
}
#[cfg(target_os = "windows")]
#[test]
fn imported_system_function_is_mocked_without_a_wrapper() {
#[link(name = "kernel32")]
unsafe extern "system" {
fn GetComputerNameW(buffer: *mut u16, size: *mut u32) -> i32;
}
let _serial = serial();
for constructor in [Session::new_global, Session::new] {
let mut session = constructor();
let hostname = mock!(
session,
GetComputerNameW,
unsafe extern "system" fn(*mut u16, *mut u32) -> i32
);
hostname.expect().once().returning(|buffer, size| {
let name: Vec<_> = "test-host\0".encode_utf16().collect();
unsafe {
assert!(*size as usize >= name.len());
std::ptr::copy_nonoverlapping(name.as_ptr(), buffer, name.len());
*size = (name.len() - 1) as u32;
}
1
});
let mut buffer = [0u16; 64];
let mut size = buffer.len() as u32;
let result = unsafe { GetComputerNameW(buffer.as_mut_ptr(), &mut size) };
assert_eq!(result, 1);
assert_eq!(
String::from_utf16(&buffer[..size as usize]).unwrap(),
"test-host"
);
session.restore();
}
}