use serde::{Deserialize, Serialize};
use std::hash::{Hash, Hasher};
use fxhash::FxHasher;
#[doc(hidden)]
pub mod __reexport
{
pub use bincode;
pub use serde;
}
pub trait Callable<Args, Output>: Send
{
fn call(&self, args: Args) -> Output;
}
#[derive(Debug)]
pub enum CallError
{
ArgsOutputMismatch,
TypeMismatch { tag: u64 },
Encode(String),
Decode(String)
}
impl std::fmt::Display for CallError
{
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { write!(f, "{:?}", self) }
}
impl std::error::Error for CallError {}
#[doc(hidden)]
pub fn moduleBase() -> usize
{
let mut info: libc::Dl_info = unsafe{ std::mem::zeroed() };
unsafe{ libc::dladdr(moduleBase as *const () as *const libc::c_void, &mut info) };
info.dli_fbase as usize
}
#[doc(hidden)]
pub fn relativeOffsetOf(absoluteAddr: usize) -> usize
{
absoluteAddr.wrapping_sub(moduleBase())
}
fn resolveRelative(offset: usize) -> usize
{
moduleBase().wrapping_add(offset)
}
#[doc(hidden)]
pub fn tagOf(sourceLocation: &str) -> u64
{
let mut hasher: FxHasher = fxhash::FxHasher::default();
sourceLocation.hash(&mut hasher);
hasher.finish()
}
fn argsOutputTagOf<Args: 'static, Output: 'static>() -> u64
{
let mut hasher: FxHasher = fxhash::FxHasher::default();
std::any::type_name::<Args>().hash(&mut hasher);
std::any::type_name::<Output>().hash(&mut hasher);
hasher.finish()
}
#[derive(Serialize, Deserialize)]
#[doc(hidden)]
pub(super) struct Envelope
{
pub relativeOffset: usize,
pub argsOutputTag: u64,
pub siteTag: u64,
pub bytes: Vec<u8>
}
pub struct Sendable<Args, Output, T: Callable<Args, Output> + Serialize>
{
relativeOffset: usize,
siteTag: u64,
value: T,
_marker: std::marker::PhantomData<(Args, Output)>
}
impl<Args: 'static, Output: 'static, T: Callable<Args, Output> + Serialize> Sendable<Args, Output, T>
{
#[doc(hidden)]
pub const fn new(relativeOffset: usize, siteTag: u64, value: T) -> Self
{
Self { relativeOffset, siteTag, value, _marker: std::marker::PhantomData }
}
pub fn call(&self, args: Args) -> Output
{
self.value.call(args)
}
pub fn encode(&self) -> Result<Vec<u8>, CallError>
{
let bytes: Vec<u8> = bincode::serde::encode_to_vec(&self.value, bincode::config::standard())
.map_err(|e| CallError::Encode(e.to_string()))?;
let envelope: Envelope = Envelope {
relativeOffset: self.relativeOffset,
argsOutputTag: argsOutputTagOf::<Args, Output>(),
siteTag: self.siteTag,
bytes
};
bincode::serde::encode_to_vec(&envelope, bincode::config::standard())
.map_err(|e| CallError::Encode(e.to_string()))
}
}
pub fn decode<Args: 'static, Output: 'static>(bytes: &[u8]) -> Result<Box<dyn Callable<Args, Output>>, CallError>
{
let (envelope, _): (Envelope, usize) = bincode::serde::decode_from_slice(bytes, bincode::config::standard())
.map_err(|e| CallError::Decode(e.to_string()))?;
if envelope.argsOutputTag != argsOutputTagOf::<Args, Output>() {
return Err(CallError::ArgsOutputMismatch);
}
type DecodeFn<Args, Output> = fn(u64, &[u8]) -> Result<Box<dyn Callable<Args, Output>>, CallError>;
let absoluteAddr: usize = resolveRelative(envelope.relativeOffset);
let decodeFn: DecodeFn<Args, Output> = unsafe{ std::mem::transmute(absoluteAddr) };
decodeFn(envelope.siteTag, &envelope.bytes)
}
#[macro_export]
macro_rules! callback
{
([$($name:ident : $ty:ty),* $(,)?] |$arg:ident : $argTy:ty| -> $retTy:ty $body:block) =>
{
{
#[derive($crate::ffi::callback::__reexport::serde::Serialize, $crate::ffi::callback::__reexport::serde::Deserialize)]
struct __CallImpl { $( $name: $ty, )* }
impl $crate::ffi::callback::Callable<$argTy, $retTy> for __CallImpl
{
fn call(&self, $arg: $argTy) -> $retTy
{
$( let $name: $ty = self.$name.clone(); )*
$body
}
}
fn __callDecode(siteTag: u64, bytes: &[u8]) -> ::std::result::Result<
::std::boxed::Box<dyn $crate::ffi::callback::Callable<$argTy, $retTy>>,
$crate::ffi::callback::CallError
>
{
let expected: u64 = $crate::ffi::callback::tagOf(
::std::concat!(::std::file!(), ":", ::std::line!(), ":", ::std::column!())
);
if siteTag != expected {
return ::std::result::Result::Err(
$crate::ffi::callback::CallError::TypeMismatch { tag: siteTag }
);
}
let (concrete, _): (__CallImpl, usize) =
$crate::ffi::callback::__reexport::bincode::serde::decode_from_slice(
bytes, $crate::ffi::callback::__reexport::bincode::config::standard()
).map_err(|e| $crate::ffi::callback::CallError::Decode(
::std::string::ToString::to_string(&e)
))?;
::std::result::Result::Ok(::std::boxed::Box::new(concrete))
}
let siteTag: u64 = $crate::ffi::callback::tagOf(
::std::concat!(::std::file!(), ":", ::std::line!(), ":", ::std::column!())
);
let relativeOffset: usize = $crate::ffi::callback::relativeOffsetOf(__callDecode as usize);
$crate::ffi::callback::Sendable::new(relativeOffset, siteTag, __CallImpl { $( $name, )* })
}
};
}
#[cfg(test)]
mod tests
{
use crate::ffi::callback::Envelope;
use crate::ffi::callback::CallError;
use crate::ffi::callback::decode;
use crate::ffi::callback::Callable;
#[test]
fn roundtrip() -> ()
{
let threshold: i32 = 5;
let compar = callback!([threshold: i32] |args: Vec<i32>| -> i32 {
args.iter().filter(|&&x| x > threshold).count() as i32
});
assert_eq!(compar.call(vec![1, 6, 9, 2]), 2);
let bytes: Vec<u8> = compar.encode().expect("encode");
let remote: Box<dyn Callable<Vec<i32>, i32>> = decode(&bytes).expect("decode");
assert_eq!(remote.call(vec![1, 6, 9, 2]), 2);
}
#[test]
fn argsOutputMismatchIsCaught() -> ()
{
let x: i32 = 1;
let c = callback!([x: i32] |args: Vec<i32>| -> i32 { args.len() as i32 + x });
let bytes: Vec<u8> = c.encode().expect("encode");
let wrong: Result<Box<dyn Callable<String, bool>>, CallError> = decode(&bytes);
assert!(matches!(wrong, Err(CallError::ArgsOutputMismatch)));
}
#[test]
fn siteTagMismatchIsCaught() -> ()
{
let x: i32 = 1;
let c = callback!([x: i32] |args: Vec<i32>| -> i32 { args.len() as i32 + x });
let mut bytes: Vec<u8> = c.encode().expect("encode");
let (mut envelope, _): (Envelope, usize) = bincode::serde::decode_from_slice(&bytes, bincode::config::standard()).unwrap();
envelope.siteTag = envelope.siteTag.wrapping_add(1);
bytes = bincode::serde::encode_to_vec(&envelope, bincode::config::standard()).unwrap();
let result: Result<Box<dyn Callable<Vec<i32>, i32>>, CallError> = decode(&bytes);
assert!(matches!(result, Err(CallError::TypeMismatch { .. })));
}
}