use crate::ffi::errors::FFIError;
use std::cell::RefCell;
use std::env;
use std::io::{self, Read, Write};
use std::os::unix::net::{UnixStream};
use std::os::fd::{AsRawFd, FromRawFd, RawFd, OwnedFd};
use std::path::PathBuf;
use std::process::{Command, Child, Stdio};
use std::sync::{Mutex, MutexGuard, OnceLock};
use std::thread;
use bincode::config::Configuration;
use fxhash::FxHashMap;
use libloading::Library;
use serde::{Serialize, Deserialize};
use crate::ffi::value::{Type, Value};
use crate::worker::executeFFI;
pub(super) const ZygoteFlag: &str = "__zygote";
#[derive(Serialize, Deserialize)]
pub(super) enum FFIRequest
{
Call { libraryPath: String, functionName: String, args: Vec<Value>, resultType: Type },
Alloc { length: usize },
Free { pointer: usize },
ReadMemory { pointer: usize, length: usize },
WriteMemory { pointer: usize, value: Value }
}
#[derive(Serialize, Deserialize)]
pub(super) enum FFIResponse
{
Ok(Value),
Err(FFIError)
}
pub(super) struct ZygoteHandle
{
process: Child,
pub(super) socket: UnixStream
}
impl Drop for ZygoteHandle
{
fn drop(&mut self) -> ()
{
let _ = self.process.kill();
}
}
pub(super) static ZygoteState: OnceLock<Mutex<ZygoteHandle>> = OnceLock::new();
pub struct ClonedZygote
{
pub pid: libc::pid_t,
pub socket: UnixStream
}
impl ClonedZygote
{
pub fn getMeClone() -> io::Result<Self>
{
let mutex: &Mutex<ZygoteHandle> = ZygoteState.get()
.ok_or_else(|| io::Error::new(io::ErrorKind::NotFound, "Zygote not initialized"))?;
let mut guard: MutexGuard<ZygoteHandle> = mutex.lock()
.map_err(|_| io::Error::other("Zygote mutex poisoned"))?;
writeMessage(&mut guard.socket, &[1u8])?;
let mut pidBytes: [u8; 4] = [0u8; 4];
guard.socket.read_exact(&mut pidBytes)?;
let pid: i32 = i32::from_le_bytes(pidBytes);
let fd: RawFd = recvFd(&mut guard.socket)?;
let socket: UnixStream = unsafe { UnixStream::from_raw_fd(fd) };
drop(guard);
Ok(Self { pid, socket })
}
pub(super) fn call(&mut self, request: FFIRequest) -> Result<FFIResponse, String>
{
let bytes: Vec<u8> = encode(&request).map_err(|e| e.to_string())?;
let responseBytes: Vec<u8> = sendAndReceive(&mut self.socket, &bytes)
.map_err(|e| format!("Zygote clone IPC failed: {}", e))?;
decode(&responseBytes).map_err(|e| e.to_string())
}
}
impl Drop for ClonedZygote
{
fn drop(&mut self) -> ()
{
unsafe {
libc::kill(self.pid, libc::SIGKILL);
}
}
}
thread_local!{
pub(crate) static ZygoteStack: RefCell<Vec<ClonedZygote>> = const { RefCell::new(Vec::new()) };
}
pub struct ZygoteGuard;
impl ZygoteGuard
{
pub fn enter(zygote: ClonedZygote) -> Self
{
ZygoteStack.with(|stack| {
stack.borrow_mut().push(zygote);
});
Self
}
}
impl Drop for ZygoteGuard
{
fn drop(&mut self) -> ()
{
ZygoteStack.with(|stack| {
stack.borrow_mut().pop();
});
}
}
pub (super) fn runAsZygote() -> !
{
let socket: UnixStream = unsafe { UnixStream::from_raw_fd(libc::STDIN_FILENO) };
zygoteLoop(socket);
}
pub(super) fn initZygote() -> io::Result<()>
{
let handle: ZygoteHandle = spawnZygote()?;
ZygoteState.set(Mutex::new(handle))
.map_err(|_| io::Error::new(io::ErrorKind::AlreadyExists, "Zygote already initialized"))?;
thread::spawn(supervisorLoop);
Ok(())
}
pub(super) fn spawnZygote() -> io::Result<ZygoteHandle>
{
let (runtimeSocket, zygoteSocket): (UnixStream, UnixStream) = UnixStream::pair()?;
let currentExe: PathBuf = env::current_exe()?;
let process: Child = Command::new(currentExe)
.arg(ZygoteFlag)
.stdin(Stdio::from(OwnedFd::from(zygoteSocket)))
.spawn()?;
Ok(ZygoteHandle{ process, socket: runtimeSocket })
}
fn zygoteLoop(mut socket: UnixStream) -> !
{
unsafe { libc::signal(libc::SIGCHLD, libc::SIG_IGN); }
loop
{
if readMessage(&mut socket).is_err() {
std::process::exit(0); }
let (runtimeSocket, cloneSocket): (UnixStream, UnixStream) = match UnixStream::pair() {
Ok(pair) => pair,
Err(_) => {
let _ = writeMessage(&mut socket, &0i32.to_le_bytes());
continue;
}
};
match unsafe { libc::fork() }
{
-1 => {
let _ = writeMessage(&mut socket, &0i32.to_le_bytes());
}
0 => {
drop(runtimeSocket);
cloneLoop(cloneSocket);
}
pid => {
drop(cloneSocket);
if socket.write_all(&pid.to_le_bytes()).is_ok() {
let _ = sendFd(&mut socket, runtimeSocket.as_raw_fd());
}
}
}
}
}
fn cloneLoop(mut socket: UnixStream) -> !
{
let mut libraryCache: FxHashMap<String, Library> = FxHashMap::default();
loop
{
let requestBytes: Vec<u8> = match readMessage(&mut socket)
{
Ok(bytes) => bytes,
Err(_) =>
std::process::exit(0)
};
let response: FFIResponse = handleRequest(&requestBytes, &mut libraryCache);
let encodedResponse: Vec<u8> = match encode(&response) {
Ok(bytes) => bytes,
Err(_) => std::process::exit(1),
};
if writeMessage(&mut socket, &encodedResponse).is_err() {
std::process::exit(0);
}
}
}
fn handleRequest(requestBytes: &[u8], cache: &mut FxHashMap<String, Library>) -> FFIResponse
{
match decode::<FFIRequest>(requestBytes)
{
Ok(request) => match executeFFI(request, cache)
{
Ok(value) => FFIResponse::Ok(value),
Err(e) => FFIResponse::Err(e)
},
Err(e) => FFIResponse::Err(e)
}
}
pub(super) fn sendAndReceive(socket: &mut UnixStream, bytes: &[u8]) -> io::Result<Vec<u8>>
{
writeMessage(socket, bytes)?;
readMessage(socket)
}
fn supervisorLoop() -> ()
{
loop
{
let pidToWait: u32 = {
let mutex: &Mutex<ZygoteHandle> = match ZygoteState.get() { Some(m) => m, None => return };
mutex.lock().unwrap_or_else(|poisoned| poisoned.into_inner()).process.id()
};
unsafe { libc::waitpid(pidToWait as libc::pid_t, std::ptr::null_mut(), 0); }
let mutex: &Mutex<ZygoteHandle> = ZygoteState.get().unwrap();
let mut guard: MutexGuard<ZygoteHandle> = mutex.lock().unwrap_or_else(|poisoned| poisoned.into_inner());
if guard.process.id() == pidToWait {
match spawnZygote()
{
Ok(newHandle) => { *guard = newHandle; }
Err(_) => { drop(guard); thread::sleep(std::time::Duration::from_millis(200)); }
}
}
}
}
fn writeMessage(socket: &mut UnixStream, data: &[u8]) -> io::Result<()>
{
socket.write_all(&(data.len() as u32).to_le_bytes())?;
socket.write_all(data)
}
fn readMessage(socket: &mut UnixStream) -> io::Result<Vec<u8>>
{
let mut lengthBuffer: [u8; 4] = [0u8; 4];
socket.read_exact(&mut lengthBuffer)?;
let mut buffer: Vec<u8> = vec![0u8; u32::from_le_bytes(lengthBuffer) as usize];
socket.read_exact(&mut buffer)?;
Ok(buffer)
}
pub(super) fn encode<T: Serialize>(value: &T) -> Result<Vec<u8>, FFIError>
{
let config: Configuration = bincode::config::standard();
bincode::serde::encode_to_vec(value, config)
.map_err(|e| FFIError::EncodeFailed(format!("Encode failed: {}", e)))
}
pub(super) fn decode<T: for<'a> Deserialize<'a>>(bytes: &[u8]) -> Result<T, FFIError>
{
let config: Configuration = bincode::config::standard();
bincode::serde::decode_from_slice(bytes, config)
.map(|(decoded, _)| decoded)
.map_err(|e| FFIError::DecodeFailed(format!("Decode failed: {}", e)))
}
fn sendFd(socket: &mut UnixStream, fd: RawFd) -> io::Result<()>
{
let mut msgHeader: libc::msghdr = unsafe { std::mem::zeroed() };
let mut dummyByte: [u8; 1] = [0u8; 1];
let mut ioVector: libc::iovec = libc::iovec {
iov_base: dummyByte.as_mut_ptr() as *mut _,
iov_len: 1,
};
let cmsgSpace: u32 = unsafe { libc::CMSG_SPACE(std::mem::size_of::<RawFd>() as u32) };
let mut cmsgBuffer: Vec<u8> = vec![0u8; cmsgSpace as usize];
msgHeader.msg_iov = &mut ioVector;
msgHeader.msg_iovlen = 1;
msgHeader.msg_control = cmsgBuffer.as_mut_ptr() as *mut _;
msgHeader.msg_controllen = cmsgBuffer.len() as _;
unsafe {
let cmsg: *mut libc::cmsghdr = libc::CMSG_FIRSTHDR(&msgHeader);
(*cmsg).cmsg_level = libc::SOL_SOCKET;
(*cmsg).cmsg_type = libc::SCM_RIGHTS;
(*cmsg).cmsg_len = libc::CMSG_LEN(std::mem::size_of::<RawFd>() as u32) as _;
let fdPtr: *mut RawFd = libc::CMSG_DATA(cmsg) as *mut RawFd;
fdPtr.write_unaligned(fd);
}
let result: libc::ssize_t = unsafe { libc::sendmsg(socket.as_raw_fd(), &msgHeader, 0) };
if result < 0 { Err(io::Error::last_os_error()) } else { Ok(()) }
}
fn recvFd(socket: &mut UnixStream) -> io::Result<RawFd>
{
let mut msgHeader: libc::msghdr = unsafe { std::mem::zeroed() };
let mut dummyByte: [u8; 1] = [0u8; 1];
let mut ioVector: libc::iovec = libc::iovec {
iov_base: dummyByte.as_mut_ptr() as *mut _,
iov_len: 1,
};
let cmsgSpace: u32 = unsafe { libc::CMSG_SPACE(std::mem::size_of::<RawFd>() as u32) };
let mut cmsgBuffer: Vec<u8> = vec![0u8; cmsgSpace as usize];
msgHeader.msg_iov = &mut ioVector;
msgHeader.msg_iovlen = 1;
msgHeader.msg_control = cmsgBuffer.as_mut_ptr() as *mut _;
msgHeader.msg_controllen = cmsgBuffer.len() as _;
let result: libc::ssize_t = unsafe { libc::recvmsg(socket.as_raw_fd(), &mut msgHeader as *mut _, 0) };
if result <= 0 { return Err(io::Error::last_os_error()); }
unsafe {
let cmsg: *mut libc::cmsghdr = libc::CMSG_FIRSTHDR(&msgHeader);
if cmsg.is_null() || (*cmsg).cmsg_type != libc::SCM_RIGHTS {
return Err(io::Error::new(io::ErrorKind::InvalidData, "No FD received"));
}
let fdPtr: *const RawFd = libc::CMSG_DATA(cmsg) as *const RawFd;
Ok(fdPtr.read_unaligned())
}
}