From a30dba668c9abbdf263872575fca23f6e2372d57 Mon Sep 17 00:00:00 2001 From: tanglearncode Date: Sun, 13 Sep 2026 13:36:53 +0800 Subject: [PATCH] Panic on failure instead of returning Result shimforge is for writing tests, where a failed setup should just fail the test. Returning Result made every test add .unwrap(), or ? with a final Ok(()). With ?, the drop check of a pending expectation could also replace the real setup error with "expected at least 1 calls, observed 0". Session::new, new_local and new_global, mock!, replace!, mock_async, replace_raw, every expectation response (returns, returning, never and the rest), and verify, checkpoint and restore now return their values directly and panic with the error's message on failure. They and the generated mock methods use #[track_caller], so a failure points at the test line. try_new_local and try_new_global still return Result, because reporting a busy session is their purpose. Internally the installation, rule and verification code still returns Result, and a hidden check helper turns an error into a panic at the public boundary. restore still removes every mock before it panics, so the drop check does not panic a second time. Tests drop their unwrap calls and ? operators. Assertions that expected an error use a panic_message helper and compare the message with the error's Display text. The README examples and tests/readme.rs no longer return Result. The compile_fail doctests no longer call unwrap, so they still fail for the reason they document. Checked locally with cargo fmt --check, cargo clippy --workspace --all-targets -- -D warnings on x86_64-pc-windows-msvc, aarch64-pc-windows-msvc, aarch64-unknown-linux-gnu, aarch64-apple-darwin and x86_64-apple-darwin, and cargo test --lib (69 passed). Integration tests and doctests run in CI. --- README.md | 152 +++++----- macros/src/lib.rs | 42 +-- src/asynchronous.rs | 69 +++-- src/asynchronous/tests.rs | 22 +- src/error.rs | 12 + src/expectation.rs | 9 +- src/expectation/tests.rs | 12 +- src/lib.rs | 137 +++++---- src/memory/tests.rs | 51 ++-- src/routing/tests.rs | 32 +-- src/tests.rs | 79 ++++-- tests/async_expectations.rs | 278 ++++++++---------- tests/async_network.rs | 25 +- tests/cruntime.rs | 32 +-- tests/expectations.rs | 534 +++++++++++++++-------------------- tests/filesystem.rs | 103 +++---- tests/generics.rs | 44 ++- tests/http_client.rs | 31 +- tests/io.rs | 52 ++-- tests/native_expectations.rs | 76 +++-- tests/parallel.rs | 58 ++-- tests/readme.rs | 140 +++++---- tests/replacement.rs | 195 +++++++------ tests/safe_api.rs | 6 +- tests/stack_safety.rs | 20 +- tests/thread_local.rs | 206 +++++++------- 26 files changed, 1158 insertions(+), 1259 deletions(-) diff --git a/README.md b/README.md index 6ea9d87..fb030b0 100644 --- a/README.md +++ b/README.md @@ -35,22 +35,21 @@ use shimforge::{Session, mock}; use std::io; #[test] -fn claim_slot_succeeds_without_a_real_directory() -> Result<(), shimforge::Error> { - let mut session = Session::new()?; +fn claim_slot_succeeds_without_a_real_directory() { + let mut session = Session::new(); let create = mock!( session, fs::create_dir_all::<&str>, fn(&str) -> io::Result<()> - )?; + ); create .expect() .with(|path| *path == "/var/run/dispatcher") .once() - .returning(|_| Ok(()))?; + .returning(|_| Ok(())); assert!(claim_slot().is_ok()); - session.verify()?; - Ok(()) + session.verify(); } ``` @@ -93,8 +92,8 @@ Import the two macros and the session type, then run `cargo test` as usual: use shimforge::{Session, mock, replace}; ``` -Every example below is a complete test. It returns `Result<(), shimforge::Error>`, -so `?` turns a shimforge error into a test failure. +Every example below is a complete test. shimforge panics when a mock cannot be +installed or an expectation is not met, so there is no `Result` to handle. ## Thread-local and global sessions @@ -114,15 +113,14 @@ fn worker_count() -> usize { } #[test] -fn other_threads_keep_the_original_function() -> Result<(), shimforge::Error> { - let mut session = Session::new()?; - let count = mock!(session, worker_count, fn() -> usize)?; - count.expect().returns(16)?; +fn other_threads_keep_the_original_function() { + let mut session = Session::new(); + let count = mock!(session, worker_count, fn() -> usize); + count.expect().returns(16); assert_eq!(worker_count(), 16); // A thread with no session of its own still calls the original. assert_eq!(std::thread::spawn(worker_count).join().unwrap(), 2); - Ok(()) } ``` @@ -144,15 +142,14 @@ fn export_state(marker: &Path) -> &'static str { } #[test] -fn a_constant_result_answers_every_call() -> Result<(), shimforge::Error> { - let mut session = Session::new()?; - let exists = mock!(session, Path::exists, fn(&Path) -> bool)?; - exists.expect().returns(true)?; +fn a_constant_result_answers_every_call() { + let mut session = Session::new(); + let exists = mock!(session, Path::exists, fn(&Path) -> bool); + exists.expect().returns(true); assert_eq!(export_state(Path::new("virtual/export.done")), "finished"); - session.restore()?; + session.restore(); assert_eq!(export_state(Path::new("virtual/export.done")), "running"); - Ok(()) } ``` @@ -174,22 +171,21 @@ fn load_port(path: &Path) -> io::Result { } #[test] -fn matching_calls_are_counted() -> Result<(), shimforge::Error> { - let mut session = Session::new()?; +fn matching_calls_are_counted() { + let mut session = Session::new(); let read = mock!( session, fs::read_to_string::<&Path>, fn(&Path) -> io::Result - )?; + ); read.expect() .with(|path| *path == Path::new("service.port")) .times(2) - .returning(|_| Ok("8080\n".to_owned()))?; + .returning(|_| Ok("8080\n".to_owned())); assert_eq!(load_port(Path::new("service.port")).unwrap(), 8080); assert_eq!(load_port(Path::new("service.port")).unwrap(), 8080); - session.verify()?; - Ok(()) + session.verify(); } ``` @@ -231,10 +227,10 @@ fn save_id(id: u64) -> bool { } #[test] -fn results_follow_the_call_order() -> Result<(), shimforge::Error> { - let mut session = Session::new()?; - let next = mock!(session, next_id, fn() -> u64)?; - let save = mock!(session, save_id, fn(u64) -> bool)?; +fn results_follow_the_call_order() { + let mut session = Session::new(); + let next = mock!(session, next_id, fn() -> u64); + let save = mock!(session, save_id, fn(u64) -> bool); let order = Sequence::new(); let mut id = 40; next.expect() @@ -243,18 +239,17 @@ fn results_follow_the_call_order() -> Result<(), shimforge::Error> { .returning(move || { id += 1; id - })?; + }); save.expect() .with(|id| *id == 42) .once() .in_sequence(&order) - .returns(true)?; + .returns(true); assert_eq!(next_id(), 41); assert_eq!(next_id(), 42); assert!(save_id(42)); - session.verify()?; - Ok(()) + session.verify(); } ``` @@ -277,9 +272,9 @@ fn fill(buffer: &mut [u8]) -> usize { } #[test] -fn a_mock_fills_output_parameters() -> Result<(), shimforge::Error> { - let mut session = Session::new()?; - let split = mock!(session, split_amount, fn(u64, &mut u64, &mut u64))?; +fn a_mock_fills_output_parameters() { + let mut session = Session::new(); + let split = mock!(session, split_amount, fn(u64, &mut u64, &mut u64)); split .expect() .with(|total, _, _| *total == 1234) @@ -287,12 +282,12 @@ fn a_mock_fills_output_parameters() -> Result<(), shimforge::Error> { .returning(|_, whole, cents| { *whole = 99; *cents = 5; - })?; - let write = mock!(session, fill, fn(&mut [u8]) -> usize)?; + }); + let write = mock!(session, fill, fn(&mut [u8]) -> usize); write.expect().once().returning(|buffer| { buffer[..2].copy_from_slice(b"ok"); 2 - })?; + }); let (mut whole, mut cents) = (0, 0); split_amount(1234, &mut whole, &mut cents); @@ -301,7 +296,6 @@ fn a_mock_fills_output_parameters() -> Result<(), shimforge::Error> { let mut buffer = [0; 8]; assert_eq!(fill(&mut buffer), 2); assert_eq!(&buffer[..2], b"ok"); - Ok(()) } ``` @@ -331,16 +325,16 @@ fn render(value: T) -> String { } #[test] -fn a_method_and_one_generic_instance_are_mocked() -> Result<(), shimforge::Error> { - let mut session = Session::new()?; - let rates = mock!(session, Cache::hit_rate, fn(&Cache, &str) -> f32)?; +fn a_method_and_one_generic_instance_are_mocked() { + let mut session = Session::new(); + let rates = mock!(session, Cache::hit_rate, fn(&Cache, &str) -> f32); rates .expect() .with(|cache, key| cache.region == "eu" && *key == "sessions") .once() - .returns(0.75)?; - let rendered = mock!(session, render::, fn(u8) -> String)?; - rendered.expect().once().returns("mocked".to_owned())?; + .returns(0.75); + let rendered = mock!(session, render::, fn(u8) -> String); + rendered.expect().once().returns("mocked".to_owned()); let cache = Cache { region: "eu".to_owned(), @@ -349,7 +343,6 @@ fn a_method_and_one_generic_instance_are_mocked() -> Result<(), shimforge::Error assert_eq!(render(7u8), "mocked"); // A different type argument is a different function. assert_eq!(render("7"), "live 7"); - Ok(()) } ``` @@ -371,15 +364,14 @@ fn fixed_checksum(_bytes: &[u8]) -> u32 { } #[test] -fn a_function_or_closure_replaces_the_original() -> Result<(), shimforge::Error> { - let mut session = Session::new()?; - replace!(session, checksum => fixed_checksum, fn(&[u8]) -> u32)?; +fn a_function_or_closure_replaces_the_original() { + let mut session = Session::new(); + replace!(session, checksum => fixed_checksum, fn(&[u8]) -> u32); assert_eq!(checksum(b"abc"), 7); - session.restore()?; + session.restore(); - replace!(session, checksum => |_| 9, fn(&[u8]) -> u32)?; + replace!(session, checksum => |_| 9, fn(&[u8]) -> u32); assert_eq!(checksum(b"abc"), 9); - Ok(()) } ``` @@ -424,26 +416,25 @@ fn ready(future: F) -> F::Output { } #[test] -fn async_functions_return_mocked_results() -> Result<(), shimforge::Error> { - let mut session = Session::new()?; - let rates = session.mock_async(exchange_rate(""))?; - rates.expect().once().returns(1.25)?; +fn async_functions_return_mocked_results() { + let mut session = Session::new(); + let rates = session.mock_async(exchange_rate("")); + rates.expect().once().returns(1.25); // A throwaway receiver is enough to name the future type. let balances = session.mock_async( Ledger { name: String::new(), } .balance(), - )?; - balances.expect().once().returns(4_200)?; + ); + balances.expect().once().returns(4_200); assert_eq!(ready(exchange_rate("EURUSD")), 1.25); let ledger = Ledger { name: "payroll".to_owned(), }; assert_eq!(ready(ledger.balance()), 4_200); - session.verify()?; - Ok(()) + session.verify(); } ``` @@ -482,17 +473,17 @@ impl Client { } #[test] -fn a_client_method_that_returns_a_boxed_future_is_mocked() -> Result<(), shimforge::Error> { - let mut session = Session::new()?; +fn a_client_method_that_returns_a_boxed_future_is_mocked() { + let mut session = Session::new(); let get = mock!( session, Client::get, for<'a> fn(&'a Client, &'a str) -> Call<'a> - )?; + ); get.expect() .with(|_, path| *path == "/health") .once() - .returning(|_, _| Box::pin(async { Ok("healthy".to_owned()) }))?; + .returning(|_, _| Box::pin(async { Ok("healthy".to_owned()) })); let client = Client { endpoint: "https://inventory.invalid".to_owned(), @@ -503,8 +494,7 @@ fn a_client_method_that_returns_a_boxed_future_is_mocked() -> Result<(), shimfor response.as_mut().poll(&mut context), Poll::Ready(Ok(body)) if body == "healthy" )); - session.verify()?; - Ok(()) + session.verify(); } ``` @@ -524,13 +514,13 @@ unsafe extern "C" { } #[test] -fn getenv_reports_a_mocked_variable() -> Result<(), shimforge::Error> { - let mut session = Session::new()?; +fn getenv_reports_a_mocked_variable() { + let mut session = Session::new(); let lookup = mock!( session, getenv, unsafe extern "C" fn(*const c_char) -> *mut c_char - )?; + ); lookup .expect() .with(|name| { @@ -539,15 +529,14 @@ fn getenv_reports_a_mocked_variable() -> Result<(), shimforge::Error> { name == c"DEPLOY_SLOT" }) .once() - .returning(|_| c"canary".as_ptr().cast_mut())?; + .returning(|_| c"canary".as_ptr().cast_mut()); // Any other variable keeps reporting that it is unset. - lookup.expect().returning(|_| std::ptr::null_mut())?; + lookup.expect().returning(|_| std::ptr::null_mut()); let key = CString::new("DEPLOY_SLOT").unwrap(); // SAFETY: the key is a valid C string and the result is only read. let slot = unsafe { CStr::from_ptr(getenv(key.as_ptr())) }; assert_eq!(slot.to_str().unwrap(), "canary"); - Ok(()) } ``` @@ -572,16 +561,15 @@ fn fake_slot_count() -> usize { } #[test] -fn a_raw_replacement_swaps_one_function_for_another() -> Result<(), shimforge::Error> { - let mut session = Session::new_global()?; +fn a_raw_replacement_swaps_one_function_for_another() { + let mut session = Session::new_global(); // SAFETY: both functions are live, share a signature, and stay loaded until the // session restores them. No thread calls them while the patch is installed. - unsafe { session.replace_raw(slot_count as *const (), fake_slot_count as *const ())? }; + unsafe { session.replace_raw(slot_count as *const (), fake_slot_count as *const ()) }; assert_eq!(slot_count(), 64); - session.restore()?; + session.restore(); assert_eq!(slot_count(), 4); - Ok(()) } ``` @@ -594,10 +582,10 @@ second panic. `session.restore()` removes the mocks early and checks them, and Each thread may hold one session. Local sessions may coexist on different threads, even for the same function. Global sessions wait for other sessions to finish, and -local sessions wait for an active global session. A nested session returns -`Error::Busy`; `try_new_local()` and `try_new_global()` return `Error::Busy` -instead of waiting. Both modes support `mock!`, `replace!`, and `mock_async`; -`replace_raw` requires a global session. +local sessions wait for an active global session. Opening a nested session panics; +`try_new_local()` and `try_new_global()` return `Err(Error::Busy)` instead of +waiting. Both modes support `mock!`, `replace!`, and `mock_async`; `replace_raw` +requires a global session. Install a local mock before calls to that target start. Only the first local installation changes code; later installs and cleanup leave the entry in place so diff --git a/macros/src/lib.rs b/macros/src/lib.rs index 613d7ff..e8367fc 100644 --- a/macros/src/lib.rs +++ b/macros/src/lib.rs @@ -86,7 +86,7 @@ fn expand_replacement(input: syn::Result) -> Tokens { #root::__install_replacement!(session, source as *const (), __dispatch::<__Site, #(#types,)* __Return> as *const (), target as *const ()) } - __install(#session, #source, #target) + #root::__private::check(__install(#session, #source, #target)) } } } @@ -224,7 +224,8 @@ fn generate(input: Input) -> Tokens { quote! { impl<__Output> __ShimforgeBuilder<__Output> where __Output: ::std::marker::Send + 'static + ::std::convert::Into<#output> { - pub fn returns(self, value: __Output) -> ::std::result::Result<#root::Expectation, #root::Error> + #[track_caller] + pub fn returns(self, value: __Output) -> #root::Expectation where __Output: ::std::clone::Clone { self.returning(move |#(#names),*| { let _ = (#(#names),*); @@ -232,14 +233,16 @@ fn generate(input: Input) -> Tokens { }) } - pub fn return_once(self, value: __Output) -> ::std::result::Result<#root::Expectation, #root::Error> { + #[track_caller] + pub fn return_once(self, value: __Output) -> #root::Expectation { self.returning_once(move |#(#names),*| { let _ = (#(#names),*); value.into() }) } - pub fn returns_default(self) -> ::std::result::Result<#root::Expectation, #root::Error> + #[track_caller] + pub fn returns_default(self) -> #root::Expectation where __Output: ::std::default::Default { self.returning(|#(#names),*| { let _ = (#(#names),*); @@ -321,9 +324,11 @@ fn generate(input: Input) -> Tokens { } } - pub fn verify(&self) -> ::std::result::Result<(), #root::Error> { self.state.verify() } + #[track_caller] + pub fn verify(&self) { #root::__private::check(self.state.verify()) } - pub fn checkpoint(&self) -> ::std::result::Result<(), #root::Error> { self.state.checkpoint() } + #[track_caller] + pub fn checkpoint(&self) { #root::__private::check(self.state.checkpoint()) } } struct __ShimforgeBuilder<__Output> { @@ -355,29 +360,34 @@ fn generate(input: Input) -> Tokens { self } - pub fn returning<__Action>(self, action: __Action) -> ::std::result::Result<#root::Expectation, #root::Error> + #[track_caller] + pub fn returning<__Action>(self, action: __Action) -> #root::Expectation where __Action: #binder ::std::ops::FnMut(#(#args),*) -> #output + ::std::marker::Send + 'static { - self.state.add(self.config, |meta| __ShimforgeRule { + #root::__private::check(self.state.add(self.config, |meta| __ShimforgeRule { meta, matcher: self.matcher, action: ::std::sync::Mutex::new(__ShimforgeAction::Repeat(::std::boxed::Box::new(action))), - }) + })) } - pub fn returning_once<__Action>(self, action: __Action) -> ::std::result::Result<#root::Expectation, #root::Error> + #[track_caller] + pub fn returning_once<__Action>(self, action: __Action) -> #root::Expectation where __Action: #binder ::std::ops::FnOnce(#(#args),*) -> #output + ::std::marker::Send + 'static { - self.state.add(self.config.for_once()?, |meta| __ShimforgeRule { + let config = #root::__private::check(self.config.for_once()); + #root::__private::check(self.state.add(config, |meta| __ShimforgeRule { meta, matcher: self.matcher, action: ::std::sync::Mutex::new(__ShimforgeAction::Once(::std::option::Option::Some(::std::boxed::Box::new(action)))), - }) + })) } - pub fn never(self) -> ::std::result::Result<#root::Expectation, #root::Error> { + #[track_caller] + pub fn never(self) -> #root::Expectation { self.times(0usize).panics("forbidden mock call") } - pub fn panics(self, message: impl ::std::convert::Into<::std::string::String>) -> ::std::result::Result<#root::Expectation, #root::Error> { + #[track_caller] + pub fn panics(self, message: impl ::std::convert::Into<::std::string::String>) -> #root::Expectation { let message = message.into(); self.returning(move |#(#names),*| { let _ = (#(#names),*); @@ -388,7 +398,7 @@ fn generate(input: Input) -> Tokens { #constants - (|| -> ::std::result::Result<__ShimforgeMock, #root::Error> { + #root::__private::check((|| -> ::std::result::Result<__ShimforgeMock, #root::Error> { let __shimforge_original = #source; #signature_check let __shimforge_source = __shimforge_original as #unsafety #abi fn(#(#infer),*) -> _; @@ -423,7 +433,7 @@ fn generate(input: Input) -> Tokens { }); #root::__install!(__shimforge_session, __shimforge_source as *const (), __shimforge_target as *const (), __shimforge_state.clone(), __shimforge_detach)?; ::std::result::Result::Ok(__ShimforgeMock { state: __shimforge_state }) - })() + })()) }} } diff --git a/src/asynchronous.rs b/src/asynchronous.rs index 94984cd..2fc5a13 100644 --- a/src/asynchronous.rs +++ b/src/asynchronous.rs @@ -1,3 +1,4 @@ +use crate::error::check; use crate::expectation::{Config, Control, Meta, Rule, State, lock}; use crate::{CallCount, Error, Expectation, Sequence, Session}; use std::any::Any; @@ -50,7 +51,12 @@ impl Session { /// The future's layout and drop code stay unchanged. /// Do not pass a boxed trait object: its poll wrapper is shared by unrelated /// futures. Mock the method returning that box with [`crate::mock!`] instead. - pub fn mock_async(&mut self, witness: F) -> Result, Error> + /// + /// # Panics + /// + /// Panics if this future type is already mocked or its poll method cannot be patched. + #[track_caller] + pub fn mock_async(&mut self, witness: F) -> AsyncMock where F: Future, F::Output: Send + 'static, @@ -59,18 +65,18 @@ impl Session { drop(witness); let state = State::new(std::any::type_name::()); let key = (source as usize, self.__thread()); - register(key, state.clone())?; + check(register(key, state.clone())); let control: Arc = state.clone(); // SAFETY: both poll functions use F's exact signature and output type. - unsafe { + check(unsafe { self.__install( source, poll_mock:: as *const (), control, Box::new(move || remove(key)), - )?; - } - Ok(AsyncMock { state }) + ) + }); + AsyncMock { state } } } @@ -170,14 +176,16 @@ impl AsyncMock { } } - /// Checks call counts and reports unexpected calls. - pub fn verify(&self) -> Result<(), Error> { - self.state.verify() + /// Checks call counts and unexpected calls, and panics if either failed. + #[track_caller] + pub fn verify(&self) { + check(self.state.verify()); } - /// Checks expectations, then clears them for the next phase. - pub fn checkpoint(&self) -> Result<(), Error> { - self.state.checkpoint() + /// Checks expectations like [`Self::verify`], then clears them for the next phase. + #[track_caller] + pub fn checkpoint(&self) { + check(self.state.checkpoint()); } } @@ -201,28 +209,29 @@ impl AsyncExpectation { } /// Calls a closure for each response. The closure may own captured values. - pub fn returning( - self, - action: impl FnMut() -> R + Send + 'static, - ) -> Result { - self.state.add(self.config, |meta| AsyncRule { + /// + /// # Panics + /// + /// Panics if the call count is invalid or the mock is no longer active. + #[track_caller] + pub fn returning(self, action: impl FnMut() -> R + Send + 'static) -> Expectation { + check(self.state.add(self.config, |meta| AsyncRule { meta, action: Mutex::new(Box::new(action)), - }) + })) } /// Calls a closure at most once, moving its captured values if needed. - pub fn returning_once( - mut self, - action: impl FnOnce() -> R + Send + 'static, - ) -> Result { - self.config = self.config.for_once()?; + #[track_caller] + pub fn returning_once(mut self, action: impl FnOnce() -> R + Send + 'static) -> Expectation { + self.config = check(self.config.for_once()); let mut action = Some(action); self.returning(move || action.take().expect("one-time response was already used")()) } /// Clones the value for each response. - pub fn returns(self, value: R) -> Result + #[track_caller] + pub fn returns(self, value: R) -> Expectation where R: Clone, { @@ -230,12 +239,14 @@ impl AsyncExpectation { } /// Moves the value into one response. It need not implement `Clone`. - pub fn return_once(self, value: R) -> Result { + #[track_caller] + pub fn return_once(self, value: R) -> Expectation { self.returning_once(move || value) } /// Creates a default value for each response. - pub fn returns_default(self) -> Result + #[track_caller] + pub fn returns_default(self) -> Expectation where R: Default, { @@ -243,13 +254,15 @@ impl AsyncExpectation { } /// Panics with this message when called. - pub fn panics(self, message: impl Into) -> Result { + #[track_caller] + pub fn panics(self, message: impl Into) -> Expectation { let message = message.into(); self.returning(move || panic!("{message}")) } /// Rejects every poll of this future type. - pub fn never(self) -> Result { + #[track_caller] + pub fn never(self) -> Expectation { self.times(0).panics("poll is forbidden") } } diff --git a/src/asynchronous/tests.rs b/src/asynchronous/tests.rs index 5dc7d83..36d1716 100644 --- a/src/asynchronous/tests.rs +++ b/src/asynchronous/tests.rs @@ -13,25 +13,21 @@ fn poll(seed: u64) -> Poll { fn local_async_duplicates_and_original_polls_are_checked() { let _serial = crate::tests::serial(); { - let mut session = Session::new_global().unwrap(); - session - .mock_async(value(1)) - .unwrap() - .expect() - .once() - .returns(30) - .unwrap(); + let mut session = Session::new_global(); + session.mock_async(value(1)).expect().once().returns(30); assert_eq!(poll(1), Poll::Ready(30)); } - let mut session = Session::new_local().unwrap(); - let mock = session.mock_async(value(1)).unwrap(); - mock.expect().once().returns(40).unwrap(); - assert!(session.mock_async(value(1)).is_err()); + let mut session = Session::new_local(); + let mock = session.mock_async(value(1)); + mock.expect().once().returns(40); + assert!( + crate::tests::panic_message(|| session.mock_async(value(1))).contains("already mocked") + ); assert_eq!(poll(1), Poll::Ready(40)); assert_eq!( std::thread::spawn(|| poll(1)).join().unwrap(), Poll::Ready(2) ); - session.restore().unwrap(); + session.restore(); assert_eq!(poll(1), Poll::Ready(2)); } diff --git a/src/error.rs b/src/error.rs index 0ff65dd..aaea1b1 100644 --- a/src/error.rs +++ b/src/error.rs @@ -62,3 +62,15 @@ impl fmt::Display for Error { } impl std::error::Error for Error {} + +/// Returns the value, or panics with the error's message. +/// +/// The panic reports the caller's location, so a failed setup points at the test. +#[doc(hidden)] +#[track_caller] +pub fn check(result: Result) -> T { + match result { + Ok(value) => value, + Err(error) => panic!("{error}"), + } +} diff --git a/src/expectation.rs b/src/expectation.rs index 7cef4b1..24b0fcb 100644 --- a/src/expectation.rs +++ b/src/expectation.rs @@ -176,8 +176,13 @@ impl Expectation { } /// Checks the call count and any order failure. - pub fn verify(&self) -> Result<(), Error> { - self.meta.verify() + /// + /// # Panics + /// + /// Panics if the expectation missed calls or received a call it rejected. + #[track_caller] + pub fn verify(&self) { + crate::error::check(self.meta.verify()); } } diff --git a/src/expectation/tests.rs b/src/expectation/tests.rs index e8aa3ba..65ba1ab 100644 --- a/src/expectation/tests.rs +++ b/src/expectation/tests.rs @@ -86,10 +86,10 @@ fn response_chains_skip_exhausted_rules() { let first = add(&state, Config::default().once(), 10); let next = add(&state, Config::default().times(1..=2), 20); let other = add(&state, Config::default(), 30); - assert!(first.verify().is_err()); + assert!(first.meta.verify().is_err()); assert!(state.verify().is_err()); assert_eq!(state.select(&|rule| rule.value != 30).value, 10); - first.verify().unwrap(); + first.verify(); assert_eq!(first.clone().calls(), 1); for _ in 0..2 { assert_eq!(state.select(&|rule| rule.value != 30).value, 20); @@ -113,11 +113,12 @@ fn missing_and_forbidden_calls_stay_failed() { let state = State::new("write"); assert!(panics(|| drop(state.select(&|_| true))).contains("no expectation")); let never = add(&state, Config::default().times(0), 1); - never.verify().unwrap(); + never.verify(); assert!(panics(|| drop(state.select(&|_| true))).contains("forbidden")); assert_eq!(never.calls(), 0); assert!( never + .meta .verify() .unwrap_err() .to_string() @@ -165,6 +166,7 @@ fn out_of_order_calls_stay_failed() { state.select(&|rule| rule.value == 2); assert!( later + .meta .verify() .unwrap_err() .to_string() @@ -234,7 +236,7 @@ fn count_overflow_is_reported_without_wrapping() { lock(&expectation.meta.progress).calls = usize::MAX; assert!(panics(|| drop(state.select(&|_| true))).contains("overflow")); assert_eq!(expectation.calls(), usize::MAX); - assert!(expectation.verify().is_err()); + assert!(expectation.meta.verify().is_err()); } #[test] @@ -254,7 +256,7 @@ fn checkpoint_clears_only_completed_expectations() { state.checkpoint().unwrap(); assert_eq!(dropped.load(Ordering::SeqCst), 1); assert_eq!(expectation.calls(), 1); - expectation.verify().unwrap(); + expectation.verify(); state.checkpoint().unwrap(); add(&state, Config::default(), 2); assert_eq!(state.select(&|_| true).value, 2); diff --git a/src/lib.rs b/src/lib.rs index c5cc689..1299c02 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -31,6 +31,7 @@ pub use shimforge_macros::{__check_signature, __mock, __replace_local}; #[doc(hidden)] pub mod __private { + pub use crate::error::check; pub use crate::expectation::{CallGuard, Config, Control, Meta, Rule, State, lock}; pub use crate::routing::route; } @@ -98,14 +99,22 @@ impl Session { } /// Opens a thread-local session. See [`Self::new_local`]. - pub fn new() -> Result { + // Opening a session takes a lock and may wait or panic, so it is not a default value. + #[allow(clippy::new_without_default)] + #[track_caller] + pub fn new() -> Self { Self::new_local() } /// Opens a global session. Mocks affect all threads. - /// Waits for other threads' sessions to end. Nested sessions return [`Error::Busy`]. - pub fn new_global() -> Result { - Self::open(false, true) + /// Waits for other threads' sessions to end. + /// + /// # Panics + /// + /// Panics if this thread already has a session. + #[track_caller] + pub fn new_global() -> Self { + error::check(Self::open(false, true)) } /// Opens a session whose mocks affect only the current thread. @@ -113,16 +122,23 @@ impl Session { /// Only one local session may be active per thread. Global sessions exclude /// local sessions, so this waits for them to end. Stop target calls during /// the first installation. Later local installs and cleanup do not patch code. - pub fn new_local() -> Result { - Self::open(true, true) + /// + /// # Panics + /// + /// Panics if this thread already has a session. + #[track_caller] + pub fn new_local() -> Self { + error::check(Self::open(true, true)) } /// Opens a local session without waiting for an active global session. + /// Returns [`Error::Busy`] instead of waiting or panicking. pub fn try_new_local() -> Result { Self::open(true, false) } /// Opens a global session without waiting for other sessions. + /// Returns [`Error::Busy`] instead of waiting or panicking. pub fn try_new_global() -> Result { Self::open(false, false) } @@ -203,11 +219,18 @@ impl Session { /// Calls must reach the source entry. Inlined or merged calls cannot be isolated. /// Do not replace functions used by shimforge or its memory and OS code. /// The replacement must meet the safety rules its callers rely on. - pub unsafe fn replace_raw( - &mut self, - source: *const (), - target: *const (), - ) -> Result<(), Error> { + /// + /// # Panics + /// + /// Panics if the session is local or the replacement cannot be installed. + #[track_caller] + pub unsafe fn replace_raw(&mut self, source: *const (), target: *const ()) { + // SAFETY: the caller follows this method's safety contract. + error::check(unsafe { self.install_raw(source, target) }); + } + + /// Installs a raw replacement under the rules of [`Self::replace_raw`]. + unsafe fn install_raw(&mut self, source: *const (), target: *const ()) -> Result<(), Error> { if self.__thread().is_some() { return Err(Error::Expectation( "use mock! or mock_async in a local session".into(), @@ -271,7 +294,16 @@ impl Session { } /// Checks all call expectations without removing the mocks. - pub fn verify(&self) -> Result<(), Error> { + /// + /// # Panics + /// + /// Panics if a mock missed expected calls or received a call it rejected. + #[track_caller] + pub fn verify(&self) { + error::check(self.check_expectations()); + } + + fn check_expectations(&self) -> Result<(), Error> { for mock in &self.mocks { mock.control.verify()?; } @@ -300,7 +332,7 @@ impl Session { routing::install(source as usize, target as usize)?; self.local_patches.push(source as usize); } else { - self.replace_raw(source, target)?; + self.install_raw(source, target)?; } } self.mocks.push(mock); @@ -309,14 +341,19 @@ impl Session { /// Removes all mocks, then checks their expectations. /// - /// Stop global target calls as required by [`Self::replace_raw`]. On error, - /// a failed patch stays in the session so you can retry. Expectation errors - /// are returned after all patches have been removed. - pub fn restore(&mut self) -> Result<(), Error> { - self.restore_patches()?; - let result = self.verify(); + /// Stop global target calls as required by [`Self::replace_raw`]. + /// + /// # Panics + /// + /// Panics if original code cannot be restored. The failed patch stays in the + /// session, so a later `restore` or the drop retries it. Also panics, after all + /// mocks are removed, if an expectation failed. + #[track_caller] + pub fn restore(&mut self) { + error::check(self.restore_patches()); + let result = self.check_expectations(); self.detach(); - result + error::check(result); } fn detach(&mut self) { @@ -345,7 +382,7 @@ impl Session { impl Drop for Session { fn drop(&mut self) { finish(self.restore_patches(), std::process::abort); - let result = self.verify(); + let result = self.check_expectations(); self.detach(); if !std::thread::panicking() { if let Err(error) = result { @@ -359,79 +396,79 @@ impl Drop for Session { /// /// ``` /// fn read_count(key: &str) -> usize { key.len() } -/// let mut session = shimforge::Session::new()?; -/// let mock = shimforge::mock!(session, read_count, fn(&str) -> usize)?; -/// mock.expect().with(|key| *key == "orders").once().returns(12)?; +/// let mut session = shimforge::Session::new(); +/// let mock = shimforge::mock!(session, read_count, fn(&str) -> usize); +/// mock.expect().with(|key| *key == "orders").once().returns(12); /// assert_eq!(read_count("orders"), 12); -/// # Ok::<(), shimforge::Error>(()) /// ``` /// /// Matchers borrow arguments. Return closures may capture owned values and must /// be `Send + 'static`. Follow the same runtime safety rules as [`replace!`]. +/// Installation panics if the function cannot be patched. /// /// The source signature must match: /// ```compile_fail -/// let mut session = shimforge::Session::new_global().unwrap(); +/// let mut session = shimforge::Session::new_global(); /// fn source(value: u64) -> u64 { value } /// shimforge::mock!(session, source, fn(u32) -> u32); /// ``` /// A replacement cannot require a longer input borrow: /// ```compile_fail -/// let mut session = shimforge::Session::new_global().unwrap(); +/// let mut session = shimforge::Session::new_global(); /// fn source(value: &str) -> usize { value.len() } /// shimforge::mock!(session, source, fn(&'static str) -> usize); /// ``` /// Type aliases do not bypass this check: /// ```compile_fail -/// let mut session = shimforge::Session::new_global().unwrap(); +/// let mut session = shimforge::Session::new_global(); /// fn source(value: &str) -> usize { value.len() } /// type Input = &'static str; /// shimforge::mock!(session, source, fn(Input) -> usize); /// ``` /// A static result cannot become a shorter borrow: /// ```compile_fail -/// let mut session = shimforge::Session::new_global().unwrap(); +/// let mut session = shimforge::Session::new_global(); /// fn source(_: &str) -> &'static str { "fixed" } /// shimforge::mock!(session, source, fn(&str) -> &str); /// ``` /// A result must stay tied to the same argument: /// ```compile_fail -/// let mut session = shimforge::Session::new_global().unwrap(); +/// let mut session = shimforge::Session::new_global(); /// fn source<'a, 'b>(left: &'a str, _: &'b str) -> &'a str { left } /// shimforge::mock!(session, source, for<'a, 'b> fn(&'a str, &'b str) -> &'b str); /// ``` /// Caller names cannot shadow the checks: /// ```compile_fail -/// let mut session = shimforge::Session::new_global().unwrap(); +/// let mut session = shimforge::Session::new_global(); /// fn __target(_: &str) -> &'static str { "fixed" } /// shimforge::mock!(session, __target, fn(&str) -> &str); /// ``` /// Captures must be safe to send between threads: /// ```compile_fail -/// let mut session = shimforge::Session::new_global().unwrap(); +/// let mut session = shimforge::Session::new_global(); /// fn source() -> usize { 1 } -/// let mock = shimforge::mock!(session, source, fn() -> usize).unwrap(); +/// let mock = shimforge::mock!(session, source, fn() -> usize); /// let value = std::rc::Rc::new(2); /// mock.expect().returning(move || *value); /// ``` /// Captures must outlive the test's stack: /// ```compile_fail -/// let mut session = shimforge::Session::new_global().unwrap(); +/// let mut session = shimforge::Session::new_global(); /// fn source() -> usize { 1 } -/// let mock = shimforge::mock!(session, source, fn() -> usize).unwrap(); +/// let mock = shimforge::mock!(session, source, fn() -> usize); /// let value = String::from("token"); /// mock.expect().returning(|| value.len()); /// ``` /// Return closures cannot create dangling references: /// ```compile_fail -/// let mut session = shimforge::Session::new_global().unwrap(); +/// let mut session = shimforge::Session::new_global(); /// fn source(value: &str) -> &str { value } -/// let mock = shimforge::mock!(session, source, fn(&str) -> &str).unwrap(); +/// let mock = shimforge::mock!(session, source, fn(&str) -> &str); /// mock.expect().returning(|_| String::from("temporary").as_str()); /// ``` /// Calling conventions must match: /// ```compile_fail -/// let mut session = shimforge::Session::new_global().unwrap(); +/// let mut session = shimforge::Session::new_global(); /// extern "C" fn source() -> usize { 1 } /// shimforge::mock!(session, source, fn() -> usize); /// ``` @@ -484,61 +521,61 @@ fn finish(result: Result<(), Error>, fatal: fn() -> !) { /// ```no_run /// # fn original(x: i32) -> i32 { x + 1 } /// # fn fake(x: i32) -> i32 { x + 10 } -/// let mut session = shimforge::Session::new_global()?; -/// shimforge::replace!(session, original => fake, fn(i32) -> i32)?; -/// # Ok::<(), shimforge::Error>(()) +/// let mut session = shimforge::Session::new_global(); +/// shimforge::replace!(session, original => fake, fn(i32) -> i32); /// ``` /// /// No `unsafe` block is needed. Follow the crate's safety rules. Lifetime checks /// are best effort; do not narrow lifetimes to force a type match. Closures -/// without captures are accepted. +/// without captures are accepted. Installation panics if the function cannot be +/// patched. /// /// Incompatible signatures are rejected: /// ```compile_fail -/// let mut session = shimforge::Session::new_global().unwrap(); +/// let mut session = shimforge::Session::new_global(); /// fn source(x: u32) -> u32 { x } /// shimforge::replace!(session, source => |x: u64| x, fn(u32) -> u32); /// ``` /// The source must also match: /// ```compile_fail -/// let mut session = shimforge::Session::new_global().unwrap(); +/// let mut session = shimforge::Session::new_global(); /// fn source(x: u64) -> u64 { x } /// shimforge::replace!(session, source => |x| x, fn(u32) -> u32); /// ``` /// Input borrows cannot be narrowed: /// ```compile_fail -/// let mut session = shimforge::Session::new_global().unwrap(); +/// let mut session = shimforge::Session::new_global(); /// fn source(value: &str) -> usize { value.len() } /// shimforge::replace!(session, source => |value| value.len(), fn(&'static str) -> usize); /// ``` /// This also applies to mutable borrows: /// ```compile_fail -/// let mut session = shimforge::Session::new_global().unwrap(); +/// let mut session = shimforge::Session::new_global(); /// fn source(value: &mut usize) { *value += 1; } /// shimforge::replace!(session, source => |_| (), fn(&'static mut usize)); /// ``` /// A static result cannot become a shorter borrow: /// ```compile_fail -/// let mut session = shimforge::Session::new_global().unwrap(); +/// let mut session = shimforge::Session::new_global(); /// fn source(_: &str) -> &'static str { "fixed" } /// shimforge::replace!(session, source => |value| value, fn(&str) -> &str); /// ``` /// A result must stay tied to the same argument: /// ```compile_fail -/// let mut session = shimforge::Session::new_global().unwrap(); +/// let mut session = shimforge::Session::new_global(); /// fn source<'a, 'b>(left: &'a str, _: &'b str) -> &'a str { left } /// shimforge::replace!(session, source => |_, right| right, /// for<'a, 'b> fn(&'a str, &'b str) -> &'b str); /// ``` /// Calling conventions must match: /// ```compile_fail -/// let mut session = shimforge::Session::new_global().unwrap(); +/// let mut session = shimforge::Session::new_global(); /// extern "C" fn source(x: u32) -> u32 { x } /// shimforge::replace!(session, source => |x| x, fn(u32) -> u32); /// ``` /// Capturing closures are rejected: /// ```compile_fail -/// let mut session = shimforge::Session::new_global().unwrap(); +/// let mut session = shimforge::Session::new_global(); /// let captured = String::from("hello"); /// fn source() -> usize { 1 } /// shimforge::replace!(session, source => || captured.len(), fn() -> usize); diff --git a/src/memory/tests.rs b/src/memory/tests.rs index 44eb68b..4fb5a8b 100644 --- a/src/memory/tests.rs +++ b/src/memory/tests.rs @@ -44,24 +44,26 @@ fn local_installation_errors_can_be_retried_and_cleanup_does_not_write_code() { fn value(seed: u64) -> u64 { seed.wrapping_add(1) } - let mut session = crate::Session::new_local().unwrap(); + let mut session = crate::Session::new_local(); for failed in [true, false] { let injection = failed.then(|| Inject::new(&[("protect memory", 1)])); - let result = crate::mock!(session, value, fn(u64) -> u64); + let result = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| { + crate::mock!(session, value, fn(u64) -> u64) + })); drop(injection); if failed { assert!(result.is_err()); assert_eq!(value(1), 2); } else { - result.unwrap().expect().once().returns(40).unwrap(); + result.unwrap().expect().once().returns(40); } } assert_eq!(value(1), 40); { let _injection = Inject::new(&[("protect memory", 1)]); - session.restore().unwrap(); + session.restore(); } - session.restore().unwrap(); + session.restore(); assert_eq!(value(1), 2); } @@ -395,7 +397,7 @@ fn installing_a_prefix_that_extends_into_an_active_patch_is_rejected() { } memory.protect(0, RX); platform::flush(memory.address, 257).unwrap(); - let mut session = crate::Session::new_global().unwrap(); + let mut session = crate::Session::new_global(); // SAFETY: Both NOP/RET functions use the same void ABI. This test owns their // memory and keeps calls stopped while patching. unsafe { @@ -403,21 +405,20 @@ fn installing_a_prefix_that_extends_into_an_active_patch_is_rejected() { (memory.address + 4) as *const (), (memory.address + 128) as *const (), ) - } - .unwrap(); + }; let installed = read(memory.address, 32).unwrap(); - assert!(matches!( + assert_eq!( // SAFETY: Both entries stay mapped and idle. The overlap is rejected before writing. - unsafe { + crate::tests::panic_message(|| unsafe { session.replace_raw( memory.address as *const (), (memory.address + 256) as *const (), ) - }, - Err(Error::Overlap) - )); + }), + Error::Overlap.to_string() + ); assert_eq!(read(memory.address, 32).unwrap(), installed); - session.restore().unwrap(); + session.restore(); memory.assert_unchanged(memory.address, 32); } @@ -436,18 +437,20 @@ fn installing_a_prefix_that_extends_into_an_active_patch_is_rejected() { let mut page = crate::executable::Executable::near(read as *const () as usize).unwrap(); page.publish(&bytes).unwrap(); let base = page.address(); - let mut session = crate::Session::new_global().unwrap(); + let mut session = crate::Session::new_global(); // SAFETY: Both NOP/RET functions use the same void ABI. This test owns their // memory and keeps calls stopped while patching. - unsafe { session.replace_raw((base + 4) as *const (), (base + 32) as *const ()) }.unwrap(); + unsafe { session.replace_raw((base + 4) as *const (), (base + 32) as *const ()) }; let installed = read(base, 32).unwrap(); - assert!(matches!( + assert_eq!( // SAFETY: Both entries stay mapped and idle. The overlap is rejected before writing. - unsafe { session.replace_raw(base as *const (), (base + 48) as *const ()) }, - Err(Error::Overlap) - )); + crate::tests::panic_message(|| unsafe { + session.replace_raw(base as *const (), (base + 48) as *const ()) + }), + Error::Overlap.to_string() + ); assert_eq!(read(base, 32).unwrap(), installed); - session.restore().unwrap(); + session.restore(); assert_eq!(read(base, bytes.len()).unwrap(), bytes); } @@ -470,13 +473,13 @@ fn a_generated_cet_function_can_execute_replace_and_restore() { // SAFETY: ENDBR64, NOPs, MOV EAX,7, RET form an extern C fn() -> u32. let source: extern "C" fn() -> u32 = unsafe { std::mem::transmute(memory.address) }; assert_eq!(source(), 7); - let mut session = crate::Session::new_global().unwrap(); + let mut session = crate::Session::new_global(); // SAFETY: Both functions match in ABI and lifetime. Only this thread can call // the generated function, and calls are stopped while patching. - unsafe { session.replace_raw(source as *const (), native_replacement as *const ()) }.unwrap(); + unsafe { session.replace_raw(source as *const (), native_replacement as *const ()) }; assert_eq!(source(), 93); assert_eq!(read(memory.address, 4).unwrap(), [0xf3, 0x0f, 0x1e, 0xfa]); - session.restore().unwrap(); + session.restore(); assert_eq!(source(), 7); assert_eq!(read(memory.address, 26).unwrap(), code); } diff --git a/src/routing/tests.rs b/src/routing/tests.rs index d49aa9d..0cd579d 100644 --- a/src/routing/tests.rs +++ b/src/routing/tests.rs @@ -92,8 +92,8 @@ fn conditional_entry_branches_execute_on_unmocked_threads() { let source: extern "C" fn(i32) -> i32 = unsafe { std::mem::transmute(page.address()) }; assert_eq!(source(0), 0); assert_eq!(source(1), 42); - let mut session = crate::Session::new().unwrap(); - crate::replace!(session, source => fake, extern "C" fn(i32) -> i32).unwrap(); + let mut session = crate::Session::new(); + crate::replace!(session, source => fake, extern "C" fn(i32) -> i32); assert_eq!(source(0), 99); std::thread::spawn(move || { assert_eq!(source(0), 0); @@ -101,13 +101,13 @@ fn conditional_entry_branches_execute_on_unmocked_threads() { }) .join() .unwrap(); - session.restore().unwrap(); + session.restore(); assert_eq!(source(1), 42); drop(session); - let mut session = crate::Session::new_global().unwrap(); - crate::replace!(session, source => fake, extern "C" fn(i32) -> i32).unwrap(); + let mut session = crate::Session::new_global(); + crate::replace!(session, source => fake, extern "C" fn(i32) -> i32); assert_eq!(source(1), 99); - session.restore().unwrap(); + session.restore(); assert_eq!(source(1), 42); std::mem::forget(page); } @@ -142,11 +142,11 @@ fn a_relocated_call_returns_to_the_original_function() { // SAFETY: the page is a complete C function with this signature. let source: extern "C" fn(i32) -> i32 = unsafe { std::mem::transmute(page.address()) }; assert_eq!(source(1), 5); - let mut session = crate::Session::new().unwrap(); - crate::replace!(session, source => replacement, extern "C" fn(i32) -> i32).unwrap(); + let mut session = crate::Session::new(); + crate::replace!(session, source => replacement, extern "C" fn(i32) -> i32); assert_eq!(source(1), 99); assert_eq!(std::thread::spawn(move || source(1)).join().unwrap(), 5); - session.restore().unwrap(); + session.restore(); assert_eq!(source(1), 5); std::mem::forget(page); } @@ -170,22 +170,22 @@ fn unsupported_prefixes_do_not_leave_routes() { #[test] fn shared_routes_outlive_the_first_owner() { let _serial = crate::tests::serial(); - let mut session = crate::Session::new_local().unwrap(); - let mock = crate::mock!(session, target, fn(u64) -> u64).unwrap(); - mock.expect().once().returns(10).unwrap(); + let mut session = crate::Session::new_local(); + let mock = crate::mock!(session, target, fn(u64) -> u64); + mock.expect().once().returns(10); let (ready_tx, ready_rx) = std::sync::mpsc::channel(); let (run_tx, run_rx) = std::sync::mpsc::channel(); let worker = std::thread::spawn(move || { - let mut session = crate::Session::new_local().unwrap(); - let mock = crate::mock!(session, target, fn(u64) -> u64).unwrap(); - mock.expect().once().returns(20).unwrap(); + let mut session = crate::Session::new_local(); + let mock = crate::mock!(session, target, fn(u64) -> u64); + mock.expect().once().returns(20); ready_tx.send(()).unwrap(); run_rx.recv().unwrap(); assert_eq!(target(1), 20); }); ready_rx.recv().unwrap(); assert_eq!(target(1), 10); - session.restore().unwrap(); + session.restore(); assert_eq!(target(1), 6); run_tx.send(()).unwrap(); worker.join().unwrap(); diff --git a/src/tests.rs b/src/tests.rs index 9a93215..7d01703 100644 --- a/src/tests.rs +++ b/src/tests.rs @@ -5,22 +5,24 @@ use std::sync::{Mutex, MutexGuard}; fn local_sessions_recover_from_global_poison_and_reject_raw_replacement() { let _serial = serial(); let _ = std::panic::catch_unwind(|| { - let mut session = Session::new_global().unwrap(); - let mock = crate::mock!(session, original, fn(u64) -> u64).unwrap(); - mock.expect().returns(4).unwrap(); + let mut session = Session::new_global(); + let mock = crate::mock!(session, original, fn(u64) -> u64); + mock.expect().returns(4); assert!(matches!(Session::try_new_local(), Err(Error::Busy))); panic!("poison the writer"); }); - let mut session = Session::new_local().unwrap(); + let mut session = Session::new_local(); assert!(matches!(Session::try_new_local(), Err(Error::Busy))); // SAFETY: valid pointers; local mode rejects raw installation before access. - let result = unsafe { session.replace_raw(original as *const (), replacement as *const ()) }; - assert!(result.is_err()); + let message = panic_message(|| unsafe { + session.replace_raw(original as *const (), replacement as *const ()) + }); + assert!(message.contains("local session")); drop(session); let result = std::panic::catch_unwind(|| { - let mut session = Session::new_local().unwrap(); - let mock = crate::mock!(session, original, fn(u64) -> u64).unwrap(); - mock.expect().once().returns(2).unwrap(); + let mut session = Session::new_local(); + let mock = crate::mock!(session, original, fn(u64) -> u64); + mock.expect().once().returns(2); }); assert!(result.is_err()); } @@ -33,13 +35,26 @@ pub(crate) fn serial() -> MutexGuard<'static, ()> { .unwrap_or_else(|poisoned| poisoned.into_inner()) } +/// Runs `action`, fails the test unless it panics, and returns the panic message. +pub(crate) fn panic_message(action: impl FnOnce() -> T) -> String { + let panic = std::panic::catch_unwind(std::panic::AssertUnwindSafe(action)) + .err() + .expect("expected a panic"); + match panic.downcast::() { + Ok(message) => *message, + Err(panic) => panic + .downcast_ref::<&str>() + .map_or_else(String::new, |message| (*message).to_owned()), + } +} + #[test] fn sessions_wait_across_threads_but_nested_sessions_fail() { let _serial = serial(); for local in [true, false] { - let session = Session::new_global().unwrap(); - assert!(matches!(Session::new(), Err(Error::Busy))); - assert!(matches!(Session::new_global(), Err(Error::Busy))); + let session = Session::new_global(); + assert_eq!(panic_message(Session::new), Error::Busy.to_string()); + assert_eq!(panic_message(Session::new_global), Error::Busy.to_string()); let (started_tx, started_rx) = std::sync::mpsc::channel(); let (done_tx, done_rx) = std::sync::mpsc::channel(); let worker = std::thread::spawn(move || { @@ -50,8 +65,7 @@ fn sessions_wait_across_threads_but_nested_sessions_fail() { Session::new() } else { Session::new_global() - } - .unwrap(); + }; done_tx.send(()).unwrap(); }); started_rx.recv().unwrap(); @@ -132,50 +146,53 @@ fn all_errors_are_useful() { fn poisoned_sessions_recover_after_unwinding() { let _serial = serial(); let panic = std::panic::catch_unwind(|| { - let mut session = Session::new_global().unwrap(); - replace!(session, original => replacement, fn(u64) -> u64).unwrap(); + let mut session = Session::new_global(); + replace!(session, original => replacement, fn(u64) -> u64); assert_eq!(original(1), 208); panic!("simulate failed test"); }); assert!(panic.is_err()); assert_eq!(original(1), 110); - let _session = Session::new_global().unwrap(); + let _session = Session::new_global(); assert!(matches!(Session::try_new_global(), Err(Error::Busy))); } #[test] fn installation_errors_do_not_change_the_target() { let _serial = serial(); - let mut session = Session::new_global().unwrap(); + let mut session = Session::new_global(); assert_eq!( - replace!(session, original => original, fn(u64) -> u64), - Err(Error::SameAddress) + panic_message(|| replace!(session, original => original, fn(u64) -> u64)), + Error::SameAddress.to_string() ); // SAFETY: Address checks reject null before reading memory. - assert!(unsafe { session.replace_raw(std::ptr::null(), replacement as *const ()) }.is_err()); + panic_message(|| unsafe { session.replace_raw(std::ptr::null(), replacement as *const ()) }); // SAFETY: The null replacement is also rejected before access. - assert!(unsafe { session.replace_raw(original as *const (), std::ptr::null()) }.is_err()); - replace!(session, original => replacement, fn(u64) -> u64).unwrap(); + panic_message(|| unsafe { session.replace_raw(original as *const (), std::ptr::null()) }); + replace!(session, original => replacement, fn(u64) -> u64); assert_eq!( - replace!(session, original => replacement, fn(u64) -> u64), - Err(Error::Overlap) + panic_message(|| replace!(session, original => replacement, fn(u64) -> u64)), + Error::Overlap.to_string() ); assert_eq!(original(4), 211); - session.restore().unwrap(); - session.restore().unwrap(); + session.restore(); + session.restore(); assert_eq!(original(4), 113); } #[test] fn failed_restoration_keeps_ownership_for_retry() { let _serial = serial(); - let mut session = Session::new_global().unwrap(); - replace!(session, retried => replacement, fn(u64) -> u64).unwrap(); + let mut session = Session::new_global(); + replace!(session, retried => replacement, fn(u64) -> u64); session.patches[0].replacement[0] ^= 1; - assert_eq!(session.restore(), Err(Error::MemoryChanged)); + assert_eq!( + panic_message(|| session.restore()), + Error::MemoryChanged.to_string() + ); assert_eq!(session.patches.len(), 1); session.patches[0].replacement[0] ^= 1; - session.restore().unwrap(); + session.restore(); assert_eq!(retried(0), 41); } diff --git a/tests/async_expectations.rs b/tests/async_expectations.rs index 8072425..1603242 100644 --- a/tests/async_expectations.rs +++ b/tests/async_expectations.rs @@ -17,6 +17,19 @@ fn serial_test() -> MutexGuard<'static, ()> { TEST_LOCK.lock().unwrap_or_else(|error| error.into_inner()) } +/// Runs `action`, fails the test unless it panics, and returns the panic message. +fn panic_message(action: impl FnOnce() -> T) -> String { + let panic = catch_unwind(AssertUnwindSafe(action)) + .err() + .expect("expected a panic"); + match panic.downcast::() { + Ok(message) => *message, + Err(panic) => panic + .downcast_ref::<&str>() + .map_or_else(String::new, |message| (*message).to_owned()), + } +} + struct ThreadWake(std::thread::Thread); impl Wake for ThreadWake { @@ -52,18 +65,20 @@ fn failed_async_installation_clears_the_registry_for_retry() { replace!( session, F::poll => |_, _| Poll::Ready("replacement".to_owned()), fn(Pin<&mut F>, &mut Context<'_>) -> Poll - ) - .unwrap(); + ); } let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); + let mut session = Session::new_global(); let witness = fetch(0); replace_poll(&mut session, &witness); - assert!(matches!(session.mock_async(witness), Err(Error::Overlap))); + assert_eq!( + panic_message(|| session.mock_async(witness)), + Error::Overlap.to_string() + ); assert_eq!(block_on(fetch(1)), "replacement"); - session.restore().unwrap(); - let mock = session.mock_async(fetch(0)).unwrap(); - mock.expect().once().returns("retry".to_owned()).unwrap(); + session.restore(); + let mock = session.mock_async(fetch(0)); + mock.expect().once().returns("retry".to_owned()); assert_eq!(block_on(fetch(2)), "retry"); } @@ -119,19 +134,15 @@ async fn issue(id: u32) -> Ticket { fn native_async_functions_return_mock_values_without_running_the_body() { let _serial = serial_test(); BODY_CALLS.store(0, Ordering::SeqCst); - let mut session = Session::new_global().unwrap(); - let fetches = session.mock_async(fetch(0)).unwrap(); - let expected = fetches - .expect() - .times(2) - .returns(String::from("cached")) - .unwrap(); + let mut session = Session::new_global(); + let fetches = session.mock_async(fetch(0)); + let expected = fetches.expect().times(2).returns(String::from("cached")); assert_eq!(block_on(load_message(5)), "message: cached"); assert_eq!(block_on(fetch(6)), "cached"); assert_eq!(expected.calls(), 2); assert_eq!(BODY_CALLS.load(Ordering::SeqCst), 0); - fetches.verify().unwrap(); - session.restore().unwrap(); + fetches.verify(); + session.restore(); assert_eq!(block_on(fetch(7)), "remote 7"); assert_eq!(BODY_CALLS.load(Ordering::SeqCst), 1); } @@ -139,13 +150,9 @@ fn native_async_functions_return_mock_values_without_running_the_body() { #[test] fn construction_and_cancellation_do_not_count_as_calls() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); - let fetches = session.mock_async(fetch(0)).unwrap(); - let expected = fetches - .expect() - .once() - .returns(String::from("ready")) - .unwrap(); + let mut session = Session::new_global(); + let fetches = session.mock_async(fetch(0)); + let expected = fetches.expect().once().returns(String::from("ready")); drop(fetch(1)); assert_eq!(expected.calls(), 0); assert_eq!(block_on(fetch(2)), "ready"); @@ -155,18 +162,14 @@ fn construction_and_cancellation_do_not_count_as_calls() { #[test] fn callbacks_capture_state_and_produce_owned_outputs() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); - let fetches = session.mock_async(fetch(0)).unwrap(); + let mut session = Session::new_global(); + let fetches = session.mock_async(fetch(0)); let prefix = String::from("page"); let mut count = 0; - fetches - .expect() - .times(2) - .returning(move || { - count += 1; - format!("{prefix} {count}") - }) - .unwrap(); + fetches.expect().times(2).returning(move || { + count += 1; + format!("{prefix} {count}") + }); assert_eq!(block_on(fetch(1)), "page 1"); assert_eq!(block_on(fetch(2)), "page 2"); } @@ -174,11 +177,11 @@ fn callbacks_capture_state_and_produce_owned_outputs() { #[test] fn once_responses_move_non_clone_outputs() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); - let issues = session.mock_async(issue(0)).unwrap(); - issues.expect().return_once(Ticket(80)).unwrap(); + let mut session = Session::new_global(); + let issues = session.mock_async(issue(0)); + issues.expect().return_once(Ticket(80)); let ticket = Ticket(81); - issues.expect().returning_once(move || ticket).unwrap(); + issues.expect().returning_once(move || ticket); assert_eq!(block_on(issue(1)), Ticket(80)); assert_eq!(block_on(issue(2)), Ticket(81)); } @@ -186,16 +189,12 @@ fn once_responses_move_non_clone_outputs() { #[test] fn default_outputs_and_checkpoints_work() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); - let fetches = session.mock_async(fetch(0)).unwrap(); - fetches.expect().once().returns_default().unwrap(); + let mut session = Session::new_global(); + let fetches = session.mock_async(fetch(0)); + fetches.expect().once().returns_default(); assert_eq!(block_on(fetch(1)), ""); - fetches.checkpoint().unwrap(); - fetches - .expect() - .times(1..=2) - .returns(String::from("next")) - .unwrap(); + fetches.checkpoint(); + fetches.expect().times(1..=2).returns(String::from("next")); assert_eq!(block_on(fetch(2)), "next"); } @@ -203,18 +202,16 @@ fn default_outputs_and_checkpoints_work() { fn async_io_can_return_errors_without_touching_disk() { let _serial = serial_test(); BODY_CALLS.store(0, Ordering::SeqCst); - let mut session = Session::new_global().unwrap(); - let reads = session.mock_async(read_config("witness.conf")).unwrap(); + let mut session = Session::new_global(); + let reads = session.mock_async(read_config("witness.conf")); reads .expect() .once() - .returning(|| Err(io::ErrorKind::PermissionDenied.into())) - .unwrap(); + .returning(|| Err(io::ErrorKind::PermissionDenied.into())); reads .expect() .once() - .returning(|| Ok(String::from("port=8080"))) - .unwrap(); + .returning(|| Ok(String::from("port=8080"))); assert_eq!( block_on(read_config("private.conf")).unwrap_err().kind(), io::ErrorKind::PermissionDenied @@ -230,15 +227,11 @@ fn borrowed_async_methods_need_no_static_receiver() { prefix: String::from("service"), }; let key = String::from("settings"); - let mut session = Session::new_global().unwrap(); - let loads = session.mock_async(client.load(&key)).unwrap(); - loads - .expect() - .once() - .returns(String::from("cached")) - .unwrap(); + let mut session = Session::new_global(); + let loads = session.mock_async(client.load(&key)); + loads.expect().once().returns(String::from("cached")); assert_eq!(block_on(client.load(&key)), "cached"); - session.restore().unwrap(); + session.restore(); assert_eq!(block_on(client.load(&key)), "service:settings"); } @@ -247,12 +240,10 @@ fn witness_and_mocked_futures_drop_their_inputs_once() { let _serial = serial_test(); BODY_CALLS.store(0, Ordering::SeqCst); let dropped = Arc::new(AtomicUsize::new(0)); - let mut session = Session::new_global().unwrap(); - let consumes = session - .mock_async(consume(DropCount(Arc::clone(&dropped)))) - .unwrap(); + let mut session = Session::new_global(); + let consumes = session.mock_async(consume(DropCount(Arc::clone(&dropped)))); assert_eq!(dropped.load(Ordering::SeqCst), 1); - consumes.expect().once().returns(7).unwrap(); + consumes.expect().once().returns(7); drop(consume(DropCount(Arc::clone(&dropped)))); assert_eq!(dropped.load(Ordering::SeqCst), 2); assert_eq!(block_on(consume(DropCount(Arc::clone(&dropped)))), 7); @@ -263,13 +254,9 @@ fn witness_and_mocked_futures_drop_their_inputs_once() { #[test] fn other_future_types_with_the_same_output_stay_unchanged() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); - let fetches = session.mock_async(fetch(0)).unwrap(); - fetches - .expect() - .once() - .returns(String::from("cached")) - .unwrap(); + let mut session = Session::new_global(); + let fetches = session.mock_async(fetch(0)); + fetches.expect().once().returns(String::from("cached")); assert_eq!(block_on(other_fetch(4)), "other 4"); assert_eq!(block_on(fetch(4)), "cached"); } @@ -277,22 +264,20 @@ fn other_future_types_with_the_same_output_stay_unchanged() { #[test] fn sequences_check_the_order_of_awaited_calls() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); - let fetches = session.mock_async(fetch(0)).unwrap(); - let issues = session.mock_async(issue(0)).unwrap(); + let mut session = Session::new_global(); + let fetches = session.mock_async(fetch(0)); + let issues = session.mock_async(issue(0)); let order = Sequence::new(); fetches .expect() .once() .in_sequence(&order) - .returns(String::from("first")) - .unwrap(); + .returns(String::from("first")); issues .expect() .once() .in_sequence(&order) - .return_once(Ticket(9)) - .unwrap(); + .return_once(Ticket(9)); assert_eq!(block_on(fetch(1)), "first"); assert_eq!(block_on(issue(1)), Ticket(9)); } @@ -300,68 +285,56 @@ fn sequences_check_the_order_of_awaited_calls() { #[test] fn missing_calls_fail_and_restore_still_restores_the_original() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); - let fetches = session.mock_async(fetch(0)).unwrap(); - fetches - .expect() - .once() - .returns(String::from("missing")) - .unwrap(); - assert!(fetches.verify().is_err()); - assert!(matches!(session.restore(), Err(Error::Expectation(_)))); + let mut session = Session::new_global(); + let fetches = session.mock_async(fetch(0)); + fetches.expect().once().returns(String::from("missing")); + panic_message(|| fetches.verify()); + panic_message(|| session.restore()); assert_eq!(block_on(fetch(1)), "remote 1"); } #[test] fn extra_calls_fail_even_when_the_panic_is_caught() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); - let fetches = session.mock_async(fetch(0)).unwrap(); - let expected = fetches - .expect() - .once() - .returns(String::from("one")) - .unwrap(); + let mut session = Session::new_global(); + let fetches = session.mock_async(fetch(0)); + let expected = fetches.expect().once().returns(String::from("one")); assert_eq!(block_on(fetch(1)), "one"); assert!(catch_unwind(|| block_on(fetch(2))).is_err()); assert_eq!(expected.calls(), 1); - assert!(fetches.verify().is_err()); - assert!(session.restore().is_err()); + panic_message(|| fetches.verify()); + panic_message(|| session.restore()); } #[test] fn never_rules_reject_awaited_calls() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); - let fetches = session.mock_async(fetch(0)).unwrap(); - fetches.expect().never().unwrap(); - fetches.verify().unwrap(); + let mut session = Session::new_global(); + let fetches = session.mock_async(fetch(0)); + fetches.expect().never(); + fetches.verify(); assert!(catch_unwind(|| block_on(fetch(1))).is_err()); - assert!(session.restore().is_err()); + panic_message(|| session.restore()); } #[test] fn configured_async_panics_count_as_calls() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); - let fetches = session.mock_async(fetch(0)).unwrap(); - let expected = fetches.expect().once().panics("offline").unwrap(); + let mut session = Session::new_global(); + let fetches = session.mock_async(fetch(0)); + let expected = fetches.expect().once().panics("offline"); assert!(catch_unwind(|| block_on(fetch(1))).is_err()); assert_eq!(expected.calls(), 1); - fetches.verify().unwrap(); + fetches.verify(); } #[test] fn unwind_restores_async_code_without_a_second_panic() { let _serial = serial_test(); let result = catch_unwind(AssertUnwindSafe(|| { - let mut session = Session::new_global().unwrap(); - let fetches = session.mock_async(fetch(0)).unwrap(); - fetches - .expect() - .once() - .returns(String::from("unused")) - .unwrap(); + let mut session = Session::new_global(); + let fetches = session.mock_async(fetch(0)); + fetches.expect().once().returns(String::from("unused")); panic!("test failed first"); })); assert!(result.is_err()); @@ -371,13 +344,9 @@ fn unwind_restores_async_code_without_a_second_panic() { #[test] fn worker_threads_share_async_expectations() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); - let fetches = session.mock_async(fetch(0)).unwrap(); - let expected = fetches - .expect() - .times(4) - .returns(String::from("shared")) - .unwrap(); + let mut session = Session::new_global(); + let fetches = session.mock_async(fetch(0)); + let expected = fetches.expect().times(4).returns(String::from("shared")); let workers: Vec<_> = (0..4) .map(|id| std::thread::spawn(move || block_on(fetch(id)))) .collect(); @@ -391,28 +360,25 @@ fn worker_threads_share_async_expectations() { fn restore_drops_callbacks_and_invalidates_old_handles() { let _serial = serial_test(); let dropped = Arc::new(AtomicUsize::new(0)); - let mut session = Session::new_global().unwrap(); - let fetches = session.mock_async(fetch(0)).unwrap(); + let mut session = Session::new_global(); + let fetches = session.mock_async(fetch(0)); let marker = DropCount(Arc::clone(&dropped)); - fetches - .expect() - .returning(move || { - let _keep = ▮ - String::from("ready") - }) - .unwrap(); - session.restore().unwrap(); + fetches.expect().returning(move || { + let _keep = ▮ + String::from("ready") + }); + session.restore(); assert_eq!(dropped.load(Ordering::SeqCst), 1); - assert!(fetches.expect().returns(String::from("late")).is_err()); + panic_message(|| fetches.expect().returns(String::from("late"))); } #[test] fn the_same_async_function_can_be_mocked_in_later_sessions() { let _serial = serial_test(); for value in ["first", "second"] { - let mut session = Session::new_global().unwrap(); - let fetches = session.mock_async(fetch(0)).unwrap(); - fetches.expect().once().returns(value.to_owned()).unwrap(); + let mut session = Session::new_global(); + let fetches = session.mock_async(fetch(0)); + fetches.expect().once().returns(value.to_owned()); assert_eq!(block_on(fetch(1)), value); } } @@ -420,16 +386,12 @@ fn the_same_async_function_can_be_mocked_in_later_sessions() { #[test] fn duplicate_async_mocks_leave_the_first_mock_usable() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); - let fetches = session.mock_async(fetch(0)).unwrap(); - fetches - .expect() - .once() - .returns(String::from("first")) - .unwrap(); - assert!(session.mock_async(fetch(1)).is_err()); + let mut session = Session::new_global(); + let fetches = session.mock_async(fetch(0)); + fetches.expect().once().returns(String::from("first")); + panic_message(|| session.mock_async(fetch(1))); assert_eq!(block_on(fetch(2)), "first"); - session.restore().unwrap(); + session.restore(); assert_eq!(block_on(fetch(2)), "remote 2"); } @@ -437,14 +399,10 @@ fn duplicate_async_mocks_leave_the_first_mock_usable() { fn futures_with_non_send_inputs_can_be_mocked() { let _serial = serial_test(); let value = Rc::new(String::from("input")); - let mut session = Session::new_global().unwrap(); - let fetches = session.mock_async(local_fetch(Rc::clone(&value))).unwrap(); + let mut session = Session::new_global(); + let fetches = session.mock_async(local_fetch(Rc::clone(&value))); assert_eq!(Rc::strong_count(&value), 1); - fetches - .expect() - .once() - .returns(String::from("cached")) - .unwrap(); + fetches.expect().once().returns(String::from("cached")); assert_eq!(block_on(local_fetch(Rc::clone(&value))), "cached"); assert_eq!(Rc::strong_count(&value), 1); } @@ -452,32 +410,30 @@ fn futures_with_non_send_inputs_can_be_mocked() { #[test] fn sync_and_async_expectations_share_sequences() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); - let fetches = session.mock_async(fetch(0)).unwrap(); - let saves = mock!(session, save, fn(&str) -> usize).unwrap(); + let mut session = Session::new_global(); + let fetches = session.mock_async(fetch(0)); + let saves = mock!(session, save, fn(&str) -> usize); let order = Sequence::new(); fetches .expect() .once() .in_sequence(&order) - .returns(String::from("value")) - .unwrap(); + .returns(String::from("value")); saves .expect() .with(|value| **value == *"value") .once() .in_sequence(&order) - .returns(5) - .unwrap(); + .returns(5); assert_eq!(save(&block_on(fetch(1))), 5); } #[test] fn recursive_async_callbacks_fail_without_deadlocking() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); - let fetches = session.mock_async(fetch(0)).unwrap(); - fetches.expect().returning(|| block_on(fetch(1))).unwrap(); + let mut session = Session::new_global(); + let fetches = session.mock_async(fetch(0)); + fetches.expect().returning(|| block_on(fetch(1))); assert!(catch_unwind(|| block_on(fetch(2))).is_err()); - assert!(session.restore().is_err()); + panic_message(|| session.restore()); } diff --git a/tests/async_network.rs b/tests/async_network.rs index 0ad5461..2548717 100644 --- a/tests/async_network.rs +++ b/tests/async_network.rs @@ -32,37 +32,35 @@ async fn fetch_banner(address: &str) -> io::Result { #[tokio::test] async fn a_connection_is_answered_without_reaching_the_address() { - let mut session = Session::new().unwrap(); - let connects = session.mock_async(TcpStream::connect("")).unwrap(); + let mut session = Session::new(); + let connects = session.mock_async(TcpStream::connect("")); connects .expect() .once() - .return_once(stream_serving(b"inventory-service ready\n")) - .unwrap(); + .return_once(stream_serving(b"inventory-service ready\n")); // 203.0.113.0/24 is reserved for documentation and is never routed. assert_eq!( fetch_banner("203.0.113.9:7000").await.unwrap(), "inventory-service ready" ); - connects.verify().unwrap(); + connects.verify(); } #[tokio::test] async fn a_refused_connection_reaches_the_caller() { - let mut session = Session::new().unwrap(); - let connects = session.mock_async(TcpStream::connect("")).unwrap(); + let mut session = Session::new(); + let connects = session.mock_async(TcpStream::connect("")); connects .expect() .once() - .returning(|| Err(io::ErrorKind::ConnectionRefused.into())) - .unwrap(); + .returning(|| Err(io::ErrorKind::ConnectionRefused.into())); assert_eq!( fetch_banner("203.0.113.9:7000").await.unwrap_err().kind(), io::ErrorKind::ConnectionRefused ); - session.restore().unwrap(); + session.restore(); } #[tokio::test] @@ -75,13 +73,12 @@ async fn other_threads_still_open_real_connections() { peer.shutdown(Shutdown::Write).unwrap(); }); - let mut session = Session::new().unwrap(); - let connects = session.mock_async(TcpStream::connect("")).unwrap(); + let mut session = Session::new(); + let connects = session.mock_async(TcpStream::connect("")); connects .expect() .once() - .return_once(stream_serving(b"mocked-peer\n")) - .unwrap(); + .return_once(stream_serving(b"mocked-peer\n")); assert_eq!( fetch_banner("203.0.113.9:7000").await.unwrap(), diff --git a/tests/cruntime.rs b/tests/cruntime.rs index 18c51aa..a10dda9 100644 --- a/tests/cruntime.rs +++ b/tests/cruntime.rs @@ -46,17 +46,13 @@ fn leading_number(text: &CStr) -> (c_long, usize) { #[test] fn an_environment_lookup_returns_a_chosen_string() { let _serial = serial(); - let mut session = Session::new().unwrap(); + let mut session = Session::new(); let lookup = mock!( session, getenv, unsafe extern "C" fn(*const c_char) -> *mut c_char - ) - .unwrap(); - lookup - .expect() - .returning(|_| c"canary".as_ptr().cast_mut()) - .unwrap(); + ); + lookup.expect().returning(|_| c"canary".as_ptr().cast_mut()); assert_eq!(deployment_slot().as_deref(), Some("canary")); } @@ -64,13 +60,12 @@ fn an_environment_lookup_returns_a_chosen_string() { #[test] fn an_environment_lookup_can_match_the_requested_key() { let _serial = serial(); - let mut session = Session::new().unwrap(); + let mut session = Session::new(); let lookup = mock!( session, getenv, unsafe extern "C" fn(*const c_char) -> *mut c_char - ) - .unwrap(); + ); lookup .expect() .with(|name| { @@ -79,29 +74,27 @@ fn an_environment_lookup_can_match_the_requested_key() { name == c"DEPLOY_SLOT" }) .once() - .returning(|_| c"blue".as_ptr().cast_mut()) - .unwrap(); + .returning(|_| c"blue".as_ptr().cast_mut()); // Every other key keeps reporting an unset variable. - lookup.expect().returning(|_| ptr::null_mut()).unwrap(); + lookup.expect().returning(|_| ptr::null_mut()); assert_eq!(deployment_slot().as_deref(), Some("blue")); let other = CString::new("PATH").unwrap(); // SAFETY: the key is a valid C string. assert!(unsafe { getenv(other.as_ptr()) }.is_null()); - session.restore().unwrap(); + session.restore(); assert_eq!(deployment_slot(), None); } #[test] fn a_parser_writes_through_its_output_pointer_and_returns_a_value() { let _serial = serial(); - let mut session = Session::new().unwrap(); + let mut session = Session::new(); let parse = mock!( session, strtol, unsafe extern "C" fn(*const c_char, *mut *mut c_char, c_int) -> c_long - ) - .unwrap(); + ); parse .expect() .with(|_, end, base| !end.is_null() && *base == 10) @@ -110,10 +103,9 @@ fn a_parser_writes_through_its_output_pointer_and_returns_a_value() { // SAFETY: the caller supplies a writable slot and a long enough string. unsafe { *end = text.cast_mut().add(4) }; 815 - }) - .unwrap(); + }); assert_eq!(leading_number(c"1234 units"), (815, 4)); - session.restore().unwrap(); + session.restore(); assert_eq!(leading_number(c"1234 units"), (1234, 4)); } diff --git a/tests/expectations.rs b/tests/expectations.rs index a8d5d58..cef7379 100644 --- a/tests/expectations.rs +++ b/tests/expectations.rs @@ -18,6 +18,19 @@ fn serial_test() -> MutexGuard<'static, ()> { TEST_LOCK.lock().unwrap_or_else(|error| error.into_inner()) } +/// Runs `action`, fails the test unless it panics, and returns the panic message. +fn panic_message(action: impl FnOnce() -> T) -> String { + let panic = catch_unwind(AssertUnwindSafe(action)) + .err() + .expect("expected a panic"); + match panic.downcast::() { + Ok(message) => *message, + Err(panic) => panic + .downcast_ref::<&str>() + .map_or_else(String::new, |message| (*message).to_owned()), + } +} + fn price(item: u32) -> u64 { u64::from(item) + 100 } @@ -123,49 +136,35 @@ fn block_on(future: F) -> F::Output { #[test] fn matches_arguments_and_checks_exact_counts() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); - let prices = mock!(session, price, fn(u32) -> u64).unwrap(); - let first = prices - .expect() - .with(|item| *item == 7) - .times(2) - .returns(50) - .unwrap(); - let second = prices - .expect() - .with(|item| *item == 8) - .once() - .returns(75) - .unwrap(); + let mut session = Session::new_global(); + let prices = mock!(session, price, fn(u32) -> u64); + let first = prices.expect().with(|item| *item == 7).times(2).returns(50); + let second = prices.expect().with(|item| *item == 8).once().returns(75); assert_eq!(price(7), 50); assert_eq!(price(8), 75); assert_eq!(price(7), 50); assert_eq!(first.calls(), 2); assert_eq!(second.calls(), 1); - first.verify().unwrap(); - second.verify().unwrap(); - prices.verify().unwrap(); - session.verify().unwrap(); - session.restore().unwrap(); + first.verify(); + second.verify(); + prices.verify(); + session.verify(); + session.restore(); assert_eq!(price(7), 107); } #[test] fn callbacks_capture_values_and_keep_mutable_state() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); - let prices = mock!(session, price, fn(u32) -> u64).unwrap(); + let mut session = Session::new_global(); + let prices = mock!(session, price, fn(u32) -> u64); let base = 20; let mut calls = 0; - prices - .expect() - .times(3) - .returning(move |item| { - calls += 1; - base + u64::from(item) + calls - }) - .unwrap(); + prices.expect().times(3).returning(move |item| { + calls += 1; + base + u64::from(item) + calls + }); assert_eq!(price(4), 25); assert_eq!(price(4), 26); assert_eq!(price(4), 27); @@ -174,13 +173,9 @@ fn callbacks_capture_values_and_keep_mutable_state() { #[test] fn constant_owned_values_are_cloned_for_each_call() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); - let titles = mock!(session, title, fn(u32) -> String).unwrap(); - titles - .expect() - .times(2) - .returns(String::from("saved")) - .unwrap(); + let mut session = Session::new_global(); + let titles = mock!(session, title, fn(u32) -> String); + titles.expect().times(2).returns(String::from("saved")); let mut first = title(1); first.push('!'); assert_eq!(first, "saved!"); @@ -190,37 +185,34 @@ fn constant_owned_values_are_cloned_for_each_call() { #[test] fn returns_non_clone_values_once() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); - let tickets = mock!(session, issue_ticket, fn(u64) -> Ticket).unwrap(); - let expected = tickets.expect().return_once(Ticket(90)).unwrap(); + let mut session = Session::new_global(); + let tickets = mock!(session, issue_ticket, fn(u64) -> Ticket); + let expected = tickets.expect().return_once(Ticket(90)); assert_eq!(issue_ticket(1), Ticket(90)); assert_eq!(expected.calls(), 1); - expected.verify().unwrap(); + expected.verify(); } #[test] fn once_callbacks_move_captured_values() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); - let titles = mock!(session, title, fn(u32) -> String).unwrap(); + let mut session = Session::new_global(); + let titles = mock!(session, title, fn(u32) -> String); let value = String::from("one owner"); - titles - .expect() - .returning_once(move |id| { - assert_eq!(id, 9); - value - }) - .unwrap(); + titles.expect().returning_once(move |id| { + assert_eq!(id, 9); + value + }); assert_eq!(title(9), "one owner"); } #[test] fn default_values_and_unbounded_counts_work() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); - let titles = mock!(session, title, fn(u32) -> String).unwrap(); - let optional = titles.expect().returns_default().unwrap(); - optional.verify().unwrap(); + let mut session = Session::new_global(); + let titles = mock!(session, title, fn(u32) -> String); + let optional = titles.expect().returns_default(); + optional.verify(); assert_eq!(title(1), ""); assert_eq!(title(2), ""); assert_eq!(optional.calls(), 2); @@ -229,44 +221,14 @@ fn default_values_and_unbounded_counts_work() { #[test] fn count_ranges_accept_their_bounds() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); - let prices = mock!(session, price, fn(u32) -> u64).unwrap(); - prices - .expect() - .with(|id| *id == 1) - .times(1..3) - .returns(11) - .unwrap(); - prices - .expect() - .with(|id| *id == 2) - .times(1..=2) - .returns(22) - .unwrap(); - prices - .expect() - .with(|id| *id == 3) - .times(1..) - .returns(33) - .unwrap(); - prices - .expect() - .with(|id| *id == 4) - .times(..2) - .returns(44) - .unwrap(); - prices - .expect() - .with(|id| *id == 5) - .times(..=1) - .returns(55) - .unwrap(); - prices - .expect() - .with(|id| *id == 6) - .times(..) - .returns(66) - .unwrap(); + let mut session = Session::new_global(); + let prices = mock!(session, price, fn(u32) -> u64); + prices.expect().with(|id| *id == 1).times(1..3).returns(11); + prices.expect().with(|id| *id == 2).times(1..=2).returns(22); + prices.expect().with(|id| *id == 3).times(1..).returns(33); + prices.expect().with(|id| *id == 4).times(..2).returns(44); + prices.expect().with(|id| *id == 5).times(..=1).returns(55); + prices.expect().with(|id| *id == 6).times(..).returns(66); for id in 1..=6 { assert_eq!(price(id), u64::from(id) * 11); } @@ -279,24 +241,24 @@ fn count_ranges_accept_their_bounds() { #[test] fn invalid_count_ranges_are_rejected_without_adding_rules() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); - let prices = mock!(session, price, fn(u32) -> u64).unwrap(); + let mut session = Session::new_global(); + let prices = mock!(session, price, fn(u32) -> u64); let start = 3; let end = 2; - assert!(prices.expect().times(start..end).returns(1).is_err()); - assert!(prices.expect().times(2).return_once(1).is_err()); - prices.expect().once().returns(6).unwrap(); + panic_message(|| prices.expect().times(start..end).returns(1)); + panic_message(|| prices.expect().times(2).return_once(1)); + prices.expect().once().returns(6); assert_eq!(price(1), 6); } #[test] fn exhausted_rules_yield_to_later_matching_rules() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); - let prices = mock!(session, price, fn(u32) -> u64).unwrap(); - prices.expect().once().returns(10).unwrap(); - prices.expect().times(2).returns(20).unwrap(); - prices.expect().returns(30).unwrap(); + let mut session = Session::new_global(); + let prices = mock!(session, price, fn(u32) -> u64); + prices.expect().once().returns(10); + prices.expect().times(2).returns(20); + prices.expect().returns(30); assert_eq!(price(1), 10); assert_eq!(price(1), 20); assert_eq!(price(1), 20); @@ -306,14 +268,13 @@ fn exhausted_rules_yield_to_later_matching_rules() { #[test] fn borrowed_arguments_and_returns_keep_their_lifetimes() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); - let words = mock!(session, first_word, for<'a> fn(&'a str) -> &'a str).unwrap(); + let mut session = Session::new_global(); + let words = mock!(session, first_word, for<'a> fn(&'a str) -> &'a str); words .expect() .with(|text| text.contains(' ')) .once() - .returning(|text| text.split_once(' ').unwrap().1) - .unwrap(); + .returning(|text| text.split_once(' ').unwrap().1); let value = String::from("first second"); assert_eq!(first_word(&value), "second"); } @@ -321,8 +282,8 @@ fn borrowed_arguments_and_returns_keep_their_lifetimes() { #[test] fn callbacks_can_write_output_parameters() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); - let reads = mock!(session, fill, fn(&mut [u8]) -> usize).unwrap(); + let mut session = Session::new_global(); + let reads = mock!(session, fill, fn(&mut [u8]) -> usize); reads .expect() .with(|buffer| buffer.len() == 4) @@ -330,8 +291,7 @@ fn callbacks_can_write_output_parameters() { .returning(|buffer| { buffer[..3].copy_from_slice(b"abc"); 3 - }) - .unwrap(); + }); let mut buffer = [9; 4]; assert_eq!(fill(&mut buffer), 3); assert_eq!(buffer, [b'a', b'b', b'c', 9]); @@ -340,62 +300,60 @@ fn callbacks_can_write_output_parameters() { #[test] fn unit_results_skip_the_original_body_and_still_count_calls() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); - let purges = mock!(session, purge, fn(u32)).unwrap(); + let mut session = Session::new_global(); + let purges = mock!(session, purge, fn(u32)); let expected = purges .expect() .with(|generation| *generation == 4) .times(2) - .returns_default() - .unwrap(); + .returns_default(); purge(4); purge(4); assert_eq!(expected.calls(), 2); - session.verify().unwrap(); + session.verify(); } #[test] fn an_extra_call_to_a_unit_function_is_rejected() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); - let purges = mock!(session, purge, fn(u32)).unwrap(); - let expected = purges.expect().once().returns_default().unwrap(); + let mut session = Session::new_global(); + let purges = mock!(session, purge, fn(u32)); + let expected = purges.expect().once().returns_default(); purge(4); assert!(catch_unwind(|| purge(4)).is_err()); assert_eq!(expected.calls(), 1); - assert!(session.restore().is_err()); + panic_message(|| session.restore()); } #[test] fn a_missing_call_to_a_unit_function_fails_verification() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); - let purges = mock!(session, purge, fn(u32)).unwrap(); - purges.expect().times(3).returns_default().unwrap(); + let mut session = Session::new_global(); + let purges = mock!(session, purge, fn(u32)); + purges.expect().times(3).returns_default(); purge(1); purge(2); - assert!(purges.verify().is_err()); - assert!(session.restore().is_err()); + panic_message(|| purges.verify()); + panic_message(|| session.restore()); } #[test] fn a_method_can_fill_an_output_parameter_without_returning() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); - let measures = mock!(session, Store::measure, fn(&Store, &str, &mut usize)).unwrap(); + let mut session = Session::new_global(); + let measures = mock!(session, Store::measure, fn(&Store, &str, &mut usize)); measures .expect() .with(|store, key, _| store.prefix == "cache" && **key == *"port") .once() - .returning(|store, key, length| *length = store.prefix.len() * key.len()) - .unwrap(); + .returning(|store, key, length| *length = store.prefix.len() * key.len()); let store = Store { prefix: String::from("cache"), }; let mut length = 0; store.measure("port", &mut length); assert_eq!(length, 20); - session.restore().unwrap(); + session.restore(); store.measure("port", &mut length); assert_eq!(length, 9); } @@ -403,14 +361,13 @@ fn a_method_can_fill_an_output_parameter_without_returning() { #[test] fn methods_match_the_receiver_and_borrowed_parameters() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); - let lookups = mock!(session, Store::lookup, fn(&Store, &str) -> String).unwrap(); + let mut session = Session::new_global(); + let lookups = mock!(session, Store::lookup, fn(&Store, &str) -> String); lookups .expect() .with(|store, key| store.prefix == "cache" && **key == *"port") .once() - .returning(|_, key| format!("mock {key}")) - .unwrap(); + .returning(|_, key| format!("mock {key}")); let store = Store { prefix: String::from("cache"), }; @@ -420,25 +377,22 @@ fn methods_match_the_receiver_and_borrowed_parameters() { #[test] fn file_reads_match_paths_and_return_io_errors() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); + let mut session = Session::new_global(); let reads = mock!( session, fs::read_to_string::<&Path>, fn(&Path) -> io::Result - ) - .unwrap(); + ); reads .expect() .with(|path| **path == *Path::new("settings.conf")) .once() - .returning(|_| Ok(String::from("port=8080"))) - .unwrap(); + .returning(|_| Ok(String::from("port=8080"))); reads .expect() .with(|path| **path == *Path::new("secret.conf")) .once() - .returning(|_| Err(io::ErrorKind::PermissionDenied.into())) - .unwrap(); + .returning(|_| Err(io::ErrorKind::PermissionDenied.into())); assert_eq!( fs::read_to_string(Path::new("settings.conf")).unwrap(), "port=8080" @@ -458,19 +412,17 @@ fn network_expectations_check_payloads_without_sending_packets() { receiver.set_nonblocking(true).unwrap(); let address = receiver.local_addr().unwrap(); let sender = UdpSocket::bind("127.0.0.1:0").unwrap(); - let mut session = Session::new_global().unwrap(); + let mut session = Session::new_global(); let sends = mock!( session, UdpSocket::send_to::, fn(&UdpSocket, &[u8], SocketAddr) -> io::Result - ) - .unwrap(); + ); sends .expect() .with(move |_, bytes, target| **bytes == *b"orders:3|c" && *target == address) .once() - .returning(|_, bytes, _| Ok(bytes.len())) - .unwrap(); + .returning(|_, bytes, _| Ok(bytes.len())); assert_eq!(sender.send_to(b"orders:3|c", address).unwrap(), 10); let mut buffer = [0; 32]; assert_eq!( @@ -482,22 +434,12 @@ fn network_expectations_check_payloads_without_sending_packets() { #[test] fn a_sequence_checks_order_across_functions() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); - let prices = mock!(session, price, fn(u32) -> u64).unwrap(); - let taxes = mock!(session, tax, fn(u64) -> u64).unwrap(); + let mut session = Session::new_global(); + let prices = mock!(session, price, fn(u32) -> u64); + let taxes = mock!(session, tax, fn(u64) -> u64); let order = Sequence::new(); - prices - .expect() - .times(2) - .in_sequence(&order) - .returns(10) - .unwrap(); - taxes - .expect() - .once() - .in_sequence(&order) - .returns(3) - .unwrap(); + prices.expect().times(2).in_sequence(&order).returns(10); + taxes.expect().once().in_sequence(&order).returns(3); assert_eq!(price(1), 10); assert_eq!(price(2), 10); assert_eq!(tax(20), 3); @@ -506,98 +448,88 @@ fn a_sequence_checks_order_across_functions() { #[test] fn a_sequence_rejects_out_of_order_calls() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); - let prices = mock!(session, price, fn(u32) -> u64).unwrap(); - let taxes = mock!(session, tax, fn(u64) -> u64).unwrap(); + let mut session = Session::new_global(); + let prices = mock!(session, price, fn(u32) -> u64); + let taxes = mock!(session, tax, fn(u64) -> u64); let order = Sequence::new(); - prices - .expect() - .once() - .in_sequence(&order) - .returns(10) - .unwrap(); - taxes - .expect() - .once() - .in_sequence(&order) - .returns(3) - .unwrap(); + prices.expect().once().in_sequence(&order).returns(10); + taxes.expect().once().in_sequence(&order).returns(3); assert!(catch_unwind(|| tax(20)).is_err()); - assert!(taxes.verify().is_err()); - assert!(matches!(session.restore(), Err(Error::Expectation(_)))); + panic_message(|| taxes.verify()); + panic_message(|| session.restore()); assert_eq!(tax(20), 2); } #[test] fn checkpoint_verifies_and_clears_finished_expectations() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); - let prices = mock!(session, price, fn(u32) -> u64).unwrap(); - prices.expect().once().returns(10).unwrap(); - assert!(prices.checkpoint().is_err()); + let mut session = Session::new_global(); + let prices = mock!(session, price, fn(u32) -> u64); + prices.expect().once().returns(10); + panic_message(|| prices.checkpoint()); assert_eq!(price(1), 10); - prices.checkpoint().unwrap(); - prices.expect().once().returns(20).unwrap(); + prices.checkpoint(); + prices.expect().once().returns(20); assert_eq!(price(1), 20); } #[test] fn missing_calls_fail_verification_and_restore_still_removes_the_patch() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); - let prices = mock!(session, price, fn(u32) -> u64).unwrap(); - let expected = prices.expect().times(2).returns(10).unwrap(); + let mut session = Session::new_global(); + let prices = mock!(session, price, fn(u32) -> u64); + let expected = prices.expect().times(2).returns(10); assert_eq!(price(1), 10); - assert!(expected.verify().is_err()); - assert!(prices.verify().is_err()); - assert!(session.verify().is_err()); - assert!(session.restore().is_err()); + panic_message(|| expected.verify()); + panic_message(|| prices.verify()); + panic_message(|| session.verify()); + panic_message(|| session.restore()); assert_eq!(price(1), 101); - session.restore().unwrap(); + session.restore(); } #[test] fn unexpected_arguments_are_reported_even_if_the_panic_is_caught() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); - let prices = mock!(session, price, fn(u32) -> u64).unwrap(); - prices.expect().with(|id| *id == 1).returns(10).unwrap(); + let mut session = Session::new_global(); + let prices = mock!(session, price, fn(u32) -> u64); + prices.expect().with(|id| *id == 1).returns(10); assert!(catch_unwind(|| price(2)).is_err()); - assert!(prices.verify().is_err()); - assert!(session.restore().is_err()); + panic_message(|| prices.verify()); + panic_message(|| session.restore()); } #[test] fn calls_beyond_the_limit_fail() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); - let prices = mock!(session, price, fn(u32) -> u64).unwrap(); - let expected = prices.expect().once().returns(10).unwrap(); + let mut session = Session::new_global(); + let prices = mock!(session, price, fn(u32) -> u64); + let expected = prices.expect().once().returns(10); assert_eq!(price(1), 10); assert!(catch_unwind(|| price(1)).is_err()); assert_eq!(expected.calls(), 1); - assert!(session.restore().is_err()); + panic_message(|| session.restore()); } #[test] fn never_rules_allow_no_calls_and_reject_matching_calls() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); - let prices = mock!(session, price, fn(u32) -> u64).unwrap(); - prices.expect().with(|id| *id == 9).never().unwrap(); - prices.expect().with(|id| *id != 9).returns(10).unwrap(); + let mut session = Session::new_global(); + let prices = mock!(session, price, fn(u32) -> u64); + prices.expect().with(|id| *id == 9).never(); + prices.expect().with(|id| *id != 9).returns(10); assert_eq!(price(1), 10); - prices.verify().unwrap(); + prices.verify(); assert!(catch_unwind(|| price(9)).is_err()); - assert!(session.restore().is_err()); + panic_message(|| session.restore()); } #[test] fn configured_panics_count_as_calls() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); - let prices = mock!(session, price, fn(u32) -> u64).unwrap(); - let expected = prices.expect().once().panics("price unavailable").unwrap(); + let mut session = Session::new_global(); + let prices = mock!(session, price, fn(u32) -> u64); + let expected = prices.expect().once().panics("price unavailable"); let panic = catch_unwind(|| price(1)).unwrap_err(); let message = panic .downcast_ref::() @@ -606,29 +538,29 @@ fn configured_panics_count_as_calls() { .unwrap(); assert!(message.contains("price unavailable")); assert_eq!(expected.calls(), 1); - prices.verify().unwrap(); + prices.verify(); } #[test] fn drop_checks_counts_and_restores_code_before_panicking() { let _serial = serial_test(); let result = catch_unwind(AssertUnwindSafe(|| { - let mut session = Session::new_global().unwrap(); - let prices = mock!(session, price, fn(u32) -> u64).unwrap(); - prices.expect().once().returns(10).unwrap(); + let mut session = Session::new_global(); + let prices = mock!(session, price, fn(u32) -> u64); + prices.expect().once().returns(10); })); assert!(result.is_err()); assert_eq!(price(1), 101); - assert!(Session::new_global().is_ok()); + drop(Session::new_global()); } #[test] fn unwinding_does_not_panic_again_for_missing_calls() { let _serial = serial_test(); let result = catch_unwind(AssertUnwindSafe(|| { - let mut session = Session::new_global().unwrap(); - let prices = mock!(session, price, fn(u32) -> u64).unwrap(); - prices.expect().once().returns(10).unwrap(); + let mut session = Session::new_global(); + let prices = mock!(session, price, fn(u32) -> u64); + prices.expect().once().returns(10); panic!("test failed first"); })); assert!(result.is_err()); @@ -638,17 +570,13 @@ fn unwinding_does_not_panic_again_for_missing_calls() { #[test] fn callbacks_are_serialized_across_threads() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); - let prices = mock!(session, price, fn(u32) -> u64).unwrap(); + let mut session = Session::new_global(); + let prices = mock!(session, price, fn(u32) -> u64); let mut total = 0; - let expected = prices - .expect() - .times(40) - .returning(move |_| { - total += 1; - total - }) - .unwrap(); + let expected = prices.expect().times(40).returning(move |_| { + total += 1; + total + }); let workers: Vec<_> = (0..4) .map(|_| std::thread::spawn(|| (0..10).map(price).collect::>())) .collect(); @@ -664,22 +592,22 @@ fn callbacks_are_serialized_across_threads() { #[test] fn recursive_calls_fail_without_deadlocking() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); - let prices = mock!(session, price, fn(u32) -> u64).unwrap(); - prices.expect().returning(|id| price(id + 1)).unwrap(); + let mut session = Session::new_global(); + let prices = mock!(session, price, fn(u32) -> u64); + prices.expect().returning(|id| price(id + 1)); assert!(catch_unwind(|| price(1)).is_err()); - assert!(session.restore().is_err()); + panic_message(|| session.restore()); } #[test] fn dropping_the_handle_keeps_the_mock_until_session_restore() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); - let prices = mock!(session, price, fn(u32) -> u64).unwrap(); - prices.expect().once().returns(42).unwrap(); + let mut session = Session::new_global(); + let prices = mock!(session, price, fn(u32) -> u64); + prices.expect().once().returns(42); drop(prices); assert_eq!(price(1), 42); - session.restore().unwrap(); + session.restore(); assert_eq!(price(1), 101); } @@ -687,20 +615,17 @@ fn dropping_the_handle_keeps_the_mock_until_session_restore() { fn restore_drops_captures_and_rejects_new_rules_on_old_handles() { let _serial = serial_test(); let dropped = Arc::new(AtomicUsize::new(0)); - let mut session = Session::new_global().unwrap(); - let prices = mock!(session, price, fn(u32) -> u64).unwrap(); + let mut session = Session::new_global(); + let prices = mock!(session, price, fn(u32) -> u64); let marker = DropCount(Arc::clone(&dropped)); - prices - .expect() - .returning(move |id| { - let _keep = ▮ - u64::from(id) - }) - .unwrap(); + prices.expect().returning(move |id| { + let _keep = ▮ + u64::from(id) + }); assert_eq!(price(3), 3); - session.restore().unwrap(); + session.restore(); assert_eq!(dropped.load(Ordering::SeqCst), 1); - assert!(prices.expect().returns(10).is_err()); + panic_message(|| prices.expect().returns(10)); drop(prices); assert_eq!(dropped.load(Ordering::SeqCst), 1); } @@ -709,9 +634,9 @@ fn restore_drops_captures_and_rejects_new_rules_on_old_handles() { fn the_same_macro_call_site_can_be_used_in_later_sessions() { let _serial = serial_test(); for value in [10, 20] { - let mut session = Session::new_global().unwrap(); - let prices = mock!(session, price, fn(u32) -> u64).unwrap(); - prices.expect().once().returns(value).unwrap(); + let mut session = Session::new_global(); + let prices = mock!(session, price, fn(u32) -> u64); + prices.expect().once().returns(value); assert_eq!(price(1), value); } assert_eq!(price(1), 101); @@ -720,32 +645,40 @@ fn the_same_macro_call_site_can_be_used_in_later_sessions() { #[test] fn duplicate_mocks_leave_the_first_mock_usable() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); - let prices = mock!(session, price, fn(u32) -> u64).unwrap(); - prices.expect().once().returns(42).unwrap(); - let duplicate = mock!(session, price, fn(u32) -> u64); - assert!(matches!(duplicate, Err(Error::Overlap))); + let mut session = Session::new_global(); + let prices = mock!(session, price, fn(u32) -> u64); + prices.expect().once().returns(42); + assert_eq!( + panic_message(|| mock!(session, price, fn(u32) -> u64)), + Error::Overlap.to_string() + ); assert_eq!(price(1), 42); - session.restore().unwrap(); + session.restore(); assert_eq!(price(1), 101); } #[test] fn failed_installation_clears_the_slot_for_the_next_attempt() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); - replace!(session, price => fixed_price, fn(u32) -> u64).unwrap(); + let mut session = Session::new_global(); + replace!(session, price => fixed_price, fn(u32) -> u64); for attempt in 0..2 { - let result = mock!(session, price, fn(u32) -> u64); + let result = catch_unwind(AssertUnwindSafe(|| mock!(session, price, fn(u32) -> u64))); if attempt == 0 { - assert!(matches!(result, Err(Error::Overlap))); + let panic = result + .err() + .expect("mocking a replaced function must panic"); + assert_eq!( + panic.downcast_ref::(), + Some(&Error::Overlap.to_string()) + ); assert_eq!(price(1), 21); } else { let prices = result.unwrap(); - prices.expect().once().returns(42).unwrap(); + prices.expect().once().returns(42); assert_eq!(price(1), 42); } - session.restore().unwrap(); + session.restore(); } assert_eq!(price(1), 101); } @@ -753,13 +686,12 @@ fn failed_installation_clears_the_slot_for_the_next_attempt() { #[test] fn boxed_future_callbacks_match_arguments_and_can_return_pending() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); + let mut session = Session::new_global(); let reads = mock!( session, boxed_read, fn(u32) -> Pin + Send>> - ) - .unwrap(); + ); let offset = 20; let expected = reads .expect() @@ -770,13 +702,12 @@ fn boxed_future_callbacks_match_arguments_and_can_return_pending() { value: id as usize + offset, ready: false, }) as Pin + Send>> - }) - .unwrap(); + }); let future = boxed_read(7); assert_eq!(expected.calls(), 1); assert_eq!(block_on(future), 27); assert_eq!(expected.calls(), 1); - session.restore().unwrap(); + session.restore(); assert_eq!(block_on(boxed_read(7)), 8); } @@ -788,11 +719,11 @@ fn macro_names_do_not_shadow_the_callers_session_or_types() { type Option = (); type String = (); let _: (Result, Box, Option, String) = ((), (), (), ()); - let mut state = Session::new_global().unwrap(); - let prices = mock!(state, price, fn(u32) -> u64).unwrap(); - prices.expect().once().returns(42).unwrap(); + let mut state = Session::new_global(); + let prices = mock!(state, price, fn(u32) -> u64); + prices.expect().once().returns(42); assert_eq!(price(1), 42); - state.restore().unwrap(); + state.restore(); assert_eq!(price(1), 101); } @@ -811,14 +742,13 @@ fn trim_left<'a>(left: &'a str, right: &str) -> &'a str { #[test] fn independent_input_lifetimes_keep_their_output_link() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); + let mut session = Session::new_global(); let trim = mock!( session, trim_left, for<'a, 'b> fn(&'a str, &'b str) -> &'a str - ) - .unwrap(); - trim.expect().once().returning(|left, _| left).unwrap(); + ); + trim.expect().once().returning(|left, _| left); let left = String::from("left"); let result; { @@ -826,37 +756,33 @@ fn independent_input_lifetimes_keep_their_output_link() { result = trim_left(&left, &right); } assert_eq!(result, "left"); - session.restore().unwrap(); + session.restore(); } #[test] fn genuine_static_inputs_and_results_stay_supported() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); - let length = mock!(session, static_length, fn(&'static str) -> usize).unwrap(); + let mut session = Session::new_global(); + let length = mock!(session, static_length, fn(&'static str) -> usize); let seen = Arc::new(Mutex::new(None)); let captured = seen.clone(); - length - .expect() - .once() - .returning(move |value| { - *captured.lock().unwrap() = Some(value); - 42 - }) - .unwrap(); - let result = mock!(session, static_result, fn(&str) -> &'static str).unwrap(); - result.expect().once().returns("mock").unwrap(); + length.expect().once().returning(move |value| { + *captured.lock().unwrap() = Some(value); + 42 + }); + let result = mock!(session, static_result, fn(&str) -> &'static str); + result.expect().once().returns("mock"); assert_eq!(static_length("saved"), 42); assert_eq!(*seen.lock().unwrap(), Some("saved")); assert_eq!(static_result(&String::from("input")), "mock"); - session.restore().unwrap(); + session.restore(); } #[test] fn signature_checks_do_not_run_source_or_target_expressions() { let _serial = serial_test(); let mut evaluated = 0; - let mut session = Session::new_global().unwrap(); + let mut session = Session::new_global(); let prices = mock!( session, { @@ -864,24 +790,22 @@ fn signature_checks_do_not_run_source_or_target_expressions() { price }, fn(u32) -> u64 - ) - .unwrap(); + ); assert_eq!(evaluated, 1); - prices.expect().once().returns(7).unwrap(); + prices.expect().once().returns(7); assert_eq!(price(1), 7); - session.restore().unwrap(); + session.restore(); replace!(session, { evaluated += 1; price } => { evaluated += 1; fixed_price }, - fn(u32) -> u64) - .unwrap(); + fn(u32) -> u64); assert_eq!(evaluated, 3); assert_eq!(price(1), 21); - session.restore().unwrap(); + session.restore(); } #[test] fn source_and_target_expressions_can_move_values_and_use_question_mark() -> Result<(), Error> { let _serial = serial_test(); - let mut session = Session::new_global()?; + let mut session = Session::new_global(); let owned = String::from("source"); let prices = mock!( session, @@ -890,19 +814,19 @@ fn source_and_target_expressions_can_move_values_and_use_question_mark() -> Resu price }, fn(u32) -> u64 - )?; - prices.expect().once().returns(9)?; + ); + prices.expect().once().returns(9); assert_eq!(price(1), 9); - session.restore()?; + session.restore(); let source: Result u64, Error> = Ok(price); - let prices = mock!(session, source?, fn(u32) -> u64)?; - prices.expect().once().returns(10)?; + let prices = mock!(session, source?, fn(u32) -> u64); + prices.expect().once().returns(10); assert_eq!(price(1), 10); - session.restore()?; + session.restore(); let owned = String::from("replacement"); let target: Result u64, Error> = Ok(fixed_price); - replace!(session, price => { drop(owned); target? }, fn(u32) -> u64)?; + replace!(session, price => { drop(owned); target? }, fn(u32) -> u64); assert_eq!(price(1), 21); - session.restore()?; + session.restore(); Ok(()) } diff --git a/tests/filesystem.rs b/tests/filesystem.rs index 9128c03..72e53c7 100644 --- a/tests/filesystem.rs +++ b/tests/filesystem.rs @@ -87,40 +87,36 @@ fn directory_creation_succeeds_without_creating_a_directory() { let _serial = serial(); let root = Path::new("virtual/cache/reports"); assert!(!root.exists()); - let mut session = Session::new().unwrap(); + let mut session = Session::new(); let create = mock!( session, fs::create_dir_all::<&Path>, fn(&Path) -> io::Result<()> - ) - .unwrap(); + ); create .expect() .with(|path| **path == *Path::new("virtual/cache/reports")) .once() - .returning(|_| Ok(())) - .unwrap(); + .returning(|_| Ok(())); assert!(prepare_cache(root).is_ok()); - session.verify().unwrap(); + session.verify(); assert!(!root.exists()); } #[test] fn directory_creation_failures_reach_the_caller() { let _serial = serial(); - let mut session = Session::new().unwrap(); + let mut session = Session::new(); let create = mock!( session, fs::create_dir_all::<&Path>, fn(&Path) -> io::Result<()> - ) - .unwrap(); + ); create .expect() .once() - .returning(|_| Err(io::ErrorKind::PermissionDenied.into())) - .unwrap(); + .returning(|_| Err(io::ErrorKind::PermissionDenied.into())); let error = prepare_cache(Path::new("virtual/cache/reports")).unwrap_err(); assert!(error.starts_with("cannot prepare cache:"), "{error}"); @@ -132,25 +128,24 @@ fn path_predicates_answer_for_paths_that_do_not_exist() { let missing = Path::new("virtual/archive"); assert_eq!(describe(missing), "missing"); { - let mut session = Session::new().unwrap(); - let is_dir = mock!(session, Path::is_dir, fn(&Path) -> bool).unwrap(); - is_dir.expect().once().returns(true).unwrap(); + let mut session = Session::new(); + let is_dir = mock!(session, Path::is_dir, fn(&Path) -> bool); + is_dir.expect().once().returns(true); assert_eq!(describe(missing), "directory"); - session.verify().unwrap(); + session.verify(); } { - let mut session = Session::new().unwrap(); - let is_dir = mock!(session, Path::is_dir, fn(&Path) -> bool).unwrap(); - is_dir.expect().once().returns(false).unwrap(); - let exists = mock!(session, Path::exists, fn(&Path) -> bool).unwrap(); + let mut session = Session::new(); + let is_dir = mock!(session, Path::is_dir, fn(&Path) -> bool); + is_dir.expect().once().returns(false); + let exists = mock!(session, Path::exists, fn(&Path) -> bool); exists .expect() .with(|path| **path == *Path::new("virtual/archive")) .once() - .returns(true) - .unwrap(); + .returns(true); assert_eq!(describe(missing), "file"); - session.verify().unwrap(); + session.verify(); } assert_eq!(describe(missing), "missing"); } @@ -158,98 +153,82 @@ fn path_predicates_answer_for_paths_that_do_not_exist() { #[test] fn a_predicate_mock_can_require_an_exact_number_of_questions() { let _serial = serial(); - let mut session = Session::new().unwrap(); - let is_dir = mock!(session, Path::is_dir, fn(&Path) -> bool).unwrap(); - let expected = is_dir.expect().times(3).returns(true).unwrap(); + let mut session = Session::new(); + let is_dir = mock!(session, Path::is_dir, fn(&Path) -> bool); + let expected = is_dir.expect().times(3).returns(true); for name in ["virtual/a", "virtual/b", "virtual/c"] { assert!(Path::new(name).is_dir()); } assert_eq!(expected.calls(), 3); - session.restore().unwrap(); + session.restore(); assert!(!Path::new("virtual/a").is_dir()); } #[test] fn a_line_is_read_from_a_handle_that_never_reaches_disk() { let _serial = serial(); - let mut session = Session::new().unwrap(); - let open = mock!(session, File::open::<&str>, fn(&str) -> io::Result).unwrap(); + let mut session = Session::new(); + let open = mock!(session, File::open::<&str>, fn(&str) -> io::Result); open.expect() .with(|path| **path == *"virtual/inventory.csv") .once() - .returning(|_| Ok(detached_handle())) - .unwrap(); + .returning(|_| Ok(detached_handle())); let read_line = mock!( session, BufReader::::read_line, fn(&mut BufReader, &mut String) -> io::Result - ) - .unwrap(); - read_line - .expect() - .once() - .returning(|_, line| { - line.push_str("sku,count,location\n"); - Ok(line.len()) - }) - .unwrap(); + ); + read_line.expect().once().returning(|_, line| { + line.push_str("sku,count,location\n"); + Ok(line.len()) + }); assert_eq!( read_header("virtual/inventory.csv").unwrap(), "sku,count,location" ); - session.verify().unwrap(); + session.verify(); } #[test] fn a_write_is_accepted_without_a_writable_handle() { let _serial = serial(); - let mut session = Session::new().unwrap(); - let open = mock!(session, File::open::<&str>, fn(&str) -> io::Result).unwrap(); - open.expect() - .once() - .returning(|_| Ok(detached_handle())) - .unwrap(); + let mut session = Session::new(); + let open = mock!(session, File::open::<&str>, fn(&str) -> io::Result); + open.expect().once().returning(|_| Ok(detached_handle())); let write_all = mock!( session, File::write_all, fn(&mut File, &[u8]) -> io::Result<()> - ) - .unwrap(); + ); write_all .expect() .with(|_, bytes| **bytes == *b"sku-77,3,aisle-2\n") .once() - .returning(|_, _| Ok(())) - .unwrap(); + .returning(|_, _| Ok(())); assert_eq!( append_entry("virtual/inventory.csv", "sku-77,3,aisle-2\n").unwrap(), 17 ); - session.verify().unwrap(); + session.verify(); } #[test] fn a_write_error_is_reported_to_the_caller() { let _serial = serial(); - let mut session = Session::new().unwrap(); - let open = mock!(session, File::open::<&str>, fn(&str) -> io::Result).unwrap(); - open.expect() - .once() - .returning(|_| Ok(detached_handle())) - .unwrap(); + let mut session = Session::new(); + let open = mock!(session, File::open::<&str>, fn(&str) -> io::Result); + open.expect().once().returning(|_| Ok(detached_handle())); let write_all = mock!( session, File::write_all, fn(&mut File, &[u8]) -> io::Result<()> - ) - .unwrap(); + ); write_all .expect() .once() - .returning(|_, _| Err(io::ErrorKind::StorageFull.into())) - .unwrap(); + .returning(|_, _| Err(io::ErrorKind::StorageFull.into())); assert_eq!( append_entry("virtual/inventory.csv", "sku-77,3,aisle-2\n") diff --git a/tests/generics.rs b/tests/generics.rs index b3bde93..e5ef7d9 100644 --- a/tests/generics.rs +++ b/tests/generics.rs @@ -45,16 +45,15 @@ fn one_instantiation_is_mocked_and_restored_with_the_session() { let target = Path::new("virtual/ledger.tar"); assert!(archive(target).is_err()); { - let mut session = Session::new().unwrap(); - let archives = mock!(session, archive::<&Path>, fn(&Path) -> io::Result<()>).unwrap(); + let mut session = Session::new(); + let archives = mock!(session, archive::<&Path>, fn(&Path) -> io::Result<()>); archives .expect() .with(|path| **path == *Path::new("virtual/ledger.tar")) .once() - .returning(|_| Ok(())) - .unwrap(); + .returning(|_| Ok(())); assert!(archive(target).is_ok()); - session.verify().unwrap(); + session.verify(); } assert_eq!( archive(target).unwrap_err().kind(), @@ -64,19 +63,17 @@ fn one_instantiation_is_mocked_and_restored_with_the_session() { #[test] fn other_instantiations_keep_the_original_body() { - let mut session = Session::new().unwrap(); + let mut session = Session::new(); let summaries = mock!( session, summarize::<&str, bool, i32>, fn(&str, bool, i32) -> String - ) - .unwrap(); + ); let expected = summaries .expect() .with(|head, flag, count| **head == *"orders" && *flag && *count == 19) .once() - .returns(String::from("mocked summary")) - .unwrap(); + .returns(String::from("mocked summary")); assert_eq!(summarize("orders", true, 19), "mocked summary"); // A different set of type arguments compiles to a different function. @@ -86,34 +83,32 @@ fn other_instantiations_keep_the_original_body() { #[test] fn a_type_parameter_inside_a_slice_can_be_named_with_a_turbofish() { - let mut session = Session::new().unwrap(); - let joins = mock!(session, joined::<&str>, fn(&str, &[&str]) -> String).unwrap(); + let mut session = Session::new(); + let joins = mock!(session, joined::<&str>, fn(&str, &[&str]) -> String); joins .expect() .with(|separator, parts| **separator == *" / " && parts.len() == 2) .once() - .returning(|separator, parts| parts.join(separator)) - .unwrap(); + .returning(|separator, parts| parts.join(separator)); assert_eq!(joined(" / ", &["north", "south"]), "north / south"); - session.verify().unwrap(); + session.verify(); } #[test] fn a_caller_that_names_its_own_lifetime_reaches_the_same_instantiation() { - let mut session = Session::new().unwrap(); + let mut session = Session::new(); let checks = mock!( session, is_stale::, fn(u16, bool, &str) -> bool - ) - .unwrap(); - checks.expect().times(2).returns(true).unwrap(); + ); + checks.expect().times(2).returns(true); assert!(is_stale(7u16, false, "12")); let count = String::from("12"); assert!(stale_with_borrowed_count(&count)); - session.verify().unwrap(); + session.verify(); } #[test] @@ -121,16 +116,15 @@ fn generic_methods_are_mocked_per_type_argument() { let queue = Queue { name: String::from("reports"), }; - let mut session = Session::new().unwrap(); - let publishes = mock!(session, Queue::publish::, fn(&Queue, u32) -> String).unwrap(); + let mut session = Session::new(); + let publishes = mock!(session, Queue::publish::, fn(&Queue, u32) -> String); publishes .expect() .with(|queue, payload| queue.name == "reports" && *payload == 5) .once() - .returning(|queue, payload| format!("{} dropped {payload}", queue.name)) - .unwrap(); + .returning(|queue, payload| format!("{} dropped {payload}", queue.name)); assert_eq!(queue.publish(5u32), "reports dropped 5"); assert_eq!(queue.publish("5"), "reports <- 5"); - session.verify().unwrap(); + session.verify(); } diff --git a/tests/http_client.rs b/tests/http_client.rs index cc859ba..67c6a97 100644 --- a/tests/http_client.rs +++ b/tests/http_client.rs @@ -61,13 +61,9 @@ fn witness() -> impl Future> { #[tokio::test] async fn a_request_is_answered_without_opening_an_outside_connection() { - let mut session = Session::new_global().unwrap(); - let connects = session.mock_async(witness()).unwrap(); - connects - .expect() - .once() - .return_once(local_http_peer()) - .unwrap(); + let mut session = Session::new_global(); + let connects = session.mock_async(witness()); + connects.expect().once().return_once(local_http_peer()); let client = Client::builder(TokioExecutor::new()).build(HttpConnector::new()); // 198.51.100.0/24 is reserved for documentation and is never routed. @@ -83,17 +79,16 @@ async fn a_request_is_answered_without_opening_an_outside_connection() { assert_eq!(response.headers()["content-type"], "application/json"); let body = response.into_body().collect().await.unwrap().to_bytes(); assert_eq!(String::from_utf8(body.to_vec()).unwrap(), BODY); - connects.verify().unwrap(); + connects.verify(); } #[tokio::test] async fn a_connect_failure_surfaces_as_a_client_error() { - let mut session = Session::new_global().unwrap(); - let connects = session.mock_async(witness()).unwrap(); + let mut session = Session::new_global(); + let connects = session.mock_async(witness()); connects .expect() - .returning(|| Err(io::ErrorKind::HostUnreachable.into())) - .unwrap(); + .returning(|| Err(io::ErrorKind::HostUnreachable.into())); let client: Client<_, String> = Client::builder(TokioExecutor::new()).build(HttpConnector::new()); @@ -105,7 +100,7 @@ async fn a_connect_failure_surfaces_as_a_client_error() { let error = client.request(request).await.unwrap_err(); assert!(error.is_connect(), "{error}"); - session.restore().unwrap(); + session.restore(); } /// An SDK-style client that hands out a boxed future for each request. @@ -138,13 +133,12 @@ impl Pipeline { #[tokio::test] async fn an_sdk_request_method_answers_without_a_transport() { - let mut session = Session::new_global().unwrap(); + let mut session = Session::new_global(); let sends = mock!( session, Pipeline::send, for<'a> fn(&'a Pipeline, &'a str) -> Call<'a> - ) - .unwrap(); + ); sends .expect() .with(|pipeline, path| { @@ -157,8 +151,7 @@ async fn an_sdk_request_method_answers_without_a_transport() { *response.status_mut() = StatusCode::OK; Ok(response) }) as Call<'_> - }) - .unwrap(); + }); let pipeline = Pipeline { endpoint: String::from("https://inventory.invalid"), @@ -166,5 +159,5 @@ async fn an_sdk_request_method_answers_without_a_transport() { let response = pipeline.send("/v1/items").await.unwrap(); assert_eq!(response.status(), StatusCode::OK); assert_eq!(response.body(), BODY); - sends.verify().unwrap(); + sends.verify(); } diff --git a/tests/io.rs b/tests/io.rs index 74a70d6..cfd1a5a 100644 --- a/tests/io.rs +++ b/tests/io.rs @@ -161,8 +161,8 @@ fn deny_datagram(socket: &UdpSocket, contents: &[u8], target: SocketAddr) -> io: fn filesystem_read_supplies_data_to_unmodified_business_code() { let _serial = serial(); READ_CALLS.store(0, Ordering::SeqCst); - let mut session = Session::new_global().unwrap(); - replace!(session, fs::read::<&Path> => member_data, fn(&Path) -> io::Result>).unwrap(); + let mut session = Session::new_global(); + replace!(session, fs::read::<&Path> => member_data, fn(&Path) -> io::Result>); assert_eq!( load_members(Path::new("virtual/members.txt")).unwrap(), ["alice", "bob"] @@ -174,9 +174,8 @@ fn filesystem_read_supplies_data_to_unmodified_business_code() { fn missing_file_uses_the_application_default() { let _serial = serial(); READ_CALLS.store(0, Ordering::SeqCst); - let mut session = Session::new_global().unwrap(); - replace!(session, fs::read::<&Path> => missing_member_data, fn(&Path) -> io::Result>) - .unwrap(); + let mut session = Session::new_global(); + replace!(session, fs::read::<&Path> => missing_member_data, fn(&Path) -> io::Result>); assert!( load_members(Path::new("virtual/members.txt")) .unwrap() @@ -189,9 +188,8 @@ fn missing_file_uses_the_application_default() { fn malformed_file_contents_are_rejected() { let _serial = serial(); READ_CALLS.store(0, Ordering::SeqCst); - let mut session = Session::new_global().unwrap(); - replace!(session, fs::read::<&Path> => invalid_member_data, fn(&Path) -> io::Result>) - .unwrap(); + let mut session = Session::new_global(); + replace!(session, fs::read::<&Path> => invalid_member_data, fn(&Path) -> io::Result>); assert_eq!( load_members(Path::new("virtual/members.txt")) .unwrap_err() @@ -206,9 +204,8 @@ fn filesystem_write_captures_the_generated_path_and_payload() { let _serial = serial(); WRITE_CALLS.store(0, Ordering::SeqCst); WRITES.lock().unwrap().clear(); - let mut session = Session::new_global().unwrap(); - replace!(session, fs::write::<&Path, &[u8]> => capture_write, fn(&Path, &[u8]) -> io::Result<()>) - .unwrap(); + let mut session = Session::new_global(); + replace!(session, fs::write::<&Path, &[u8]> => capture_write, fn(&Path, &[u8]) -> io::Result<()>); save_members(Path::new("virtual/members.txt"), &["alice", "bob"]).unwrap(); assert_eq!(WRITE_CALLS.load(Ordering::SeqCst), 1); let writes = WRITES.lock().unwrap(); @@ -221,9 +218,8 @@ fn filesystem_write_captures_the_generated_path_and_payload() { fn filesystem_write_permission_errors_reach_the_caller() { let _serial = serial(); WRITE_CALLS.store(0, Ordering::SeqCst); - let mut session = Session::new_global().unwrap(); - replace!(session, fs::write::<&Path, &[u8]> => deny_write, fn(&Path, &[u8]) -> io::Result<()>) - .unwrap(); + let mut session = Session::new_global(); + replace!(session, fs::write::<&Path, &[u8]> => deny_write, fn(&Path, &[u8]) -> io::Result<()>); let error = save_members(Path::new("virtual/readonly/members.txt"), &["alice", "bob"]).unwrap_err(); assert_eq!(error.kind(), io::ErrorKind::PermissionDenied); @@ -234,8 +230,8 @@ fn filesystem_write_permission_errors_reach_the_caller() { fn file_open_can_report_a_missing_session_without_touching_disk() { let _serial = serial(); OPEN_CALLS.store(0, Ordering::SeqCst); - let mut session = Session::new_global().unwrap(); - replace!(session, File::open::<&str> => missing_session, fn(&str) -> io::Result).unwrap(); + let mut session = Session::new_global(); + replace!(session, File::open::<&str> => missing_session, fn(&str) -> io::Result); assert!(!has_saved_session("virtual/session.bin").unwrap()); assert_eq!(OPEN_CALLS.load(Ordering::SeqCst), 1); } @@ -244,8 +240,8 @@ fn file_open_can_report_a_missing_session_without_touching_disk() { fn file_open_permission_errors_are_not_treated_as_missing_data() { let _serial = serial(); OPEN_CALLS.store(0, Ordering::SeqCst); - let mut session = Session::new_global().unwrap(); - replace!(session, File::open::<&str> => deny_session, fn(&str) -> io::Result).unwrap(); + let mut session = Session::new_global(); + replace!(session, File::open::<&str> => deny_session, fn(&str) -> io::Result); assert_eq!( has_saved_session("virtual/session.bin").unwrap_err().kind(), io::ErrorKind::PermissionDenied @@ -257,9 +253,8 @@ fn file_open_permission_errors_are_not_treated_as_missing_data() { fn tcp_timeout_is_handled_without_connecting() { let _serial = serial(); CONNECT_CALLS.store(0, Ordering::SeqCst); - let mut session = Session::new_global().unwrap(); - replace!(session, TcpStream::connect_timeout => connection_timeout, fn(&SocketAddr, Duration) -> io::Result) - .unwrap(); + let mut session = Session::new_global(); + replace!(session, TcpStream::connect_timeout => connection_timeout, fn(&SocketAddr, Duration) -> io::Result); assert!(!upstream_is_available(SocketAddr::from(([192, 0, 2, 10], 443))).unwrap()); assert_eq!(CONNECT_CALLS.load(Ordering::SeqCst), 1); } @@ -268,9 +263,8 @@ fn tcp_timeout_is_handled_without_connecting() { fn tcp_connection_errors_reach_the_caller_without_connecting() { let _serial = serial(); CONNECT_CALLS.store(0, Ordering::SeqCst); - let mut session = Session::new_global().unwrap(); - replace!(session, TcpStream::connect_timeout => connection_refused, fn(&SocketAddr, Duration) -> io::Result) - .unwrap(); + let mut session = Session::new_global(); + replace!(session, TcpStream::connect_timeout => connection_refused, fn(&SocketAddr, Duration) -> io::Result); let error = upstream_is_available(SocketAddr::from(([192, 0, 2, 10], 443))).unwrap_err(); assert_eq!(error.kind(), io::ErrorKind::ConnectionRefused); assert_eq!(CONNECT_CALLS.load(Ordering::SeqCst), 1); @@ -284,9 +278,8 @@ fn udp_send_captures_the_metric_without_delivering_a_packet() { let sender = UdpSocket::bind("127.0.0.1:0").unwrap(); let receiver = UdpSocket::bind("127.0.0.1:0").unwrap(); let target = receiver.local_addr().unwrap(); - let mut session = Session::new_global().unwrap(); - replace!(session, UdpSocket::send_to:: => capture_datagram, fn(&UdpSocket, &[u8], SocketAddr) -> io::Result) - .unwrap(); + let mut session = Session::new_global(); + replace!(session, UdpSocket::send_to:: => capture_datagram, fn(&UdpSocket, &[u8], SocketAddr) -> io::Result); assert_eq!( emit_counter(&sender, target, "jobs.completed", 12).unwrap(), 19 @@ -311,9 +304,8 @@ fn udp_send_errors_reach_the_caller_without_delivering_a_packet() { let sender = UdpSocket::bind("127.0.0.1:0").unwrap(); let receiver = UdpSocket::bind("127.0.0.1:0").unwrap(); let target = receiver.local_addr().unwrap(); - let mut session = Session::new_global().unwrap(); - replace!(session, UdpSocket::send_to:: => deny_datagram, fn(&UdpSocket, &[u8], SocketAddr) -> io::Result) - .unwrap(); + let mut session = Session::new_global(); + replace!(session, UdpSocket::send_to:: => deny_datagram, fn(&UdpSocket, &[u8], SocketAddr) -> io::Result); assert_eq!( emit_counter(&sender, target, "jobs.completed", 12) .unwrap_err() diff --git a/tests/native_expectations.rs b/tests/native_expectations.rs index f341574..e515707 100644 --- a/tests/native_expectations.rs +++ b/tests/native_expectations.rs @@ -26,44 +26,43 @@ extern "C-unwind" fn unwind_sum(a: i64, b: i64) -> i64 { #[test] fn native_calls_match_arguments_and_count_captured_responses() { let _serial = serial(); - let mut session = Session::new_global().unwrap(); - let sum = mock!(session, native_sum, extern "C" fn(i64, i64) -> i64).unwrap(); + 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()) - .unwrap(); + .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().unwrap(); + 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().unwrap(); - let sum = mock!(session, unchecked_sum, unsafe fn(i64, i64) -> i64).unwrap(); - sum.expect().once().returns(42).unwrap(); + let mut session = Session::new_global(); + let sum = mock!(session, unchecked_sum, unsafe fn(i64, i64) -> i64); + sum.expect().once().returns(42); // SAFETY: the function accepts all i64 values. assert_eq!(unsafe { unchecked_sum(6, 7) }, 42); - let system = mock!(session, system_sum, extern "system" fn(i64, i64) -> i64).unwrap(); - system.expect().once().returning(|a, b| a * b).unwrap(); + 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().unwrap(); + session.verify(); } #[test] fn unwind_abi_keeps_rust_panic_behavior() { let _serial = serial(); - let mut session = Session::new_global().unwrap(); - let sum = mock!(session, unwind_sum, extern "C-unwind" fn(i64, i64) -> i64).unwrap(); - sum.expect().once().panics("chosen failure").unwrap(); + 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().unwrap(); + session.restore(); assert_eq!(unwind_sum(6, 7), 13); } @@ -71,9 +70,9 @@ fn unwind_abi_keeps_rust_panic_behavior() { 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().unwrap(); - let sum = mock!(session, native_sum, extern "C" fn(i64, i64) -> i64).unwrap(); - sum.expect().panics("native callback failed").unwrap(); + 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); } @@ -98,13 +97,12 @@ fn native_panic_does_not_cross_the_abi_boundary() { fn imported_system_function_is_mocked_without_a_wrapper() { let _serial = serial(); for constructor in [Session::new_global, Session::new] { - let mut session = constructor().unwrap(); + let mut session = constructor(); let hostname = mock!( session, libc::gethostname, unsafe extern "C" fn(*mut libc::c_char, usize) -> libc::c_int - ) - .unwrap(); + ); hostname .expect() .with(|_, size| *size == 64) @@ -115,14 +113,13 @@ fn imported_system_function_is_mocked_without_a_wrapper() { // SAFETY: the caller provides a writable buffer of this size. unsafe { std::ptr::copy_nonoverlapping(name.as_ptr().cast(), buffer, name.len()) }; 0 - }) - .unwrap(); + }); let mut buffer = [0u8; 64]; // SAFETY: buffer is writable for all 64 bytes. 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().unwrap(); + session.restore(); } } @@ -135,27 +132,22 @@ fn imported_system_function_is_mocked_without_a_wrapper() { } let _serial = serial(); for constructor in [Session::new_global, Session::new] { - let mut session = constructor().unwrap(); + let mut session = constructor(); let hostname = mock!( session, GetComputerNameW, unsafe extern "system" fn(*mut u16, *mut u32) -> i32 - ) - .unwrap(); - hostname - .expect() - .once() - .returning(|buffer, size| { - let name: Vec<_> = "test-host\0".encode_utf16().collect(); - // SAFETY: the caller supplies a valid size and a writable buffer. - unsafe { - assert!(*size as usize >= name.len()); - std::ptr::copy_nonoverlapping(name.as_ptr(), buffer, name.len()); - *size = (name.len() - 1) as u32; - } - 1 - }) - .unwrap(); + ); + hostname.expect().once().returning(|buffer, size| { + let name: Vec<_> = "test-host\0".encode_utf16().collect(); + // SAFETY: the caller supplies a valid size and a writable buffer. + 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; // SAFETY: both pointers are valid; size describes the writable buffer. @@ -165,6 +157,6 @@ fn imported_system_function_is_mocked_without_a_wrapper() { String::from_utf16(&buffer[..size as usize]).unwrap(), "test-host" ); - session.restore().unwrap(); + session.restore(); } } diff --git a/tests/parallel.rs b/tests/parallel.rs index d245374..bb45cf8 100644 --- a/tests/parallel.rs +++ b/tests/parallel.rs @@ -28,17 +28,16 @@ impl Meter { fn isolated(seed: u64) { for _ in 0..100 { - let mut session = Session::new().unwrap(); - let mock = mock!(session, value, fn(u64) -> u64).unwrap(); + let mut session = Session::new(); + let mock = mock!(session, value, fn(u64) -> u64); mock.expect() .with(move |arg| *arg == seed) .times(3) - .returns(seed + 20) - .unwrap(); + .returns(seed + 20); for _ in 0..3 { assert_eq!(value(seed), seed + 20); } - session.restore().unwrap(); + session.restore(); assert_eq!(value(seed), seed + 1); } } @@ -86,11 +85,11 @@ fn calls_from_another_thread_are_safe_during_setup_and_teardown() { // execute the entry while those bytes are written. Install once here so the // reader below never races that write; later installs reuse the same entry. { - let mut session = Session::new().unwrap(); - let mock = mock!(session, value, fn(u64) -> u64).unwrap(); - mock.expect().once().returns(0).unwrap(); + let mut session = Session::new(); + let mock = mock!(session, value, fn(u64) -> u64); + mock.expect().once().returns(0); assert_eq!(value(41), 0); - session.restore().unwrap(); + session.restore(); } let stop = Arc::new(AtomicBool::new(false)); let reader = { @@ -106,11 +105,11 @@ fn calls_from_another_thread_are_safe_during_setup_and_teardown() { }) }; for round in 0..200 { - let mut session = Session::new().unwrap(); - let mock = mock!(session, value, fn(u64) -> u64).unwrap(); - mock.expect().once().returns(round).unwrap(); + let mut session = Session::new(); + let mock = mock!(session, value, fn(u64) -> u64); + mock.expect().once().returns(round); assert_eq!(value(41), round); - session.restore().unwrap(); + session.restore(); } stop.store(true, Ordering::Relaxed); assert!(reader.join().unwrap() > 0); @@ -126,15 +125,15 @@ fn struct_methods_keep_one_result_per_thread() { .map(|result| { let start = Arc::clone(&start); std::thread::spawn(move || { - let mut session = Session::new().unwrap(); - let readings = mock!(session, Meter::reading, fn(&Meter, u64) -> u64).unwrap(); - readings.expect().times(50).returns(result).unwrap(); + let mut session = Session::new(); + let readings = mock!(session, Meter::reading, fn(&Meter, u64) -> u64); + readings.expect().times(50).returns(result); start.wait(); let local = Meter { offset: 10 }; for _ in 0..50 { assert_eq!(local.reading(5), result); } - session.restore().unwrap(); + session.restore(); }) }) .collect(); @@ -155,20 +154,16 @@ fn unit_functions_keep_one_behavior_per_thread() { let observer = { let start = Arc::clone(&start); std::thread::spawn(move || { - let mut session = Session::new().unwrap(); - let discards = mock!(session, discard, fn(u64)).unwrap(); - discards - .expect() - .times(20) - .returning(|_| { - CALLS.fetch_add(1, Ordering::SeqCst); - }) - .unwrap(); + let mut session = Session::new(); + let discards = mock!(session, discard, fn(u64)); + discards.expect().times(20).returning(|_| { + CALLS.fetch_add(1, Ordering::SeqCst); + }); start.wait(); for _ in 0..20 { discard(0); } - session.restore().unwrap(); + session.restore(); }) }; start.wait(); @@ -186,18 +181,17 @@ fn owned_results_keep_one_value_per_thread() { let worker = { let start = Arc::clone(&start); std::thread::spawn(move || { - let mut session = Session::new().unwrap(); - let labels = mock!(session, label, fn(u64) -> String).unwrap(); + let mut session = Session::new(); + let labels = mock!(session, label, fn(u64) -> String); labels .expect() .times(30) - .returning(|seed| format!("worker {seed}")) - .unwrap(); + .returning(|seed| format!("worker {seed}")); start.wait(); for _ in 0..30 { assert_eq!(label(8), "worker 8"); } - session.restore().unwrap(); + session.restore(); }) }; start.wait(); diff --git a/tests/readme.rs b/tests/readme.rs index 703b2b2..9feae85 100644 --- a/tests/readme.rs +++ b/tests/readme.rs @@ -21,22 +21,21 @@ mod why_shimforge { use std::io; #[test] - fn claim_slot_succeeds_without_a_real_directory() -> Result<(), shimforge::Error> { - let mut session = Session::new()?; + fn claim_slot_succeeds_without_a_real_directory() { + let mut session = Session::new(); let create = mock!( session, fs::create_dir_all::<&str>, fn(&str) -> io::Result<()> - )?; + ); create .expect() .with(|path| *path == "/var/run/dispatcher") .once() - .returning(|_| Ok(()))?; + .returning(|_| Ok(())); assert!(claim_slot().is_ok()); - session.verify()?; - Ok(()) + session.verify(); } } @@ -53,15 +52,14 @@ mod thread_local_and_global_sessions { } #[test] - fn other_threads_keep_the_original_function() -> Result<(), shimforge::Error> { - let mut session = Session::new()?; - let count = mock!(session, worker_count, fn() -> usize)?; - count.expect().returns(16)?; + fn other_threads_keep_the_original_function() { + let mut session = Session::new(); + let count = mock!(session, worker_count, fn() -> usize); + count.expect().returns(16); assert_eq!(worker_count(), 16); // A thread with no session of its own still calls the original. assert_eq!(std::thread::spawn(worker_count).join().unwrap(), 2); - Ok(()) } } @@ -78,15 +76,14 @@ mod constant_results { } #[test] - fn a_constant_result_answers_every_call() -> Result<(), shimforge::Error> { - let mut session = Session::new()?; - let exists = mock!(session, Path::exists, fn(&Path) -> bool)?; - exists.expect().returns(true)?; + fn a_constant_result_answers_every_call() { + let mut session = Session::new(); + let exists = mock!(session, Path::exists, fn(&Path) -> bool); + exists.expect().returns(true); assert_eq!(export_state(Path::new("virtual/export.done")), "finished"); - session.restore()?; + session.restore(); assert_eq!(export_state(Path::new("virtual/export.done")), "running"); - Ok(()) } } @@ -102,22 +99,21 @@ mod matching_arguments_and_counting_calls { } #[test] - fn matching_calls_are_counted() -> Result<(), shimforge::Error> { - let mut session = Session::new()?; + fn matching_calls_are_counted() { + let mut session = Session::new(); let read = mock!( session, fs::read_to_string::<&Path>, fn(&Path) -> io::Result - )?; + ); read.expect() .with(|path| *path == Path::new("service.port")) .times(2) - .returning(|_| Ok("8080\n".to_owned()))?; + .returning(|_| Ok("8080\n".to_owned())); assert_eq!(load_port(Path::new("service.port")).unwrap(), 8080); assert_eq!(load_port(Path::new("service.port")).unwrap(), 8080); - session.verify()?; - Ok(()) + session.verify(); } } @@ -133,10 +129,10 @@ mod return_values_and_call_order { } #[test] - fn results_follow_the_call_order() -> Result<(), shimforge::Error> { - let mut session = Session::new()?; - let next = mock!(session, next_id, fn() -> u64)?; - let save = mock!(session, save_id, fn(u64) -> bool)?; + fn results_follow_the_call_order() { + let mut session = Session::new(); + let next = mock!(session, next_id, fn() -> u64); + let save = mock!(session, save_id, fn(u64) -> bool); let order = Sequence::new(); let mut id = 40; next.expect() @@ -145,18 +141,17 @@ mod return_values_and_call_order { .returning(move || { id += 1; id - })?; + }); save.expect() .with(|id| *id == 42) .once() .in_sequence(&order) - .returns(true)?; + .returns(true); assert_eq!(next_id(), 41); assert_eq!(next_id(), 42); assert!(save_id(42)); - session.verify()?; - Ok(()) + session.verify(); } } @@ -174,9 +169,9 @@ mod writing_through_reference_parameters { } #[test] - fn a_mock_fills_output_parameters() -> Result<(), shimforge::Error> { - let mut session = Session::new()?; - let split = mock!(session, split_amount, fn(u64, &mut u64, &mut u64))?; + fn a_mock_fills_output_parameters() { + let mut session = Session::new(); + let split = mock!(session, split_amount, fn(u64, &mut u64, &mut u64)); split .expect() .with(|total, _, _| *total == 1234) @@ -184,12 +179,12 @@ mod writing_through_reference_parameters { .returning(|_, whole, cents| { *whole = 99; *cents = 5; - })?; - let write = mock!(session, fill, fn(&mut [u8]) -> usize)?; + }); + let write = mock!(session, fill, fn(&mut [u8]) -> usize); write.expect().once().returning(|buffer| { buffer[..2].copy_from_slice(b"ok"); 2 - })?; + }); let (mut whole, mut cents) = (0, 0); split_amount(1234, &mut whole, &mut cents); @@ -198,7 +193,6 @@ mod writing_through_reference_parameters { let mut buffer = [0; 8]; assert_eq!(fill(&mut buffer), 2); assert_eq!(&buffer[..2], b"ok"); - Ok(()) } } @@ -222,16 +216,16 @@ mod methods_and_generic_functions { } #[test] - fn a_method_and_one_generic_instance_are_mocked() -> Result<(), shimforge::Error> { - let mut session = Session::new()?; - let rates = mock!(session, Cache::hit_rate, fn(&Cache, &str) -> f32)?; + fn a_method_and_one_generic_instance_are_mocked() { + let mut session = Session::new(); + let rates = mock!(session, Cache::hit_rate, fn(&Cache, &str) -> f32); rates .expect() .with(|cache, key| cache.region == "eu" && *key == "sessions") .once() - .returns(0.75)?; - let rendered = mock!(session, render::, fn(u8) -> String)?; - rendered.expect().once().returns("mocked".to_owned())?; + .returns(0.75); + let rendered = mock!(session, render::, fn(u8) -> String); + rendered.expect().once().returns("mocked".to_owned()); let cache = Cache { region: "eu".to_owned(), @@ -240,7 +234,6 @@ mod methods_and_generic_functions { assert_eq!(render(7u8), "mocked"); // A different type argument is a different function. assert_eq!(render("7"), "live 7"); - Ok(()) } } @@ -256,15 +249,14 @@ mod replacing_a_whole_function { } #[test] - fn a_function_or_closure_replaces_the_original() -> Result<(), shimforge::Error> { - let mut session = Session::new()?; - replace!(session, checksum => fixed_checksum, fn(&[u8]) -> u32)?; + fn a_function_or_closure_replaces_the_original() { + let mut session = Session::new(); + replace!(session, checksum => fixed_checksum, fn(&[u8]) -> u32); assert_eq!(checksum(b"abc"), 7); - session.restore()?; + session.restore(); - replace!(session, checksum => |_| 9, fn(&[u8]) -> u32)?; + replace!(session, checksum => |_| 9, fn(&[u8]) -> u32); assert_eq!(checksum(b"abc"), 9); - Ok(()) } } @@ -302,26 +294,25 @@ mod async_functions { } #[test] - fn async_functions_return_mocked_results() -> Result<(), shimforge::Error> { - let mut session = Session::new()?; - let rates = session.mock_async(exchange_rate(""))?; - rates.expect().once().returns(1.25)?; + fn async_functions_return_mocked_results() { + let mut session = Session::new(); + let rates = session.mock_async(exchange_rate("")); + rates.expect().once().returns(1.25); // A throwaway receiver is enough to name the future type. let balances = session.mock_async( Ledger { name: String::new(), } .balance(), - )?; - balances.expect().once().returns(4_200)?; + ); + balances.expect().once().returns(4_200); assert_eq!(ready(exchange_rate("EURUSD")), 1.25); let ledger = Ledger { name: "payroll".to_owned(), }; assert_eq!(ready(ledger.balance()), 4_200); - session.verify()?; - Ok(()) + session.verify(); } } @@ -350,17 +341,17 @@ mod client_libraries_that_return_boxed_futures { } #[test] - fn a_client_method_that_returns_a_boxed_future_is_mocked() -> Result<(), shimforge::Error> { - let mut session = Session::new()?; + fn a_client_method_that_returns_a_boxed_future_is_mocked() { + let mut session = Session::new(); let get = mock!( session, Client::get, for<'a> fn(&'a Client, &'a str) -> Call<'a> - )?; + ); get.expect() .with(|_, path| *path == "/health") .once() - .returning(|_, _| Box::pin(async { Ok("healthy".to_owned()) }))?; + .returning(|_, _| Box::pin(async { Ok("healthy".to_owned()) })); let client = Client { endpoint: "https://inventory.invalid".to_owned(), @@ -371,8 +362,7 @@ mod client_libraries_that_return_boxed_futures { response.as_mut().poll(&mut context), Poll::Ready(Ok(body)) if body == "healthy" )); - session.verify()?; - Ok(()) + session.verify(); } } @@ -385,13 +375,13 @@ mod system_and_c_runtime_functions { } #[test] - fn getenv_reports_a_mocked_variable() -> Result<(), shimforge::Error> { - let mut session = Session::new()?; + fn getenv_reports_a_mocked_variable() { + let mut session = Session::new(); let lookup = mock!( session, getenv, unsafe extern "C" fn(*const c_char) -> *mut c_char - )?; + ); lookup .expect() .with(|name| { @@ -400,15 +390,14 @@ mod system_and_c_runtime_functions { name == c"DEPLOY_SLOT" }) .once() - .returning(|_| c"canary".as_ptr().cast_mut())?; + .returning(|_| c"canary".as_ptr().cast_mut()); // Any other variable keeps reporting that it is unset. - lookup.expect().returning(|_| std::ptr::null_mut())?; + lookup.expect().returning(|_| std::ptr::null_mut()); let key = CString::new("DEPLOY_SLOT").unwrap(); // SAFETY: the key is a valid C string and the result is only read. let slot = unsafe { CStr::from_ptr(getenv(key.as_ptr())) }; assert_eq!(slot.to_str().unwrap(), "canary"); - Ok(()) } } @@ -424,16 +413,15 @@ mod raw_replacement_without_signature_checks { } #[test] - fn a_raw_replacement_swaps_one_function_for_another() -> Result<(), shimforge::Error> { - let mut session = Session::new_global()?; + fn a_raw_replacement_swaps_one_function_for_another() { + let mut session = Session::new_global(); // SAFETY: both functions are live, share a signature, and stay loaded until the // session restores them. No thread calls them while the patch is installed. - unsafe { session.replace_raw(slot_count as *const (), fake_slot_count as *const ())? }; + unsafe { session.replace_raw(slot_count as *const (), fake_slot_count as *const ()) }; assert_eq!(slot_count(), 64); - session.restore()?; + session.restore(); assert_eq!(slot_count(), 4); - Ok(()) } } diff --git a/tests/replacement.rs b/tests/replacement.rs index 95640c2..dc6a0ca 100644 --- a/tests/replacement.rs +++ b/tests/replacement.rs @@ -8,6 +8,19 @@ fn serial_test() -> MutexGuard<'static, ()> { TEST_LOCK.lock().unwrap_or_else(|error| error.into_inner()) } +/// Runs `action`, fails the test unless it panics, and returns the panic message. +fn panic_message(action: impl FnOnce() -> T) -> String { + let panic = catch_unwind(AssertUnwindSafe(action)) + .err() + .expect("expected a panic"); + match panic.downcast::() { + Ok(message) => *message, + Err(panic) => panic + .downcast_ref::<&str>() + .map_or_else(String::new, |message| (*message).to_owned()), + } +} + fn increment(value: i64) -> i64 { value + 1 } @@ -33,8 +46,8 @@ fn direct_calls_are_replaced_and_restored_on_drop() { let _serial = serial_test(); assert_eq!(increment(7), 8); { - let mut session = Session::new_global().unwrap(); - replace!(session, increment => add_ten, fn(i64) -> i64).unwrap(); + let mut session = Session::new_global(); + replace!(session, increment => add_ten, fn(i64) -> i64); assert_eq!(increment(7), 17); assert_eq!(add_ten(7), 17); } @@ -44,13 +57,13 @@ fn direct_calls_are_replaced_and_restored_on_drop() { #[test] fn multiple_replacements_restore_together() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); - replace!(session, increment => add_ten, fn(i64) -> i64).unwrap(); - replace!(session, multiply => triple, fn(i64) -> i64).unwrap(); + let mut session = Session::new_global(); + replace!(session, increment => add_ten, fn(i64) -> i64); + replace!(session, multiply => triple, fn(i64) -> i64); assert_eq!(increment(7), 17); assert_eq!(multiply(7), 21); - session.restore().unwrap(); + session.restore(); assert_eq!(increment(7), 8); assert_eq!(multiply(7), 14); } @@ -58,17 +71,17 @@ fn multiple_replacements_restore_together() { #[test] fn explicit_restore_is_idempotent_and_session_can_be_reused() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); - session.restore().unwrap(); - replace!(session, increment => add_ten, fn(i64) -> i64).unwrap(); + let mut session = Session::new_global(); + session.restore(); + replace!(session, increment => add_ten, fn(i64) -> i64); assert_eq!(increment(3), 13); - session.restore().unwrap(); - session.restore().unwrap(); + session.restore(); + session.restore(); assert_eq!(increment(3), 4); - replace!(session, increment => add_hundred, fn(i64) -> i64).unwrap(); + replace!(session, increment => add_hundred, fn(i64) -> i64); assert_eq!(increment(3), 103); - session.restore().unwrap(); + session.restore(); assert_eq!(increment(3), 4); } @@ -76,83 +89,81 @@ fn explicit_restore_is_idempotent_and_session_can_be_reused() { fn panic_unwinding_restores_and_releases_session() { let _serial = serial_test(); let result = catch_unwind(AssertUnwindSafe(|| { - let mut session = Session::new_global().unwrap(); - replace!(session, increment => add_ten, fn(i64) -> i64).unwrap(); + let mut session = Session::new_global(); + replace!(session, increment => add_ten, fn(i64) -> i64); assert_eq!(increment(2), 12); panic!("exercise session cleanup"); })); assert!(result.is_err()); assert_eq!(increment(2), 3); - let mut session = Session::new_global().unwrap(); - replace!(session, increment => add_hundred, fn(i64) -> i64).unwrap(); + let mut session = Session::new_global(); + replace!(session, increment => add_hundred, fn(i64) -> i64); assert_eq!(increment(2), 102); } #[test] fn nested_sessions_fail_without_blocking() { let _serial = serial_test(); - let session = Session::new_global().unwrap(); + let session = Session::new_global(); assert!(matches!(Session::try_new_global(), Err(Error::Busy))); drop(session); - assert!(Session::new_global().is_ok()); + drop(Session::new_global()); } #[test] fn session_exclusivity_extends_to_other_threads() { let _serial = serial_test(); - let session = Session::new_global().unwrap(); + let session = Session::new_global(); assert!( std::thread::spawn(|| matches!(Session::try_new_global(), Err(Error::Busy))) .join() .unwrap() ); drop(session); - assert!( - std::thread::spawn(|| Session::new_global().is_ok()) - .join() - .unwrap() - ); + std::thread::spawn(|| drop(Session::new_global())) + .join() + .unwrap(); } #[test] fn installed_replacement_is_visible_to_a_worker_thread() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); - replace!(session, increment => add_ten, fn(i64) -> i64).unwrap(); + let mut session = Session::new_global(); + replace!(session, increment => add_ten, fn(i64) -> i64); // Start workers after installation and join them before restoration. let answer = std::thread::spawn(|| increment(31)).join().unwrap(); assert_eq!(answer, 41); - session.restore().unwrap(); + session.restore(); assert_eq!(increment(31), 32); } #[test] fn duplicate_replacement_fails_without_losing_original() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); - replace!(session, increment => add_ten, fn(i64) -> i64).unwrap(); + let mut session = Session::new_global(); + replace!(session, increment => add_ten, fn(i64) -> i64); assert_eq!( - replace!(session, increment => add_hundred, fn(i64) -> i64), - Err(Error::Overlap) + panic_message(|| replace!(session, increment => add_hundred, fn(i64) -> i64)), + Error::Overlap.to_string() ); assert_eq!(increment(9), 19); - session.restore().unwrap(); + session.restore(); assert_eq!(increment(9), 10); } #[test] fn replacement_cycles_are_rejected_without_disturbing_existing_patches() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); - replace!(session, increment => add_ten, fn(i64) -> i64).unwrap(); + let mut session = Session::new_global(); + replace!(session, increment => add_ten, fn(i64) -> i64); assert_eq!( - replace!(session, add_ten => increment, fn(i64) -> i64), - Err(Error::Overlap) + panic_message(|| replace!(session, add_ten => increment, fn(i64) -> i64)), + Error::Overlap.to_string() ); assert_eq!(increment(9), 19); assert_eq!(add_ten(9), 19); - session.restore().unwrap(); + session.restore(); assert_eq!(increment(9), 10); assert_eq!(add_ten(9), 19); } @@ -165,8 +176,8 @@ fn panic_replacement(value: i64) -> i64 { fn replacement_can_unwind_through_the_original_caller() { let _serial = serial_test(); let result = catch_unwind(AssertUnwindSafe(|| { - let mut session = Session::new_global().unwrap(); - replace!(session, increment => panic_replacement, fn(i64) -> i64).unwrap(); + let mut session = Session::new_global(); + replace!(session, increment => panic_replacement, fn(i64) -> i64); increment(42) })); let payload = result.unwrap_err(); @@ -175,16 +186,16 @@ fn replacement_can_unwind_through_the_original_caller() { "replacement rejected 42" ); assert_eq!(increment(42), 43); - assert!(Session::new_global().is_ok()); + drop(Session::new_global()); } #[test] fn replacing_a_function_with_itself_is_rejected() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); + let mut session = Session::new_global(); assert_eq!( - replace!(session, increment => increment, fn(i64) -> i64), - Err(Error::SameAddress) + panic_message(|| replace!(session, increment => increment, fn(i64) -> i64)), + Error::SameAddress.to_string() ); assert_eq!(increment(9), 10); } @@ -200,10 +211,10 @@ fn fake_greeting(name: &str) -> String { #[test] fn owned_return_values_preserve_ownership() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); - replace!(session, greeting => fake_greeting, fn(&str) -> String).unwrap(); + let mut session = Session::new_global(); + replace!(session, greeting => fake_greeting, fn(&str) -> String); let answer = greeting("Rust"); - session.restore().unwrap(); + session.restore(); assert_eq!(answer, "welcome Rust"); assert_eq!(greeting("Rust"), "hello Rust"); } @@ -231,8 +242,8 @@ fn fake_aggregate(seed: u64, label: &str) -> Aggregate { #[test] fn large_aggregate_return_keeps_the_hidden_return_pointer() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); - replace!(session, aggregate => fake_aggregate, fn(u64, &str) -> Aggregate).unwrap(); + let mut session = Session::new_global(); + replace!(session, aggregate => fake_aggregate, fn(u64, &str) -> Aggregate); assert_eq!( aggregate(25, "payload"), Aggregate { @@ -240,7 +251,7 @@ fn large_aggregate_return_keeps_the_hidden_return_pointer() { label: "fake payload".to_owned(), } ); - session.restore().unwrap(); + session.restore(); assert_eq!(aggregate(25, "payload").values, [25; 12]); } @@ -286,18 +297,17 @@ fn fake_many_arguments( #[test] fn replacement_preserves_integer_float_and_stack_arguments() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); + let mut session = Session::new_global(); replace!( session, many_arguments => fake_many_arguments, fn(u64, u64, u64, u64, u64, u64, u64, u64, f64, f64, f64, f64, f64) -> f64 - ) - .unwrap(); + ); assert_eq!( many_arguments(1, 2, 3, 4, 5, 6, 7, 8, 0.5, 1.0, 1.5, 2.0, 2.5), 231.5 ); - session.restore().unwrap(); + session.restore(); assert_eq!( many_arguments(1, 2, 3, 4, 5, 6, 7, 8, 0.5, 1.0, 1.5, 2.0, 2.5), 43.5 @@ -307,19 +317,18 @@ fn replacement_preserves_integer_float_and_stack_arguments() { #[test] fn expectations_preserve_integer_float_and_stack_arguments() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); + let mut session = Session::new_global(); let call = shimforge::mock!( session, many_arguments, fn(u64, u64, u64, u64, u64, u64, u64, u64, f64, f64, f64, f64, f64) -> f64 - ) - .unwrap(); - call.expect().once().returning(fake_many_arguments).unwrap(); + ); + call.expect().once().returning(fake_many_arguments); assert_eq!( many_arguments(1, 2, 3, 4, 5, 6, 7, 8, 0.5, 1.0, 1.5, 2.0, 2.5), 231.5 ); - session.restore().unwrap(); + session.restore(); assert_eq!( many_arguments(1, 2, 3, 4, 5, 6, 7, 8, 0.5, 1.0, 1.5, 2.0, 2.5), 43.5 @@ -338,12 +347,12 @@ fn borrowed_identity(value: &str) -> &str { fn replacement_preserves_the_borrowed_argument_lifetime() { let _serial = serial_test(); let owned = String::from("borrowed data"); - let mut session = Session::new_global().unwrap(); - replace!(session, borrowed => borrowed_identity, for<'a> fn(&'a str) -> &'a str).unwrap(); + let mut session = Session::new_global(); + replace!(session, borrowed => borrowed_identity, for<'a> fn(&'a str) -> &'a str); let answer = borrowed(owned.as_str()); assert_eq!(answer, owned); assert_eq!(answer.as_ptr(), owned.as_ptr()); - session.restore().unwrap(); + session.restore(); assert_eq!(answer, "borrowed data"); assert_eq!(borrowed(owned.as_str()), "b"); } @@ -368,11 +377,10 @@ impl Counter { fn methods_with_mutable_receivers_are_supported() { let _serial = serial_test(); let mut counter = Counter { value: 5 }; - let mut session = Session::new_global().unwrap(); - replace!(session, Counter::advance => Counter::fake_advance, fn(&mut Counter, i64) -> i64) - .unwrap(); + let mut session = Session::new_global(); + replace!(session, Counter::advance => Counter::fake_advance, fn(&mut Counter, i64) -> i64); assert_eq!(counter.advance(3), 35); - session.restore().unwrap(); + session.restore(); assert_eq!(counter.advance(2), 37); } @@ -387,11 +395,11 @@ fn fake_generic_size(seed: usize) -> usize { #[test] fn replacing_one_generic_instantiation_preserves_another() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); - replace!(session, generic_size:: => fake_generic_size, fn(usize) -> usize).unwrap(); + let mut session = Session::new_global(); + replace!(session, generic_size:: => fake_generic_size, fn(usize) -> usize); assert_eq!(generic_size::(1), 101); assert_eq!(generic_size::(1), 9); - session.restore().unwrap(); + session.restore(); assert_eq!(generic_size::(1), 5); assert_eq!(generic_size::(1), 9); } @@ -407,20 +415,20 @@ extern "C" fn native_product(a: i64, b: i64) -> i64 { #[test] fn native_abi_functions_are_supported() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); - replace!(session, native_sum => native_product, extern "C" fn(i64, i64) -> i64).unwrap(); + let mut session = Session::new_global(); + replace!(session, native_sum => native_product, extern "C" fn(i64, i64) -> i64); assert_eq!(native_sum(6, 7), 42); - session.restore().unwrap(); + session.restore(); assert_eq!(native_sum(6, 7), 13); } #[test] fn noncapturing_closure_can_be_a_replacement() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); - replace!(session, increment => |value| value - 2, fn(i64) -> i64).unwrap(); + let mut session = Session::new_global(); + replace!(session, increment => |value| value - 2, fn(i64) -> i64); assert_eq!(increment(42), 40); - session.restore().unwrap(); + session.restore(); assert_eq!(increment(42), 43); } @@ -436,11 +444,11 @@ fn fake_touch(value: &mut usize) { fn unit_return_functions_are_supported() { let _serial = serial_test(); let mut value = 0; - let mut session = Session::new_global().unwrap(); - replace!(session, touch => fake_touch, fn(&mut usize)).unwrap(); + let mut session = Session::new_global(); + replace!(session, touch => fake_touch, fn(&mut usize)); touch(&mut value); assert_eq!(value, 5); - session.restore().unwrap(); + session.restore(); touch(&mut value); assert_eq!(value, 6); } @@ -456,11 +464,11 @@ unsafe fn unsafe_add_ten(value: i64) -> i64 { #[test] fn unsafe_function_installation_does_not_require_an_unsafe_block() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); - replace!(session, unsafe_increment => unsafe_add_ten, unsafe fn(i64) -> i64).unwrap(); + let mut session = Session::new_global(); + replace!(session, unsafe_increment => unsafe_add_ten, unsafe fn(i64) -> i64); // SAFETY: Both test functions accept any i64. assert_eq!(unsafe { unsafe_increment(5) }, 15); - session.restore().unwrap(); + session.restore(); // SAFETY: Both test functions accept any i64. assert_eq!(unsafe { unsafe_increment(5) }, 6); } @@ -484,14 +492,13 @@ unsafe extern "C" fn unsafe_native_product(a: i64, b: i64) -> i64 { #[test] fn system_abi_and_unsafe_native_functions_are_supported() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); - replace!(session, system_sum => system_product, extern "system" fn(i64, i64) -> i64).unwrap(); - replace!(session, unsafe_native_sum => unsafe_native_product, unsafe extern "C" fn(i64, i64) -> i64) - .unwrap(); + let mut session = Session::new_global(); + replace!(session, system_sum => system_product, extern "system" fn(i64, i64) -> i64); + replace!(session, unsafe_native_sum => unsafe_native_product, unsafe extern "C" fn(i64, i64) -> i64); assert_eq!(system_sum(6, 7), 42); // SAFETY: Both test functions accept any pair of i64 values. assert_eq!(unsafe { unsafe_native_sum(6, 7) }, 42); - session.restore().unwrap(); + session.restore(); assert_eq!(system_sum(6, 7), 13); // SAFETY: Both test functions accept any pair of i64 values. assert_eq!(unsafe { unsafe_native_sum(6, 7) }, 13); @@ -508,26 +515,26 @@ fn scaled_twice(value: i64) -> i64 { #[test] fn raw_replacement_skips_the_signature_check_and_reaches_every_thread() { let _serial = serial_test(); - let mut session = Session::new_global().unwrap(); + let mut session = Session::new_global(); // SAFETY: both functions are live, share a signature, and stay loaded while // the session holds the patch. No thread calls them during installation. unsafe { - session - .replace_raw(scaled as *const (), scaled_twice as *const ()) - .unwrap(); + session.replace_raw(scaled as *const (), scaled_twice as *const ()); } assert_eq!(scaled(5), 20); assert_eq!(std::thread::spawn(|| scaled(5)).join().unwrap(), 20); - session.restore().unwrap(); + session.restore(); assert_eq!(scaled(5), 10); } #[test] fn raw_replacement_rejects_a_local_session() { let _serial = serial_test(); - let mut session = Session::new_local().unwrap(); + let mut session = Session::new_local(); // SAFETY: both pointers are live functions; local mode refuses before access. - let result = unsafe { session.replace_raw(scaled as *const (), scaled_twice as *const ()) }; - assert!(matches!(result, Err(Error::Expectation(_)))); + let message = panic_message(|| unsafe { + session.replace_raw(scaled as *const (), scaled_twice as *const ()) + }); + assert!(message.contains("local session")); assert_eq!(scaled(5), 10); } diff --git a/tests/safe_api.rs b/tests/safe_api.rs index b9dc068..935239d 100644 --- a/tests/safe_api.rs +++ b/tests/safe_api.rs @@ -8,9 +8,9 @@ fn original(value: u64) -> u64 { #[test] fn public_api_works_when_the_calling_crate_forbids_unsafe_code() { - let mut session = Session::new_global().unwrap(); - replace!(session, original => |value| value + 10, fn(u64) -> u64).unwrap(); + let mut session = Session::new_global(); + replace!(session, original => |value| value + 10, fn(u64) -> u64); assert_eq!(original(2), 12); - session.restore().unwrap(); + session.restore(); assert_eq!(original(2), 3); } diff --git a/tests/stack_safety.rs b/tests/stack_safety.rs index 905fd2c..77b9b97 100644 --- a/tests/stack_safety.rs +++ b/tests/stack_safety.rs @@ -39,9 +39,9 @@ fn descend(depth: u32) -> u64 { #[test] fn deep_recursion_still_works_after_a_mock_is_installed() { - let mut session = Session::new().unwrap(); - let limit = mock!(session, retry_limit, fn() -> u32).unwrap(); - limit.expect().returns(9).unwrap(); + let mut session = Session::new(); + let limit = mock!(session, retry_limit, fn() -> u32); + limit.expect().returns(9); assert_eq!(retry_limit(), 9); assert_eq!(descend(1024), 1); @@ -58,21 +58,21 @@ fn threads_that_patch_at_the_same_time_keep_their_stacks() { thread::Builder::new() .stack_size(1 << 20) .spawn(move || { - let mut session = Session::new().unwrap(); + let mut session = Session::new(); match index % 3 { 0 => { - let limit = mock!(session, retry_limit, fn() -> u32).unwrap(); - limit.expect().returns(11).unwrap(); + let limit = mock!(session, retry_limit, fn() -> u32); + limit.expect().returns(11); assert_eq!(retry_limit(), 11); } 1 => { - let size = mock!(session, batch_size, fn() -> u32).unwrap(); - size.expect().returns(256).unwrap(); + let size = mock!(session, batch_size, fn() -> u32); + size.expect().returns(256); assert_eq!(batch_size(), 256); } _ => { - let tracing = mock!(session, tracing_enabled, fn() -> bool).unwrap(); - tracing.expect().returns(true).unwrap(); + let tracing = mock!(session, tracing_enabled, fn() -> bool); + tracing.expect().returns(true); assert!(tracing_enabled()); } } diff --git a/tests/thread_local.rs b/tests/thread_local.rs index de795ad..11fc338 100644 --- a/tests/thread_local.rs +++ b/tests/thread_local.rs @@ -13,6 +13,19 @@ fn serial_test() -> MutexGuard<'static, ()> { TEST_LOCK.lock().unwrap_or_else(|error| error.into_inner()) } +/// Runs `action`, fails the test unless it panics, and returns the panic message. +fn panic_message(action: impl FnOnce() -> T) -> String { + let panic = catch_unwind(AssertUnwindSafe(action)) + .err() + .expect("expected a panic"); + match panic.downcast::() { + Ok(message) => *message, + Err(panic) => panic + .downcast_ref::<&str>() + .map_or_else(String::new, |message| (*message).to_owned()), + } +} + fn increment(value: i64) -> i64 { value + 1 } @@ -30,22 +43,22 @@ extern "C" fn native(value: i64) -> i64 { } fn install_shared(session: &mut Session, result: i64) { - let mock = mock!(session, increment, fn(i64) -> i64).unwrap(); - mock.expect().once().returns(result).unwrap(); + let mock = mock!(session, increment, fn(i64) -> i64); + mock.expect().once().returns(result); } fn two_threads( install: impl FnOnce(&mut Session), worker_install: impl FnOnce(&mut Session) + Send + 'static, ) { - let mut session = Session::new_local().unwrap(); + let mut session = Session::new_local(); install(&mut session); let (ready_tx, ready_rx) = mpsc::channel(); let (run_tx, run_rx) = mpsc::channel(); let (done_tx, done_rx) = mpsc::channel(); let (drop_tx, drop_rx) = mpsc::channel(); let worker = std::thread::spawn(move || { - let mut session = Session::new_local().unwrap(); + let mut session = Session::new_local(); worker_install(&mut session); ready_tx.send(()).unwrap(); run_rx.recv().unwrap(); @@ -59,18 +72,18 @@ fn two_threads( done_rx.recv().unwrap(); drop_tx.send(()).unwrap(); worker.join().unwrap(); - session.restore().unwrap(); + session.restore(); assert_eq!(increment(10), 11); } #[test] fn unmocked_threads_call_the_original() { let _serial = serial_test(); - let mut session = Session::new_local().unwrap(); + let mut session = Session::new_local(); install_shared(&mut session, 90); assert_eq!(increment(3), 90); assert_eq!(std::thread::spawn(|| increment(3)).join().unwrap(), 4); - session.restore().unwrap(); + session.restore(); assert_eq!(increment(3), 4); } @@ -79,12 +92,12 @@ fn different_macro_sites_can_mock_the_same_function() { let _serial = serial_test(); two_threads( |session| { - let mock = mock!(session, increment, fn(i64) -> i64).unwrap(); - mock.expect().once().returns(100).unwrap(); + let mock = mock!(session, increment, fn(i64) -> i64); + mock.expect().once().returns(100); }, |session| { - let mock = mock!(session, increment, fn(i64) -> i64).unwrap(); - mock.expect().once().returns(200).unwrap(); + let mock = mock!(session, increment, fn(i64) -> i64); + mock.expect().once().returns(200); }, ); } @@ -101,13 +114,13 @@ fn one_macro_site_can_serve_parallel_sessions() { #[test] fn dropping_the_first_session_keeps_the_second_mock_alive() { let _serial = serial_test(); - let mut session = Session::new_local().unwrap(); + let mut session = Session::new_local(); install_shared(&mut session, 100); assert_eq!(increment(4), 100); let (ready_tx, ready_rx) = mpsc::channel(); let (run_tx, run_rx) = mpsc::channel(); let worker = std::thread::spawn(move || { - let mut session = Session::new_local().unwrap(); + let mut session = Session::new_local(); install_shared(&mut session, 200); ready_tx.send(()).unwrap(); run_rx.recv().unwrap(); @@ -125,14 +138,14 @@ fn dropping_the_first_session_keeps_the_second_mock_alive() { fn panic_restores_the_original_and_releases_the_local_session() { let _serial = serial_test(); let result = catch_unwind(AssertUnwindSafe(|| { - let mut session = Session::new_local().unwrap(); - let mock = mock!(session, increment, fn(i64) -> i64).unwrap(); - mock.expect().once().panics("local failure").unwrap(); + let mut session = Session::new_local(); + let mock = mock!(session, increment, fn(i64) -> i64); + mock.expect().once().panics("local failure"); increment(1); })); assert!(result.is_err()); assert_eq!(increment(1), 2); - let mut session = Session::new_local().unwrap(); + let mut session = Session::new_local(); install_shared(&mut session, 40); assert_eq!(increment(1), 40); } @@ -140,18 +153,24 @@ fn panic_restores_the_original_and_releases_the_local_session() { #[test] fn duplicate_installation_can_be_retried_after_restore() { let _serial = serial_test(); - let mut session = Session::new_local().unwrap(); - let first = mock!(session, increment, fn(i64) -> i64).unwrap(); - first.expect().once().returns(80).unwrap(); + let mut session = Session::new_local(); + let first = mock!(session, increment, fn(i64) -> i64); + first.expect().once().returns(80); for should_succeed in [false, true] { - let result = mock!(session, increment, fn(i64) -> i64); + let result = catch_unwind(AssertUnwindSafe(|| { + mock!(session, increment, fn(i64) -> i64) + })); if should_succeed { - result.unwrap().expect().once().returns(90).unwrap(); + result.unwrap().expect().once().returns(90); assert_eq!(increment(1), 90); } else { - assert!(matches!(result, Err(Error::Overlap))); + let panic = result.err().expect("a duplicate mock must panic"); + assert_eq!( + panic.downcast_ref::(), + Some(&Error::Overlap.to_string()) + ); assert_eq!(increment(1), 80); - session.restore().unwrap(); + session.restore(); } } } @@ -159,17 +178,17 @@ fn duplicate_installation_can_be_retried_after_restore() { #[test] fn global_and_local_sessions_exclude_each_other() { let _serial = serial_test(); - let local = Session::new_local().unwrap(); + let local = Session::new_local(); assert!(matches!(Session::try_new_local(), Err(Error::Busy))); assert!(matches!(Session::try_new_global(), Err(Error::Busy))); std::thread::spawn(|| { assert!(matches!(Session::try_new_global(), Err(Error::Busy))); - assert!(Session::new_local().is_ok()); + drop(Session::new_local()); }) .join() .unwrap(); drop(local); - let global = Session::new_global().unwrap(); + let global = Session::new_global(); assert!(matches!(Session::try_new_local(), Err(Error::Busy))); std::thread::spawn(|| { assert!(matches!(Session::try_new_local(), Err(Error::Busy))); @@ -177,17 +196,17 @@ fn global_and_local_sessions_exclude_each_other() { .join() .unwrap(); drop(global); - assert!(Session::new_local().is_ok()); + drop(Session::new_local()); } #[test] fn replacements_use_local_mode_by_default() { let _serial = serial_test(); - let mut session = Session::new().unwrap(); - shimforge::replace!(session, increment => |x| x + 10, fn(i64) -> i64).unwrap(); + let mut session = Session::new(); + shimforge::replace!(session, increment => |x| x + 10, fn(i64) -> i64); assert_eq!(increment(1), 11); assert_eq!(std::thread::spawn(|| increment(1)).join().unwrap(), 2); - session.restore().unwrap(); + session.restore(); assert_eq!(increment(1), 2); } @@ -199,33 +218,32 @@ fn a_function_can_switch_between_local_and_global_modes() { Session::new_global() } else { Session::new() - } - .unwrap(); + }; install_shared(&mut session, 70); assert_eq!(increment(1), 70); - session.restore().unwrap(); + session.restore(); assert_eq!(increment(1), 2); } - let mut session = Session::new_global().unwrap(); - shimforge::replace!(session, increment => |value| value + 10, fn(i64) -> i64).unwrap(); + let mut session = Session::new_global(); + shimforge::replace!(session, increment => |value| value + 10, fn(i64) -> i64); assert_eq!(increment(1), 11); assert_eq!(std::thread::spawn(|| increment(1)).join().unwrap(), 11); - assert!(shimforge::replace!(session, increment => |value| value + 20, fn(i64) -> i64).is_err()); - session.restore().unwrap(); + panic_message(|| shimforge::replace!(session, increment => |value| value + 20, fn(i64) -> i64)); + session.restore(); assert_eq!(increment(1), 2); } #[test] fn local_replacements_keep_borrows_and_native_arguments() { let _serial = serial_test(); - let mut session = Session::new().unwrap(); - shimforge::replace!(session, borrowed => |value| &value[1..], fn(&str) -> &str).unwrap(); + let mut session = Session::new(); + shimforge::replace!(session, borrowed => |value| &value[1..], fn(&str) -> &str); assert_eq!(borrowed("word"), "ord"); assert_eq!( std::thread::spawn(|| borrowed("word")).join().unwrap(), "word" ); - shimforge::replace!(session, native => native_replacement, extern "C" fn(i64) -> i64).unwrap(); + shimforge::replace!(session, native => native_replacement, extern "C" fn(i64) -> i64); assert_eq!(native(1), 31); assert_eq!(std::thread::spawn(|| native(1)).join().unwrap(), 3); } @@ -244,9 +262,9 @@ fn failing_callee() -> usize { #[test] fn original_calls_can_unwind_through_the_saved_entry() { let _serial = serial_test(); - let mut session = Session::new().unwrap(); - let mock = mock!(session, first_call, fn() -> usize).unwrap(); - mock.expect().once().returns(20).unwrap(); + let mut session = Session::new(); + let mock = mock!(session, first_call, fn() -> usize); + mock.expect().once().returns(20); assert_eq!(first_call(), 20); std::thread::spawn(|| { let error = catch_unwind(first_call).unwrap_err(); @@ -254,7 +272,7 @@ fn original_calls_can_unwind_through_the_saved_entry() { }) .join() .unwrap(); - session.restore().unwrap(); + session.restore(); assert!(catch_unwind(first_call).is_err()); } @@ -266,15 +284,14 @@ fn filesystem_replacements_need_no_wrappers() { io, path::Path, }; - let mut session = Session::new().unwrap(); - shimforge::replace!(session, fs::read::<&Path> => |_| Ok(b"file contents".to_vec()), fn(&Path) -> io::Result>).unwrap(); + let mut session = Session::new(); + shimforge::replace!(session, fs::read::<&Path> => |_| Ok(b"file contents".to_vec()), fn(&Path) -> io::Result>); shimforge::replace!(session, fs::write::<&Path, &[u8]> => |path, data| { assert_eq!(path, Path::new("virtual/output")); assert_eq!(data, b"data"); Ok(()) - }, fn(&Path, &[u8]) -> io::Result<()>) - .unwrap(); - shimforge::replace!(session, File::open::<&str> => |_| Err(io::ErrorKind::PermissionDenied.into()), fn(&str) -> io::Result).unwrap(); + }, fn(&Path, &[u8]) -> io::Result<()>); + shimforge::replace!(session, File::open::<&str> => |_| Err(io::ErrorKind::PermissionDenied.into()), fn(&str) -> io::Result); assert_eq!( fs::read(Path::new("virtual/input")).unwrap(), b"file contents" @@ -291,23 +308,21 @@ fn file_mocks_do_not_intercept_memory_inspection() { let _serial = serial_test(); use std::{fs, io, path::Path}; for constructor in [Session::new, Session::new_global] { - let mut session = constructor().unwrap(); + let mut session = constructor(); let read = mock!( session, fs::read_to_string::<&str>, fn(&str) -> io::Result - ) - .unwrap(); + ); read.expect() .once() - .returning(|_| Ok("mock contents".to_owned())) - .unwrap(); - let exists = mock!(session, Path::exists, fn(&Path) -> bool).unwrap(); - exists.expect().once().returns(false).unwrap(); + .returning(|_| Ok("mock contents".to_owned())); + let exists = mock!(session, Path::exists, fn(&Path) -> bool); + exists.expect().once().returns(false); assert_eq!(fs::read_to_string("virtual/file").unwrap(), "mock contents"); let executable = std::env::current_exe().unwrap(); assert!(!executable.exists()); - session.restore().unwrap(); + session.restore(); assert!(executable.exists()); } } @@ -320,19 +335,15 @@ fn native_async_mocks_can_switch_modes() { Session::new_global() } else { Session::new() - } - .unwrap(); - let mock = session.mock_async(fetch(1)).unwrap(); - mock.expect() - .times(if global { 2 } else { 1 }) - .returns(80) - .unwrap(); + }; + let mock = session.mock_async(fetch(1)); + mock.expect().times(if global { 2 } else { 1 }).returns(80); assert_eq!(ready(fetch(1)), 80); assert_eq!( std::thread::spawn(|| ready(fetch(1))).join().unwrap(), if global { 80 } else { 2 } ); - session.restore().unwrap(); + session.restore(); assert_eq!(ready(fetch(1)), 2); } } @@ -342,12 +353,12 @@ fn local_sessions_recover_after_a_global_panic() { let _serial = serial_test(); assert!( catch_unwind(|| { - let _session = Session::new_global().unwrap(); + let _session = Session::new_global(); panic!("global failure"); }) .is_err() ); - let mut session = Session::new_local().unwrap(); + let mut session = Session::new_local(); install_shared(&mut session, 40); assert_eq!(increment(1), 40); } @@ -359,9 +370,9 @@ fn aggregate(seed: u64, gain: f64) -> [u64; 12] { #[test] fn original_fallback_preserves_large_returns_and_float_arguments() { let _serial = serial_test(); - let mut session = Session::new_local().unwrap(); - let mock = mock!(session, aggregate, fn(u64, f64) -> [u64; 12]).unwrap(); - mock.expect().once().returns([40; 12]).unwrap(); + let mut session = Session::new_local(); + let mock = mock!(session, aggregate, fn(u64, f64) -> [u64; 12]); + mock.expect().once().returns([40; 12]); assert_eq!(aggregate(3, 4.0), [40; 12]); assert_eq!( std::thread::spawn(|| aggregate(3, 4.0)).join().unwrap(), @@ -377,9 +388,9 @@ fn may_panic(value: i64) -> i64 { #[test] fn original_fallback_can_unwind() { let _serial = serial_test(); - let mut session = Session::new_local().unwrap(); - let mock = mock!(session, may_panic, fn(i64) -> i64).unwrap(); - mock.expect().once().returns(10).unwrap(); + let mut session = Session::new_local(); + let mock = mock!(session, may_panic, fn(i64) -> i64); + mock.expect().once().returns(10); std::thread::spawn(|| assert!(catch_unwind(|| may_panic(0)).is_err())) .join() .unwrap(); @@ -389,14 +400,13 @@ fn original_fallback_can_unwind() { #[test] fn borrowed_and_owned_results_keep_their_original_fallback() { let _serial = serial_test(); - let mut session = Session::new_local().unwrap(); - let mock = mock!(session, borrowed, fn(&str) -> &str).unwrap(); - mock.expect().once().returning(|value| value).unwrap(); - let mock = mock!(session, owned, fn(String) -> String).unwrap(); + let mut session = Session::new_local(); + let mock = mock!(session, borrowed, fn(&str) -> &str); + mock.expect().once().returning(|value| value); + let mock = mock!(session, owned, fn(String) -> String); mock.expect() .once() - .returning(|value| format!("mock {value}")) - .unwrap(); + .returning(|value| format!("mock {value}")); assert_eq!(borrowed(" value "), " value "); assert_eq!(owned("value".to_owned()), "mock value"); std::thread::spawn(|| { @@ -411,9 +421,9 @@ fn borrowed_and_owned_results_keep_their_original_fallback() { #[test] fn native_abi_mocks_are_thread_local() { let _serial = serial_test(); - let mut session = Session::new_local().unwrap(); - let mock = mock!(session, native, extern "C" fn(i64) -> i64).unwrap(); - mock.expect().once().returns(70).unwrap(); + let mut session = Session::new_local(); + let mock = mock!(session, native, extern "C" fn(i64) -> i64); + mock.expect().once().returns(70); assert_eq!(native(5), 70); assert_eq!(std::thread::spawn(|| native(5)).join().unwrap(), 7); } @@ -433,27 +443,27 @@ async fn fetch(value: i64) -> i64 { #[test] fn async_mocks_leave_other_threads_unchanged() { let _serial = serial_test(); - let mut session = Session::new_local().unwrap(); - let mock = session.mock_async(fetch(0)).unwrap(); - mock.expect().once().returns(60).unwrap(); + let mut session = Session::new_local(); + let mock = session.mock_async(fetch(0)); + mock.expect().once().returns(60); assert_eq!(ready(fetch(5)), 60); assert_eq!(std::thread::spawn(|| ready(fetch(5))).join().unwrap(), 6); - session.restore().unwrap(); + session.restore(); assert_eq!(ready(fetch(5)), 6); } #[test] fn async_sessions_can_share_a_poll_function_and_drop_in_either_order() { let _serial = serial_test(); - let mut session = Session::new_local().unwrap(); - let mock = session.mock_async(fetch(0)).unwrap(); - mock.expect().once().returns(100).unwrap(); + let mut session = Session::new_local(); + let mock = session.mock_async(fetch(0)); + mock.expect().once().returns(100); let (ready_tx, ready_rx) = mpsc::channel(); let (run_tx, run_rx) = mpsc::channel(); let worker = std::thread::spawn(move || { - let mut session = Session::new_local().unwrap(); - let mock = session.mock_async(fetch(0)).unwrap(); - mock.expect().once().returns(200).unwrap(); + let mut session = Session::new_local(); + let mock = session.mock_async(fetch(0)); + mock.expect().once().returns(200); ready_tx.send(()).unwrap(); run_rx.recv().unwrap(); assert_eq!(ready(fetch(4)), 200); @@ -486,17 +496,15 @@ fn file_reads_are_mocked_only_on_the_current_thread() { .unwrap(); file.write_all(b"disk content").unwrap(); drop(file); - let mut session = Session::new_local().unwrap(); + let mut session = Session::new_local(); let mock = mock!( session, std::fs::read_to_string::<&std::path::Path>, fn(&std::path::Path) -> std::io::Result - ) - .unwrap(); + ); mock.expect() .once() - .return_once(Ok("mock content".to_owned())) - .unwrap(); + .return_once(Ok("mock content".to_owned())); assert_eq!( std::fs::read_to_string(path.as_path()).unwrap(), "mock content" @@ -508,6 +516,6 @@ fn file_reads_are_mocked_only_on_the_current_thread() { .unwrap(), "disk content" ); - session.restore().unwrap(); + session.restore(); std::fs::remove_file(path).unwrap(); }