Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
18 changes: 10 additions & 8 deletions src/injector_core/internal.rs
Original file line number Diff line number Diff line change
Expand Up @@ -4,10 +4,10 @@ use super::patch_trait::PatchTrait;

#[cfg(target_arch = "x86_64")]
use super::patch_amd64::PatchAmd64;
#[cfg(target_arch = "aarch64")]
use super::patch_arm64::PatchArm64;
#[cfg(target_arch = "arm")]
use super::patch_arm::PatchArm;
#[cfg(target_arch = "aarch64")]
use super::patch_arm64::PatchArm64;

#[cfg(any(target_arch = "x86_64", target_arch = "aarch64", target_arch = "arm"))]
use super::thread_local_registry;
Expand Down Expand Up @@ -49,10 +49,7 @@ impl WhenCalled {
/// The original function is patched to a dispatcher that routes calls
/// to per-thread replacement functions.
#[cfg(any(target_arch = "x86_64", target_arch = "aarch64", target_arch = "arm"))]
pub(crate) fn will_execute_thread_local(
self,
target: FuncPtrInternal,
) -> ThreadRegistration {
pub(crate) fn will_execute_thread_local(self, target: FuncPtrInternal) -> ThreadRegistration {
let replacement_addr = target.as_ptr() as usize;
thread_local_registry::register_replacement(&self.func_ptr, replacement_addr, None)
}
Expand All @@ -64,8 +61,13 @@ impl WhenCalled {
#[cfg(target_arch = "x86_64")]
let (jit_size, asm_code_vec) = {
let code: [u8; 8] = [
0x48, 0xC7, 0xC0, // mov rax, imm32
value as u8, 0x00, 0x00, 0x00, // imm32
0x48,
0xC7,
0xC0, // mov rax, imm32
value as u8,
0x00,
0x00,
0x00, // imm32
0xC3, // ret
];
(8usize, code.to_vec())
Expand Down
166 changes: 93 additions & 73 deletions src/injector_core/thread_local_registry.rs

Large diffs are not rendered by default.

63 changes: 36 additions & 27 deletions src/interface/injector.rs
Original file line number Diff line number Diff line change
Expand Up @@ -347,9 +347,10 @@ impl InjectorPP {
F: Future<Output = T>,
{
let poll_fn: fn(Pin<&mut F>, &mut Context<'_>) -> Poll<T> = <F as Future>::poll;
let when = WhenCalled::new(unsafe {
FuncPtr::new(poll_fn as *const (), std::any::type_name_of_val(&poll_fn))
}.func_ptr_internal);
let when = WhenCalled::new(
unsafe { FuncPtr::new(poll_fn as *const (), std::any::type_name_of_val(&poll_fn)) }
.func_ptr_internal,
);

let signature = fake_pair.1;
WhenCalledBuilderAsync {
Expand Down Expand Up @@ -407,9 +408,10 @@ impl InjectorPP {
F: Future<Output = T>,
{
let poll_fn: fn(Pin<&mut F>, &mut Context<'_>) -> Poll<T> = <F as Future>::poll;
let when = WhenCalled::new(unsafe {
FuncPtr::new(poll_fn as *const (), std::any::type_name_of_val(&poll_fn))
}.func_ptr_internal);
let when = WhenCalled::new(
unsafe { FuncPtr::new(poll_fn as *const (), std::any::type_name_of_val(&poll_fn)) }
.func_ptr_internal,
);

WhenCalledBuilderAsync {
lib: self,
Expand Down Expand Up @@ -506,16 +508,16 @@ impl WhenCalledBuilder<'_> {
self.expected_signature, target.signature
);
}
(None, _) | (_, None) => {
(None, _) | (_, None)
if normalize_signature(target.signature)
!= normalize_signature(self.expected_signature)
{
panic!(
"Signature mismatch: expected {:?} but got {:?}",
self.expected_signature, target.signature
);
}
!= normalize_signature(self.expected_signature) =>
{
panic!(
"Signature mismatch: expected {:?} but got {:?}",
self.expected_signature, target.signature
);
}

_ => {}
}

Expand All @@ -525,7 +527,9 @@ impl WhenCalledBuilder<'_> {
} else {
#[cfg(any(target_arch = "x86_64", target_arch = "aarch64", target_arch = "arm"))]
{
let reg = self.when.will_execute_thread_local(target.func_ptr_internal);
let reg = self
.when
.will_execute_thread_local(target.func_ptr_internal);
self.lib.registrations.push(reg);
}

Expand Down Expand Up @@ -597,7 +601,9 @@ impl WhenCalledBuilder<'_> {
} else {
#[cfg(any(target_arch = "x86_64", target_arch = "aarch64", target_arch = "arm"))]
{
let reg = self.when.will_execute_thread_local(target.func_ptr_internal);
let reg = self
.when
.will_execute_thread_local(target.func_ptr_internal);
self.lib.registrations.push(reg);
}

Expand Down Expand Up @@ -737,16 +743,16 @@ impl WhenCalledBuilderAsync<'_> {
self.expected_signature, target.signature
);
}
(None, _) | (_, None) => {
(None, _) | (_, None)
if normalize_signature(target.signature)
!= normalize_signature(self.expected_signature)
{
panic!(
"Signature mismatch: expected {:?} but got {:?}",
self.expected_signature, target.signature
);
}
!= normalize_signature(self.expected_signature) =>
{
panic!(
"Signature mismatch: expected {:?} but got {:?}",
self.expected_signature, target.signature
);
}

_ => {}
}

Expand All @@ -756,7 +762,9 @@ impl WhenCalledBuilderAsync<'_> {
} else {
#[cfg(any(target_arch = "x86_64", target_arch = "aarch64", target_arch = "arm"))]
{
let reg = self.when.will_execute_thread_local(target.func_ptr_internal);
let reg = self
.when
.will_execute_thread_local(target.func_ptr_internal);
self.lib.registrations.push(reg);
}

Expand Down Expand Up @@ -806,7 +814,9 @@ impl WhenCalledBuilderAsync<'_> {
} else {
#[cfg(any(target_arch = "x86_64", target_arch = "aarch64", target_arch = "arm"))]
{
let reg = self.when.will_execute_thread_local(target.func_ptr_internal);
let reg = self
.when
.will_execute_thread_local(target.func_ptr_internal);
self.lib.registrations.push(reg);
}

Expand All @@ -818,4 +828,3 @@ impl WhenCalledBuilderAsync<'_> {
}
}
}

10 changes: 8 additions & 2 deletions tests/global.rs
Original file line number Diff line number Diff line change
Expand Up @@ -86,7 +86,10 @@ fn test_global_fake_closure_cross_thread() {
let mut injector = InjectorPP::new_global();
injector
.when_called(injectorpp::func!(fn (global_multiply)(i32, i32) -> i32))
.will_execute_raw(injectorpp::closure!(|_a: i32, _b: i32| -> i32 { 777 }, fn(i32, i32) -> i32));
.will_execute_raw(injectorpp::closure!(
|_a: i32, _b: i32| -> i32 { 777 },
fn(i32, i32) -> i32
));

assert_eq!(global_multiply(3, 4), 777);

Expand Down Expand Up @@ -193,7 +196,10 @@ fn test_thread_local_mode_not_visible_from_spawned_thread() {
let mut injector = InjectorPP::new();
injector
.when_called(injectorpp::func!(fn (global_add)(i32, i32) -> i32))
.will_execute_raw(injectorpp::closure!(|_a: i32, _b: i32| -> i32 { 9999 }, fn(i32, i32) -> i32));
.will_execute_raw(injectorpp::closure!(
|_a: i32, _b: i32| -> i32 { 9999 },
fn(i32, i32) -> i32
));

// Test thread sees the fake
assert_eq!(global_add(1, 2), 9999);
Expand Down
3 changes: 1 addition & 2 deletions tests/hyper.rs
Original file line number Diff line number Diff line change
Expand Up @@ -36,8 +36,7 @@ fn make_tcp_with_http_response() -> std::io::Result<TcpStream> {
}
}

let body =
r#"{"status": "ok", "message": "mock response", "headers": {"User-Agent": "hyper-test/1.0"}}"#;
let body = r#"{"status": "ok", "message": "mock response", "headers": {"User-Agent": "hyper-test/1.0"}}"#;
let response = format!(
"HTTP/1.1 200 OK\r\n\
Content-Type: application/json\r\n\
Expand Down
53 changes: 41 additions & 12 deletions tests/lifetime_safety.rs
Original file line number Diff line number Diff line change
Expand Up @@ -14,18 +14,38 @@ use injectorpp::interface::injector::*;
/// Helper: tries to compile a source file and returns whether it succeeded.
/// Returns None if build artifacts can't be found (e.g. cross-compilation).
fn try_compile(source_path: &str) -> Option<bool> {
let rlib = find_file(&["target/debug/deps", "target/debug"], "libinjectorpp", ".rlib")?;
let ext = if cfg!(windows) { ".dll" } else if cfg!(target_os = "macos") { ".dylib" } else { ".so" };
let proc_dylib = find_file(&["target/debug/deps", "target/debug"], "injectorpp_macros", ext)?;
let rlib = find_file(
&["target/debug/deps", "target/debug"],
"libinjectorpp",
".rlib",
)?;
let ext = if cfg!(windows) {
".dll"
} else if cfg!(target_os = "macos") {
".dylib"
} else {
".so"
};
let proc_dylib = find_file(
&["target/debug/deps", "target/debug"],
"injectorpp_macros",
ext,
)?;

let output = std::process::Command::new("rustc")
.args([
"--edition", "2021",
"--crate-type", "bin",
"-L", "target/debug/deps",
"--extern", &format!("injectorpp={}", rlib),
"--extern", &format!("injectorpp_macros={}", proc_dylib),
"-o", if cfg!(windows) { "NUL" } else { "/dev/null" },
"--edition",
"2021",
"--crate-type",
"bin",
"-L",
"target/debug/deps",
"--extern",
&format!("injectorpp={}", rlib),
"--extern",
&format!("injectorpp_macros={}", proc_dylib),
"-o",
if cfg!(windows) { "NUL" } else { "/dev/null" },
source_path,
])
.output()
Expand All @@ -50,23 +70,32 @@ fn find_file(dirs: &[&str], prefix: &str, suffix: &str) -> Option<String> {
#[test]
fn static_str_coerced_to_bare_ref_must_not_compile() {
match try_compile("tests/compile_fail/static_str_coerced_to_bare_ref.rs") {
Some(compiled) => assert!(!compiled, "expected compile error: &'static str coerced to bare &str should be rejected"),
Some(compiled) => assert!(
!compiled,
"expected compile error: &'static str coerced to bare &str should be rejected"
),
None => eprintln!("skipped: build artifacts not found"),
}
}

#[test]
fn static_slice_coerced_to_bare_ref_must_not_compile() {
match try_compile("tests/compile_fail/static_slice_coerced_to_bare_ref.rs") {
Some(compiled) => assert!(!compiled, "expected compile error: &'static [u8] coerced to bare &[u8] should be rejected"),
Some(compiled) => assert!(
!compiled,
"expected compile error: &'static [u8] coerced to bare &[u8] should be rejected"
),
None => eprintln!("skipped: build artifacts not found"),
}
}

#[test]
fn func_info_prefix_lifetime_mismatch_must_not_compile() {
match try_compile("tests/compile_fail/func_info_prefix_lifetime_mismatch.rs") {
Some(compiled) => assert!(!compiled, "expected compile error: lifetime mismatch with func_info: prefix should be rejected"),
Some(compiled) => assert!(
!compiled,
"expected compile error: lifetime mismatch with func_info: prefix should be rejected"
),
None => eprintln!("skipped: build artifacts not found"),
}
}
Expand Down
10 changes: 8 additions & 2 deletions tests/thread_safety.rs
Original file line number Diff line number Diff line change
Expand Up @@ -859,7 +859,10 @@ fn test_string_return_thread_isolation() {
let mut injector = InjectorPP::new();
injector
.when_called(injectorpp::func!(fn(get_greeting)() -> String))
.will_execute_raw(injectorpp::closure!(|| { "from_thread_1".to_string() }, fn() -> String));
.will_execute_raw(injectorpp::closure!(
|| { "from_thread_1".to_string() },
fn() -> String
));
b1.wait();
if get_greeting() != "from_thread_1" {
e1.fetch_add(1, Ordering::SeqCst);
Expand All @@ -873,7 +876,10 @@ fn test_string_return_thread_isolation() {
let mut injector = InjectorPP::new();
injector
.when_called(injectorpp::func!(fn(get_greeting)() -> String))
.will_execute_raw(injectorpp::closure!(|| { "from_thread_2".to_string() }, fn() -> String));
.will_execute_raw(injectorpp::closure!(
|| { "from_thread_2".to_string() },
fn() -> String
));
b2.wait();
if get_greeting() != "from_thread_2" {
e2.fetch_add(1, Ordering::SeqCst);
Expand Down
3 changes: 2 additions & 1 deletion tests/will_execute.rs
Original file line number Diff line number Diff line change
Expand Up @@ -160,7 +160,8 @@ fn test_will_execute_when_fake_no_return_function_over_called_should_panic() {

let message = result.unwrap_err();
let message_str = message
.downcast_ref::<&str>().copied()
.downcast_ref::<&str>()
.copied()
.or_else(|| message.downcast_ref::<String>().map(|s| s.as_str()))
.unwrap();

Expand Down
Loading