use crate::{
Result,
connection::socket::{self, Socket},
};
use std::os::{
fd::{AsFd, BorrowedFd},
unix::net::UnixStream as StdUnixStream,
};
use tokio::net::{UnixStream, unix};
use zlink_core::connection::socket::ReadResult;
pub type Connection = crate::Connection<Stream>;
pub async fn connect<P>(path: P) -> Result<Connection>
where
P: AsRef<std::path::Path>,
{
UnixStream::connect(path)
.await
.map_err(Into::into)
.and_then(TryInto::try_into)
.map(Connection::new)
}
#[derive(Debug)]
pub struct Stream(UnixStream);
impl Socket for Stream {
type ReadHalf = ReadHalf;
type WriteHalf = WriteHalf;
const CAN_TRANSFER_FDS: bool = true;
fn split(self) -> (Self::ReadHalf, Self::WriteHalf) {
let (read, write) = self.0.into_split();
(ReadHalf(read), WriteHalf(write))
}
}
impl TryFrom<UnixStream> for Stream {
type Error = crate::Error;
fn try_from(stream: UnixStream) -> Result<Self> {
#[cfg(target_os = "linux")]
zlink_core::unix_utils::enable_passcred(&stream)?;
Ok(Self(stream))
}
}
impl TryFrom<StdUnixStream> for Stream {
type Error = crate::Error;
fn try_from(stream: StdUnixStream) -> Result<Self> {
stream.set_nonblocking(true)?;
UnixStream::from_std(stream)
.map_err(Into::into)
.and_then(TryInto::try_into)
}
}
impl socket::UnixSocket for Stream {}
impl AsFd for Stream {
fn as_fd(&self) -> BorrowedFd<'_> {
self.0.as_fd()
}
}
#[derive(Debug)]
pub struct ReadHalf(unix::OwnedReadHalf);
impl socket::ReadHalf for ReadHalf {
async fn read(&mut self, buf: &mut [u8]) -> Result<ReadResult> {
use std::{future::poll_fn, task::Poll};
poll_fn(|cx| {
loop {
let stream: &UnixStream = self.0.as_ref();
match stream.try_io(tokio::io::Interest::READABLE, || {
crate::unix_utils::recvmsg(stream, buf)
}) {
Err(ref e) if e.kind() == std::io::ErrorKind::WouldBlock => {
match stream.poll_read_ready(cx) {
Poll::Pending => return Poll::Pending,
Poll::Ready(res) => res?,
}
}
v => return Poll::Ready(v.map_err(Into::into)),
}
}
})
.await
}
}
impl AsFd for ReadHalf {
fn as_fd(&self) -> BorrowedFd<'_> {
let stream: &UnixStream = self.0.as_ref();
stream.as_fd()
}
}
impl socket::UnixSocket for ReadHalf {}
#[derive(Debug)]
pub struct WriteHalf(unix::OwnedWriteHalf);
impl socket::WriteHalf for WriteHalf {
async fn write(
&mut self,
buf: &[u8],
fds: &[impl AsFd],
#[cfg(target_os = "linux")] creds: Option<&crate::connection::PassedCredentials>,
) -> Result<()> {
use std::{future::poll_fn, task::Poll};
let borrowed_fds: Vec<BorrowedFd<'_>> = fds.iter().map(|f| f.as_fd()).collect();
let mut pos = 0;
while pos < buf.len() {
let fds_to_send = if pos == 0 { &borrowed_fds[..] } else { &[] };
let n: usize = poll_fn(|cx| {
loop {
let stream: &UnixStream = self.0.as_ref();
match stream.try_io(tokio::io::Interest::WRITABLE, || {
crate::unix_utils::sendmsg(
stream,
&buf[pos..],
fds_to_send,
#[cfg(target_os = "linux")]
creds,
)
}) {
Ok(bytes_sent) => return Poll::Ready(Ok::<_, crate::Error>(bytes_sent)),
Err(ref e) if e.kind() == std::io::ErrorKind::WouldBlock => {
match stream.poll_write_ready(cx) {
Poll::Pending => return Poll::Pending,
Poll::Ready(res) => res?,
}
}
Err(e) => return Poll::Ready(Err(e.into())),
}
}
})
.await?;
pos += n;
}
Ok(())
}
}
impl AsFd for WriteHalf {
fn as_fd(&self) -> BorrowedFd<'_> {
let stream: &UnixStream = self.0.as_ref();
stream.as_fd()
}
}
impl socket::UnixSocket for WriteHalf {}
#[cfg(test)]
mod tests {
use super::*;
use std::{
io::Write,
os::{
fd::{FromRawFd, IntoRawFd},
unix::net::UnixStream as StdUnixStream,
},
};
#[tokio::test]
async fn fd_passing_across_separate_connections() {
let (std_a, std_b) = StdUnixStream::pair().unwrap();
std_a.set_nonblocking(true).unwrap();
std_b.set_nonblocking(true).unwrap();
let conn_a = Connection::new(UnixStream::from_std(std_a).unwrap().try_into().unwrap());
let conn_b = Connection::new(UnixStream::from_std(std_b).unwrap().try_into().unwrap());
let (_, mut write_a) = conn_a.split();
let (mut read_b, _) = conn_b.split();
for name in ["alpha", "beta", "gamma"] {
let (r, mut w) = StdUnixStream::pair().unwrap();
w.write_all(name.as_bytes()).unwrap();
drop(w);
let reply = crate::Reply::new(Some(name.to_string())).set_continues(Some(false));
write_a.send_reply(&reply, vec![r.into()]).await.unwrap();
}
for name in ["alpha", "beta", "gamma"] {
let (reply, fds) = read_b.receive_reply::<String, ()>().await.unwrap();
let params = reply.unwrap().into_parameters().unwrap();
assert_eq!(params, name);
assert_eq!(fds.len(), 1);
let recv_fd = fds.into_iter().next().unwrap();
let mut stream = unsafe { StdUnixStream::from_raw_fd(recv_fd.into_raw_fd()) };
let mut buf = String::new();
std::io::Read::read_to_string(&mut stream, &mut buf).unwrap();
assert_eq!(buf, name, "FD data mismatch for {name:?}");
}
}
}